Compare commits

...
Author SHA1 Message Date
sgugger 73fb43f32b Remove defaults to None from diff 2020-09-09 10:17:54 -04:00
sgugger 07d77d8f70 Present scripting alternative in a clean diff 2020-09-09 10:15:01 -04:00
3 changed files with 749 additions and 584 deletions
File diff suppressed because it is too large Load Diff
+67 -87
View File
@@ -17,9 +17,8 @@
import inspect import inspect
import os import os
import re import re
import warnings
from dataclasses import dataclass from dataclasses import dataclass
from typing import Any, Callable, Dict, List, Optional, Set, Tuple, Union from typing import Callable, Dict, List, Optional, Set, Tuple, Union
import torch import torch
from torch import Tensor, device, dtype, nn from torch import Tensor, device, dtype, nn
@@ -46,6 +45,7 @@ from .utils import logging
logger = logging.get_logger(__name__) logger = logging.get_logger(__name__)
try: try:
from torch.nn import Identity from torch.nn import Identity
except ImportError: except ImportError:
@@ -91,6 +91,20 @@ class ModuleUtilsMixin:
A few utilities for :obj:`torch.nn.Modules`, to be used as a mixin. A few utilities for :obj:`torch.nn.Modules`, to be used as a mixin.
""" """
def num_parameters(self, only_trainable: bool = False) -> int:
"""
Get the number of (optionally, trainable) parameters in the model.
Args:
only_trainable (:obj:`bool`, `optional`, defaults to :obj:`False`):
Whether or not to return only the number of trainable parameters
Returns:
:obj:`int`: The number of parameters.
"""
params = filter(lambda x: x.requires_grad, self.parameters()) if only_trainable else self.parameters()
return sum(p.numel() for p in params)
@staticmethod @staticmethod
def _hook_rss_memory_pre_forward(module, *args, **kwargs): def _hook_rss_memory_pre_forward(module, *args, **kwargs):
try: try:
@@ -158,10 +172,10 @@ class ModuleUtilsMixin:
first_tuple = next(gen) first_tuple = next(gen)
return first_tuple[1].device return first_tuple[1].device
@property # TorchScript does not support, so add non-property option
def dtype(self) -> dtype: def get_dtype(self) -> dtype:
""" """
:obj:`torch.dtype`: The dtype of the module (assuming that all the module parameters have the same dtype). Get torch.dtype from module, assuming that the whole module has one dtype.
""" """
try: try:
return next(self.parameters()).dtype return next(self.parameters()).dtype
@@ -176,6 +190,10 @@ class ModuleUtilsMixin:
first_tuple = next(gen) first_tuple = next(gen)
return first_tuple[1].dtype return first_tuple[1].dtype
@property
def dtype(self):
return self.get_dtype()
def invert_attention_mask(self, encoder_attention_mask: Tensor) -> Tensor: def invert_attention_mask(self, encoder_attention_mask: Tensor) -> Tensor:
""" """
Invert an attention mask (e.g., switches 0. and 1.). Invert an attention mask (e.g., switches 0. and 1.).
@@ -186,31 +204,43 @@ class ModuleUtilsMixin:
Returns: Returns:
:obj:`torch.Tensor`: The inverted attention mask. :obj:`torch.Tensor`: The inverted attention mask.
""" """
encoder_extended_attention_mask: Optional[Tensor] = None
if encoder_attention_mask.dim() == 3: if encoder_attention_mask.dim() == 3:
encoder_extended_attention_mask = encoder_attention_mask[:, None, :, :] encoder_extended_attention_mask = encoder_attention_mask[:, None, :, :]
if encoder_attention_mask.dim() == 2: if encoder_attention_mask.dim() == 2:
encoder_extended_attention_mask = encoder_attention_mask[:, None, None, :] encoder_extended_attention_mask = encoder_attention_mask[:, None, None, :]
assert encoder_extended_attention_mask is not None
# T5 has a mask that can compare sequence ids, we can simulate this here with this transposition # T5 has a mask that can compare sequence ids, we can simulate this here with this transposition
# Cf. https://github.com/tensorflow/mesh/blob/8d2465e9bc93129b913b5ccc6a59aa97abd96ec6/mesh_tensorflow # Cf. https://github.com/tensorflow/mesh/blob/8d2465e9bc93129b913b5ccc6a59aa97abd96ec6/mesh_tensorflow
# /transformer/transformer_layers.py#L270 # /transformer/transformer_layers.py#L270
# encoder_extended_attention_mask = (encoder_extended_attention_mask == # encoder_extended_attention_mask = (encoder_extended_attention_mask ==
# encoder_extended_attention_mask.transpose(-1, -2)) # encoder_extended_attention_mask.transpose(-1, -2))
encoder_extended_attention_mask = encoder_extended_attention_mask.to(dtype=self.dtype) # fp16 compatibility encoder_extended_attention_mask = encoder_extended_attention_mask.to(
dtype=self.get_dtype()
) # fp16 compatibility
if self.dtype == torch.float16: if self.get_dtype() == torch.float16:
encoder_extended_attention_mask = (1.0 - encoder_extended_attention_mask) * -1e4 encoder_extended_attention_mask = (1.0 - encoder_extended_attention_mask) * -1e4
elif self.dtype == torch.float32: elif self.get_dtype() == torch.float32:
encoder_extended_attention_mask = (1.0 - encoder_extended_attention_mask) * -1e9 encoder_extended_attention_mask = (1.0 - encoder_extended_attention_mask) * -1e9
else: else:
raise ValueError( raise ValueError(
"{} not recognized. `dtype` should be set to either `torch.float32` or `torch.float16`".format( "{} not recognized. `dtype` should be set to either `torch.float32` or `torch.float16`".format(
self.dtype self.get_dtype()
) )
) )
return encoder_extended_attention_mask return encoder_extended_attention_mask
def get_extended_attention_mask(self, attention_mask: Tensor, input_shape: Tuple[int], device: device) -> Tensor: def get_is_decoder(self):
if hasattr(self, "is_decoder"):
return self.is_decoder
else:
return self.config.is_decoder
def get_extended_attention_mask(
self, attention_mask: Tensor, input_shape: Tuple[int, int], device: device
) -> Tensor:
""" """
Makes broadcastable attention and causal masks so that future and masked tokens are ignored. Makes broadcastable attention and causal masks so that future and masked tokens are ignored.
@@ -233,7 +263,7 @@ class ModuleUtilsMixin:
# Provided a padding mask of dimensions [batch_size, seq_length] # Provided a padding mask of dimensions [batch_size, seq_length]
# - if the model is a decoder, apply a causal mask in addition to the padding mask # - if the model is a decoder, apply a causal mask in addition to the padding mask
# - if the model is an encoder, make the mask broadcastable to [batch_size, num_heads, seq_length, seq_length] # - if the model is an encoder, make the mask broadcastable to [batch_size, num_heads, seq_length, seq_length]
if self.config.is_decoder: if self.get_is_decoder():
batch_size, seq_length = input_shape batch_size, seq_length = input_shape
seq_ids = torch.arange(seq_length, device=device) seq_ids = torch.arange(seq_length, device=device)
causal_mask = seq_ids[None, None, :].repeat(batch_size, seq_length, 1) <= seq_ids[None, :, None] causal_mask = seq_ids[None, None, :].repeat(batch_size, seq_length, 1) <= seq_ids[None, :, None]
@@ -254,7 +284,7 @@ class ModuleUtilsMixin:
# positions we want to attend and -10000.0 for masked positions. # positions we want to attend and -10000.0 for masked positions.
# Since we are adding it to the raw scores before the softmax, this is # Since we are adding it to the raw scores before the softmax, this is
# effectively the same as removing these entirely. # effectively the same as removing these entirely.
extended_attention_mask = extended_attention_mask.to(dtype=self.dtype) # fp16 compatibility extended_attention_mask = extended_attention_mask.to(dtype=self.get_dtype()) # fp16 compatibility
extended_attention_mask = (1.0 - extended_attention_mask) * -10000.0 extended_attention_mask = (1.0 - extended_attention_mask) * -10000.0
return extended_attention_mask return extended_attention_mask
@@ -276,16 +306,34 @@ class ModuleUtilsMixin:
:obj:`torch.Tensor` with shape :obj:`[num_hidden_layers x batch x num_heads x seq_length x seq_length]` :obj:`torch.Tensor` with shape :obj:`[num_hidden_layers x batch x num_heads x seq_length x seq_length]`
or list with :obj:`[None]` for each layer. or list with :obj:`[None]` for each layer.
""" """
if head_mask is not None: if head_mask is None:
return [None] * num_hidden_layers
else:
return self.get_scriptable_head_mask(head_mask, num_hidden_layers, is_attention_chunked)
def get_scriptable_head_mask(
self, head_mask: Optional[Tensor], num_hidden_layers: int, is_attention_chunked: bool = False
) -> Optional[Tensor]:
"""
# Prepare head mask if needed
# 1.0 in head_mask indicate we keep the head
attention_probs has shape bsz x n_heads x N x N
Arguments:
head_mask: torch.Tensor or None: has shape [num_heads] or [num_hidden_layers x num_heads]
num_hidden_layers: int
Returns:
Tensor of shape shape [num_hidden_layers x batch x num_heads x seq_length x seq_length]
or None
"""
if head_mask is None:
return None
else:
head_mask = self._convert_head_mask_to_5d(head_mask, num_hidden_layers) head_mask = self._convert_head_mask_to_5d(head_mask, num_hidden_layers)
if is_attention_chunked is True: if is_attention_chunked is True:
head_mask = head_mask.unsqueeze(-1) head_mask = head_mask.unsqueeze(-1)
else: return head_mask
head_mask = [None] * num_hidden_layers
return head_mask def _convert_head_mask_to_5d(self, head_mask, num_hidden_layers: int):
def _convert_head_mask_to_5d(self, head_mask, num_hidden_layers):
"""-> [num_hidden_layers x batch x num_heads x seq_length x seq_length]""" """-> [num_hidden_layers x batch x num_heads x seq_length x seq_length]"""
if head_mask.dim() == 1: if head_mask.dim() == 1:
head_mask = head_mask.unsqueeze(0).unsqueeze(0).unsqueeze(-1).unsqueeze(-1) head_mask = head_mask.unsqueeze(0).unsqueeze(0).unsqueeze(-1).unsqueeze(-1)
@@ -293,77 +341,9 @@ class ModuleUtilsMixin:
elif head_mask.dim() == 2: elif head_mask.dim() == 2:
head_mask = head_mask.unsqueeze(1).unsqueeze(-1).unsqueeze(-1) # We can specify head_mask for each layer head_mask = head_mask.unsqueeze(1).unsqueeze(-1).unsqueeze(-1) # We can specify head_mask for each layer
assert head_mask.dim() == 5, f"head_mask.dim != 5, instead {head_mask.dim()}" assert head_mask.dim() == 5, f"head_mask.dim != 5, instead {head_mask.dim()}"
head_mask = head_mask.to(dtype=self.dtype) # switch to float if need + fp16 compatibility head_mask = head_mask.to(dtype=self.get_dtype()) # switch to fload if need + fp16 compatibility
return head_mask return head_mask
def num_parameters(self, only_trainable: bool = False, exclude_embeddings: bool = False) -> int:
"""
Get number of (optionally, trainable or non-embeddings) parameters in the module.
Args:
only_trainable (:obj:`bool`, `optional`, defaults to :obj:`False`):
Whether or not to return only the number of trainable parameters
exclude_embeddings (:obj:`bool`, `optional`, defaults to :obj:`False`):
Whether or not to return only the number of non-embeddings parameters
Returns:
:obj:`int`: The number of parameters.
"""
def parameter_filter(x):
return (x.requires_grad or not only_trainable) and not (
isinstance(x, torch.nn.Embedding) and exclude_embeddings
)
params = filter(parameter_filter, self.parameters()) if only_trainable else self.parameters()
return sum(p.numel() for p in params)
def estimate_tokens(self, input_dict: Dict[str, Union[torch.Tensor, Any]]) -> int:
"""
Helper function to estimate the total number of tokens from the model inputs.
Args:
inputs (:obj:`dict`): The model inputs.
Returns:
:obj:`int`: The total number of tokens.
"""
token_inputs = [tensor for key, tensor in input_dict.items() if "input" in key]
if token_inputs:
return sum([token_input.numel() for token_input in token_inputs])
else:
warnings.warn(
"Could not estimate the number of tokens of the input, floating-point operations will not be computed"
)
return 0
def floating_point_ops(
self, input_dict: Dict[str, Union[torch.Tensor, Any]], exclude_embeddings: bool = True
) -> int:
"""
Get number of (optionally, non-embeddings) floating-point operations for the forward and backward passes of a
batch with this transformer model. Default approximation neglects the quadratic dependency on the number of
tokens (valid if :obj:`12 * d_model << sequence_length`) as laid out in `this paper <https://arxiv.org/pdf/2001.08361.pdf>`__ section
2.1. Should be overriden for transformers with parameter re-use e.g. Albert or Universal Transformers, or
if doing long-range modeling with very high sequence lengths.
Args:
batch_size (:obj:`int`):
The batch size for the forward pass.
sequence_length (:obj:`int`):
The number of tokens in each line of the batch.
exclude_embeddings (:obj:`bool`, `optional`, defaults to :obj:`True`):
Whether or not to count embedding and softmax operations.
Returns:
:obj:`int`: The number of floating-point operations.
"""
return 6 * self.estimate_tokens(input_dict) * self.num_parameters(exclude_embeddings=exclude_embeddings)
class PreTrainedModel(nn.Module, ModuleUtilsMixin, GenerationMixin): class PreTrainedModel(nn.Module, ModuleUtilsMixin, GenerationMixin):
r""" r"""
+77
View File
@@ -453,6 +453,83 @@ class BertModelTest(ModelTesterMixin, unittest.TestCase):
config_and_inputs = self.model_tester.prepare_config_and_inputs() config_and_inputs = self.model_tester.prepare_config_and_inputs()
self.model_tester.create_and_check_for_token_classification(*config_and_inputs) self.model_tester.create_and_check_for_token_classification(*config_and_inputs)
# Copied from test_modeling_common.test_torchscript, but using jit.script, not jit.trace
def test_full_torchscript(self):
import copy
import os
import tempfile
if is_torch_available():
import torch
from transformers.modeling_bert import (
BertScriptableForMultipleChoice,
BertScriptableForNextSentencePrediction,
BertScriptableForPreTraining,
BertScriptableForQuestionAnswering,
BertScriptableForSequenceClassification,
BertScriptableForTokenClassification,
BertScriptableModel,
)
config, unused_inputs_dict = self.model_tester.prepare_config_and_inputs_for_common()
config.return_dict = False
configs_no_init = copy.deepcopy(config)
for key in configs_no_init.__dict__.keys():
if "_range" in key or "_std" in key or "initializer_factor" in key:
setattr(configs_no_init, key, 1e-10)
scriptable_model_classes = (
BertScriptableModel,
BertScriptableForMultipleChoice,
BertScriptableForNextSentencePrediction,
BertScriptableForPreTraining,
BertScriptableForQuestionAnswering,
BertScriptableForSequenceClassification,
BertScriptableForTokenClassification,
)
for model_class in scriptable_model_classes:
model = model_class(config=configs_no_init)
model.to(torch_device)
model.eval()
try:
scripted = torch.jit.script(model)
except RuntimeError:
self.fail("Couldn't script module.")
with tempfile.TemporaryDirectory() as tmp_dir_name:
pt_file_name = os.path.join(tmp_dir_name, "scripted_model.pt")
try:
torch.jit.save(scripted, pt_file_name)
except Exception:
self.fail("Couldn't save scripted module.")
try:
loaded_model = torch.jit.load(pt_file_name)
except Exception:
self.fail("Couldn't load scripted module.")
model.to(torch_device)
model.eval()
loaded_model.to(torch_device)
loaded_model.eval()
model_state_dict = model.state_dict()
loaded_model_state_dict = loaded_model.state_dict()
self.assertEqual(set(model_state_dict.keys()), set(loaded_model_state_dict.keys()))
models_equal = True
for layer_name, p1 in model_state_dict.items():
p2 = loaded_model_state_dict[layer_name]
if p1.data.ne(p2.data).sum() > 0:
models_equal = False
self.assertTrue(models_equal)
@slow @slow
def test_model_from_pretrained(self): def test_model_from_pretrained(self):
for model_name in BERT_PRETRAINED_MODEL_ARCHIVE_LIST[:1]: for model_name in BERT_PRETRAINED_MODEL_ARCHIVE_LIST[:1]: