This commit is contained in:
yjernite
2020-06-25 11:23:54 -04:00
parent ca51fdc79f
commit 8d36410be2
4 changed files with 8 additions and 8 deletions
+2 -2
View File
@@ -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):
+2 -2
View File
@@ -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):
+2 -2
View File
@@ -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
+2 -2
View File
@@ -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