Compare commits

...
Author SHA1 Message Date
Lysandre b7e417ab94 Docstring 2020-06-30 19:40:22 -04:00
Lysandre 4aa3c76528 Address comment 2020-06-30 19:37:36 -04:00
Lysandre 826a082c84 Imports 2020-06-30 17:42:59 -04:00
Lysandre 998721e1dc BERT & ALBERT poc 2020-06-30 17:35:18 -04:00
7 changed files with 67 additions and 25 deletions
+1
View File
@@ -61,6 +61,7 @@ class PretrainedConfig(object):
self.torchscript = kwargs.pop("torchscript", False) # Only used by PyTorch models
self.use_bfloat16 = kwargs.pop("use_bfloat16", False)
self.pruned_heads = kwargs.pop("pruned_heads", {})
self.gradient_checkpointing = kwargs.pop("gradient_checkpointing", False)
# Is decoder is used in encoder-decoder models to differentiate encoder from decoder
self.is_encoder_decoder = kwargs.pop("is_encoder_decoder", False)
+4
View File
@@ -327,6 +327,7 @@ class AlbertTransformer(nn.Module):
hidden_states = self.embedding_hidden_mapping_in(hidden_states)
all_attentions = ()
output_attentions = torch.tensor(output_attentions, dtype=torch.bool)
if output_hidden_states:
all_hidden_states = (hidden_states,)
@@ -462,6 +463,9 @@ class AlbertModel(AlbertPreTrainedModel):
def set_input_embeddings(self, value):
self.embeddings.word_embeddings = value
def get_layers(self):
return [layer for layer_group in self.encoder.albert_layer_groups for layer in layer_group.albert_layers]
def _resize_token_embeddings(self, new_num_tokens):
old_embeddings = self.embeddings.word_embeddings
new_embeddings = self._get_resized_embeddings(old_embeddings, new_num_tokens)
+12 -25
View File
@@ -409,35 +409,19 @@ class BertEncoder(nn.Module):
):
all_hidden_states = ()
all_attentions = ()
output_attentions = torch.tensor(output_attentions, dtype=torch.bool)
for i, layer_module in enumerate(self.layer):
if output_hidden_states:
all_hidden_states = all_hidden_states + (hidden_states,)
if getattr(self.config, "gradient_checkpointing", False):
def create_custom_forward(module):
def custom_forward(*inputs):
return module(*inputs, output_attentions)
return custom_forward
layer_outputs = torch.utils.checkpoint.checkpoint(
create_custom_forward(layer_module),
hidden_states,
attention_mask,
head_mask[i],
encoder_hidden_states,
encoder_attention_mask,
)
else:
layer_outputs = layer_module(
hidden_states,
attention_mask,
head_mask[i],
encoder_hidden_states,
encoder_attention_mask,
output_attentions,
)
layer_outputs = layer_module(
hidden_states,
attention_mask,
head_mask[i],
encoder_hidden_states,
encoder_attention_mask,
output_attentions,
)
hidden_states = layer_outputs[0]
if output_attentions:
@@ -657,6 +641,9 @@ class BertModel(BertPreTrainedModel):
def set_input_embeddings(self, value):
self.embeddings.word_embeddings = value
def get_layers(self):
return self.encoder.layer
def _prune_heads(self, heads_to_prune):
""" Prunes heads of the model.
heads_to_prune: dict of {layer_num: list of heads to prune in this layer}
+20
View File
@@ -14,6 +14,7 @@
# See the License for the specific language governing permissions and
# limitations under the License.
import functools
import inspect
import logging
import os
@@ -345,6 +346,20 @@ class PreTrainedModel(nn.Module, ModuleUtilsMixin, GenerationMixin):
"""
return None # Overwrite for models with output embeddings
def get_layers(self):
"""
Returns the model's transformer layers.
Returns:
:obj:`List` or :obj:`nn.ModuleList`:
A list or :obj:`nn.ModuleList` containing all transformer layers.
"""
base_model = getattr(self, self.base_model_prefix, self)
if base_model is not self:
return base_model.get_layers()
else:
raise NotImplementedError
def tie_weights(self):
"""
Tie the weights between the input embeddings and the output embeddings.
@@ -456,6 +471,11 @@ class PreTrainedModel(nn.Module, ModuleUtilsMixin, GenerationMixin):
# Tie weights if needed
self.tie_weights()
# Gradient accumulation if needed
if self.config.gradient_checkpointing:
for layer in self.get_layers():
layer.forward = functools.partial(torch.utils.checkpoint.checkpoint, layer.forward)
def prune_heads(self, heads_to_prune: Dict):
""" Prunes heads of the base model.
+2
View File
@@ -260,6 +260,8 @@ class AlbertModelTester:
@require_torch
class AlbertModelTest(ModelTesterMixin, unittest.TestCase):
test_gradient_checkpointing = True
all_model_classes = (
(
AlbertModel,
+2
View File
@@ -431,6 +431,8 @@ class BertModelTester:
@require_torch
class BertModelTest(ModelTesterMixin, unittest.TestCase):
test_gradient_checkpointing = True
all_model_classes = (
(
BertModel,
+26
View File
@@ -62,6 +62,7 @@ class ModelTesterMixin:
test_head_masking = True
test_missing_keys = True
is_encoder_decoder = False
test_gradient_checkpointing = False
def _prepare_for_class(self, inputs_dict, model_class):
if model_class in MODEL_FOR_MULTIPLE_CHOICE_MAPPING.values():
@@ -676,6 +677,31 @@ class ModelTesterMixin:
with torch.no_grad():
model(**inputs)
def test_model_gradient_checkpointing(self):
if not self.test_gradient_checkpointing:
return
for model_class in self.all_model_classes:
print(model_class)
config, inputs_dict = self.model_tester.prepare_config_and_inputs_for_common()
model = model_class(config)
model.to(torch_device)
model.eval()
outputs_no_checkpointing = model(**self._prepare_for_class(inputs_dict, model_class))
config.gradient_checkpointing = True
model_with_gc = model_class(config)
model_with_gc.load_state_dict(model.state_dict())
model_with_gc.to(torch_device)
model_with_gc.eval()
outputs_with_checkpointing = model_with_gc(**self._prepare_for_class(inputs_dict, model_class))
for output_no_checkpointing, output_with_checkpointing in zip(
outputs_no_checkpointing, outputs_with_checkpointing
):
if isinstance(output_with_checkpointing, torch.Tensor):
self.assertTrue(torch.allclose(output_no_checkpointing, output_with_checkpointing))
def test_lm_head_model_random_no_beam_search_generate(self):
config, inputs_dict = self.model_tester.prepare_config_and_inputs_for_common()
input_ids = inputs_dict["input_ids"] if "input_ids" in inputs_dict else inputs_dict["inputs"]