Compare commits
4
Commits
| Author | SHA1 | Date | |
|---|---|---|---|
|
|
b7e417ab94 | ||
|
|
4aa3c76528 | ||
|
|
826a082c84 | ||
|
|
998721e1dc |
@@ -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)
|
||||
|
||||
@@ -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)
|
||||
|
||||
@@ -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}
|
||||
|
||||
@@ -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.
|
||||
|
||||
|
||||
@@ -260,6 +260,8 @@ class AlbertModelTester:
|
||||
@require_torch
|
||||
class AlbertModelTest(ModelTesterMixin, unittest.TestCase):
|
||||
|
||||
test_gradient_checkpointing = True
|
||||
|
||||
all_model_classes = (
|
||||
(
|
||||
AlbertModel,
|
||||
|
||||
@@ -431,6 +431,8 @@ class BertModelTester:
|
||||
@require_torch
|
||||
class BertModelTest(ModelTesterMixin, unittest.TestCase):
|
||||
|
||||
test_gradient_checkpointing = True
|
||||
|
||||
all_model_classes = (
|
||||
(
|
||||
BertModel,
|
||||
|
||||
@@ -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"]
|
||||
|
||||
Reference in New Issue
Block a user