renamed
This commit is contained in:
@@ -25,9 +25,9 @@ from torch.nn import functional as F
|
||||
logger = logging.getLogger(__name__)
|
||||
|
||||
|
||||
class ModuleGenerationMixin:
|
||||
class GenerationMixin:
|
||||
"""
|
||||
A class contraining all of the functions supporting generation, to be used as a mixin.
|
||||
A class contraining all of the functions supporting generation, to be used as a mixin in PreTrainedModel.
|
||||
"""
|
||||
|
||||
def prepare_inputs_for_generation(self, input_ids, **kwargs):
|
||||
|
||||
@@ -23,9 +23,9 @@ import tensorflow as tf
|
||||
logger = logging.getLogger(__name__)
|
||||
|
||||
|
||||
class TFModelGenerationMixin:
|
||||
class TFGenerationMixin:
|
||||
"""
|
||||
A class contraining all of the functions supporting generation, to be used as a mixin.
|
||||
A class contraining all of the functions supporting generation, to be used as a mixin in TFPreTrainedModel.
|
||||
"""
|
||||
|
||||
def prepare_inputs_for_generation(self, inputs, **kwargs):
|
||||
|
||||
@@ -25,7 +25,7 @@ from tensorflow.python.keras.saving import hdf5_format
|
||||
|
||||
from .configuration_utils import PretrainedConfig
|
||||
from .file_utils import DUMMY_INPUTS, TF2_WEIGHTS_NAME, WEIGHTS_NAME, cached_path, hf_bucket_url, is_remote_url
|
||||
from .modeling_tf_generation import TFModelGenerationMixin, shape_list
|
||||
from .modeling_tf_generation import TFGenerationMixin, shape_list
|
||||
from .modeling_tf_pytorch_utils import load_pytorch_checkpoint_in_tf2_model
|
||||
|
||||
|
||||
@@ -145,7 +145,7 @@ class TFSequenceClassificationLoss:
|
||||
TFMultipleChoiceLoss = TFSequenceClassificationLoss
|
||||
|
||||
|
||||
class TFPreTrainedModel(tf.keras.Model, TFModelUtilsMixin, TFModelGenerationMixin):
|
||||
class TFPreTrainedModel(tf.keras.Model, TFModelUtilsMixin, TFGenerationMixin):
|
||||
r""" Base class for all TF models.
|
||||
|
||||
:class:`~transformers.TFPreTrainedModel` takes care of storing the configuration of the models and handles methods for loading/downloading/saving models
|
||||
|
||||
@@ -35,7 +35,7 @@ from .file_utils import (
|
||||
hf_bucket_url,
|
||||
is_remote_url,
|
||||
)
|
||||
from .modeling_generation import ModuleGenerationMixin
|
||||
from .modeling_generation import GenerationMixin
|
||||
|
||||
|
||||
logger = logging.getLogger(__name__)
|
||||
@@ -262,7 +262,7 @@ class ModuleUtilsMixin:
|
||||
return head_mask
|
||||
|
||||
|
||||
class PreTrainedModel(nn.Module, ModuleUtilsMixin, ModuleGenerationMixin):
|
||||
class PreTrainedModel(nn.Module, ModuleUtilsMixin, GenerationMixin):
|
||||
r""" Base class for all models.
|
||||
|
||||
:class:`~transformers.PreTrainedModel` takes care of storing the configuration of the models and handles methods for loading/downloading/saving models
|
||||
|
||||
Reference in New Issue
Block a user