Adding tests that don't require downloading models + conversation can be
fully created from static state.
This commit is contained in:
@@ -12,6 +12,7 @@
|
||||
# See the License for the specific language governing permissions and
|
||||
# limitations under the License.
|
||||
|
||||
from string import ascii_lowercase
|
||||
from typing import List, Optional
|
||||
from unittest import mock
|
||||
|
||||
@@ -256,6 +257,7 @@ class DummyTok:
|
||||
def __init__(self, **kwargs):
|
||||
for name, v in kwargs.items():
|
||||
setattr(self, name, v)
|
||||
self.index = 0
|
||||
|
||||
def __call__(self, inputs, **kwargs):
|
||||
if kwargs.get("return_tensors", "") == "pt":
|
||||
@@ -304,4 +306,9 @@ class DummyTok:
|
||||
return list(input_.encode("utf-8"))
|
||||
|
||||
def decode(self, sequence, **kwargs):
|
||||
return "D" * len(sequence)
|
||||
string = ""
|
||||
for i in range(len(sequence)):
|
||||
string += ascii_lowercase[self.index]
|
||||
self.index += 1
|
||||
self.index %= len(ascii_lowercase)
|
||||
return string
|
||||
|
||||
@@ -25,22 +25,34 @@ DEFAULT_DEVICE_NUM = -1 if torch_device == "cpu" else 0
|
||||
|
||||
|
||||
class SimpleConversationPipelineTests(unittest.TestCase):
|
||||
@require_torch
|
||||
def test_integration_torch_conversation(self):
|
||||
def get_pipeline(self):
|
||||
# When
|
||||
config = GPT2Config(
|
||||
vocab_size=257, n_ctx=64, n_embd=64, n_layer=1, n_head=8, eos_token_id=0, bos_token_id=0, pad_token_id=0
|
||||
vocab_size=257,
|
||||
n_ctx=64,
|
||||
max_length=64,
|
||||
n_embd=64,
|
||||
n_layer=1,
|
||||
n_head=8,
|
||||
eos_token_id=0,
|
||||
bos_token_id=0,
|
||||
pad_token_id=0,
|
||||
)
|
||||
model = GPT2LMHeadModel(config)
|
||||
tokenizer = DummyTok()
|
||||
nlp = pipeline(task="conversational", device=DEFAULT_DEVICE_NUM, model=model, tokenizer=tokenizer)
|
||||
return nlp
|
||||
|
||||
@require_torch
|
||||
def test_integration_torch_conversation(self):
|
||||
nlp = self.get_pipeline()
|
||||
conversation_1 = Conversation("Going to the movies tonight - any suggestions?")
|
||||
conversation_2 = Conversation("What's the last book you have read?")
|
||||
# Then
|
||||
self.assertEqual(len(conversation_1.past_user_inputs), 0)
|
||||
self.assertEqual(len(conversation_2.past_user_inputs), 0)
|
||||
# When
|
||||
result = nlp([conversation_1, conversation_2], do_sample=False, max_length=1000)
|
||||
result = nlp([conversation_1, conversation_2], do_sample=False)
|
||||
# Then
|
||||
self.assertEqual(result, [conversation_1, conversation_2])
|
||||
self.assertEqual(
|
||||
@@ -49,17 +61,17 @@ class SimpleConversationPipelineTests(unittest.TestCase):
|
||||
Conversation(
|
||||
None,
|
||||
past_user_inputs=["Going to the movies tonight - any suggestions?"],
|
||||
generated_responses=["D"],
|
||||
generated_responses=["a"],
|
||||
),
|
||||
Conversation(
|
||||
None, past_user_inputs=["What's the last book you have read?"], generated_responses=["D"]
|
||||
None, past_user_inputs=["What's the last book you have read?"], generated_responses=["b"]
|
||||
),
|
||||
],
|
||||
)
|
||||
|
||||
# When
|
||||
conversation_2.add_user_input("Why do you recommend it?")
|
||||
result = nlp(conversation_2, do_sample=False, max_length=1000)
|
||||
result = nlp(conversation_2, do_sample=False)
|
||||
# Then
|
||||
self.assertEqual(result, conversation_2)
|
||||
self.assertEqual(
|
||||
@@ -67,10 +79,63 @@ class SimpleConversationPipelineTests(unittest.TestCase):
|
||||
Conversation(
|
||||
None,
|
||||
past_user_inputs=["What's the last book you have read?", "Why do you recommend it?"],
|
||||
generated_responses=["D", "D"],
|
||||
generated_responses=["b", "c"],
|
||||
),
|
||||
)
|
||||
|
||||
def test_history_cache(self):
|
||||
nlp = self.get_pipeline()
|
||||
conversation = Conversation(
|
||||
"Why do you recommend it?",
|
||||
past_user_inputs=["What's the last book you have read?"],
|
||||
generated_responses=["b"],
|
||||
)
|
||||
_ = nlp(conversation)
|
||||
self.assertEquals(conversation._index, 1)
|
||||
self.assertEquals(
|
||||
conversation._history,
|
||||
[
|
||||
87,
|
||||
104,
|
||||
97,
|
||||
116,
|
||||
39,
|
||||
115,
|
||||
32,
|
||||
116,
|
||||
104,
|
||||
101,
|
||||
32,
|
||||
108,
|
||||
97,
|
||||
115,
|
||||
116,
|
||||
32,
|
||||
98,
|
||||
111,
|
||||
111,
|
||||
107,
|
||||
32,
|
||||
121,
|
||||
111,
|
||||
117,
|
||||
32,
|
||||
104,
|
||||
97,
|
||||
118,
|
||||
101,
|
||||
32,
|
||||
114,
|
||||
101,
|
||||
97,
|
||||
100,
|
||||
63,
|
||||
0,
|
||||
98,
|
||||
0,
|
||||
],
|
||||
)
|
||||
|
||||
|
||||
class ConversationalPipelineTests(MonoInputPipelineCommonMixin, unittest.TestCase):
|
||||
pipeline_task = "conversational"
|
||||
|
||||
@@ -51,7 +51,7 @@ class SimpleSummarizationPipelineTests(unittest.TestCase):
|
||||
|
||||
output = nlp("This is a test", truncation=TruncationStrategy.ONLY_FIRST)
|
||||
|
||||
self.assertEqual(output, [{"summary_text": "DDDD"}])
|
||||
self.assertEqual(output, [{"summary_text": "abcd"}])
|
||||
|
||||
|
||||
class SummarizationPipelineTests(MonoInputPipelineCommonMixin, unittest.TestCase):
|
||||
|
||||
Reference in New Issue
Block a user