chendl/compositional_test
1
1# coding=utf-82# Copyright 2019 HuggingFace Inc.3#4# Licensed under the Apache License, Version 2.0 (the "License");5# you may not use this file except in compliance with the License.6# You may obtain a copy of the License at7#8# http://www.apache.org/licenses/LICENSE-2.09#10# Unless required by applicable law or agreed to in writing, software11# distributed under the License is distributed on an "AS IS" BASIS,12# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.13# See the License for the specific language governing permissions and14# limitations under the License.15 16import copy17import gc18import glob19import inspect20import json21import os22import os.path23import pickle24import random25import sys26import tempfile27import unittest28import unittest.mock as mock29import warnings30from collections import defaultdict31from pathlib import Path32from typing import Dict, List, Tuple33 34import numpy as np35from huggingface_hub import HfFolder, delete_repo36from huggingface_hub.file_download import http_get37from pytest import mark38from requests.exceptions import HTTPError39 40import transformers41from transformers import (42 AutoConfig,43 AutoModel,44 AutoModelForSequenceClassification,45 PretrainedConfig,46 is_torch_available,47 logging,48)49from transformers.models.auto import get_values50from transformers.models.auto.modeling_auto import (51 MODEL_FOR_AUDIO_CLASSIFICATION_MAPPING_NAMES,52 MODEL_FOR_AUDIO_XVECTOR_MAPPING_NAMES,53 MODEL_FOR_BACKBONE_MAPPING_NAMES,54 MODEL_FOR_CAUSAL_IMAGE_MODELING_MAPPING_NAMES,55 MODEL_FOR_CAUSAL_LM_MAPPING_NAMES,56 MODEL_FOR_DOCUMENT_QUESTION_ANSWERING_MAPPING_NAMES,57 MODEL_FOR_IMAGE_CLASSIFICATION_MAPPING_NAMES,58 MODEL_FOR_MASKED_IMAGE_MODELING_MAPPING_NAMES,59 MODEL_FOR_MASKED_LM_MAPPING_NAMES,60 MODEL_FOR_MULTIPLE_CHOICE_MAPPING_NAMES,61 MODEL_FOR_NEXT_SENTENCE_PREDICTION_MAPPING_NAMES,62 MODEL_FOR_QUESTION_ANSWERING_MAPPING_NAMES,63 MODEL_FOR_SEMANTIC_SEGMENTATION_MAPPING_NAMES,64 MODEL_FOR_SEQ_TO_SEQ_CAUSAL_LM_MAPPING_NAMES,65 MODEL_FOR_SEQUENCE_CLASSIFICATION_MAPPING_NAMES,66 MODEL_FOR_TOKEN_CLASSIFICATION_MAPPING_NAMES,67 MODEL_FOR_VIDEO_CLASSIFICATION_MAPPING_NAMES,68 MODEL_MAPPING_NAMES,69)70from transformers.testing_utils import (71 TOKEN,72 USER,73 CaptureLogger,74 TestCasePlus,75 is_pt_flax_cross_test,76 is_pt_tf_cross_test,77 is_staging_test,78 require_accelerate,79 require_safetensors,80 require_torch,81 require_torch_gpu,82 require_torch_multi_gpu,83 require_usr_bin_time,84 slow,85 torch_device,86)87from transformers.utils import (88 CONFIG_NAME,89 GENERATION_CONFIG_NAME,90 SAFE_WEIGHTS_INDEX_NAME,91 SAFE_WEIGHTS_NAME,92 WEIGHTS_INDEX_NAME,93 WEIGHTS_NAME,94 is_accelerate_available,95 is_flax_available,96 is_tf_available,97 is_torch_fx_available,98)99from transformers.utils.generic import ModelOutput100 101 102sys.path.append(str(Path(__file__).parent.parent / "utils"))103 104from test_module.custom_configuration import CustomConfig, NoSuperInitConfig # noqa E402105 106 107if is_accelerate_available():108 from accelerate.utils import compute_module_sizes109 110 111if is_torch_available():112 import torch113 from test_module.custom_modeling import CustomModel, NoSuperInitModel114 from torch import nn115 116 from transformers import (117 BERT_PRETRAINED_MODEL_ARCHIVE_LIST,118 MODEL_MAPPING,119 AdaptiveEmbedding,120 AutoModelForCausalLM,121 AutoTokenizer,122 BertConfig,123 BertModel,124 CLIPTextModel,125 PreTrainedModel,126 T5Config,127 T5ForConditionalGeneration,128 )129 from transformers.modeling_utils import shard_checkpoint130 131 # Fake pretrained models for tests132 class BaseModel(PreTrainedModel):133 config_class = PretrainedConfig134 135 def __init__(self, config):136 super().__init__(config)137 self.linear = nn.Linear(4, 5)138 self.linear_2 = nn.Linear(5, 6)139 140 def forward(self, x):141 return self.linear_2(self.linear(x))142 143 class ModelWithHead(PreTrainedModel):144 base_model_prefix = "base"145 config_class = PretrainedConfig146 147 def _init_weights(self, module):148 pass149 150 def __init__(self, config):151 super().__init__(config)152 self.base = BaseModel(config)153 # linear is a common name between Base and Head on purpose.154 self.linear = nn.Linear(6, 3)155 self.linear2 = nn.Linear(3, 5)156 157 def forward(self, x):158 return self.linear2(self.linear(self.base(x)))159 160 161if is_tf_available():162 import tensorflow as tf163 164if is_flax_available():165 import jax.numpy as jnp166 167 from transformers.modeling_flax_pytorch_utils import (168 convert_pytorch_state_dict_to_flax,169 load_flax_weights_in_pytorch_model,170 )171 172if is_torch_fx_available():173 from transformers.utils.fx import symbolic_trace174 175 176def _config_zero_init(config):177 configs_no_init = copy.deepcopy(config)178 for key in configs_no_init.__dict__.keys():179 if "_range" in key or "_std" in key or "initializer_factor" in key or "layer_scale" in key:180 setattr(configs_no_init, key, 1e-10)181 if isinstance(getattr(configs_no_init, key, None), PretrainedConfig):182 no_init_subconfig = _config_zero_init(getattr(configs_no_init, key))183 setattr(configs_no_init, key, no_init_subconfig)184 return configs_no_init185 186 187TINY_T5 = "patrickvonplaten/t5-tiny-random"188TINY_BERT_FOR_TOKEN_CLASSIFICATION = "hf-internal-testing/tiny-bert-for-token-classification"189 190 191def _mock_init_weights(self, module):192 for name, param in module.named_parameters(recurse=False):193 # Use the first letter of the name to get a value and go from a <> -13 to z <> 12194 value = ord(name[0].lower()) - 110195 param.data.fill_(value)196 197 198def _mock_all_init_weights(self):199 # Prune heads if needed200 if self.config.pruned_heads:201 self.prune_heads(self.config.pruned_heads)202 203 import transformers.modeling_utils204 205 if transformers.modeling_utils._init_weights:206 for module in self.modules():207 module._is_hf_initialized = False208 # Initialize weights209 self.apply(self._initialize_weights)210 211 # Tie weights should be skipped when not initializing all weights212 # since from_pretrained(...) calls tie weights anyways213 self.tie_weights()214 215 216@require_torch217class ModelTesterMixin:218 model_tester = None219 all_model_classes = ()220 all_generative_model_classes = ()221 fx_compatible = False222 test_torchscript = True223 test_pruning = True224 test_resize_embeddings = True225 test_resize_position_embeddings = False226 test_head_masking = True227 test_mismatched_shapes = True228 test_missing_keys = True229 test_model_parallel = False230 is_encoder_decoder = False231 has_attentions = True232 model_split_percents = [0.5, 0.7, 0.9]233 234 def _prepare_for_class(self, inputs_dict, model_class, return_labels=False):235 inputs_dict = copy.deepcopy(inputs_dict)236 if model_class.__name__ in get_values(MODEL_FOR_MULTIPLE_CHOICE_MAPPING_NAMES):237 inputs_dict = {238 k: v.unsqueeze(1).expand(-1, self.model_tester.num_choices, -1).contiguous()239 if isinstance(v, torch.Tensor) and v.ndim > 1240 else v241 for k, v in inputs_dict.items()242 }243 elif model_class.__name__ in get_values(MODEL_FOR_AUDIO_XVECTOR_MAPPING_NAMES):244 inputs_dict.pop("attention_mask")245 246 if return_labels:247 if model_class.__name__ in get_values(MODEL_FOR_MULTIPLE_CHOICE_MAPPING_NAMES):248 inputs_dict["labels"] = torch.ones(self.model_tester.batch_size, dtype=torch.long, device=torch_device)249 elif model_class.__name__ in [250 *get_values(MODEL_FOR_QUESTION_ANSWERING_MAPPING_NAMES),251 *get_values(MODEL_FOR_DOCUMENT_QUESTION_ANSWERING_MAPPING_NAMES),252 ]:253 inputs_dict["start_positions"] = torch.zeros(254 self.model_tester.batch_size, dtype=torch.long, device=torch_device255 )256 inputs_dict["end_positions"] = torch.zeros(257 self.model_tester.batch_size, dtype=torch.long, device=torch_device258 )259 elif model_class.__name__ in [260 *get_values(MODEL_FOR_SEQUENCE_CLASSIFICATION_MAPPING_NAMES),261 *get_values(MODEL_FOR_NEXT_SENTENCE_PREDICTION_MAPPING_NAMES),262 *get_values(MODEL_FOR_IMAGE_CLASSIFICATION_MAPPING_NAMES),263 *get_values(MODEL_FOR_VIDEO_CLASSIFICATION_MAPPING_NAMES),264 *get_values(MODEL_FOR_AUDIO_CLASSIFICATION_MAPPING_NAMES),265 ]:266 inputs_dict["labels"] = torch.zeros(267 self.model_tester.batch_size, dtype=torch.long, device=torch_device268 )269 elif model_class.__name__ in [270 *get_values(MODEL_FOR_TOKEN_CLASSIFICATION_MAPPING_NAMES),271 *get_values(MODEL_FOR_CAUSAL_LM_MAPPING_NAMES),272 *get_values(MODEL_FOR_CAUSAL_IMAGE_MODELING_MAPPING_NAMES),273 *get_values(MODEL_FOR_MASKED_LM_MAPPING_NAMES),274 *get_values(MODEL_FOR_SEQ_TO_SEQ_CAUSAL_LM_MAPPING_NAMES),275 ]:276 inputs_dict["labels"] = torch.zeros(277 (self.model_tester.batch_size, self.model_tester.seq_length), dtype=torch.long, device=torch_device278 )279 elif model_class.__name__ in get_values(MODEL_FOR_MASKED_IMAGE_MODELING_MAPPING_NAMES):280 num_patches = self.model_tester.image_size // self.model_tester.patch_size281 inputs_dict["bool_masked_pos"] = torch.zeros(282 (self.model_tester.batch_size, num_patches**2), dtype=torch.long, device=torch_device283 )284 elif model_class.__name__ in get_values(MODEL_FOR_SEMANTIC_SEGMENTATION_MAPPING_NAMES):285 batch_size, num_channels, height, width = inputs_dict["pixel_values"].shape286 inputs_dict["labels"] = torch.zeros(287 [self.model_tester.batch_size, height, width], device=torch_device288 ).long()289 290 return inputs_dict291 292 def test_save_load(self):293 config, inputs_dict = self.model_tester.prepare_config_and_inputs_for_common()294 295 def check_save_load(out1, out2):296 # make sure we don't have nans297 out_2 = out2.cpu().numpy()298 out_2[np.isnan(out_2)] = 0299 300 out_1 = out1.cpu().numpy()301 out_1[np.isnan(out_1)] = 0302 max_diff = np.amax(np.abs(out_1 - out_2))303 self.assertLessEqual(max_diff, 1e-5)304 305 for model_class in self.all_model_classes:306 model = model_class(config)307 model.to(torch_device)308 model.eval()309 with torch.no_grad():310 first = model(**self._prepare_for_class(inputs_dict, model_class))[0]311 312 with tempfile.TemporaryDirectory() as tmpdirname:313 model.save_pretrained(tmpdirname)314 315 # the config file (and the generation config file, if it can generate) should be saved316 self.assertTrue(os.path.exists(os.path.join(tmpdirname, CONFIG_NAME)))317 self.assertEqual(318 model.can_generate(), os.path.exists(os.path.join(tmpdirname, GENERATION_CONFIG_NAME))319 )320 321 model = model_class.from_pretrained(tmpdirname)322 model.to(torch_device)323 with torch.no_grad():324 second = model(**self._prepare_for_class(inputs_dict, model_class))[0]325 326 if isinstance(first, tuple) and isinstance(second, tuple):327 for tensor1, tensor2 in zip(first, second):328 check_save_load(tensor1, tensor2)329 else:330 check_save_load(first, second)331 332 def test_from_pretrained_no_checkpoint(self):333 config, _ = self.model_tester.prepare_config_and_inputs_for_common()334 for model_class in self.all_model_classes:335 model = model_class(config)336 state_dict = model.state_dict()337 338 new_model = model_class.from_pretrained(339 pretrained_model_name_or_path=None, config=config, state_dict=state_dict340 )341 for p1, p2 in zip(model.parameters(), new_model.parameters()):342 self.assertTrue(torch.equal(p1, p2))343 344 def test_save_load_keys_to_ignore_on_save(self):345 config, inputs_dict = self.model_tester.prepare_config_and_inputs_for_common()346 347 for model_class in self.all_model_classes:348 model = model_class(config)349 _keys_to_ignore_on_save = getattr(model, "_keys_to_ignore_on_save", None)350 if _keys_to_ignore_on_save is None:351 continue352 353 # check the keys are in the original state_dict354 for k in _keys_to_ignore_on_save:355 self.assertIn(k, model.state_dict().keys(), "\n".join(model.state_dict().keys()))356 357 # check that certain keys didn't get saved with the model358 with tempfile.TemporaryDirectory() as tmpdirname:359 model.save_pretrained(tmpdirname)360 output_model_file = os.path.join(tmpdirname, WEIGHTS_NAME)361 state_dict_saved = torch.load(output_model_file)362 for k in _keys_to_ignore_on_save:363 self.assertNotIn(k, state_dict_saved.keys(), "\n".join(state_dict_saved.keys()))364 365 # Test we can load the state dict in the model, necessary for the checkpointing API in Trainer.366 load_result = model.load_state_dict(state_dict_saved, strict=False)367 self.assertTrue(368 len(load_result.missing_keys) == 0369 or set(load_result.missing_keys) == set(model._keys_to_ignore_on_save)370 )371 self.assertTrue(len(load_result.unexpected_keys) == 0)372 373 def test_gradient_checkpointing_backward_compatibility(self):374 config, inputs_dict = self.model_tester.prepare_config_and_inputs_for_common()375 376 for model_class in self.all_model_classes:377 if not model_class.supports_gradient_checkpointing:378 continue379 380 config.gradient_checkpointing = True381 model = model_class(config)382 self.assertTrue(model.is_gradient_checkpointing)383 384 def test_gradient_checkpointing_enable_disable(self):385 config, inputs_dict = self.model_tester.prepare_config_and_inputs_for_common()386 387 for model_class in self.all_model_classes:388 if not model_class.supports_gradient_checkpointing:389 continue390 391 # at init model should have gradient checkpointing disabled392 model = model_class(config)393 self.assertFalse(model.is_gradient_checkpointing)394 395 # check enable works396 model.gradient_checkpointing_enable()397 self.assertTrue(model.is_gradient_checkpointing)398 399 # check disable works400 model.gradient_checkpointing_disable()401 self.assertFalse(model.is_gradient_checkpointing)402 403 def test_save_load_fast_init_from_base(self):404 config, inputs_dict = self.model_tester.prepare_config_and_inputs_for_common()405 if config.__class__ not in MODEL_MAPPING:406 return407 base_class = MODEL_MAPPING[config.__class__]408 409 if isinstance(base_class, tuple):410 base_class = base_class[0]411 412 for model_class in self.all_model_classes:413 if model_class == base_class:414 continue415 416 # make a copy of model class to not break future tests417 # from https://stackoverflow.com/questions/9541025/how-to-copy-a-python-class418 class CopyClass(model_class):419 pass420 421 model_class_copy = CopyClass422 423 # make sure that all keys are expected for test424 model_class_copy._keys_to_ignore_on_load_missing = []425 426 # make init deterministic, but make sure that427 # non-initialized weights throw errors nevertheless428 model_class_copy._init_weights = _mock_init_weights429 model_class_copy.init_weights = _mock_all_init_weights430 431 model = base_class(config)432 state_dict = model.state_dict()433 434 # this will often delete a single weight of a multi-weight module435 # to test an edge case436 random_key_to_del = random.choice(list(state_dict.keys()))437 del state_dict[random_key_to_del]438 439 # check that certain keys didn't get saved with the model440 with tempfile.TemporaryDirectory() as tmpdirname:441 model.save_pretrained(tmpdirname)442 torch.save(state_dict, os.path.join(tmpdirname, "pytorch_model.bin"))443 444 model_fast_init = model_class_copy.from_pretrained(tmpdirname)445 model_slow_init = model_class_copy.from_pretrained(tmpdirname, _fast_init=False)446 # Before we test anything447 448 for key in model_fast_init.state_dict().keys():449 if isinstance(model_slow_init.state_dict()[key], torch.BoolTensor):450 max_diff = (model_slow_init.state_dict()[key] ^ model_fast_init.state_dict()[key]).sum().item()451 else:452 max_diff = (model_slow_init.state_dict()[key] - model_fast_init.state_dict()[key]).sum().item()453 self.assertLessEqual(max_diff, 1e-3, msg=f"{key} not identical")454 455 def test_save_load_fast_init_to_base(self):456 config, inputs_dict = self.model_tester.prepare_config_and_inputs_for_common()457 if config.__class__ not in MODEL_MAPPING:458 return459 base_class = MODEL_MAPPING[config.__class__]460 461 if isinstance(base_class, tuple):462 base_class = base_class[0]463 464 for model_class in self.all_model_classes:465 if model_class == base_class:466 continue467 468 # make a copy of model class to not break future tests469 # from https://stackoverflow.com/questions/9541025/how-to-copy-a-python-class470 class CopyClass(base_class):471 pass472 473 base_class_copy = CopyClass474 475 # make sure that all keys are expected for test476 base_class_copy._keys_to_ignore_on_load_missing = []477 478 # make init deterministic, but make sure that479 # non-initialized weights throw errors nevertheless480 base_class_copy._init_weights = _mock_init_weights481 base_class_copy.init_weights = _mock_all_init_weights482 483 model = model_class(config)484 state_dict = model.state_dict()485 486 # this will often delete a single weight of a multi-weight module487 # to test an edge case488 random_key_to_del = random.choice(list(state_dict.keys()))489 del state_dict[random_key_to_del]490 491 # check that certain keys didn't get saved with the model492 with tempfile.TemporaryDirectory() as tmpdirname:493 model.config.save_pretrained(tmpdirname)494 torch.save(state_dict, os.path.join(tmpdirname, "pytorch_model.bin"))495 496 model_fast_init = base_class_copy.from_pretrained(tmpdirname)497 model_slow_init = base_class_copy.from_pretrained(tmpdirname, _fast_init=False)498 499 for key in model_fast_init.state_dict().keys():500 if isinstance(model_slow_init.state_dict()[key], torch.BoolTensor):501 max_diff = torch.max(502 model_slow_init.state_dict()[key] ^ model_fast_init.state_dict()[key]503 ).item()504 else:505 max_diff = torch.max(506 torch.abs(model_slow_init.state_dict()[key] - model_fast_init.state_dict()[key])507 ).item()508 self.assertLessEqual(max_diff, 1e-3, msg=f"{key} not identical")509 510 def test_initialization(self):511 config, inputs_dict = self.model_tester.prepare_config_and_inputs_for_common()512 513 configs_no_init = _config_zero_init(config)514 for model_class in self.all_model_classes:515 model = model_class(config=configs_no_init)516 for name, param in model.named_parameters():517 if param.requires_grad:518 self.assertIn(519 ((param.data.mean() * 1e9).round() / 1e9).item(),520 [0.0, 1.0],521 msg=f"Parameter {name} of model {model_class} seems not properly initialized",522 )523 524 def test_determinism(self):525 config, inputs_dict = self.model_tester.prepare_config_and_inputs_for_common()526 527 def check_determinism(first, second):528 out_1 = first.cpu().numpy()529 out_2 = second.cpu().numpy()530 out_1 = out_1[~np.isnan(out_1)]531 out_2 = out_2[~np.isnan(out_2)]532 max_diff = np.amax(np.abs(out_1 - out_2))533 self.assertLessEqual(max_diff, 1e-5)534 535 for model_class in self.all_model_classes:536 model = model_class(config)537 model.to(torch_device)538 model.eval()539 with torch.no_grad():540 first = model(**self._prepare_for_class(inputs_dict, model_class))[0]541 second = model(**self._prepare_for_class(inputs_dict, model_class))[0]542 543 if isinstance(first, tuple) and isinstance(second, tuple):544 for tensor1, tensor2 in zip(first, second):545 check_determinism(tensor1, tensor2)546 else:547 check_determinism(first, second)548 549 def test_forward_signature(self):550 config, _ = self.model_tester.prepare_config_and_inputs_for_common()551 552 for model_class in self.all_model_classes:553 model = model_class(config)554 signature = inspect.signature(model.forward)555 # signature.parameters is an OrderedDict => so arg_names order is deterministic556 arg_names = [*signature.parameters.keys()]557 558 if model.config.is_encoder_decoder:559 expected_arg_names = [560 "input_ids",561 "attention_mask",562 "decoder_input_ids",563 "decoder_attention_mask",564 ]565 expected_arg_names.extend(566 ["head_mask", "decoder_head_mask", "cross_attn_head_mask", "encoder_outputs"]567 if "head_mask" and "decoder_head_mask" and "cross_attn_head_mask" in arg_names568 else ["encoder_outputs"]569 )570 self.assertListEqual(arg_names[: len(expected_arg_names)], expected_arg_names)571 else:572 expected_arg_names = ["input_ids"]573 self.assertListEqual(arg_names[:1], expected_arg_names)574 575 def test_training(self):576 if not self.model_tester.is_training:577 return578 579 for model_class in self.all_model_classes:580 config, inputs_dict = self.model_tester.prepare_config_and_inputs_for_common()581 config.return_dict = True582 583 if model_class.__name__ in [584 *get_values(MODEL_MAPPING_NAMES),585 *get_values(MODEL_FOR_BACKBONE_MAPPING_NAMES),586 ]:587 continue588 589 model = model_class(config)590 model.to(torch_device)591 model.train()592 inputs = self._prepare_for_class(inputs_dict, model_class, return_labels=True)593 loss = model(**inputs).loss594 loss.backward()595 596 def test_training_gradient_checkpointing(self):597 if not self.model_tester.is_training:598 return599 600 for model_class in self.all_model_classes:601 config, inputs_dict = self.model_tester.prepare_config_and_inputs_for_common()602 config.use_cache = False603 config.return_dict = True604 605 if (606 model_class.__name__607 in [*get_values(MODEL_MAPPING_NAMES), *get_values(MODEL_FOR_BACKBONE_MAPPING_NAMES)]608 or not model_class.supports_gradient_checkpointing609 ):610 continue611 model = model_class(config)612 model.to(torch_device)613 model.gradient_checkpointing_enable()614 model.train()615 inputs = self._prepare_for_class(inputs_dict, model_class, return_labels=True)616 loss = model(**inputs).loss617 loss.backward()618 619 def test_attention_outputs(self):620 if not self.has_attentions:621 self.skipTest(reason="Model does not output attentions")622 623 config, inputs_dict = self.model_tester.prepare_config_and_inputs_for_common()624 config.return_dict = True625 626 seq_len = getattr(self.model_tester, "seq_length", None)627 decoder_seq_length = getattr(self.model_tester, "decoder_seq_length", seq_len)628 encoder_seq_length = getattr(self.model_tester, "encoder_seq_length", seq_len)629 decoder_key_length = getattr(self.model_tester, "decoder_key_length", decoder_seq_length)630 encoder_key_length = getattr(self.model_tester, "key_length", encoder_seq_length)631 chunk_length = getattr(self.model_tester, "chunk_length", None)632 if chunk_length is not None and hasattr(self.model_tester, "num_hashes"):633 encoder_seq_length = encoder_seq_length * self.model_tester.num_hashes634 635 for model_class in self.all_model_classes:636 inputs_dict["output_attentions"] = True637 inputs_dict["output_hidden_states"] = False638 config.return_dict = True639 model = model_class(config)640 model.to(torch_device)641 model.eval()642 with torch.no_grad():643 outputs = model(**self._prepare_for_class(inputs_dict, model_class))644 attentions = outputs.encoder_attentions if config.is_encoder_decoder else outputs.attentions645 self.assertEqual(len(attentions), self.model_tester.num_hidden_layers)646 647 # check that output_attentions also work using config648 del inputs_dict["output_attentions"]649 config.output_attentions = True650 model = model_class(config)651 model.to(torch_device)652 model.eval()653 with torch.no_grad():654 outputs = model(**self._prepare_for_class(inputs_dict, model_class))655 attentions = outputs.encoder_attentions if config.is_encoder_decoder else outputs.attentions656 self.assertEqual(len(attentions), self.model_tester.num_hidden_layers)657 658 if chunk_length is not None:659 self.assertListEqual(660 list(attentions[0].shape[-4:]),661 [self.model_tester.num_attention_heads, encoder_seq_length, chunk_length, encoder_key_length],662 )663 else:664 self.assertListEqual(665 list(attentions[0].shape[-3:]),666 [self.model_tester.num_attention_heads, encoder_seq_length, encoder_key_length],667 )668 out_len = len(outputs)669 670 if self.is_encoder_decoder:671 correct_outlen = 5672 673 # loss is at first position674 if "labels" in inputs_dict:675 correct_outlen += 1 # loss is added to beginning676 # Question Answering model returns start_logits and end_logits677 if model_class.__name__ in [678 *get_values(MODEL_FOR_QUESTION_ANSWERING_MAPPING_NAMES),679 *get_values(MODEL_FOR_DOCUMENT_QUESTION_ANSWERING_MAPPING_NAMES),680 ]:681 correct_outlen += 1 # start_logits and end_logits instead of only 1 output682 if "past_key_values" in outputs:683 correct_outlen += 1 # past_key_values have been returned684 685 self.assertEqual(out_len, correct_outlen)686 687 # decoder attentions688 decoder_attentions = outputs.decoder_attentions689 self.assertIsInstance(decoder_attentions, (list, tuple))690 self.assertEqual(len(decoder_attentions), self.model_tester.num_hidden_layers)691 self.assertListEqual(692 list(decoder_attentions[0].shape[-3:]),693 [self.model_tester.num_attention_heads, decoder_seq_length, decoder_key_length],694 )695 696 # cross attentions697 cross_attentions = outputs.cross_attentions698 self.assertIsInstance(cross_attentions, (list, tuple))699 self.assertEqual(len(cross_attentions), self.model_tester.num_hidden_layers)700 self.assertListEqual(701 list(cross_attentions[0].shape[-3:]),702 [703 self.model_tester.num_attention_heads,704 decoder_seq_length,705 encoder_key_length,706 ],707 )708 709 # Check attention is always last and order is fine710 inputs_dict["output_attentions"] = True711 inputs_dict["output_hidden_states"] = True712 model = model_class(config)713 model.to(torch_device)714 model.eval()715 with torch.no_grad():716 outputs = model(**self._prepare_for_class(inputs_dict, model_class))717 718 if hasattr(self.model_tester, "num_hidden_states_types"):719 added_hidden_states = self.model_tester.num_hidden_states_types720 elif self.is_encoder_decoder:721 added_hidden_states = 2722 else:723 added_hidden_states = 1724 self.assertEqual(out_len + added_hidden_states, len(outputs))725 726 self_attentions = outputs.encoder_attentions if config.is_encoder_decoder else outputs.attentions727 728 self.assertEqual(len(self_attentions), self.model_tester.num_hidden_layers)729 if chunk_length is not None:730 self.assertListEqual(731 list(self_attentions[0].shape[-4:]),732 [self.model_tester.num_attention_heads, encoder_seq_length, chunk_length, encoder_key_length],733 )734 else:735 self.assertListEqual(736 list(self_attentions[0].shape[-3:]),737 [self.model_tester.num_attention_heads, encoder_seq_length, encoder_key_length],738 )739 740 @slow741 def test_torchscript_simple(self):742 config, inputs_dict = self.model_tester.prepare_config_and_inputs_for_common()743 self._create_and_check_torchscript(config, inputs_dict)744 745 @slow746 def test_torchscript_output_attentions(self):747 config, inputs_dict = self.model_tester.prepare_config_and_inputs_for_common()748 config.output_attentions = True749 self._create_and_check_torchscript(config, inputs_dict)750 751 @slow752 def test_torchscript_output_hidden_state(self):753 config, inputs_dict = self.model_tester.prepare_config_and_inputs_for_common()754 config.output_hidden_states = True755 self._create_and_check_torchscript(config, inputs_dict)756 757 # This is copied from `torch/testing/_internal/jit_utils.py::clear_class_registry`758 def clear_torch_jit_class_registry(self):759 torch._C._jit_clear_class_registry()760 torch.jit._recursive.concrete_type_store = torch.jit._recursive.ConcreteTypeStore()761 # torch 1.8 has no `_clear_class_state` in `torch.jit._state`762 if hasattr(torch.jit._state, "_clear_class_state"):763 torch.jit._state._clear_class_state()764 765 def _create_and_check_torchscript(self, config, inputs_dict):766 if not self.test_torchscript:767 return768 769 configs_no_init = _config_zero_init(config) # To be sure we have no Nan770 configs_no_init.torchscript = True771 for model_class in self.all_model_classes:772 model = model_class(config=configs_no_init)773 model.to(torch_device)774 model.eval()775 inputs = self._prepare_for_class(inputs_dict, model_class)776 777 main_input_name = model_class.main_input_name778 779 try:780 if model.config.is_encoder_decoder:781 model.config.use_cache = False # FSTM still requires this hack -> FSTM should probably be refactored similar to BART afterward782 main_input = inputs[main_input_name]783 attention_mask = inputs["attention_mask"]784 decoder_input_ids = inputs["decoder_input_ids"]785 decoder_attention_mask = inputs["decoder_attention_mask"]786 model(main_input, attention_mask, decoder_input_ids, decoder_attention_mask)787 traced_model = torch.jit.trace(788 model, (main_input, attention_mask, decoder_input_ids, decoder_attention_mask)789 )790 elif "bbox" in inputs and "image" in inputs: # LayoutLMv2 requires additional inputs791 input_ids = inputs["input_ids"]792 bbox = inputs["bbox"]793 image = inputs["image"].tensor794 model(input_ids, bbox, image)795 traced_model = torch.jit.trace(796 model, (input_ids, bbox, image), check_trace=False797 ) # when traced model is checked, an error is produced due to name mangling798 else:799 main_input = inputs[main_input_name]800 model(main_input)801 traced_model = torch.jit.trace(model, main_input)802 except RuntimeError:803 self.fail("Couldn't trace module.")804 805 with tempfile.TemporaryDirectory() as tmp_dir_name:806 pt_file_name = os.path.join(tmp_dir_name, "traced_model.pt")807 808 try:809 torch.jit.save(traced_model, pt_file_name)810 except Exception:811 self.fail("Couldn't save module.")812 813 try:814 loaded_model = torch.jit.load(pt_file_name)815 except Exception:816 self.fail("Couldn't load module.")817 818 model.to(torch_device)819 model.eval()820 821 loaded_model.to(torch_device)822 loaded_model.eval()823 824 model_state_dict = model.state_dict()825 loaded_model_state_dict = loaded_model.state_dict()826 827 non_persistent_buffers = {}828 for key in loaded_model_state_dict.keys():829 if key not in model_state_dict.keys():830 non_persistent_buffers[key] = loaded_model_state_dict[key]831 832 loaded_model_state_dict = {833 key: value for key, value in loaded_model_state_dict.items() if key not in non_persistent_buffers834 }835 836 self.assertEqual(set(model_state_dict.keys()), set(loaded_model_state_dict.keys()))837 838 model_buffers = list(model.buffers())839 for non_persistent_buffer in non_persistent_buffers.values():840 found_buffer = False841 for i, model_buffer in enumerate(model_buffers):842 if torch.equal(non_persistent_buffer, model_buffer):843 found_buffer = True844 break845 846 self.assertTrue(found_buffer)847 model_buffers.pop(i)848 849 models_equal = True850 for layer_name, p1 in model_state_dict.items():851 if layer_name in loaded_model_state_dict:852 p2 = loaded_model_state_dict[layer_name]853 if p1.data.ne(p2.data).sum() > 0:854 models_equal = False855 856 self.assertTrue(models_equal)857 858 # Avoid memory leak. Without this, each call increase RAM usage by ~20MB.859 # (Even with this call, there are still memory leak by ~0.04MB)860 self.clear_torch_jit_class_registry()861 862 def test_torch_fx(self):863 config, inputs_dict = self.model_tester.prepare_config_and_inputs_for_common()864 self._create_and_check_torch_fx_tracing(config, inputs_dict)865 866 def test_torch_fx_output_loss(self):867 config, inputs_dict = self.model_tester.prepare_config_and_inputs_for_common()868 self._create_and_check_torch_fx_tracing(config, inputs_dict, output_loss=True)869 870 def _create_and_check_torch_fx_tracing(self, config, inputs_dict, output_loss=False):871 if not is_torch_fx_available() or not self.fx_compatible:872 return873 874 configs_no_init = _config_zero_init(config) # To be sure we have no Nan875 configs_no_init.return_dict = False876 877 for model_class in self.all_model_classes:878 model = model_class(config=configs_no_init)879 model.to(torch_device)880 model.eval()881 inputs = self._prepare_for_class(inputs_dict, model_class, return_labels=output_loss)882 883 try:884 if model.config.is_encoder_decoder:885 model.config.use_cache = False # FSTM still requires this hack -> FSTM should probably be refactored similar to BART afterward886 labels = inputs.get("labels", None)887 input_names = [888 "attention_mask",889 "decoder_attention_mask",890 "decoder_input_ids",891 "input_features",892 "input_ids",893 "input_values",894 ]895 if labels is not None:896 input_names.append("labels")897 898 filtered_inputs = {k: v for (k, v) in inputs.items() if k in input_names}899 input_names = list(filtered_inputs.keys())900 901 model_output = model(**filtered_inputs)902 903 traced_model = symbolic_trace(model, input_names)904 traced_output = traced_model(**filtered_inputs)905 else:906 input_names = [907 "attention_mask",908 "bbox",909 "input_features",910 "input_ids",911 "input_values",912 "pixel_values",913 "token_type_ids",914 "visual_feats",915 "visual_pos",916 ]917 918 labels = inputs.get("labels", None)919 start_positions = inputs.get("start_positions", None)920 end_positions = inputs.get("end_positions", None)921 if labels is not None:922 input_names.append("labels")923 if start_positions is not None:924 input_names.append("start_positions")925 if end_positions is not None:926 input_names.append("end_positions")927 928 filtered_inputs = {k: v for (k, v) in inputs.items() if k in input_names}929 input_names = list(filtered_inputs.keys())930 931 if model.__class__.__name__ in set(MODEL_FOR_SEQUENCE_CLASSIFICATION_MAPPING_NAMES.values()) and (932 not hasattr(model.config, "problem_type") or model.config.problem_type is None933 ):934 model.config.problem_type = "single_label_classification"935 936 traced_model = symbolic_trace(model, input_names)937 traced_output = traced_model(**filtered_inputs)938 model_output = model(**filtered_inputs)939 940 except Exception as e:941 self.fail(f"Couldn't trace module: {e}")942 943 def flatten_output(output):944 flatten = []945 for x in output:946 if isinstance(x, (tuple, list)):947 flatten += flatten_output(x)948 elif not isinstance(x, torch.Tensor):949 continue950 else:951 flatten.append(x)952 return flatten953 954 model_output = flatten_output(model_output)955 traced_output = flatten_output(traced_output)956 num_outputs = len(model_output)957 958 for i in range(num_outputs):959 self.assertTrue(960 torch.allclose(model_output[i], traced_output[i]),961 f"traced {i}th output doesn't match model {i}th output for {model_class}",962 )963 964 # Test that the model can be serialized and restored properly965 with tempfile.TemporaryDirectory() as tmp_dir_name:966 pkl_file_name = os.path.join(tmp_dir_name, "model.pkl")967 try:968 with open(pkl_file_name, "wb") as f:969 pickle.dump(traced_model, f)970 with open(pkl_file_name, "rb") as f:971 loaded = pickle.load(f)972 except Exception as e:973 self.fail(f"Couldn't serialize / deserialize the traced model: {e}")974 975 loaded_output = loaded(**filtered_inputs)976 loaded_output = flatten_output(loaded_output)977 978 for i in range(num_outputs):979 self.assertTrue(980 torch.allclose(model_output[i], loaded_output[i]),981 f"serialized model {i}th output doesn't match model {i}th output for {model_class}",982 )983 984 # Avoid memory leak. Without this, each call increase RAM usage by ~20MB.985 # (Even with this call, there are still memory leak by ~0.04MB)986 self.clear_torch_jit_class_registry()987 988 def test_headmasking(self):989 if not self.test_head_masking:990 return991 992 global_rng.seed(42)993 config, inputs_dict = self.model_tester.prepare_config_and_inputs_for_common()994 global_rng.seed()995 996 inputs_dict["output_attentions"] = True997 config.output_hidden_states = True998 configs_no_init = _config_zero_init(config) # To be sure we have no Nan999 for model_class in self.all_model_classes:1000 model = model_class(config=configs_no_init)1001 model.to(torch_device)1002 model.eval()1003 1004 # Prepare head_mask1005 # Set require_grad after having prepared the tensor to avoid error (leaf variable has been moved into the graph interior)1006 head_mask = torch.ones(1007 self.model_tester.num_hidden_layers,1008 self.model_tester.num_attention_heads,1009 device=torch_device,1010 )1011 head_mask[0, 0] = 01012 head_mask[-1, :-1] = 01013 head_mask.requires_grad_(requires_grad=True)1014 inputs = self._prepare_for_class(inputs_dict, model_class).copy()1015 inputs["head_mask"] = head_mask1016 if model.config.is_encoder_decoder:1017 signature = inspect.signature(model.forward)1018 arg_names = [*signature.parameters.keys()]1019 if "decoder_head_mask" in arg_names: # necessary diferentiation because of T5 model1020 inputs["decoder_head_mask"] = head_mask1021 if "cross_attn_head_mask" in arg_names:1022 inputs["cross_attn_head_mask"] = head_mask1023 outputs = model(**inputs, return_dict=True)1024 1025 # Test that we can get a gradient back for importance score computation1026 output = sum(t.sum() for t in outputs[0])1027 output = output.sum()1028 output.backward()1029 multihead_outputs = head_mask.grad1030 1031 self.assertIsNotNone(multihead_outputs)1032 self.assertEqual(len(multihead_outputs), self.model_tester.num_hidden_layers)1033 1034 def check_attentions_validity(attentions):1035 # Remove Nan1036 for t in attentions:1037 self.assertLess(1038 torch.sum(torch.isnan(t)), t.numel() / 41039 ) # Check we don't have more than 25% nans (arbitrary)1040 attentions = [1041 t.masked_fill(torch.isnan(t), 0.0) for t in attentions1042 ] # remove them (the test is less complete)1043 1044 self.assertAlmostEqual(attentions[0][..., 0, :, :].flatten().sum().item(), 0.0)1045 self.assertNotEqual(attentions[0][..., -1, :, :].flatten().sum().item(), 0.0)1046 if len(attentions) > 2: # encoder-decoder models have only 2 layers in each module1047 self.assertNotEqual(attentions[1][..., 0, :, :].flatten().sum().item(), 0.0)1048 self.assertAlmostEqual(attentions[-1][..., -2, :, :].flatten().sum().item(), 0.0)1049 self.assertNotEqual(attentions[-1][..., -1, :, :].flatten().sum().item(), 0.0)1050 1051 if model.config.is_encoder_decoder:1052 check_attentions_validity(outputs.encoder_attentions)1053 check_attentions_validity(outputs.decoder_attentions)1054 check_attentions_validity(outputs.cross_attentions)1055 else:1056 check_attentions_validity(outputs.attentions)1057 1058 def test_head_pruning(self):1059 if not self.test_pruning:1060 return1061 1062 for model_class in self.all_model_classes:1063 (1064 config,1065 inputs_dict,1066 ) = self.model_tester.prepare_config_and_inputs_for_common()1067 1068 if "head_mask" in inputs_dict:1069 del inputs_dict["head_mask"]1070 1071 inputs_dict["output_attentions"] = True1072 config.output_hidden_states = False1073 model = model_class(config=config)1074 model.to(torch_device)1075 model.eval()1076 heads_to_prune = {1077 0: list(range(1, self.model_tester.num_attention_heads)),1078 -1: [0],1079 }1080 model.prune_heads(heads_to_prune)1081 with torch.no_grad():1082 outputs = model(**self._prepare_for_class(inputs_dict, model_class))1083 1084 attentions = outputs[-1]1085 1086 self.assertEqual(attentions[0].shape[-3], 1)1087 self.assertEqual(attentions[1].shape[-3], self.model_tester.num_attention_heads)1088 self.assertEqual(attentions[-1].shape[-3], self.model_tester.num_attention_heads - 1)1089 1090 def test_head_pruning_save_load_from_pretrained(self):1091 if not self.test_pruning:1092 return1093 1094 for model_class in self.all_model_classes:1095 (1096 config,1097 inputs_dict,1098 ) = self.model_tester.prepare_config_and_inputs_for_common()1099 1100 if "head_mask" in inputs_dict:1101 del inputs_dict["head_mask"]1102 1103 inputs_dict["output_attentions"] = True1104 config.output_hidden_states = False1105 model = model_class(config=config)1106 model.to(torch_device)1107 model.eval()1108 heads_to_prune = {1109 0: list(range(1, self.model_tester.num_attention_heads)),1110 -1: [0],1111 }1112 model.prune_heads(heads_to_prune)1113 1114 with tempfile.TemporaryDirectory() as temp_dir_name:1115 model.save_pretrained(temp_dir_name)1116 model = model_class.from_pretrained(temp_dir_name)1117 model.to(torch_device)1118 1119 with torch.no_grad():1120 outputs = model(**self._prepare_for_class(inputs_dict, model_class))1121 attentions = outputs[-1]1122 self.assertEqual(attentions[0].shape[-3], 1)1123 self.assertEqual(attentions[1].shape[-3], self.model_tester.num_attention_heads)1124 self.assertEqual(attentions[-1].shape[-3], self.model_tester.num_attention_heads - 1)1125 1126 def test_head_pruning_save_load_from_config_init(self):1127 if not self.test_pruning:1128 return1129 1130 for model_class in self.all_model_classes:1131 (1132 config,1133 inputs_dict,1134 ) = self.model_tester.prepare_config_and_inputs_for_common()1135 1136 if "head_mask" in inputs_dict:1137 del inputs_dict["head_mask"]1138 1139 inputs_dict["output_attentions"] = True1140 config.output_hidden_states = False1141 1142 heads_to_prune = {1143 0: list(range(1, self.model_tester.num_attention_heads)),1144 -1: [0],1145 }1146 config.pruned_heads = heads_to_prune1147 1148 model = model_class(config=config)1149 model.to(torch_device)1150 model.eval()1151 1152 with torch.no_grad():1153 outputs = model(**self._prepare_for_class(inputs_dict, model_class))1154 attentions = outputs[-1]1155 1156 self.assertEqual(attentions[0].shape[-3], 1)1157 self.assertEqual(attentions[1].shape[-3], self.model_tester.num_attention_heads)1158 self.assertEqual(attentions[-1].shape[-3], self.model_tester.num_attention_heads - 1)1159 1160 def test_head_pruning_integration(self):1161 if not self.test_pruning:1162 return1163 1164 for model_class in self.all_model_classes:1165 (1166 config,1167 inputs_dict,1168 ) = self.model_tester.prepare_config_and_inputs_for_common()1169 1170 if "head_mask" in inputs_dict:1171 del inputs_dict["head_mask"]1172 1173 inputs_dict["output_attentions"] = True1174 config.output_hidden_states = False1175 1176 heads_to_prune = {0: [0], 1: [1, 2]}1177 config.pruned_heads = heads_to_prune1178 1179 model = model_class(config=config)1180 model.to(torch_device)1181 model.eval()1182 1183 with torch.no_grad():1184 outputs = model(**self._prepare_for_class(inputs_dict, model_class))1185 attentions = outputs[-1]1186 1187 self.assertEqual(attentions[0].shape[-3], self.model_tester.num_attention_heads - 1)1188 self.assertEqual(attentions[1].shape[-3], self.model_tester.num_attention_heads - 2)1189 self.assertEqual(attentions[2].shape[-3], self.model_tester.num_attention_heads)1190 self.assertEqual(attentions[3].shape[-3], self.model_tester.num_attention_heads)1191 1192 with tempfile.TemporaryDirectory() as temp_dir_name:1193 model.save_pretrained(temp_dir_name)1194 model = model_class.from_pretrained(temp_dir_name)1195 model.to(torch_device)1196 1197 with torch.no_grad():1198 outputs = model(**self._prepare_for_class(inputs_dict, model_class))1199 attentions = outputs[-1]1200 