CoolFace
Apppublic

chendl/compositional_test

sourceHugging Faceupdated 3y agoView on Hugging Face
1likes
test_modeling_common.py3609 linesDownload Raw Back to tests
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 

Showing the first 1,200 of 3609 lines. Download the file for the rest.