Compare commits
1
Commits
| Author | SHA1 | Date | |
|---|---|---|---|
|
|
16a10cb6c2 |
@@ -3,7 +3,7 @@ from typing import List, Optional
|
|||||||
|
|
||||||
from transformers import is_tf_available, is_torch_available, pipeline
|
from transformers import is_tf_available, is_torch_available, pipeline
|
||||||
from transformers.pipelines import DefaultArgumentHandler, Pipeline
|
from transformers.pipelines import DefaultArgumentHandler, Pipeline
|
||||||
from transformers.testing_utils import _run_slow_tests, is_pipeline_test, require_tf, require_torch, slow
|
from transformers.testing_utils import _run_slow_tests, is_pipeline_test, require_tf, require_tokenizers, require_torch, slow
|
||||||
|
|
||||||
|
|
||||||
VALID_INPUTS = ["A simple string", ["list of strings"]]
|
VALID_INPUTS = ["A simple string", ["list of strings"]]
|
||||||
@@ -79,9 +79,26 @@ class CustomInputPipelineCommonMixin:
|
|||||||
nlp = pipeline(task=self.pipeline_task, model=model_name, tokenizer=model_name, framework="tf")
|
nlp = pipeline(task=self.pipeline_task, model=model_name, tokenizer=model_name, framework="tf")
|
||||||
self._test_pipeline(nlp)
|
self._test_pipeline(nlp)
|
||||||
|
|
||||||
|
@require_torch
|
||||||
|
def test_fast_tokenizer_equivalence_torch_small(self):
|
||||||
|
for model_name in self.small_models:
|
||||||
|
fast_nlp = pipeline(task=self.pipeline_task, model=model_name, tokenizer=model_name, framework="pt", use_fast=True)
|
||||||
|
slow_nlp = pipeline(task=self.pipeline_task, model=model_name, tokenizer=model_name, framework="pt", use_fast=False)
|
||||||
|
self._test_equivalence(fast_nlp, slow_nlp)
|
||||||
|
|
||||||
|
@require_tf
|
||||||
|
def test_fast_tokenizer_equivalence_tf_small(self):
|
||||||
|
for model_name in self.small_models:
|
||||||
|
fast_nlp = pipeline(task=self.pipeline_task, model=model_name, tokenizer=model_name, framework="tf", use_fast=True)
|
||||||
|
slow_nlp = pipeline(task=self.pipeline_task, model=model_name, tokenizer=model_name, framework="tf", use_fast=False)
|
||||||
|
self._test_equivalence(fast_nlp, slow_nlp)
|
||||||
|
|
||||||
def _test_pipeline(self, nlp: Pipeline):
|
def _test_pipeline(self, nlp: Pipeline):
|
||||||
raise NotImplementedError
|
raise NotImplementedError
|
||||||
|
|
||||||
|
def _test_equivalence(self, fast_nlp: Pipeline, slow_nlp: Pipeline):
|
||||||
|
raise NotImplementedError
|
||||||
|
|
||||||
|
|
||||||
@is_pipeline_test
|
@is_pipeline_test
|
||||||
class MonoInputPipelineCommonMixin:
|
class MonoInputPipelineCommonMixin:
|
||||||
@@ -165,6 +182,45 @@ class MonoInputPipelineCommonMixin:
|
|||||||
)
|
)
|
||||||
self._test_pipeline(nlp)
|
self._test_pipeline(nlp)
|
||||||
|
|
||||||
|
@require_tokenizers
|
||||||
|
@require_torch
|
||||||
|
def test_fast_tokenizer_equivalence_torch_small(self):
|
||||||
|
for model_name in self.small_models:
|
||||||
|
fast_nlp = pipeline(task=self.pipeline_task, model=model_name, tokenizer=model_name, framework="pt", use_fast=True, **self.pipeline_loading_kwargs)
|
||||||
|
slow_nlp = pipeline(task=self.pipeline_task, model=model_name, tokenizer=model_name, framework="pt", use_fast=False, **self.pipeline_loading_kwargs)
|
||||||
|
self._test_equivalence(fast_nlp, slow_nlp)
|
||||||
|
|
||||||
|
@require_tokenizers
|
||||||
|
@require_tf
|
||||||
|
def test_fast_tokenizer_equivalence_tf_small(self):
|
||||||
|
for model_name in self.small_models:
|
||||||
|
fast_nlp = pipeline(task=self.pipeline_task, model=model_name, tokenizer=model_name, framework="tf", use_fast=True)
|
||||||
|
slow_nlp = pipeline(task=self.pipeline_task, model=model_name, tokenizer=model_name, framework="tf", use_fast=False)
|
||||||
|
self._test_equivalence(fast_nlp, slow_nlp)
|
||||||
|
|
||||||
|
def _test_equivalence(self, fast_nlp: Pipeline, slow_nlp: Pipeline):
|
||||||
|
mono_result_fast = fast_nlp(self.valid_inputs[0], **self.pipeline_running_kwargs)
|
||||||
|
mono_result_slow = slow_nlp(self.valid_inputs[0], **self.pipeline_running_kwargs)
|
||||||
|
|
||||||
|
if isinstance(mono_result_fast[0], list):
|
||||||
|
mono_result_fast = mono_result_fast[0]
|
||||||
|
mono_result_slow = mono_result_slow[0]
|
||||||
|
|
||||||
|
for key in mono_result_fast[0]:
|
||||||
|
self.assertEqual(mono_result_fast[0][key], mono_result_slow[0][key])
|
||||||
|
|
||||||
|
multi_result_fast = [fast_nlp(input, **self.pipeline_running_kwargs) for input in self.valid_inputs]
|
||||||
|
multi_result_slow = [slow_nlp(input, **self.pipeline_running_kwargs) for input in self.valid_inputs]
|
||||||
|
|
||||||
|
if self.expected_multi_result is not None:
|
||||||
|
for result_fast, result_slow, expect in zip(multi_result_fast, multi_result_slow, self.expected_multi_result):
|
||||||
|
for key in result_fast:
|
||||||
|
self.assertEqual(
|
||||||
|
set([o[key] for o in result_fast]),
|
||||||
|
set([o[key] for o in result_slow]),
|
||||||
|
set([o[key] for o in expect]),
|
||||||
|
)
|
||||||
|
|
||||||
def _test_pipeline(self, nlp: Pipeline):
|
def _test_pipeline(self, nlp: Pipeline):
|
||||||
self.assertIsNotNone(nlp)
|
self.assertIsNotNone(nlp)
|
||||||
|
|
||||||
|
|||||||
Reference in New Issue
Block a user