Compare commits
7
Commits
| Author | SHA1 | Date | |
|---|---|---|---|
|
|
68dc37d599 | ||
|
|
6175976506 | ||
|
|
e9343aa36c | ||
|
|
59182f841c | ||
|
|
775eebb39d | ||
|
|
329ec23b1c | ||
|
|
fdc819729d |
Executable
+428
@@ -0,0 +1,428 @@
|
||||
# flake8: noqa
|
||||
# There's no way to ignore "F401 '...' imported but unused" warnings in this
|
||||
# module, but to preserve other warnings. So, don't check this module at all.
|
||||
|
||||
# coding=utf-8
|
||||
# Copyright 2018 The HuggingFace Inc. team.
|
||||
#
|
||||
# Licensed under the Apache License, Version 2.0 (the "License");
|
||||
# you may not use this file except in compliance with the License.
|
||||
# You may obtain a copy of the License at
|
||||
#
|
||||
# http://www.apache.org/licenses/LICENSE-2.0
|
||||
#
|
||||
# Unless required by applicable law or agreed to in writing, software
|
||||
# distributed under the License is distributed on an "AS IS" BASIS,
|
||||
# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
|
||||
# See the License for the specific language governing permissions and
|
||||
# limitations under the License.
|
||||
import warnings
|
||||
from typing import TYPE_CHECKING, Any, Dict, Optional, Tuple, Union
|
||||
|
||||
from transformers.configuration_utils import PretrainedConfig
|
||||
from transformers.file_utils import is_tf_available, is_torch_available
|
||||
from transformers.modelcard import ModelCard
|
||||
from transformers.models.auto.tokenization_auto import AutoTokenizer
|
||||
from transformers.pipelines.base import (
|
||||
ArgumentHandler,
|
||||
Conversation,
|
||||
ConversationalPipeline,
|
||||
CsvPipelineDataFormat,
|
||||
JsonPipelineDataFormat,
|
||||
NerPipeline,
|
||||
PipedPipelineDataFormat,
|
||||
Pipeline,
|
||||
PipelineDataFormat,
|
||||
PipelineException,
|
||||
QuestionAnsweringArgumentHandler,
|
||||
QuestionAnsweringPipeline,
|
||||
SummarizationPipeline,
|
||||
TableQuestionAnsweringArgumentHandler,
|
||||
TableQuestionAnsweringPipeline,
|
||||
Text2TextGenerationPipeline,
|
||||
TokenClassificationArgumentHandler,
|
||||
TokenClassificationPipeline,
|
||||
TranslationPipeline,
|
||||
get_default_model,
|
||||
get_framework,
|
||||
)
|
||||
from transformers.pipelines.feature_extraction import FeatureExtractionPipeline
|
||||
from transformers.pipelines.fill_mask import FillMaskPipeline
|
||||
from transformers.pipelines.text_classification import TextClassificationPipeline
|
||||
from transformers.pipelines.text_generation import TextGenerationPipeline
|
||||
from transformers.pipelines.zero_shot_classification import (
|
||||
ZeroShotClassificationArgumentHandler,
|
||||
ZeroShotClassificationPipeline,
|
||||
)
|
||||
from transformers.tokenization_utils import PreTrainedTokenizer
|
||||
from transformers.utils import logging
|
||||
|
||||
|
||||
if is_tf_available():
|
||||
import tensorflow as tf
|
||||
|
||||
from transformers.models.auto.modeling_tf_auto import (
|
||||
TF_MODEL_FOR_QUESTION_ANSWERING_MAPPING,
|
||||
TF_MODEL_FOR_SEQ_TO_SEQ_CAUSAL_LM_MAPPING,
|
||||
TF_MODEL_FOR_SEQUENCE_CLASSIFICATION_MAPPING,
|
||||
TF_MODEL_FOR_TOKEN_CLASSIFICATION_MAPPING,
|
||||
TF_MODEL_WITH_LM_HEAD_MAPPING,
|
||||
TFAutoModel,
|
||||
TFAutoModelForCausalLM,
|
||||
TFAutoModelForMaskedLM,
|
||||
TFAutoModelForQuestionAnswering,
|
||||
TFAutoModelForSeq2SeqLM,
|
||||
TFAutoModelForSequenceClassification,
|
||||
TFAutoModelForTokenClassification,
|
||||
)
|
||||
|
||||
if is_torch_available():
|
||||
import torch
|
||||
|
||||
from transformers.models.auto.modeling_auto import (
|
||||
MODEL_FOR_MASKED_LM_MAPPING,
|
||||
MODEL_FOR_QUESTION_ANSWERING_MAPPING,
|
||||
MODEL_FOR_SEQ_TO_SEQ_CAUSAL_LM_MAPPING,
|
||||
MODEL_FOR_SEQUENCE_CLASSIFICATION_MAPPING,
|
||||
MODEL_FOR_TABLE_QUESTION_ANSWERING_MAPPING,
|
||||
MODEL_FOR_TOKEN_CLASSIFICATION_MAPPING,
|
||||
AutoModel,
|
||||
AutoModelForCausalLM,
|
||||
AutoModelForMaskedLM,
|
||||
AutoModelForQuestionAnswering,
|
||||
AutoModelForSeq2SeqLM,
|
||||
AutoModelForSequenceClassification,
|
||||
AutoModelForTableQuestionAnswering,
|
||||
AutoModelForTokenClassification,
|
||||
)
|
||||
if TYPE_CHECKING:
|
||||
from transformers.modeling_tf_utils import TFPreTrainedModel
|
||||
from transformers.modeling_utils import PreTrainedModel
|
||||
|
||||
logger = logging.get_logger(__name__)
|
||||
|
||||
|
||||
# Register all the supported tasks here
|
||||
SUPPORTED_TASKS = {
|
||||
"feature-extraction": {
|
||||
"impl": FeatureExtractionPipeline,
|
||||
"tf": TFAutoModel if is_tf_available() else None,
|
||||
"pt": AutoModel if is_torch_available() else None,
|
||||
"default": {"model": {"pt": "distilbert-base-cased", "tf": "distilbert-base-cased"}},
|
||||
},
|
||||
"sentiment-analysis": {
|
||||
"impl": TextClassificationPipeline,
|
||||
"tf": TFAutoModelForSequenceClassification if is_tf_available() else None,
|
||||
"pt": AutoModelForSequenceClassification if is_torch_available() else None,
|
||||
"default": {
|
||||
"model": {
|
||||
"pt": "distilbert-base-uncased-finetuned-sst-2-english",
|
||||
"tf": "distilbert-base-uncased-finetuned-sst-2-english",
|
||||
},
|
||||
},
|
||||
},
|
||||
"ner": {
|
||||
"impl": TokenClassificationPipeline,
|
||||
"tf": TFAutoModelForTokenClassification if is_tf_available() else None,
|
||||
"pt": AutoModelForTokenClassification if is_torch_available() else None,
|
||||
"default": {
|
||||
"model": {
|
||||
"pt": "dbmdz/bert-large-cased-finetuned-conll03-english",
|
||||
"tf": "dbmdz/bert-large-cased-finetuned-conll03-english",
|
||||
},
|
||||
},
|
||||
},
|
||||
"question-answering": {
|
||||
"impl": QuestionAnsweringPipeline,
|
||||
"tf": TFAutoModelForQuestionAnswering if is_tf_available() else None,
|
||||
"pt": AutoModelForQuestionAnswering if is_torch_available() else None,
|
||||
"default": {
|
||||
"model": {"pt": "distilbert-base-cased-distilled-squad", "tf": "distilbert-base-cased-distilled-squad"},
|
||||
},
|
||||
},
|
||||
"table-question-answering": {
|
||||
"impl": TableQuestionAnsweringPipeline,
|
||||
"pt": AutoModelForTableQuestionAnswering if is_torch_available() else None,
|
||||
"tf": None,
|
||||
"default": {
|
||||
"model": {
|
||||
"pt": "nielsr/tapas-base-finetuned-wtq",
|
||||
"tokenizer": "nielsr/tapas-base-finetuned-wtq",
|
||||
"tf": "nielsr/tapas-base-finetuned-wtq",
|
||||
},
|
||||
},
|
||||
},
|
||||
"fill-mask": {
|
||||
"impl": FillMaskPipeline,
|
||||
"tf": TFAutoModelForMaskedLM if is_tf_available() else None,
|
||||
"pt": AutoModelForMaskedLM if is_torch_available() else None,
|
||||
"default": {"model": {"pt": "distilroberta-base", "tf": "distilroberta-base"}},
|
||||
},
|
||||
"summarization": {
|
||||
"impl": SummarizationPipeline,
|
||||
"tf": TFAutoModelForSeq2SeqLM if is_tf_available() else None,
|
||||
"pt": AutoModelForSeq2SeqLM if is_torch_available() else None,
|
||||
"default": {"model": {"pt": "sshleifer/distilbart-cnn-12-6", "tf": "t5-small"}},
|
||||
},
|
||||
# This task is a special case as it's parametrized by SRC, TGT languages.
|
||||
"translation": {
|
||||
"impl": TranslationPipeline,
|
||||
"tf": TFAutoModelForSeq2SeqLM if is_tf_available() else None,
|
||||
"pt": AutoModelForSeq2SeqLM if is_torch_available() else None,
|
||||
"default": {
|
||||
("en", "fr"): {"model": {"pt": "t5-base", "tf": "t5-base"}},
|
||||
("en", "de"): {"model": {"pt": "t5-base", "tf": "t5-base"}},
|
||||
("en", "ro"): {"model": {"pt": "t5-base", "tf": "t5-base"}},
|
||||
},
|
||||
},
|
||||
"text2text-generation": {
|
||||
"impl": Text2TextGenerationPipeline,
|
||||
"tf": TFAutoModelForSeq2SeqLM if is_tf_available() else None,
|
||||
"pt": AutoModelForSeq2SeqLM if is_torch_available() else None,
|
||||
"default": {"model": {"pt": "t5-base", "tf": "t5-base"}},
|
||||
},
|
||||
"text-generation": {
|
||||
"impl": TextGenerationPipeline,
|
||||
"tf": TFAutoModelForCausalLM if is_tf_available() else None,
|
||||
"pt": AutoModelForCausalLM if is_torch_available() else None,
|
||||
"default": {"model": {"pt": "gpt2", "tf": "gpt2"}},
|
||||
},
|
||||
"zero-shot-classification": {
|
||||
"impl": ZeroShotClassificationPipeline,
|
||||
"tf": TFAutoModelForSequenceClassification if is_tf_available() else None,
|
||||
"pt": AutoModelForSequenceClassification if is_torch_available() else None,
|
||||
"default": {
|
||||
"model": {"pt": "facebook/bart-large-mnli", "tf": "roberta-large-mnli"},
|
||||
"config": {"pt": "facebook/bart-large-mnli", "tf": "roberta-large-mnli"},
|
||||
"tokenizer": {"pt": "facebook/bart-large-mnli", "tf": "roberta-large-mnli"},
|
||||
},
|
||||
},
|
||||
"conversational": {
|
||||
"impl": ConversationalPipeline,
|
||||
"tf": TFAutoModelForCausalLM if is_tf_available() else None,
|
||||
"pt": AutoModelForCausalLM if is_torch_available() else None,
|
||||
"default": {"model": {"pt": "microsoft/DialoGPT-medium", "tf": "microsoft/DialoGPT-medium"}},
|
||||
},
|
||||
}
|
||||
|
||||
|
||||
def check_task(task: str) -> Tuple[Dict, Any]:
|
||||
"""
|
||||
Checks an incoming task string, to validate it's correct and return the default Pipeline and Model classes, and
|
||||
default models if they exist.
|
||||
|
||||
Args:
|
||||
task (:obj:`str`):
|
||||
The task defining which pipeline will be returned. Currently accepted tasks are:
|
||||
|
||||
- :obj:`"feature-extraction"`
|
||||
- :obj:`"sentiment-analysis"`
|
||||
- :obj:`"ner"`
|
||||
- :obj:`"question-answering"`
|
||||
- :obj:`"fill-mask"`
|
||||
- :obj:`"summarization"`
|
||||
- :obj:`"translation_xx_to_yy"`
|
||||
- :obj:`"translation"`
|
||||
- :obj:`"text-generation"`
|
||||
- :obj:`"conversational"`
|
||||
|
||||
Returns:
|
||||
(task_defaults:obj:`dict`, task_options: (:obj:`tuple`, None)) The actual dictionary required to initialize the
|
||||
pipeline and some extra task options for parametrized tasks like "translation_XX_to_YY"
|
||||
|
||||
|
||||
"""
|
||||
if task in SUPPORTED_TASKS:
|
||||
targeted_task = SUPPORTED_TASKS[task]
|
||||
return targeted_task, None
|
||||
|
||||
if task.startswith("translation"):
|
||||
tokens = task.split("_")
|
||||
if len(tokens) == 4 and tokens[0] == "translation" and tokens[2] == "to":
|
||||
targeted_task = SUPPORTED_TASKS["translation"]
|
||||
return targeted_task, (tokens[1], tokens[3])
|
||||
raise KeyError("Invalid translation task {}, use 'translation_XX_to_YY' format".format(task))
|
||||
|
||||
raise KeyError(
|
||||
"Unknown task {}, available tasks are {}".format(task, list(SUPPORTED_TASKS.keys()) + ["translation_XX_to_YY"])
|
||||
)
|
||||
|
||||
|
||||
def pipeline(
|
||||
task: str,
|
||||
model: Optional = None,
|
||||
config: Optional[Union[str, PretrainedConfig]] = None,
|
||||
tokenizer: Optional[Union[str, PreTrainedTokenizer]] = None,
|
||||
framework: Optional[str] = None,
|
||||
revision: Optional[str] = None,
|
||||
use_fast: bool = True,
|
||||
**kwargs
|
||||
) -> Pipeline:
|
||||
"""
|
||||
Utility factory method to build a :class:`~transformers.Pipeline`.
|
||||
|
||||
Pipelines are made of:
|
||||
|
||||
- A :doc:`tokenizer <tokenizer>` in charge of mapping raw textual input to token.
|
||||
- A :doc:`model <model>` to make predictions from the inputs.
|
||||
- Some (optional) post processing for enhancing model's output.
|
||||
|
||||
Args:
|
||||
task (:obj:`str`):
|
||||
The task defining which pipeline will be returned. Currently accepted tasks are:
|
||||
|
||||
- :obj:`"feature-extraction"`: will return a :class:`~transformers.FeatureExtractionPipeline`.
|
||||
- :obj:`"sentiment-analysis"`: will return a :class:`~transformers.TextClassificationPipeline`.
|
||||
- :obj:`"ner"`: will return a :class:`~transformers.TokenClassificationPipeline`.
|
||||
- :obj:`"question-answering"`: will return a :class:`~transformers.QuestionAnsweringPipeline`.
|
||||
- :obj:`"fill-mask"`: will return a :class:`~transformers.FillMaskPipeline`.
|
||||
- :obj:`"summarization"`: will return a :class:`~transformers.SummarizationPipeline`.
|
||||
- :obj:`"translation_xx_to_yy"`: will return a :class:`~transformers.TranslationPipeline`.
|
||||
- :obj:`"text2text-generation"`: will return a :class:`~transformers.Text2TextGenerationPipeline`.
|
||||
- :obj:`"text-generation"`: will return a :class:`~transformers.TextGenerationPipeline`.
|
||||
- :obj:`"zero-shot-classification:`: will return a :class:`~transformers.ZeroShotClassificationPipeline`.
|
||||
- :obj:`"conversation"`: will return a :class:`~transformers.ConversationalPipeline`.
|
||||
model (:obj:`str` or :obj:`~transformers.PreTrainedModel` or :obj:`~transformers.TFPreTrainedModel`, `optional`):
|
||||
The model that will be used by the pipeline to make predictions. This can be a model identifier or an
|
||||
actual instance of a pretrained model inheriting from :class:`~transformers.PreTrainedModel` (for PyTorch)
|
||||
or :class:`~transformers.TFPreTrainedModel` (for TensorFlow).
|
||||
|
||||
If not provided, the default for the :obj:`task` will be loaded.
|
||||
config (:obj:`str` or :obj:`~transformers.PretrainedConfig`, `optional`):
|
||||
The configuration that will be used by the pipeline to instantiate the model. This can be a model
|
||||
identifier or an actual pretrained model configuration inheriting from
|
||||
:class:`~transformers.PretrainedConfig`.
|
||||
|
||||
If not provided, the default configuration file for the requested model will be used. That means that if
|
||||
:obj:`model` is given, its default configuration will be used. However, if :obj:`model` is not supplied,
|
||||
this :obj:`task`'s default model's config is used instead.
|
||||
tokenizer (:obj:`str` or :obj:`~transformers.PreTrainedTokenizer`, `optional`):
|
||||
The tokenizer that will be used by the pipeline to encode data for the model. This can be a model
|
||||
identifier or an actual pretrained tokenizer inheriting from :class:`~transformers.PreTrainedTokenizer`.
|
||||
|
||||
If not provided, the default tokenizer for the given :obj:`model` will be loaded (if it is a string). If
|
||||
:obj:`model` is not specified or not a string, then the default tokenizer for :obj:`config` is loaded (if
|
||||
it is a string). However, if :obj:`config` is also not given or not a string, then the default tokenizer
|
||||
for the given :obj:`task` will be loaded.
|
||||
framework (:obj:`str`, `optional`):
|
||||
The framework to use, either :obj:`"pt"` for PyTorch or :obj:`"tf"` for TensorFlow. The specified framework
|
||||
must be installed.
|
||||
|
||||
If no framework is specified, will default to the one currently installed. If no framework is specified and
|
||||
both frameworks are installed, will default to the framework of the :obj:`model`, or to PyTorch if no model
|
||||
is provided.
|
||||
revision(:obj:`str`, `optional`, defaults to :obj:`"main"`):
|
||||
When passing a task name or a string model identifier: The specific model version to use. It can be a
|
||||
branch name, a tag name, or a commit id, since we use a git-based system for storing models and other
|
||||
artifacts on huggingface.co, so ``revision`` can be any identifier allowed by git.
|
||||
use_fast (:obj:`bool`, `optional`, defaults to :obj:`True`):
|
||||
Whether or not to use a Fast tokenizer if possible (a :class:`~transformers.PreTrainedTokenizerFast`).
|
||||
kwargs:
|
||||
Additional keyword arguments passed along to the specific pipeline init (see the documentation for the
|
||||
corresponding pipeline class for possible values).
|
||||
|
||||
Returns:
|
||||
:class:`~transformers.Pipeline`: A suitable pipeline for the task.
|
||||
|
||||
Examples::
|
||||
|
||||
>>> from transformers import pipeline, AutoModelForTokenClassification, AutoTokenizer
|
||||
|
||||
>>> # Sentiment analysis pipeline
|
||||
>>> pipeline('sentiment-analysis')
|
||||
|
||||
>>> # Question answering pipeline, specifying the checkpoint identifier
|
||||
>>> pipeline('question-answering', model='distilbert-base-cased-distilled-squad', tokenizer='bert-base-cased')
|
||||
|
||||
>>> # Named entity recognition pipeline, passing in a specific model and tokenizer
|
||||
>>> model = AutoModelForTokenClassification.from_pretrained("dbmdz/bert-large-cased-finetuned-conll03-english")
|
||||
>>> tokenizer = AutoTokenizer.from_pretrained("bert-base-cased")
|
||||
>>> pipeline('ner', model=model, tokenizer=tokenizer)
|
||||
"""
|
||||
# Retrieve the task
|
||||
targeted_task, task_options = check_task(task)
|
||||
|
||||
# Use default model/config/tokenizer for the task if no model is provided
|
||||
if model is None:
|
||||
# At that point framework might still be undetermined
|
||||
model = get_default_model(targeted_task, framework, task_options)
|
||||
|
||||
framework = framework or get_framework(model)
|
||||
|
||||
task_class, model_class = targeted_task["impl"], targeted_task[framework]
|
||||
|
||||
# Try to infer tokenizer from model or config name (if provided as str)
|
||||
if tokenizer is None:
|
||||
if isinstance(model, str):
|
||||
tokenizer = model
|
||||
elif isinstance(config, str):
|
||||
tokenizer = config
|
||||
else:
|
||||
# Impossible to guest what is the right tokenizer here
|
||||
raise Exception(
|
||||
"Impossible to guess which tokenizer to use. "
|
||||
"Please provided a PretrainedTokenizer class or a path/identifier to a pretrained tokenizer."
|
||||
)
|
||||
|
||||
modelcard = None
|
||||
# Try to infer modelcard from model or config name (if provided as str)
|
||||
if isinstance(model, str):
|
||||
modelcard = model
|
||||
elif isinstance(config, str):
|
||||
modelcard = config
|
||||
|
||||
# Instantiate tokenizer if needed
|
||||
if isinstance(tokenizer, (str, tuple)):
|
||||
if isinstance(tokenizer, tuple):
|
||||
# For tuple we have (tokenizer name, {kwargs})
|
||||
use_fast = tokenizer[1].pop("use_fast", use_fast)
|
||||
tokenizer = AutoTokenizer.from_pretrained(
|
||||
tokenizer[0], use_fast=use_fast, revision=revision, **tokenizer[1]
|
||||
)
|
||||
else:
|
||||
tokenizer = AutoTokenizer.from_pretrained(tokenizer, revision=revision, use_fast=use_fast)
|
||||
|
||||
# Instantiate config if needed
|
||||
if isinstance(config, str):
|
||||
config = AutoConfig.from_pretrained(config, revision=revision)
|
||||
|
||||
# Instantiate modelcard if needed
|
||||
if isinstance(modelcard, str):
|
||||
modelcard = ModelCard.from_pretrained(modelcard, revision=revision)
|
||||
|
||||
# Instantiate model if needed
|
||||
if isinstance(model, str):
|
||||
# Handle transparent TF/PT model conversion
|
||||
model_kwargs = {}
|
||||
if framework == "pt" and model.endswith(".h5"):
|
||||
model_kwargs["from_tf"] = True
|
||||
logger.warning(
|
||||
"Model might be a TensorFlow model (ending with `.h5`) but TensorFlow is not available. "
|
||||
"Trying to load the model with PyTorch."
|
||||
)
|
||||
elif framework == "tf" and model.endswith(".bin"):
|
||||
model_kwargs["from_pt"] = True
|
||||
logger.warning(
|
||||
"Model might be a PyTorch model (ending with `.bin`) but PyTorch is not available. "
|
||||
"Trying to load the model with Tensorflow."
|
||||
)
|
||||
|
||||
if model_class is None:
|
||||
raise ValueError(
|
||||
f"Pipeline using {framework} framework, but this framework is not supported by this pipeline."
|
||||
)
|
||||
|
||||
model = model_class.from_pretrained(model, config=config, revision=revision, **model_kwargs)
|
||||
if task == "translation" and model.config.task_specific_params:
|
||||
for key in model.config.task_specific_params:
|
||||
if key.startswith("translation"):
|
||||
task = key
|
||||
warnings.warn(
|
||||
'"translation" task was used, instead of "translation_XX_to_YY", defaulting to "{}"'.format(
|
||||
task
|
||||
),
|
||||
UserWarning,
|
||||
)
|
||||
break
|
||||
|
||||
return task_class(model=model, tokenizer=tokenizer, modelcard=modelcard, framework=framework, task=task, **kwargs)
|
||||
Executable → Regular
+11
-1011
File diff suppressed because it is too large
Load Diff
@@ -0,0 +1,83 @@
|
||||
from typing import TYPE_CHECKING, Optional, Union
|
||||
|
||||
from transformers.modelcard import ModelCard
|
||||
from transformers.tokenization_utils import PreTrainedTokenizer
|
||||
|
||||
from .base import ArgumentHandler, Pipeline
|
||||
|
||||
|
||||
if TYPE_CHECKING:
|
||||
from transformers.modeling_tf_utils import TFPreTrainedModel
|
||||
from transformers.modeling_utils import PreTrainedModel
|
||||
|
||||
|
||||
# Can't use @add_end_docstrings(PIPELINE_INIT_ARGS) here because this one does not accept `binary_output`
|
||||
class FeatureExtractionPipeline(Pipeline):
|
||||
"""
|
||||
Feature extraction pipeline using no model head. This pipeline extracts the hidden states from the base
|
||||
transformer, which can be used as features in downstream tasks.
|
||||
|
||||
This feature extraction pipeline can currently be loaded from :func:`~transformers.pipeline` using the task
|
||||
identifier: :obj:`"feature-extraction"`.
|
||||
|
||||
All models may be used for this pipeline. See a list of all models, including community-contributed models on
|
||||
`huggingface.co/models <https://huggingface.co/models>`__.
|
||||
|
||||
Arguments:
|
||||
model (:obj:`~transformers.PreTrainedModel` or :obj:`~transformers.TFPreTrainedModel`):
|
||||
The model that will be used by the pipeline to make predictions. This needs to be a model inheriting from
|
||||
:class:`~transformers.PreTrainedModel` for PyTorch and :class:`~transformers.TFPreTrainedModel` for
|
||||
TensorFlow.
|
||||
tokenizer (:obj:`~transformers.PreTrainedTokenizer`):
|
||||
The tokenizer that will be used by the pipeline to encode data for the model. This object inherits from
|
||||
:class:`~transformers.PreTrainedTokenizer`.
|
||||
modelcard (:obj:`str` or :class:`~transformers.ModelCard`, `optional`):
|
||||
Model card attributed to the model for this pipeline.
|
||||
framework (:obj:`str`, `optional`):
|
||||
The framework to use, either :obj:`"pt"` for PyTorch or :obj:`"tf"` for TensorFlow. The specified framework
|
||||
must be installed.
|
||||
|
||||
If no framework is specified, will default to the one currently installed. If no framework is specified and
|
||||
both frameworks are installed, will default to the framework of the :obj:`model`, or to PyTorch if no model
|
||||
is provided.
|
||||
task (:obj:`str`, defaults to :obj:`""`):
|
||||
A task-identifier for the pipeline.
|
||||
args_parser (:class:`~transformers.pipelines.ArgumentHandler`, `optional`):
|
||||
Reference to the object in charge of parsing supplied pipeline parameters.
|
||||
device (:obj:`int`, `optional`, defaults to -1):
|
||||
Device ordinal for CPU/GPU supports. Setting this to -1 will leverage CPU, a positive will run the model on
|
||||
the associated CUDA device id.
|
||||
"""
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
model: Union["PreTrainedModel", "TFPreTrainedModel"],
|
||||
tokenizer: PreTrainedTokenizer,
|
||||
modelcard: Optional[ModelCard] = None,
|
||||
framework: Optional[str] = None,
|
||||
args_parser: ArgumentHandler = None,
|
||||
device: int = -1,
|
||||
task: str = "",
|
||||
):
|
||||
super().__init__(
|
||||
model=model,
|
||||
tokenizer=tokenizer,
|
||||
modelcard=modelcard,
|
||||
framework=framework,
|
||||
args_parser=args_parser,
|
||||
device=device,
|
||||
binary_output=True,
|
||||
task=task,
|
||||
)
|
||||
|
||||
def __call__(self, *args, **kwargs):
|
||||
"""
|
||||
Extract the features of the input(s).
|
||||
|
||||
Args:
|
||||
args (:obj:`str` or :obj:`List[str]`): One or several texts (or one list of texts) to get the features of.
|
||||
|
||||
Return:
|
||||
A nested list of :obj:`float`: The features computed by the model.
|
||||
"""
|
||||
return super().__call__(*args, **kwargs).tolist()
|
||||
@@ -0,0 +1,195 @@
|
||||
from typing import TYPE_CHECKING, Optional, Union
|
||||
|
||||
import numpy as np
|
||||
|
||||
from transformers.file_utils import add_end_docstrings, is_tf_available, is_torch_available
|
||||
from transformers.modelcard import ModelCard
|
||||
from transformers.tokenization_utils import PreTrainedTokenizer
|
||||
from transformers.utils import logging
|
||||
|
||||
from .base import PIPELINE_INIT_ARGS, ArgumentHandler, Pipeline, PipelineException
|
||||
|
||||
|
||||
if TYPE_CHECKING:
|
||||
from transformers.modeling_tf_utils import TFPreTrainedModel
|
||||
from transformers.modeling_utils import PreTrainedModel
|
||||
|
||||
if is_tf_available():
|
||||
import tensorflow as tf
|
||||
|
||||
from transformers.models.auto.modeling_tf_auto import TF_MODEL_WITH_LM_HEAD_MAPPING
|
||||
|
||||
if is_torch_available():
|
||||
import torch
|
||||
|
||||
from transformers.models.auto.modeling_auto import MODEL_FOR_MASKED_LM_MAPPING
|
||||
|
||||
|
||||
logger = logging.get_logger(__name__)
|
||||
|
||||
|
||||
@add_end_docstrings(
|
||||
PIPELINE_INIT_ARGS,
|
||||
r"""
|
||||
top_k (:obj:`int`, defaults to 5): The number of predictions to return.
|
||||
""",
|
||||
)
|
||||
class FillMaskPipeline(Pipeline):
|
||||
"""
|
||||
Masked language modeling prediction pipeline using any :obj:`ModelWithLMHead`. See the `masked language modeling
|
||||
examples <../task_summary.html#masked-language-modeling>`__ for more information.
|
||||
|
||||
This mask filling pipeline can currently be loaded from :func:`~transformers.pipeline` using the following task
|
||||
identifier: :obj:`"fill-mask"`.
|
||||
|
||||
The models that this pipeline can use are models that have been trained with a masked language modeling objective,
|
||||
which includes the bi-directional models in the library. See the up-to-date list of available models on
|
||||
`huggingface.co/models <https://huggingface.co/models?filter=masked-lm>`__.
|
||||
|
||||
.. note::
|
||||
|
||||
This pipeline only works for inputs with exactly one token masked.
|
||||
"""
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
model: Union["PreTrainedModel", "TFPreTrainedModel"],
|
||||
tokenizer: PreTrainedTokenizer,
|
||||
modelcard: Optional[ModelCard] = None,
|
||||
framework: Optional[str] = None,
|
||||
args_parser: ArgumentHandler = None,
|
||||
device: int = -1,
|
||||
top_k=5,
|
||||
task: str = "",
|
||||
):
|
||||
super().__init__(
|
||||
model=model,
|
||||
tokenizer=tokenizer,
|
||||
modelcard=modelcard,
|
||||
framework=framework,
|
||||
args_parser=args_parser,
|
||||
device=device,
|
||||
binary_output=True,
|
||||
task=task,
|
||||
)
|
||||
|
||||
self.check_model_type(TF_MODEL_WITH_LM_HEAD_MAPPING if self.framework == "tf" else MODEL_FOR_MASKED_LM_MAPPING)
|
||||
self.top_k = top_k
|
||||
|
||||
def ensure_exactly_one_mask_token(self, masked_index: np.ndarray):
|
||||
numel = np.prod(masked_index.shape)
|
||||
if numel > 1:
|
||||
raise PipelineException(
|
||||
"fill-mask",
|
||||
self.model.base_model_prefix,
|
||||
f"More than one mask_token ({self.tokenizer.mask_token}) is not supported",
|
||||
)
|
||||
elif numel < 1:
|
||||
raise PipelineException(
|
||||
"fill-mask",
|
||||
self.model.base_model_prefix,
|
||||
f"No mask_token ({self.tokenizer.mask_token}) found on the input",
|
||||
)
|
||||
|
||||
def __call__(self, *args, targets=None, top_k: Optional[int] = None, **kwargs):
|
||||
"""
|
||||
Fill the masked token in the text(s) given as inputs.
|
||||
|
||||
Args:
|
||||
args (:obj:`str` or :obj:`List[str]`):
|
||||
One or several texts (or one list of prompts) with masked tokens.
|
||||
targets (:obj:`str` or :obj:`List[str]`, `optional`):
|
||||
When passed, the model will return the scores for the passed token or tokens rather than the top k
|
||||
predictions in the entire vocabulary. If the provided targets are not in the model vocab, they will be
|
||||
tokenized and the first resulting token will be used (with a warning).
|
||||
top_k (:obj:`int`, `optional`):
|
||||
When passed, overrides the number of predictions to return.
|
||||
|
||||
Return:
|
||||
A list or a list of list of :obj:`dict`: Each result comes as list of dictionaries with the following keys:
|
||||
|
||||
- **sequence** (:obj:`str`) -- The corresponding input with the mask token prediction.
|
||||
- **score** (:obj:`float`) -- The corresponding probability.
|
||||
- **token** (:obj:`int`) -- The predicted token id (to replace the masked one).
|
||||
- **token** (:obj:`str`) -- The predicted token (to replace the masked one).
|
||||
"""
|
||||
inputs = self._parse_and_tokenize(*args, **kwargs)
|
||||
outputs = self._forward(inputs, return_tensors=True)
|
||||
|
||||
results = []
|
||||
batch_size = outputs.shape[0] if self.framework == "tf" else outputs.size(0)
|
||||
|
||||
if targets is not None:
|
||||
if len(targets) == 0 or len(targets[0]) == 0:
|
||||
raise ValueError("At least one target must be provided when passed.")
|
||||
if isinstance(targets, str):
|
||||
targets = [targets]
|
||||
|
||||
targets_proc = []
|
||||
for target in targets:
|
||||
target_enc = self.tokenizer.tokenize(target)
|
||||
if len(target_enc) > 1 or target_enc[0] == self.tokenizer.unk_token:
|
||||
logger.warning(
|
||||
"The specified target token `{}` does not exist in the model vocabulary. Replacing with `{}`.".format(
|
||||
target, target_enc[0]
|
||||
)
|
||||
)
|
||||
targets_proc.append(target_enc[0])
|
||||
target_inds = np.array(self.tokenizer.convert_tokens_to_ids(targets_proc))
|
||||
|
||||
for i in range(batch_size):
|
||||
input_ids = inputs["input_ids"][i]
|
||||
result = []
|
||||
|
||||
if self.framework == "tf":
|
||||
masked_index = tf.where(input_ids == self.tokenizer.mask_token_id).numpy()
|
||||
|
||||
# Fill mask pipeline supports only one ${mask_token} per sample
|
||||
self.ensure_exactly_one_mask_token(masked_index)
|
||||
|
||||
logits = outputs[i, masked_index.item(), :]
|
||||
probs = tf.nn.softmax(logits)
|
||||
if targets is None:
|
||||
topk = tf.math.top_k(probs, k=top_k if top_k is not None else self.top_k)
|
||||
values, predictions = topk.values.numpy(), topk.indices.numpy()
|
||||
else:
|
||||
values = tf.gather_nd(probs, tf.reshape(target_inds, (-1, 1)))
|
||||
sort_inds = tf.reverse(tf.argsort(values), [0])
|
||||
values = tf.gather_nd(values, tf.reshape(sort_inds, (-1, 1))).numpy()
|
||||
predictions = target_inds[sort_inds.numpy()]
|
||||
else:
|
||||
masked_index = torch.nonzero(input_ids == self.tokenizer.mask_token_id, as_tuple=False)
|
||||
|
||||
# Fill mask pipeline supports only one ${mask_token} per sample
|
||||
self.ensure_exactly_one_mask_token(masked_index.numpy())
|
||||
|
||||
logits = outputs[i, masked_index.item(), :]
|
||||
probs = logits.softmax(dim=0)
|
||||
if targets is None:
|
||||
values, predictions = probs.topk(top_k if top_k is not None else self.top_k)
|
||||
else:
|
||||
values = probs[..., target_inds]
|
||||
sort_inds = list(reversed(values.argsort(dim=-1)))
|
||||
values = values[..., sort_inds]
|
||||
predictions = target_inds[sort_inds]
|
||||
|
||||
for v, p in zip(values.tolist(), predictions.tolist()):
|
||||
tokens = input_ids.numpy()
|
||||
tokens[masked_index] = p
|
||||
# Filter padding out:
|
||||
tokens = tokens[np.where(tokens != self.tokenizer.pad_token_id)]
|
||||
result.append(
|
||||
{
|
||||
"sequence": self.tokenizer.decode(tokens),
|
||||
"score": v,
|
||||
"token": p,
|
||||
"token_str": self.tokenizer.convert_ids_to_tokens(p),
|
||||
}
|
||||
)
|
||||
|
||||
# Append
|
||||
results += [result]
|
||||
|
||||
if len(results) == 1:
|
||||
return results[0]
|
||||
return results
|
||||
@@ -0,0 +1,80 @@
|
||||
import numpy as np
|
||||
|
||||
from transformers.file_utils import add_end_docstrings, is_tf_available, is_torch_available
|
||||
|
||||
from .base import PIPELINE_INIT_ARGS, Pipeline
|
||||
|
||||
|
||||
if is_tf_available():
|
||||
from transformers.models.auto.modeling_tf_auto import TF_MODEL_FOR_SEQUENCE_CLASSIFICATION_MAPPING
|
||||
|
||||
if is_torch_available():
|
||||
from transformers.models.auto.modeling_auto import MODEL_FOR_SEQUENCE_CLASSIFICATION_MAPPING
|
||||
|
||||
|
||||
@add_end_docstrings(
|
||||
PIPELINE_INIT_ARGS,
|
||||
r"""
|
||||
return_all_scores (:obj:`bool`, `optional`, defaults to :obj:`False`):
|
||||
Whether to return all prediction scores or just the one of the predicted class.
|
||||
""",
|
||||
)
|
||||
class TextClassificationPipeline(Pipeline):
|
||||
"""
|
||||
Text classification pipeline using any :obj:`ModelForSequenceClassification`. See the `sequence classification
|
||||
examples <../task_summary.html#sequence-classification>`__ for more information.
|
||||
|
||||
This text classification pipeline can currently be loaded from :func:`~transformers.pipeline` using the following
|
||||
task identifier: :obj:`"sentiment-analysis"` (for classifying sequences according to positive or negative
|
||||
sentiments).
|
||||
|
||||
If multiple classification labels are available (:obj:`model.config.num_labels >= 2`), the pipeline will run a
|
||||
softmax over the results. If there is a single label, the pipeline will run a sigmoid over the result.
|
||||
|
||||
The models that this pipeline can use are models that have been fine-tuned on a sequence classification task. See
|
||||
the up-to-date list of available models on `huggingface.co/models
|
||||
<https://huggingface.co/models?filter=text-classification>`__.
|
||||
"""
|
||||
|
||||
def __init__(self, return_all_scores: bool = False, **kwargs):
|
||||
super().__init__(**kwargs)
|
||||
|
||||
self.check_model_type(
|
||||
TF_MODEL_FOR_SEQUENCE_CLASSIFICATION_MAPPING
|
||||
if self.framework == "tf"
|
||||
else MODEL_FOR_SEQUENCE_CLASSIFICATION_MAPPING
|
||||
)
|
||||
|
||||
self.return_all_scores = return_all_scores
|
||||
|
||||
def __call__(self, *args, **kwargs):
|
||||
"""
|
||||
Classify the text(s) given as inputs.
|
||||
|
||||
Args:
|
||||
args (:obj:`str` or :obj:`List[str]`):
|
||||
One or several texts (or one list of prompts) to classify.
|
||||
|
||||
Return:
|
||||
A list or a list of list of :obj:`dict`: Each result comes as list of dictionaries with the following keys:
|
||||
|
||||
- **label** (:obj:`str`) -- The label predicted.
|
||||
- **score** (:obj:`float`) -- The corresponding probability.
|
||||
|
||||
If ``self.return_all_scores=True``, one such dictionary is returned per label.
|
||||
"""
|
||||
outputs = super().__call__(*args, **kwargs)
|
||||
|
||||
if self.model.config.num_labels == 1:
|
||||
scores = 1.0 / (1.0 + np.exp(-outputs))
|
||||
else:
|
||||
scores = np.exp(outputs) / np.exp(outputs).sum(-1, keepdims=True)
|
||||
if self.return_all_scores:
|
||||
return [
|
||||
[{"label": self.model.config.id2label[i], "score": score.item()} for i, score in enumerate(item)]
|
||||
for item in scores
|
||||
]
|
||||
else:
|
||||
return [
|
||||
{"label": self.model.config.id2label[item.argmax()], "score": item.max().item()} for item in scores
|
||||
]
|
||||
@@ -0,0 +1,190 @@
|
||||
from transformers.file_utils import add_end_docstrings
|
||||
|
||||
from .base import PIPELINE_INIT_ARGS, Pipeline
|
||||
|
||||
|
||||
@add_end_docstrings(PIPELINE_INIT_ARGS)
|
||||
class TextGenerationPipeline(Pipeline):
|
||||
"""
|
||||
Language generation pipeline using any :obj:`ModelWithLMHead`. This pipeline predicts the words that will follow a
|
||||
specified text prompt.
|
||||
|
||||
This language generation pipeline can currently be loaded from :func:`~transformers.pipeline` using the following
|
||||
task identifier: :obj:`"text-generation"`.
|
||||
|
||||
The models that this pipeline can use are models that have been trained with an autoregressive language modeling
|
||||
objective, which includes the uni-directional models in the library (e.g. gpt2). See the list of available models
|
||||
on `huggingface.co/models <https://huggingface.co/models?filter=causal-lm>`__.
|
||||
"""
|
||||
|
||||
# Prefix text to help Transformer-XL and XLNet with short prompts as proposed by Aman Rusia
|
||||
# in https://github.com/rusiaaman/XLNet-gen#methodology
|
||||
# and https://medium.com/@amanrusia/xlnet-speaks-comparison-to-gpt-2-ea1a4e9ba39e
|
||||
|
||||
XL_PREFIX = """
|
||||
In 1991, the remains of Russian Tsar Nicholas II and his family (except for Alexei and Maria) are discovered. The
|
||||
voice of Nicholas's young son, Tsarevich Alexei Nikolaevich, narrates the remainder of the story. 1883 Western
|
||||
Siberia, a young Grigori Rasputin is asked by his father and a group of men to perform magic. Rasputin has a vision
|
||||
and denounces one of the men as a horse thief. Although his father initially slaps him for making such an
|
||||
accusation, Rasputin watches as the man is chased outside and beaten. Twenty years later, Rasputin sees a vision of
|
||||
the Virgin Mary, prompting him to become a priest. Rasputin quickly becomes famous, with people, even a bishop,
|
||||
begging for his blessing. <eod> </s> <eos>
|
||||
"""
|
||||
|
||||
ALLOWED_MODELS = [
|
||||
"XLNetLMHeadModel",
|
||||
"TransfoXLLMHeadModel",
|
||||
"ReformerModelWithLMHead",
|
||||
"GPT2LMHeadModel",
|
||||
"OpenAIGPTLMHeadModel",
|
||||
"CTRLLMHeadModel",
|
||||
"TFXLNetLMHeadModel",
|
||||
"TFTransfoXLLMHeadModel",
|
||||
"TFGPT2LMHeadModel",
|
||||
"TFOpenAIGPTLMHeadModel",
|
||||
"TFCTRLLMHeadModel",
|
||||
]
|
||||
|
||||
def __init__(self, *args, **kwargs):
|
||||
super().__init__(*args, **kwargs)
|
||||
|
||||
self.check_model_type(self.ALLOWED_MODELS)
|
||||
|
||||
# overriding _parse_and_tokenize to allow for unusual language-modeling tokenizer arguments
|
||||
|
||||
def _parse_and_tokenize(self, inputs, padding=True, add_special_tokens=True, **kwargs):
|
||||
"""
|
||||
Parse arguments and tokenize
|
||||
"""
|
||||
# Parse arguments
|
||||
if self.model.__class__.__name__ in ["TransfoXLLMHeadModel"]:
|
||||
tokenizer_kwargs = {"add_space_before_punct_symbol": True}
|
||||
else:
|
||||
tokenizer_kwargs = {}
|
||||
inputs = self.tokenizer(
|
||||
inputs,
|
||||
add_special_tokens=add_special_tokens,
|
||||
return_tensors=self.framework,
|
||||
padding=padding,
|
||||
**tokenizer_kwargs,
|
||||
)
|
||||
|
||||
return inputs
|
||||
|
||||
def __call__(
|
||||
self,
|
||||
text_inputs,
|
||||
return_tensors=False,
|
||||
return_text=True,
|
||||
clean_up_tokenization_spaces=False,
|
||||
prefix=None,
|
||||
**generate_kwargs
|
||||
):
|
||||
"""
|
||||
Complete the prompt(s) given as inputs.
|
||||
|
||||
Args:
|
||||
args (:obj:`str` or :obj:`List[str]`):
|
||||
One or several prompts (or one list of prompts) to complete.
|
||||
return_tensors (:obj:`bool`, `optional`, defaults to :obj:`False`):
|
||||
Whether or not to include the tensors of predictions (as token indices) in the outputs.
|
||||
return_text (:obj:`bool`, `optional`, defaults to :obj:`True`):
|
||||
Whether or not to include the decoded texts in the outputs.
|
||||
clean_up_tokenization_spaces (:obj:`bool`, `optional`, defaults to :obj:`False`):
|
||||
Whether or not to clean up the potential extra spaces in the text output.
|
||||
prefix (:obj:`str`, `optional`):
|
||||
Prefix added to prompt.
|
||||
generate_kwargs:
|
||||
Additional keyword arguments to pass along to the generate method of the model (see the generate method
|
||||
corresponding to your framework `here <./model.html#generative-models>`__).
|
||||
|
||||
Return:
|
||||
A list or a list of list of :obj:`dict`: Each result comes as a dictionary with the following keys:
|
||||
|
||||
- **generated_text** (:obj:`str`, present when ``return_text=True``) -- The generated text.
|
||||
- **generated_token_ids** (:obj:`torch.Tensor` or :obj:`tf.Tensor`, present when ``return_tensors=True``)
|
||||
-- The token ids of the generated text.
|
||||
"""
|
||||
|
||||
if isinstance(text_inputs, str):
|
||||
text_inputs = [text_inputs]
|
||||
results = []
|
||||
for prompt_text in text_inputs:
|
||||
# Manage correct placement of the tensors
|
||||
with self.device_placement():
|
||||
prefix = prefix if prefix is not None else self.model.config.prefix
|
||||
if prefix is None and self.model.__class__.__name__ in [
|
||||
"XLNetLMHeadModel",
|
||||
"TransfoXLLMHeadModel",
|
||||
"TFXLNetLMHeadModel",
|
||||
"TFTransfoXLLMHeadModel",
|
||||
]:
|
||||
# For XLNet and TransformerXL we add an article to the prompt to give more state to the model.
|
||||
prefix = self.XL_PREFIX
|
||||
|
||||
if prefix:
|
||||
prefix_inputs = self._parse_and_tokenize(prefix, padding=False, add_special_tokens=False)
|
||||
# This impacts max_length and min_length argument that need adjusting.
|
||||
prefix_length = prefix_inputs["input_ids"].shape[-1]
|
||||
if generate_kwargs.get("max_length", None) is not None:
|
||||
generate_kwargs["max_length"] += prefix_length
|
||||
if generate_kwargs.get("min_length", None) is not None:
|
||||
generate_kwargs["min_length"] += prefix_length
|
||||
|
||||
prefix = prefix or ""
|
||||
inputs = self._parse_and_tokenize(prefix + prompt_text, padding=False, add_special_tokens=False)
|
||||
|
||||
# set input_ids to None to allow empty prompt
|
||||
if inputs["input_ids"].shape[-1] == 0:
|
||||
inputs["input_ids"] = None
|
||||
inputs["attention_mask"] = None
|
||||
|
||||
if self.framework == "pt" and inputs["input_ids"] is not None:
|
||||
inputs = self.ensure_tensor_on_device(**inputs)
|
||||
|
||||
input_ids = inputs["input_ids"]
|
||||
|
||||
# Ensure that batch size = 1 (batch generation not allowed for now)
|
||||
assert (
|
||||
input_ids is None or input_ids.shape[0] == 1
|
||||
), "Batch generation is currently not supported. See https://github.com/huggingface/transformers/issues/3021 for more information."
|
||||
|
||||
output_sequences = self.model.generate(input_ids=input_ids, **generate_kwargs) # BS x SL
|
||||
|
||||
result = []
|
||||
for generated_sequence in output_sequences:
|
||||
if self.framework == "pt" and generated_sequence is not None:
|
||||
generated_sequence = generated_sequence.cpu()
|
||||
generated_sequence = generated_sequence.numpy().tolist()
|
||||
record = {}
|
||||
if return_tensors:
|
||||
record["generated_token_ids"] = generated_sequence
|
||||
if return_text:
|
||||
# Decode text
|
||||
text = self.tokenizer.decode(
|
||||
generated_sequence,
|
||||
skip_special_tokens=True,
|
||||
clean_up_tokenization_spaces=clean_up_tokenization_spaces,
|
||||
)
|
||||
|
||||
# Remove PADDING prompt of the sequence if XLNet or Transfo-XL model is used
|
||||
if input_ids is None:
|
||||
prompt_length = 0
|
||||
else:
|
||||
prompt_length = len(
|
||||
self.tokenizer.decode(
|
||||
input_ids[0],
|
||||
skip_special_tokens=True,
|
||||
clean_up_tokenization_spaces=clean_up_tokenization_spaces,
|
||||
)
|
||||
)
|
||||
|
||||
record["generated_text"] = prompt_text + text[prompt_length:]
|
||||
|
||||
result.append(record)
|
||||
results += [result]
|
||||
|
||||
if len(results) == 1:
|
||||
return results[0]
|
||||
|
||||
return results
|
||||
@@ -0,0 +1,171 @@
|
||||
from typing import List, Union
|
||||
|
||||
import numpy as np
|
||||
|
||||
from transformers.file_utils import add_end_docstrings
|
||||
from transformers.utils import logging
|
||||
|
||||
from .base import PIPELINE_INIT_ARGS, ArgumentHandler, Pipeline
|
||||
|
||||
|
||||
logger = logging.get_logger(__name__)
|
||||
|
||||
|
||||
class ZeroShotClassificationArgumentHandler(ArgumentHandler):
|
||||
"""
|
||||
Handles arguments for zero-shot for text classification by turning each possible label into an NLI
|
||||
premise/hypothesis pair.
|
||||
"""
|
||||
|
||||
def _parse_labels(self, labels):
|
||||
if isinstance(labels, str):
|
||||
labels = [label.strip() for label in labels.split(",")]
|
||||
return labels
|
||||
|
||||
def __call__(self, sequences, labels, hypothesis_template):
|
||||
if len(labels) == 0 or len(sequences) == 0:
|
||||
raise ValueError("You must include at least one label and at least one sequence.")
|
||||
if hypothesis_template.format(labels[0]) == hypothesis_template:
|
||||
raise ValueError(
|
||||
(
|
||||
'The provided hypothesis_template "{}" was not able to be formatted with the target labels. '
|
||||
"Make sure the passed template includes formatting syntax such as {{}} where the label should go."
|
||||
).format(hypothesis_template)
|
||||
)
|
||||
|
||||
if isinstance(sequences, str):
|
||||
sequences = [sequences]
|
||||
labels = self._parse_labels(labels)
|
||||
|
||||
sequence_pairs = []
|
||||
for sequence in sequences:
|
||||
sequence_pairs.extend([[sequence, hypothesis_template.format(label)] for label in labels])
|
||||
|
||||
return sequence_pairs
|
||||
|
||||
|
||||
@add_end_docstrings(PIPELINE_INIT_ARGS)
|
||||
class ZeroShotClassificationPipeline(Pipeline):
|
||||
"""
|
||||
NLI-based zero-shot classification pipeline using a :obj:`ModelForSequenceClassification` trained on NLI (natural
|
||||
language inference) tasks.
|
||||
|
||||
Any combination of sequences and labels can be passed and each combination will be posed as a premise/hypothesis
|
||||
pair and passed to the pretrained model. Then, the logit for `entailment` is taken as the logit for the candidate
|
||||
label being valid. Any NLI model can be used, but the id of the `entailment` label must be included in the model
|
||||
config's :attr:`~transformers.PretrainedConfig.label2id`.
|
||||
|
||||
This NLI pipeline can currently be loaded from :func:`~transformers.pipeline` using the following task identifier:
|
||||
:obj:`"zero-shot-classification"`.
|
||||
|
||||
The models that this pipeline can use are models that have been fine-tuned on an NLI task. See the up-to-date list
|
||||
of available models on `huggingface.co/models <https://huggingface.co/models?search=nli>`__.
|
||||
"""
|
||||
|
||||
def __init__(self, args_parser=ZeroShotClassificationArgumentHandler(), *args, **kwargs):
|
||||
super().__init__(*args, **kwargs)
|
||||
self._args_parser = args_parser
|
||||
if self.entailment_id == -1:
|
||||
logger.warning(
|
||||
"Failed to determine 'entailment' label id from the label2id mapping in the model config. Setting to "
|
||||
"-1. Define a descriptive label2id mapping in the model config to ensure correct outputs."
|
||||
)
|
||||
|
||||
@property
|
||||
def entailment_id(self):
|
||||
for label, ind in self.model.config.label2id.items():
|
||||
if label.lower().startswith("entail"):
|
||||
return ind
|
||||
return -1
|
||||
|
||||
def _parse_and_tokenize(
|
||||
self, sequences, candidate_labels, hypothesis_template, padding=True, add_special_tokens=True, **kwargs
|
||||
):
|
||||
"""
|
||||
Parse arguments and tokenize only_first so that hypothesis (label) is not truncated
|
||||
"""
|
||||
sequence_pairs = self._args_parser(sequences, candidate_labels, hypothesis_template)
|
||||
inputs = self.tokenizer(
|
||||
sequence_pairs,
|
||||
add_special_tokens=add_special_tokens,
|
||||
return_tensors=self.framework,
|
||||
padding=padding,
|
||||
truncation="only_first",
|
||||
)
|
||||
|
||||
return inputs
|
||||
|
||||
def __call__(
|
||||
self,
|
||||
sequences: Union[str, List[str]],
|
||||
candidate_labels,
|
||||
hypothesis_template="This example is {}.",
|
||||
multi_class=False,
|
||||
):
|
||||
"""
|
||||
Classify the sequence(s) given as inputs. See the :obj:`~transformers.ZeroShotClassificationPipeline`
|
||||
documentation for more information.
|
||||
|
||||
Args:
|
||||
sequences (:obj:`str` or :obj:`List[str]`):
|
||||
The sequence(s) to classify, will be truncated if the model input is too large.
|
||||
candidate_labels (:obj:`str` or :obj:`List[str]`):
|
||||
The set of possible class labels to classify each sequence into. Can be a single label, a string of
|
||||
comma-separated labels, or a list of labels.
|
||||
hypothesis_template (:obj:`str`, `optional`, defaults to :obj:`"This example is {}."`):
|
||||
The template used to turn each label into an NLI-style hypothesis. This template must include a {} or
|
||||
similar syntax for the candidate label to be inserted into the template. For example, the default
|
||||
template is :obj:`"This example is {}."` With the candidate label :obj:`"sports"`, this would be fed
|
||||
into the model like :obj:`"<cls> sequence to classify <sep> This example is sports . <sep>"`. The
|
||||
default template works well in many cases, but it may be worthwhile to experiment with different
|
||||
templates depending on the task setting.
|
||||
multi_class (:obj:`bool`, `optional`, defaults to :obj:`False`):
|
||||
Whether or not multiple candidate labels can be true. If :obj:`False`, the scores are normalized such
|
||||
that the sum of the label likelihoods for each sequence is 1. If :obj:`True`, the labels are considered
|
||||
independent and probabilities are normalized for each candidate by doing a softmax of the entailment
|
||||
score vs. the contradiction score.
|
||||
|
||||
Return:
|
||||
A :obj:`dict` or a list of :obj:`dict`: Each result comes as a dictionary with the following keys:
|
||||
|
||||
- **sequence** (:obj:`str`) -- The sequence for which this is the output.
|
||||
- **labels** (:obj:`List[str]`) -- The labels sorted by order of likelihood.
|
||||
- **scores** (:obj:`List[float]`) -- The probabilities for each of the labels.
|
||||
"""
|
||||
if sequences and isinstance(sequences, str):
|
||||
sequences = [sequences]
|
||||
|
||||
outputs = super().__call__(sequences, candidate_labels, hypothesis_template)
|
||||
num_sequences = len(sequences)
|
||||
candidate_labels = self._args_parser._parse_labels(candidate_labels)
|
||||
reshaped_outputs = outputs.reshape((num_sequences, len(candidate_labels), -1))
|
||||
|
||||
if len(candidate_labels) == 1:
|
||||
multi_class = True
|
||||
|
||||
if not multi_class:
|
||||
# softmax the "entailment" logits over all candidate labels
|
||||
entail_logits = reshaped_outputs[..., self.entailment_id]
|
||||
scores = np.exp(entail_logits) / np.exp(entail_logits).sum(-1, keepdims=True)
|
||||
else:
|
||||
# softmax over the entailment vs. contradiction dim for each label independently
|
||||
entailment_id = self.entailment_id
|
||||
contradiction_id = -1 if entailment_id == 0 else 0
|
||||
entail_contr_logits = reshaped_outputs[..., [contradiction_id, entailment_id]]
|
||||
scores = np.exp(entail_contr_logits) / np.exp(entail_contr_logits).sum(-1, keepdims=True)
|
||||
scores = scores[..., 1]
|
||||
|
||||
result = []
|
||||
for iseq in range(num_sequences):
|
||||
top_inds = list(reversed(scores[iseq].argsort()))
|
||||
result.append(
|
||||
{
|
||||
"sequence": sequences if isinstance(sequences, str) else sequences[iseq],
|
||||
"labels": [candidate_labels[i] for i in top_inds],
|
||||
"scores": scores[iseq][top_inds].tolist(),
|
||||
}
|
||||
)
|
||||
|
||||
if len(result) == 1:
|
||||
return result[0]
|
||||
return result
|
||||
Reference in New Issue
Block a user