Cleaning up conversation tests.
This commit is contained in:
@@ -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)
|
||||
|
||||
@@ -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
|
||||
|
||||
@@ -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
|
||||
|
||||
Reference in New Issue
Block a user