Cleaning up conversation tests.

This commit is contained in:
Nicolas Patry
2020-12-25 14:20:35 +01:00
parent 41e55cb75a
commit ba0252b7f9
4 changed files with 124 additions and 37 deletions
+63
View File
@@ -18,9 +18,14 @@ from unittest import mock
from transformers import is_tf_available, is_torch_available, pipeline
from transformers.pipelines import Pipeline
from transformers.testing_utils import _run_slow_tests, is_pipeline_test, require_tf, require_torch, slow
from transformers.tokenization_utils import TruncationStrategy
from transformers.tokenization_utils_base import to_py_obj
if is_torch_available():
import torch
VALID_INPUTS = ["A simple string", ["list of strings"]]
@@ -242,3 +247,61 @@ class MonoInputPipelineCommonMixin(CustomInputPipelineCommonMixin):
self.assertIn(key, result)
self.assertRaises(Exception, nlp, self.invalid_inputs)
class DummyTok:
pad_token_id = 0
eos_token_id = 0
def __init__(self, **kwargs):
for name, v in kwargs.items():
setattr(self, name, v)
def __call__(self, inputs, **kwargs):
if kwargs.get("return_tensors", "") == "pt":
return self.encode_pt(inputs, **kwargs)
else:
return self.encode_list(inputs, **kwargs)
def encode_list(self, inputs, **kwargs):
unwrap = False
if isinstance(inputs, str):
unwrap = True
inputs = [inputs]
assert isinstance(inputs, list)
input_ids = [self.encode(input_) for input_ in inputs]
if unwrap:
input_ids = input_ids[0]
return {"input_ids": input_ids}
def encode_pt(self, inputs, **kwargs):
if isinstance(inputs, str):
input_ids = torch.LongTensor(self.encode(inputs)).unsqueeze(0)
else:
input_ids = self._pad([self.encode(input_) for input_ in inputs])
return self.finalize_pt(input_ids, **kwargs)
def finalize_pt(self, input_ids, **kwargs):
if kwargs.get("truncation", TruncationStrategy.DO_NOT_TRUNCATE) == TruncationStrategy.ONLY_FIRST:
input_ids = input_ids[:, : self.model_max_length]
attention_mask = torch.zeros_like(input_ids).long() + 1
return {"input_ids": input_ids, "attention_mask": attention_mask}
def _pad(self, inputs):
return torch.nn.utils.rnn.pad_sequence(
[torch.LongTensor(input_) for input_ in inputs],
padding_value=self.pad_token_id,
).transpose(1, 0)
def pad(self, inputs, **kwargs):
input_ids = self._pad(inputs["input_ids"])
return self.finalize_pt(input_ids, **kwargs)
def encode(self, input_):
return list(input_.encode("utf-8"))
def decode(self, sequence, **kwargs):
return "D" * len(sequence)
+50 -1
View File
@@ -15,14 +15,63 @@
import unittest
from transformers import AutoModelForSeq2SeqLM, AutoTokenizer, Conversation, ConversationalPipeline, pipeline
from transformers.models.gpt2 import GPT2Config, GPT2LMHeadModel
from transformers.testing_utils import require_torch, slow, torch_device
from .test_pipelines_common import MonoInputPipelineCommonMixin
from .test_pipelines_common import DummyTok, MonoInputPipelineCommonMixin
DEFAULT_DEVICE_NUM = -1 if torch_device == "cpu" else 0
class SimpleConversationPipelineTests(unittest.TestCase):
@require_torch
def test_integration_torch_conversation(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
)
model = GPT2LMHeadModel(config)
tokenizer = DummyTok()
nlp = pipeline(task="conversational", device=DEFAULT_DEVICE_NUM, model=model, tokenizer=tokenizer)
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)
# Then
self.assertEqual(result, [conversation_1, conversation_2])
self.assertEqual(
result,
[
Conversation(
None,
past_user_inputs=["Going to the movies tonight - any suggestions?"],
generated_responses=["D"],
),
Conversation(
None, past_user_inputs=["What's the last book you have read?"], generated_responses=["D"]
),
],
)
# When
conversation_2.add_user_input("Why do you recommend it?")
result = nlp(conversation_2, do_sample=False, max_length=1000)
# Then
self.assertEqual(result, conversation_2)
self.assertEqual(
result,
Conversation(
None,
past_user_inputs=["What's the last book you have read?", "Why do you recommend it?"],
generated_responses=["D", "D"],
),
)
class ConversationalPipelineTests(MonoInputPipelineCommonMixin, unittest.TestCase):
pipeline_task = "conversational"
small_models = [] # Models tested without the @slow decorator
+2 -33
View File
@@ -14,47 +14,16 @@
import unittest
from transformers import is_torch_available, pipeline
from transformers import pipeline
from transformers.models.bart import BartConfig, BartForConditionalGeneration
from transformers.testing_utils import require_torch, slow, torch_device
from transformers.tokenization_utils import TruncationStrategy
from .test_pipelines_common import MonoInputPipelineCommonMixin
from .test_pipelines_common import DummyTok, MonoInputPipelineCommonMixin
DEFAULT_DEVICE_NUM = -1 if torch_device == "cpu" else 0
if is_torch_available():
import torch
class DummyTok:
pad_token_id = 0
def __init__(self, **kwargs):
for name, v in kwargs.items():
setattr(self, name, v)
def __call__(self, inputs, **kwargs):
if isinstance(inputs, str):
input_ids = self.encode(inputs).unsqueeze(0)
else:
input_ids = torch.nn.utils.rnn.pad_sequence(
[self.encode(input_) for input_ in inputs],
padding_value=self.pad_token_id,
)
if kwargs.get("truncation", TruncationStrategy.DO_NOT_TRUNCATE) == TruncationStrategy.ONLY_FIRST:
input_ids = input_ids[:, : self.model_max_length]
attention_mask = torch.zeros_like(input_ids).long() + 1
return {"input_ids": input_ids, "attention_mask": attention_mask}
def encode(self, input_):
return torch.LongTensor(list(input_.encode("utf-8")))
def decode(self, sequence, **kwargs):
return "D" * len(sequence)
class SimpleSummarizationPipelineTests(unittest.TestCase):
@require_torch