From 1b32c8e3e8379fe52be7ead5fabd1f08b3a28558 Mon Sep 17 00:00:00 2001 From: Marc Sun Date: Mon, 5 Jun 2023 19:30:40 +0000 Subject: [PATCH 1/5] Add check for tied parameters --- src/transformers/modeling_utils.py | 8 ++++++++ 1 file changed, 8 insertions(+) diff --git a/src/transformers/modeling_utils.py b/src/transformers/modeling_utils.py index fb142432863b..e7415e15f251 100644 --- a/src/transformers/modeling_utils.py +++ b/src/transformers/modeling_utils.py @@ -90,6 +90,7 @@ offload_weight, save_offload_index, set_module_tensor_to_device, + check_tied_parameters_on_same_device ) if version.parse(accelerate_version) > version.parse("0.11.0"): @@ -2824,6 +2825,12 @@ def from_pretrained(cls, pretrained_model_name_or_path: Optional[Union[str, os.P ) del device_map_without_lm_head + if device_map is not None: + model.tie_weights() + tied_params = find_tied_parameters(model) + # check if we don't have tied param in different devices + check_tied_parameters_on_same_device(tied_params,device_map) + if from_tf: if resolved_archive_file.endswith(".index"): # Load from a TensorFlow 1.X checkpoint - provided by original authors @@ -3015,6 +3022,7 @@ def _fix_key(key): unexpected_keys = list(set(loaded_keys) - set(expected_keys)) if find_tied_parameters is not None: + model.tie_weights() tied_params = find_tied_parameters(model) else: tied_params = [] From 9509e02c6f1970ca355ff81f9c358608774663ca Mon Sep 17 00:00:00 2001 From: Marc Sun Date: Mon, 5 Jun 2023 19:36:51 +0000 Subject: [PATCH 2/5] Fix style --- .../research_projects/deebert/src/modeling_highway_bert.py | 5 +---- .../movement-pruning/emmental/modeling_bert_masked.py | 5 +---- src/transformers/modeling_utils.py | 4 ++-- src/transformers/models/reformer/modeling_reformer.py | 7 +------ 4 files changed, 5 insertions(+), 16 deletions(-) diff --git a/examples/research_projects/deebert/src/modeling_highway_bert.py b/examples/research_projects/deebert/src/modeling_highway_bert.py index 2a881decbbd5..37d81248ed45 100644 --- a/examples/research_projects/deebert/src/modeling_highway_bert.py +++ b/examples/research_projects/deebert/src/modeling_highway_bert.py @@ -229,10 +229,7 @@ def forward( sequence_output = encoder_outputs[0] pooled_output = self.pooler(sequence_output) - outputs = ( - sequence_output, - pooled_output, - ) + encoder_outputs[ + outputs = (sequence_output, pooled_output,) + encoder_outputs[ 1: ] # add hidden_states and attentions if they are here return outputs # sequence_output, pooled_output, (hidden_states), (attentions), highway exits diff --git a/examples/research_projects/movement-pruning/emmental/modeling_bert_masked.py b/examples/research_projects/movement-pruning/emmental/modeling_bert_masked.py index d404bf49aaa6..4228050fe123 100644 --- a/examples/research_projects/movement-pruning/emmental/modeling_bert_masked.py +++ b/examples/research_projects/movement-pruning/emmental/modeling_bert_masked.py @@ -649,10 +649,7 @@ def forward( sequence_output = encoder_outputs[0] pooled_output = self.pooler(sequence_output) - outputs = ( - sequence_output, - pooled_output, - ) + encoder_outputs[ + outputs = (sequence_output, pooled_output,) + encoder_outputs[ 1: ] # add hidden_states and attentions if they are here return outputs # sequence_output, pooled_output, (hidden_states), (attentions) diff --git a/src/transformers/modeling_utils.py b/src/transformers/modeling_utils.py index e7415e15f251..ca4cb58a55ff 100644 --- a/src/transformers/modeling_utils.py +++ b/src/transformers/modeling_utils.py @@ -90,7 +90,7 @@ offload_weight, save_offload_index, set_module_tensor_to_device, - check_tied_parameters_on_same_device + check_tied_parameters_on_same_device, ) if version.parse(accelerate_version) > version.parse("0.11.0"): @@ -2829,7 +2829,7 @@ def from_pretrained(cls, pretrained_model_name_or_path: Optional[Union[str, os.P model.tie_weights() tied_params = find_tied_parameters(model) # check if we don't have tied param in different devices - check_tied_parameters_on_same_device(tied_params,device_map) + check_tied_parameters_on_same_device(tied_params, device_map) if from_tf: if resolved_archive_file.endswith(".index"): diff --git a/src/transformers/models/reformer/modeling_reformer.py b/src/transformers/models/reformer/modeling_reformer.py index 4bd29e78eeba..94339573f10e 100755 --- a/src/transformers/models/reformer/modeling_reformer.py +++ b/src/transformers/models/reformer/modeling_reformer.py @@ -891,12 +891,7 @@ def _get_relevant_hid_states_and_buckets( bucket_idx = _stable_argsort(concat_buckets, dim=-1) # bucket_idx has shape: BatchSize x NumAttnHeads x NumHashes x SequenceLength - assert bucket_idx.shape == ( - batch_size, - self.num_attention_heads, - num_hashes, - sequence_length, - ), ( + assert bucket_idx.shape == (batch_size, self.num_attention_heads, num_hashes, sequence_length,), ( f"bucket_idx should have shape {(batch_size, self.num_attention_heads, num_hashes, sequence_length)}, but" f" has shape {bucket_idx.shape}." ) From ac48bfeb0f57428b25297850390f9ce558bacb6f Mon Sep 17 00:00:00 2001 From: Marc Sun Date: Mon, 5 Jun 2023 19:43:47 +0000 Subject: [PATCH 3/5] fix style --- .../research_projects/deebert/src/modeling_highway_bert.py | 5 ++++- .../movement-pruning/emmental/modeling_bert_masked.py | 5 ++++- src/transformers/modeling_utils.py | 2 +- src/transformers/models/reformer/modeling_reformer.py | 7 ++++++- 4 files changed, 15 insertions(+), 4 deletions(-) diff --git a/examples/research_projects/deebert/src/modeling_highway_bert.py b/examples/research_projects/deebert/src/modeling_highway_bert.py index 37d81248ed45..2a881decbbd5 100644 --- a/examples/research_projects/deebert/src/modeling_highway_bert.py +++ b/examples/research_projects/deebert/src/modeling_highway_bert.py @@ -229,7 +229,10 @@ def forward( sequence_output = encoder_outputs[0] pooled_output = self.pooler(sequence_output) - outputs = (sequence_output, pooled_output,) + encoder_outputs[ + outputs = ( + sequence_output, + pooled_output, + ) + encoder_outputs[ 1: ] # add hidden_states and attentions if they are here return outputs # sequence_output, pooled_output, (hidden_states), (attentions), highway exits diff --git a/examples/research_projects/movement-pruning/emmental/modeling_bert_masked.py b/examples/research_projects/movement-pruning/emmental/modeling_bert_masked.py index 4228050fe123..d404bf49aaa6 100644 --- a/examples/research_projects/movement-pruning/emmental/modeling_bert_masked.py +++ b/examples/research_projects/movement-pruning/emmental/modeling_bert_masked.py @@ -649,7 +649,10 @@ def forward( sequence_output = encoder_outputs[0] pooled_output = self.pooler(sequence_output) - outputs = (sequence_output, pooled_output,) + encoder_outputs[ + outputs = ( + sequence_output, + pooled_output, + ) + encoder_outputs[ 1: ] # add hidden_states and attentions if they are here return outputs # sequence_output, pooled_output, (hidden_states), (attentions) diff --git a/src/transformers/modeling_utils.py b/src/transformers/modeling_utils.py index ca4cb58a55ff..6d4b82228210 100644 --- a/src/transformers/modeling_utils.py +++ b/src/transformers/modeling_utils.py @@ -85,12 +85,12 @@ from accelerate import __version__ as accelerate_version from accelerate import dispatch_model, infer_auto_device_map, init_empty_weights from accelerate.utils import ( + check_tied_parameters_on_same_device, find_tied_parameters, load_offloaded_weights, offload_weight, save_offload_index, set_module_tensor_to_device, - check_tied_parameters_on_same_device, ) if version.parse(accelerate_version) > version.parse("0.11.0"): diff --git a/src/transformers/models/reformer/modeling_reformer.py b/src/transformers/models/reformer/modeling_reformer.py index 94339573f10e..4bd29e78eeba 100755 --- a/src/transformers/models/reformer/modeling_reformer.py +++ b/src/transformers/models/reformer/modeling_reformer.py @@ -891,7 +891,12 @@ def _get_relevant_hid_states_and_buckets( bucket_idx = _stable_argsort(concat_buckets, dim=-1) # bucket_idx has shape: BatchSize x NumAttnHeads x NumHashes x SequenceLength - assert bucket_idx.shape == (batch_size, self.num_attention_heads, num_hashes, sequence_length,), ( + assert bucket_idx.shape == ( + batch_size, + self.num_attention_heads, + num_hashes, + sequence_length, + ), ( f"bucket_idx should have shape {(batch_size, self.num_attention_heads, num_hashes, sequence_length)}, but" f" has shape {bucket_idx.shape}." ) From d724054bc5ac2e991cd5d057ee13d633f7a7c4f0 Mon Sep 17 00:00:00 2001 From: Marc Sun Date: Mon, 5 Jun 2023 20:44:56 +0000 Subject: [PATCH 4/5] Fix versioning --- src/transformers/modeling_utils.py | 8 ++++++-- 1 file changed, 6 insertions(+), 2 deletions(-) diff --git a/src/transformers/modeling_utils.py b/src/transformers/modeling_utils.py index 6d4b82228210..0e3c6b19ba72 100644 --- a/src/transformers/modeling_utils.py +++ b/src/transformers/modeling_utils.py @@ -85,7 +85,6 @@ from accelerate import __version__ as accelerate_version from accelerate import dispatch_model, infer_auto_device_map, init_empty_weights from accelerate.utils import ( - check_tied_parameters_on_same_device, find_tied_parameters, load_offloaded_weights, offload_weight, @@ -97,6 +96,10 @@ from accelerate.utils import get_balanced_memory else: get_balanced_memory = None + if version.parse(accelerate_version) > version.parse("0.19.0"): + from accelerate.utils import check_tied_parameters_on_same_device + else: + check_tied_parameters_on_same_device = None else: find_tied_parameters = None @@ -2829,7 +2832,8 @@ def from_pretrained(cls, pretrained_model_name_or_path: Optional[Union[str, os.P model.tie_weights() tied_params = find_tied_parameters(model) # check if we don't have tied param in different devices - check_tied_parameters_on_same_device(tied_params, device_map) + if check_tied_parameters_on_same_device is not None: + check_tied_parameters_on_same_device(tied_params, device_map) if from_tf: if resolved_archive_file.endswith(".index"): From b65a007d75b5eac7e0e8ed19c8aef08333b84512 Mon Sep 17 00:00:00 2001 From: Marc Sun Date: Mon, 5 Jun 2023 21:05:33 +0000 Subject: [PATCH 5/5] Change if to elif --- src/transformers/modeling_utils.py | 2 +- 1 file changed, 1 insertion(+), 1 deletion(-) diff --git a/src/transformers/modeling_utils.py b/src/transformers/modeling_utils.py index 0e3c6b19ba72..4002f188e9ee 100644 --- a/src/transformers/modeling_utils.py +++ b/src/transformers/modeling_utils.py @@ -2828,7 +2828,7 @@ def from_pretrained(cls, pretrained_model_name_or_path: Optional[Union[str, os.P ) del device_map_without_lm_head - if device_map is not None: + elif device_map is not None: model.tie_weights() tied_params = find_tied_parameters(model) # check if we don't have tied param in different devices