Adding tests that don't require downloading models + conversation can be

fully created from static state.
This commit is contained in:
Nicolas Patry
2020-12-28 14:00:36 +01:00
parent ba0252b7f9
commit 1b9edfee3e
4 changed files with 116 additions and 27 deletions
+8 -1
View File
@@ -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
+73 -8
View File
@@ -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"
+1 -1
View File
@@ -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):