Compare commits
3
Commits
| Author | SHA1 | Date | |
|---|---|---|---|
|
|
7cfdbb24b7 | ||
|
|
a53e7b2e87 | ||
|
|
d8043889c1 |
@@ -120,27 +120,27 @@ from .modeling_tf_pytorch_utils import (
|
||||
)
|
||||
|
||||
# Pipelines
|
||||
from .pipelines import (
|
||||
Conversation,
|
||||
ConversationalPipeline,
|
||||
CsvPipelineDataFormat,
|
||||
FeatureExtractionPipeline,
|
||||
FillMaskPipeline,
|
||||
JsonPipelineDataFormat,
|
||||
NerPipeline,
|
||||
PipedPipelineDataFormat,
|
||||
Pipeline,
|
||||
PipelineDataFormat,
|
||||
QuestionAnsweringPipeline,
|
||||
SummarizationPipeline,
|
||||
Text2TextGenerationPipeline,
|
||||
TextClassificationPipeline,
|
||||
TextGenerationPipeline,
|
||||
TokenClassificationPipeline,
|
||||
TranslationPipeline,
|
||||
ZeroShotClassificationPipeline,
|
||||
pipeline,
|
||||
)
|
||||
# from .pipelines import (
|
||||
# Conversation,
|
||||
# ConversationalPipeline,
|
||||
# CsvPipelineDataFormat,
|
||||
# FeatureExtractionPipeline,
|
||||
# FillMaskPipeline,
|
||||
# JsonPipelineDataFormat,
|
||||
# NerPipeline,
|
||||
# PipedPipelineDataFormat,
|
||||
# Pipeline,
|
||||
# PipelineDataFormat,
|
||||
# QuestionAnsweringPipeline,
|
||||
# SummarizationPipeline,
|
||||
# Text2TextGenerationPipeline,
|
||||
# TextClassificationPipeline,
|
||||
# TextGenerationPipeline,
|
||||
# TokenClassificationPipeline,
|
||||
# TranslationPipeline,
|
||||
# ZeroShotClassificationPipeline,
|
||||
# pipeline,
|
||||
# )
|
||||
|
||||
# Retriever
|
||||
from .retrieval_rag import RagRetriever
|
||||
|
||||
@@ -0,0 +1,13 @@
|
||||
from pathlib import Path
|
||||
from typing import Any, Union
|
||||
|
||||
from transformers import PreTrainedTokenizer
|
||||
|
||||
from .base import PipelineConfigType, PipelineInputType, PipelineIntermediateType, MaybeBatch, ModelType, \
|
||||
PipelineOutputType, Pipeline, PipelineConfig
|
||||
from .configs import TokenClassificationConfig
|
||||
from .outputs import TokenClassificationOutput
|
||||
|
||||
|
||||
def pipeline(task: str, model: Union[str, Path, Any], tokenizer: Union[str, Path, PreTrainedTokenizer]):
|
||||
pass
|
||||
@@ -0,0 +1,168 @@
|
||||
from abc import ABC, abstractmethod
|
||||
from dataclasses import dataclass, field
|
||||
from typing import Any, ClassVar, Dict, Generic, Mapping, Optional, TypeVar, Union, List
|
||||
|
||||
from transformers import BatchEncoding, PreTrainedTokenizer
|
||||
|
||||
|
||||
# Syntactic sugar to indicate it can take multiple inputs of type T at once.
|
||||
T = TypeVar("T")
|
||||
MaybeBatch = Union[T, List[T]]
|
||||
|
||||
|
||||
@dataclass
|
||||
class PipelineConfig:
|
||||
tokenizer_kwargs: Dict[str, Any] = field(default_factory=dict)
|
||||
model_kwargs: Dict[str, Any] = field(default_factory=dict)
|
||||
|
||||
def __post_init__(self):
|
||||
self.model_kwargs.setdefault("return_dict", True)
|
||||
|
||||
|
||||
# Define all the generic types for a Pipeline
|
||||
PipelineInputType = TypeVar("PipelineInputType") # Pipeline input type
|
||||
PipelineOutputType = TypeVar("PipelineOutputType") # Pipeline output type
|
||||
PipelineIntermediateType = TypeVar("PipelineIntermediateType") # Model output type (i.e. after forward())
|
||||
PipelineConfigType = TypeVar("PipelineConfigType", bound=PipelineConfig) # Pipeline configuration type
|
||||
ModelType = TypeVar("ModelType") # Pipeline model type
|
||||
|
||||
|
||||
class Pipeline(ABC, Generic[PipelineConfigType, ModelType, PipelineInputType, PipelineIntermediateType, PipelineOutputType]):
|
||||
""""""
|
||||
|
||||
__slots__ = ["_tokenizer", "_model", "_configs"]
|
||||
|
||||
# TODO: Check if the $task variable is required and won't introduce dependency on it
|
||||
task: ClassVar[str]
|
||||
default_config: ClassVar[PipelineConfigType]
|
||||
|
||||
def __init__(self, tokenizer: PreTrainedTokenizer, model: ModelType):
|
||||
self._tokenizer = tokenizer
|
||||
self._model = model
|
||||
|
||||
self._configs: Dict[str, PipelineConfigType] = dict()
|
||||
self._configs["default"] = self.default_config
|
||||
|
||||
@property
|
||||
def tokenizer(self) -> PreTrainedTokenizer:
|
||||
"""
|
||||
Tokenizer associated with this pipeline
|
||||
Returns:
|
||||
PreTrainedTokenizer instance
|
||||
"""
|
||||
return self._tokenizer
|
||||
|
||||
@property
|
||||
def model(self) -> ModelType:
|
||||
"""
|
||||
Model associated with this pipeline
|
||||
Returns:
|
||||
Pretrained model instance
|
||||
"""
|
||||
raise self._model
|
||||
|
||||
@property
|
||||
def configs(self) -> Mapping[str, PipelineConfigType]:
|
||||
"""
|
||||
Return all registered configuration's mapping on the pipeline
|
||||
Returns:
|
||||
Mapping[str, ConfigType]
|
||||
|
||||
"""
|
||||
return self._configs
|
||||
|
||||
@property
|
||||
def default_config(self) -> PipelineConfigType:
|
||||
return self.default_config
|
||||
|
||||
def get_config(self, name: str) -> Optional[PipelineConfigType]:
|
||||
"""
|
||||
Attempt to retrieve a configuration from its registered name.
|
||||
If no configuration matches the requested name, None is returned.
|
||||
Args:
|
||||
name (:obj:`str`): Name of the configuration to look for
|
||||
|
||||
Returns:
|
||||
Instance of ConfigType if the configuration associated with the requested name is present in the registry
|
||||
None if the requested named cannot be found in the registry
|
||||
"""
|
||||
return self._configs.get(name, None)
|
||||
|
||||
def register_config(self, name: str, config: PipelineConfigType) -> PipelineConfigType:
|
||||
"""
|
||||
Register a configuration with the specified name uniquely identifying a configuration.
|
||||
Args:
|
||||
name (:obj:`str`): Name used to reference the configuration
|
||||
config (:obj:`ConfigType`): The configuration instance to associate to the specified name
|
||||
Returns:
|
||||
The provided ConfigType instance
|
||||
"""
|
||||
self._configs[name] = config
|
||||
return config
|
||||
|
||||
def delete_config(self, name: str) -> Optional[PipelineConfigType]:
|
||||
"""
|
||||
Attempt to delete a configuration if registered on the pipeline.
|
||||
The ConfigType instance associated with the provided name is returned if present on the pipeline,
|
||||
None is returned otherwise.
|
||||
Args:
|
||||
name (:obj:`str`): Name identifying the configuration to delete
|
||||
|
||||
Returns:
|
||||
ConfigType instance if the name identifier was found on the pipeline.
|
||||
None if the name identifier is unknown on the pipeline.
|
||||
"""
|
||||
return self._configs.pop(name, None)
|
||||
|
||||
@abstractmethod
|
||||
def __call__(
|
||||
self, inputs: MaybeBatch[PipelineInputType], config: Optional[Union[str, PipelineConfigType]], **kwargs
|
||||
) -> MaybeBatch[PipelineIntermediateType]:
|
||||
|
||||
"""
|
||||
Args:
|
||||
inputs (:obj:`InputType`):
|
||||
config (:obj:`ConfigType`)
|
||||
**kwargs:
|
||||
|
||||
Returns:
|
||||
"""
|
||||
raise NotImplementedError()
|
||||
|
||||
@abstractmethod
|
||||
def preprocess(self, inputs: MaybeBatch[PipelineInputType], config: PipelineConfigType) -> BatchEncoding:
|
||||
"""
|
||||
Process the raw inputs to generate a BatchEncoding representation through the use of a tokenizer.
|
||||
Args:
|
||||
inputs (:obj:`InputType`):
|
||||
config (:obj:`ConfigType`)
|
||||
|
||||
Returns:
|
||||
"""
|
||||
# TODO: preprocess while calling batch_encode_plus might support truncation and thus all the remaining
|
||||
# steps [forward(), postprocess()] should supports inference and input reconstruction.
|
||||
raise NotImplementedError()
|
||||
|
||||
@abstractmethod
|
||||
def forward(self, encodings: BatchEncoding, config: PipelineConfigType) -> PipelineIntermediateType:
|
||||
"""
|
||||
|
||||
Args:
|
||||
encodings (:obj:`BatchEncoding`):
|
||||
config (:obj:`ConfigType`)
|
||||
Returns:
|
||||
"""
|
||||
raise NotImplementedError()
|
||||
|
||||
@abstractmethod
|
||||
def postprocess(self, encodings: BatchEncoding, model_output: PipelineIntermediateType, config: PipelineConfigType) -> MaybeBatch[PipelineOutputType]:
|
||||
"""
|
||||
|
||||
Args:
|
||||
encodings:
|
||||
model_output (:obj:`PipelineIntermediateType`):
|
||||
config (:obj:`ConfigType`)
|
||||
|
||||
Returns:
|
||||
"""
|
||||
raise NotImplementedError()
|
||||
@@ -0,0 +1,11 @@
|
||||
from dataclasses import dataclass, field
|
||||
from typing import List
|
||||
|
||||
from transformers.pipelines import PipelineConfig
|
||||
|
||||
|
||||
@dataclass
|
||||
class TokenClassificationConfig(PipelineConfig):
|
||||
group_entities: bool = False
|
||||
ignore_labels: List[int] = field(default_factory=list)
|
||||
|
||||
@@ -0,0 +1,13 @@
|
||||
from dataclasses import dataclass
|
||||
from typing import Union, List, NamedTuple
|
||||
|
||||
from transformers import BatchEncoding
|
||||
|
||||
Entity = NamedTuple("Entity", [("token", str), ("index", int), ("label", str), ("score", float)])
|
||||
|
||||
|
||||
@dataclass
|
||||
class TokenClassificationOutput:
|
||||
input: Union[str, List[str]]
|
||||
encodings: BatchEncoding
|
||||
entities: List[Entity]
|
||||
@@ -0,0 +1,8 @@
|
||||
from .pipeline_utils import PreTrainedPipeline
|
||||
from .token_classification import (
|
||||
TokenClassificationConfig,
|
||||
TokenClassificationInput,
|
||||
TokenClassificationOutput,
|
||||
TokenClassificationPipeline,
|
||||
TokenClassifierOutput,
|
||||
)
|
||||
@@ -0,0 +1,55 @@
|
||||
from abc import ABC
|
||||
from typing import Optional, Union
|
||||
|
||||
import torch
|
||||
|
||||
from transformers import PreTrainedModel, BatchEncoding, TensorType
|
||||
from transformers.pipelines import MaybeBatch, Pipeline
|
||||
from transformers.pipelines.base import PipelineConfigType, PipelineInputType, PipelineOutputType, \
|
||||
PipelineIntermediateType
|
||||
|
||||
|
||||
class PreTrainedPipeline(Pipeline[PipelineConfigType, PreTrainedModel, PipelineInputType, PipelineIntermediateType, PipelineOutputType], ABC):
|
||||
|
||||
@property
|
||||
def device(self) -> torch.device:
|
||||
return self._model.device
|
||||
|
||||
def cpu(self) -> "PreTrainedPipeline":
|
||||
self._model.cpu()
|
||||
return self
|
||||
|
||||
def gpu(self) -> "PreTrainedPipeline":
|
||||
self._model.gpu()
|
||||
return self
|
||||
|
||||
def preprocess(self, inputs: MaybeBatch[PipelineInputType], config: PipelineConfigType) -> BatchEncoding:
|
||||
return self._tokenizer(
|
||||
inputs,
|
||||
return_tensors=TensorType.PYTORCH,
|
||||
**(config.tokenizer_kwargs or {}),
|
||||
)
|
||||
|
||||
def __call__(self, inputs: MaybeBatch[PipelineInputType], config: Optional[Union[str, PipelineConfigType]], **kwargs) -> MaybeBatch[PipelineIntermediateType]:
|
||||
# Retrieve the appropriate config if identifier is provided or None
|
||||
if isinstance(config, str):
|
||||
config = self.get_config(config)
|
||||
elif not config:
|
||||
config = self.default_config
|
||||
|
||||
# Preprocess the input
|
||||
encodings = self.preprocess(inputs, config)
|
||||
|
||||
# Forward the encoding through the model
|
||||
model_output = self.forward(encodings, config)
|
||||
|
||||
# Apply any postprocessing steps required
|
||||
return self.postprocess(encodings, model_output, config)
|
||||
|
||||
def forward(self, encodings, config) -> PipelineIntermediateType:
|
||||
return self._model(**encodings, **config.model_kwargs)
|
||||
|
||||
|
||||
|
||||
|
||||
|
||||
@@ -0,0 +1,48 @@
|
||||
from transformers import BatchEncoding
|
||||
|
||||
from transformers.modeling_outputs import TokenClassifierOutput
|
||||
from transformers.pipelines import MaybeBatch, TokenClassificationConfig, TokenClassificationOutput
|
||||
from transformers.pipelines.outputs import Entity
|
||||
from transformers.pipelines.torch import PreTrainedPipeline
|
||||
|
||||
|
||||
# Token Classification
|
||||
|
||||
TokenClassificationInput = str
|
||||
|
||||
|
||||
class TokenClassificationPipeline(
|
||||
PreTrainedPipeline[
|
||||
TokenClassificationConfig, TokenClassificationInput, TokenClassifierOutput, TokenClassificationOutput
|
||||
]
|
||||
):
|
||||
task = "token-classification"
|
||||
default_config = TokenClassificationConfig()
|
||||
|
||||
def postprocess(self, encodings: BatchEncoding, model_output: TokenClassifierOutput, config: TokenClassificationConfig) -> MaybeBatch[TokenClassificationOutput]:
|
||||
# Retrieve labels
|
||||
labels = self._model.config.id2label
|
||||
|
||||
# Convert to probabilities
|
||||
probs = model_output.logits.softmax(dim=-1)
|
||||
|
||||
# Retrieve the argmax for each token over sequence axis
|
||||
batch_entities_labels = probs.argmax(dim=-1)
|
||||
|
||||
batch_entities = []
|
||||
for batch_idx, entities_labels in enumerate(batch_entities_labels):
|
||||
entities = []
|
||||
for token_idx, label_idx in enumerate(entities_labels):
|
||||
label_idx = label_idx.item()
|
||||
if label_idx not in config.ignore_labels:
|
||||
entities.append(Entity(
|
||||
token=self._tokenizer.convert_ids_to_tokens([encodings["input_ids"][batch_idx][token_idx]])[0],
|
||||
index=token_idx,
|
||||
label=labels[label_idx],
|
||||
score=probs[batch_idx, token_idx, label_idx].item(),
|
||||
))
|
||||
batch_entities.append(entities)
|
||||
|
||||
return batch_entities
|
||||
|
||||
|
||||
Reference in New Issue
Block a user