Team Ai
Apppublic

HarryLee/eCommerceImageCaptioning

sourceHugging Faceupdated 4y agoView on Hugging Face
2likes
checkpoint_utils.py876 linesDownload Raw Back to utils
1# Copyright (c) Facebook, Inc. and its affiliates.2#3# This source code is licensed under the MIT license found in the4# LICENSE file in the root directory of this source tree.5 6import ast7import collections8import contextlib9import logging10import numpy as np11import os12import re13import time14import traceback15import math16from collections import OrderedDict17from typing import Any, Dict, Optional, Union18 19import torch20from fairseq.dataclass.configs import CheckpointConfig21from fairseq.dataclass.utils import (22    convert_namespace_to_omegaconf,23    overwrite_args_by_name,24)25from fairseq.distributed.fully_sharded_data_parallel import FSDP, has_FSDP26from fairseq.file_io import PathManager27from fairseq.models import FairseqDecoder, FairseqEncoder28from omegaconf import DictConfig, open_dict, OmegaConf29 30from data import data_utils31 32logger = logging.getLogger(__name__)33 34 35def save_checkpoint(cfg: CheckpointConfig, trainer, epoch_itr, val_loss):36    from fairseq import meters37 38    # only one worker should attempt to create the required dir39    if trainer.data_parallel_rank == 0:40        os.makedirs(cfg.save_dir, exist_ok=True)41 42    prev_best = getattr(save_checkpoint, "best", val_loss)43    if val_loss is not None:44        best_function = max if cfg.maximize_best_checkpoint_metric else min45        save_checkpoint.best = best_function(val_loss, prev_best)46 47    if cfg.no_save:48        return49 50    trainer.consolidate_optimizer()  # TODO(SS): do we need this if no_save_optimizer_state51 52    if not trainer.should_save_checkpoint_on_current_rank:53        if trainer.always_call_state_dict_during_save_checkpoint:54            trainer.state_dict()55        return56 57    write_timer = meters.StopwatchMeter()58    write_timer.start()59 60    epoch = epoch_itr.epoch61    end_of_epoch = epoch_itr.end_of_epoch()62    updates = trainer.get_num_updates()63 64    logger.info(f"Preparing to save checkpoint for epoch {epoch} @ {updates} updates")65 66    def is_better(a, b):67        return a >= b if cfg.maximize_best_checkpoint_metric else a <= b68 69    suffix = trainer.checkpoint_suffix70    checkpoint_conds = collections.OrderedDict()71    checkpoint_conds["checkpoint{}{}.pt".format(epoch, suffix)] = (72        end_of_epoch and not cfg.no_epoch_checkpoints and epoch % cfg.save_interval == 073    )74    checkpoint_conds["checkpoint_{}_{}{}.pt".format(epoch, updates, suffix)] = (75        not end_of_epoch76        and cfg.save_interval_updates > 077        and updates % cfg.save_interval_updates == 078    )79    checkpoint_conds["checkpoint_best{}.pt".format(suffix)] = val_loss is not None and (80        not hasattr(save_checkpoint, "best")81        or is_better(val_loss, save_checkpoint.best)82    )83    if val_loss is not None and cfg.keep_best_checkpoints > 0:84        worst_best = getattr(save_checkpoint, "best", None)85        chkpts = checkpoint_paths(86            cfg.save_dir,87            pattern=r"checkpoint\.best_{}_(\d+\.?\d*){}\.pt".format(88                cfg.best_checkpoint_metric, suffix89            ),90        )91        if len(chkpts) > 0:92            p = chkpts[-1] if cfg.maximize_best_checkpoint_metric else chkpts[0]93            worst_best = float(p.rsplit("_")[-1].replace("{}.pt".format(suffix), ""))94        # add random digits to resolve ties95        with data_utils.numpy_seed(epoch, updates, val_loss):96            rand_sfx = np.random.randint(0, cfg.keep_best_checkpoints)97 98        checkpoint_conds[99            "checkpoint.best_{}_{:.3f}{}{}.pt".format(100                cfg.best_checkpoint_metric,101                val_loss,102                rand_sfx,103                suffix104            )105        ] = worst_best is None or is_better(val_loss, worst_best)106    checkpoint_conds[107        "checkpoint_last{}.pt".format(suffix)108    ] = not cfg.no_last_checkpoints109 110    extra_state = {"train_iterator": epoch_itr.state_dict(), "val_loss": val_loss}111    if hasattr(save_checkpoint, "best"):112        extra_state.update({"best": save_checkpoint.best})113 114    checkpoints = [115        os.path.join(cfg.save_dir, fn) for fn, cond in checkpoint_conds.items() if cond116    ]117    if len(checkpoints) > 0:118        trainer.save_checkpoint(checkpoints[0], extra_state)119        for cp in checkpoints[1:]:120            if cfg.write_checkpoints_asynchronously:121                # TODO[ioPath]: Need to implement a delayed asynchronous122                # file copying/moving feature.123                logger.warning(124                    f"ioPath is not copying {checkpoints[0]} to {cp} "125                    "since async write mode is on."126                )127            else:128                assert PathManager.copy(129                    checkpoints[0], cp, overwrite=True130                ), f"Failed to copy {checkpoints[0]} to {cp}"131 132        write_timer.stop()133        logger.info(134            "Saved checkpoint {} (epoch {} @ {} updates, score {}) (writing took {} seconds)".format(135                checkpoints[0], epoch, updates, val_loss, write_timer.sum136            )137        )138 139    if not end_of_epoch and cfg.keep_interval_updates > 0:140        # remove old checkpoints; checkpoints are sorted in descending order141        if cfg.keep_interval_updates_pattern == -1:142            checkpoints = checkpoint_paths(143                cfg.save_dir, pattern=r"checkpoint_\d+_(\d+){}\.pt".format(suffix)144            )145        else:146            checkpoints = checkpoint_paths(147                cfg.save_dir,148                pattern=r"checkpoint_\d+_(\d+){}\.pt".format(suffix),149                keep_match=True,150            )151            checkpoints = [152                x[0]153                for x in checkpoints154                if x[1] % cfg.keep_interval_updates_pattern != 0155            ]156 157        for old_chk in checkpoints[cfg.keep_interval_updates :]:158            if os.path.lexists(old_chk):159                os.remove(old_chk)160            elif PathManager.exists(old_chk):161                PathManager.rm(old_chk)162 163    if cfg.keep_last_epochs > 0:164        # remove old epoch checkpoints; checkpoints are sorted in descending order165        checkpoints = checkpoint_paths(166            cfg.save_dir, pattern=r"checkpoint(\d+){}\.pt".format(suffix)167        )168        for old_chk in checkpoints[cfg.keep_last_epochs :]:169            if os.path.lexists(old_chk):170                os.remove(old_chk)171            elif PathManager.exists(old_chk):172                PathManager.rm(old_chk)173 174    if cfg.keep_best_checkpoints > 0:175        # only keep the best N checkpoints according to validation metric176        checkpoints = checkpoint_paths(177            cfg.save_dir,178            pattern=r"checkpoint\.best_{}_(\d+\.?\d*){}\.pt".format(179                cfg.best_checkpoint_metric, suffix180            ),181        )182        if not cfg.maximize_best_checkpoint_metric:183            checkpoints = checkpoints[::-1]184        for old_chk in checkpoints[cfg.keep_best_checkpoints :]:185            if os.path.lexists(old_chk):186                os.remove(old_chk)187            elif PathManager.exists(old_chk):188                PathManager.rm(old_chk)189 190 191def load_checkpoint(cfg: CheckpointConfig, trainer, **passthrough_args):192    """193    Load a checkpoint and restore the training iterator.194 195    *passthrough_args* will be passed through to196    ``trainer.get_train_iterator``.197    """198 199    reset_optimizer = cfg.reset_optimizer200    reset_lr_scheduler = cfg.reset_lr_scheduler201    optimizer_overrides = ast.literal_eval(cfg.optimizer_overrides)202    reset_meters = cfg.reset_meters203    reset_dataloader = cfg.reset_dataloader204 205    if cfg.finetune_from_model is not None and (206        reset_optimizer or reset_lr_scheduler or reset_meters or reset_dataloader207    ):208        raise ValueError(209            "--finetune-from-model can not be set together with either --reset-optimizer"210            " or reset_lr_scheduler or reset_meters or reset_dataloader"211        )212 213    suffix = trainer.checkpoint_suffix214    if (215        cfg.restore_file == "checkpoint_last.pt"216    ):  # default value of restore_file is 'checkpoint_last.pt'217        checkpoint_path = os.path.join(218            cfg.save_dir, "checkpoint_last{}.pt".format(suffix)219        )220        first_launch = not PathManager.exists(checkpoint_path)221        if cfg.finetune_from_model is not None and first_launch:222            # if there is no last checkpoint to restore, start the finetune from pretrained model223            # else just use usual logic to load checkpoint, e.g. restart from last checkpoint and etc.224            if PathManager.exists(cfg.finetune_from_model):225                checkpoint_path = cfg.finetune_from_model226                reset_optimizer = True227                reset_lr_scheduler = True228                reset_meters = True229                reset_dataloader = True230                logger.info(231                    f"loading pretrained model from {checkpoint_path}: "232                    "optimizer, lr scheduler, meters, dataloader will be reset"233                )234            else:235                raise ValueError(236                    f"--funetune-from-model {cfg.finetune_from_model} does not exist"237                )238    elif suffix is not None:239        checkpoint_path = cfg.restore_file.replace(".pt", suffix + ".pt")240    else:241        checkpoint_path = cfg.restore_file242 243    if cfg.restore_file != "checkpoint_last.pt" and cfg.finetune_from_model:244        raise ValueError(245            "--finetune-from-model and --restore-file (non-default value) "246            "can not be specified together: " + str(cfg)247        )248 249    extra_state = trainer.load_checkpoint(250        checkpoint_path,251        reset_optimizer,252        reset_lr_scheduler,253        optimizer_overrides,254        reset_meters=reset_meters,255    )256 257    if (258        extra_state is not None259        and "best" in extra_state260        and not reset_optimizer261        and not reset_meters262    ):263        save_checkpoint.best = extra_state["best"]264 265    if extra_state is not None and not reset_dataloader:266        # restore iterator from checkpoint267        itr_state = extra_state["train_iterator"]268        epoch_itr = trainer.get_train_iterator(269            epoch=itr_state["epoch"], load_dataset=True, **passthrough_args270        )271        epoch_itr.load_state_dict(itr_state)272        _n = itr_state['iterations_in_epoch']273        offset = sum(len(_) for _ in epoch_itr.batch_sampler[:_n])274        epoch_itr.dataset.dataset._seek(offset=offset)275        true_num = int(math.ceil(len(epoch_itr.dataset) / 8)) * 8276        another_offset = ((epoch_itr.epoch - 1) * true_num + offset) // 8277        if hasattr(epoch_itr.dataset, 'pure_text_dataset'):278            text_offset = (2 * another_offset) % len(epoch_itr.dataset.pure_text_dataset)279            epoch_itr.dataset.pure_text_dataset._seek(offset=text_offset)280        if hasattr(epoch_itr.dataset, 'pure_image_dataset'):281            image_offset = another_offset % len(epoch_itr.dataset.pure_image_dataset)282            epoch_itr.dataset.pure_image_dataset._seek(offset=image_offset)283        if hasattr(epoch_itr.dataset, 'detection_dataset'):284            detection_offset = another_offset % len(epoch_itr.dataset.detection_dataset)285            epoch_itr.dataset.detection_dataset._seek(offset=detection_offset)286    else:287        epoch_itr = trainer.get_train_iterator(288            epoch=1, load_dataset=True, **passthrough_args289        )290 291    trainer.lr_step(epoch_itr.epoch)292 293    return extra_state, epoch_itr294 295 296def load_checkpoint_to_cpu(path, arg_overrides=None, load_on_all_ranks=False):297    """Loads a checkpoint to CPU (with upgrading for backward compatibility).298 299    If doing single-GPU training or if the checkpoint is only being loaded by at300    most one process on each node (current default behavior is for only rank 0301    to read the checkpoint from disk), load_on_all_ranks should be False to302    avoid errors from torch.distributed not having been initialized or303    torch.distributed.barrier() hanging.304 305    If all processes on each node may be loading the checkpoint306    simultaneously, load_on_all_ranks should be set to True to avoid I/O307    conflicts.308 309    There's currently no support for > 1 but < all processes loading the310    checkpoint on each node.311    """312    local_path = PathManager.get_local_path(path)313    # The locally cached file returned by get_local_path() may be stale for314    # remote files that are periodically updated/overwritten (ex:315    # checkpoint_last.pt) - so we remove the local copy, sync across processes316    # (if needed), and then download a fresh copy.317    if local_path != path and PathManager.path_requires_pathmanager(path):318        try:319            os.remove(local_path)320        except FileNotFoundError:321            # With potentially multiple processes removing the same file, the322            # file being missing is benign (missing_ok isn't available until323            # Python 3.8).324            pass325        if load_on_all_ranks:326            torch.distributed.barrier()327        local_path = PathManager.get_local_path(path)328 329    with open(local_path, "rb") as f:330        state = torch.load(f, map_location=torch.device("cpu"))331 332    if "args" in state and state["args"] is not None and arg_overrides is not None:333        args = state["args"]334        for arg_name, arg_val in arg_overrides.items():335            setattr(args, arg_name, arg_val)336 337    if "cfg" in state and state["cfg"] is not None:338 339        # hack to be able to set Namespace in dict config. this should be removed when we update to newer340        # omegaconf version that supports object flags, or when we migrate all existing models341        from omegaconf import _utils342 343        old_primitive = _utils.is_primitive_type344        _utils.is_primitive_type = lambda _: True345 346        state["cfg"] = OmegaConf.create(state["cfg"])347 348        _utils.is_primitive_type = old_primitive349        OmegaConf.set_struct(state["cfg"], True)350 351        if arg_overrides is not None:352            overwrite_args_by_name(state["cfg"], arg_overrides)353 354    state = _upgrade_state_dict(state)355    return state356 357 358def load_model_ensemble(359    filenames,360    arg_overrides: Optional[Dict[str, Any]] = None,361    task=None,362    strict=True,363    suffix="",364    num_shards=1,365    state=None,366):367    """Loads an ensemble of models.368 369    Args:370        filenames (List[str]): checkpoint files to load371        arg_overrides (Dict[str,Any], optional): override model args that372            were used during model training373        task (fairseq.tasks.FairseqTask, optional): task to use for loading374    """375    assert not (376        strict and num_shards > 1377    ), "Cannot load state dict with strict=True and checkpoint shards > 1"378    ensemble, args, _task = load_model_ensemble_and_task(379        filenames,380        arg_overrides,381        task,382        strict,383        suffix,384        num_shards,385        state,386    )387    return ensemble, args388 389 390def get_maybe_sharded_checkpoint_filename(391    filename: str, suffix: str, shard_idx: int, num_shards: int392) -> str:393    orig_filename = filename394    filename = filename.replace(".pt", suffix + ".pt")395    fsdp_filename = filename[:-3] + f"-shard{shard_idx}.pt"396    model_parallel_filename = orig_filename[:-3] + f"_part{shard_idx}.pt"397    if PathManager.exists(fsdp_filename):398        return fsdp_filename399    elif num_shards > 1:400        return model_parallel_filename401    else:402        return filename403 404 405def load_model_ensemble_and_task(406    filenames,407    arg_overrides: Optional[Dict[str, Any]] = None,408    task=None,409    strict=True,410    suffix="",411    num_shards=1,412    state=None,413):414    assert state is None or len(filenames) == 1415 416    from fairseq import tasks417 418    assert not (419        strict and num_shards > 1420    ), "Cannot load state dict with strict=True and checkpoint shards > 1"421    ensemble = []422    cfg = None423    for filename in filenames:424        orig_filename = filename425        model_shard_state = {"shard_weights": [], "shard_metadata": []}426        assert num_shards > 0427        st = time.time()428        for shard_idx in range(num_shards):429            filename = get_maybe_sharded_checkpoint_filename(430                orig_filename, suffix, shard_idx, num_shards431            )432 433            if not PathManager.exists(filename):434                raise IOError("Model file not found: {}".format(filename))435            if state is None:436                state = load_checkpoint_to_cpu(filename, arg_overrides)437            if "args" in state and state["args"] is not None:438                cfg = convert_namespace_to_omegaconf(state["args"])439            elif "cfg" in state and state["cfg"] is not None:440                cfg = state["cfg"]441            else:442                raise RuntimeError(443                    f"Neither args nor cfg exist in state keys = {state.keys()}"444                )445 446            if task is None:447                task = tasks.setup_task(cfg.task)448 449            if "task_state" in state:450                task.load_state_dict(state["task_state"])451 452            if "fsdp_metadata" in state and num_shards > 1:453                model_shard_state["shard_weights"].append(state["model"])454                model_shard_state["shard_metadata"].append(state["fsdp_metadata"])455                # check FSDP import before the code goes too far456                if not has_FSDP:457                    raise ImportError(458                        "Cannot find FullyShardedDataParallel. "459                        "Please install fairscale with: pip install fairscale"460                    )461                if shard_idx == num_shards - 1:462                    consolidated_model_state = FSDP.consolidate_shard_weights(463                        shard_weights=model_shard_state["shard_weights"],464                        shard_metadata=model_shard_state["shard_metadata"],465                    )466                    model = task.build_model(cfg.model)467                    model.load_state_dict(468                        consolidated_model_state, strict=strict, model_cfg=cfg.model469                    )470            else:471                # model parallel checkpoint or unsharded checkpoint472                model = task.build_model(cfg.model)473                model.load_state_dict(474                    state["model"], strict=strict, model_cfg=cfg.model475                )476 477            # reset state so it gets loaded for the next model in ensemble478            state = None479            if shard_idx % 10 == 0 and shard_idx > 0:480                elapsed = time.time() - st481                logger.info(482                    f"Loaded {shard_idx} shards in {elapsed:.2f}s, {elapsed / (shard_idx+1):.2f}s/shard"483                )484 485        # build model for ensemble486        ensemble.append(model)487    return ensemble, cfg, task488 489 490def checkpoint_paths(path, pattern=r"checkpoint(\d+)\.pt", keep_match=False):491    """Retrieves all checkpoints found in `path` directory.492 493    Checkpoints are identified by matching filename to the specified pattern. If494    the pattern contains groups, the result will be sorted by the first group in495    descending order.496    """497    pt_regexp = re.compile(pattern)498    files = PathManager.ls(path)499 500    entries = []501    for i, f in enumerate(files):502        m = pt_regexp.fullmatch(f)503        if m is not None:504            idx = float(m.group(1)) if len(m.groups()) > 0 else i505            entries.append((idx, m.group(0)))506    if keep_match:507        return [(os.path.join(path, x[1]), x[0]) for x in sorted(entries, reverse=True)]508    else:509        return [os.path.join(path, x[1]) for x in sorted(entries, reverse=True)]510 511 512def torch_persistent_save(obj, filename, async_write: bool = False):513    if async_write:514        with PathManager.opena(filename, "wb") as f:515            _torch_persistent_save(obj, f)516    else:517        with PathManager.open(filename, "wb") as f:518            _torch_persistent_save(obj, f)519        # if PathManager.supports_rename(filename):520        #     # do atomic save521        #     with PathManager.open(filename + ".tmp", "wb") as f:522        #         _torch_persistent_save(obj, f)523        #     PathManager.rename(filename + ".tmp", filename)524        # else:525        #     # fallback to non-atomic save526        #     with PathManager.open(filename, "wb") as f:527        #         _torch_persistent_save(obj, f)528 529 530def _torch_persistent_save(obj, f):531    if isinstance(f, str):532        with PathManager.open(f, "wb") as h:533            torch_persistent_save(obj, h)534        return535    for i in range(3):536        try:537            return torch.save(obj, f)538        except Exception:539            if i == 2:540                logger.error(traceback.format_exc())541                raise542 543 544def _upgrade_state_dict(state):545    """Helper for upgrading old model checkpoints."""546 547    # add optimizer_history548    if "optimizer_history" not in state:549        state["optimizer_history"] = [550            {"criterion_name": "CrossEntropyCriterion", "best_loss": state["best_loss"]}551        ]552        state["last_optimizer_state"] = state["optimizer"]553        del state["optimizer"]554        del state["best_loss"]555    # move extra_state into sub-dictionary556    if "epoch" in state and "extra_state" not in state:557        state["extra_state"] = {558            "epoch": state["epoch"],559            "batch_offset": state["batch_offset"],560            "val_loss": state["val_loss"],561        }562        del state["epoch"]563        del state["batch_offset"]564        del state["val_loss"]565    # reduce optimizer history's memory usage (only keep the last state)566    if "optimizer" in state["optimizer_history"][-1]:567        state["last_optimizer_state"] = state["optimizer_history"][-1]["optimizer"]568        for optim_hist in state["optimizer_history"]:569            del optim_hist["optimizer"]570    # record the optimizer class name571    if "optimizer_name" not in state["optimizer_history"][-1]:572        state["optimizer_history"][-1]["optimizer_name"] = "FairseqNAG"573    # move best_loss into lr_scheduler_state574    if "lr_scheduler_state" not in state["optimizer_history"][-1]:575        state["optimizer_history"][-1]["lr_scheduler_state"] = {576            "best": state["optimizer_history"][-1]["best_loss"]577        }578        del state["optimizer_history"][-1]["best_loss"]579    # keep track of number of updates580    if "num_updates" not in state["optimizer_history"][-1]:581        state["optimizer_history"][-1]["num_updates"] = 0582    # old model checkpoints may not have separate source/target positions583    if (584        "args" in state585        and hasattr(state["args"], "max_positions")586        and not hasattr(state["args"], "max_source_positions")587    ):588        state["args"].max_source_positions = state["args"].max_positions589        state["args"].max_target_positions = state["args"].max_positions590    # use stateful training data iterator591    if "train_iterator" not in state["extra_state"]:592        state["extra_state"]["train_iterator"] = {593            "epoch": state["extra_state"]["epoch"],594            "iterations_in_epoch": state["extra_state"].get("batch_offset", 0),595        }596 597    # backward compatibility, cfg updates598    if "args" in state and state["args"] is not None:599        # default to translation task600        if not hasattr(state["args"], "task"):601            state["args"].task = "translation"602        # --raw-text and --lazy-load are deprecated603        if getattr(state["args"], "raw_text", False):604            state["args"].dataset_impl = "raw"605        elif getattr(state["args"], "lazy_load", False):606            state["args"].dataset_impl = "lazy"607        # epochs start at 1608        if state["extra_state"]["train_iterator"] is not None:609            state["extra_state"]["train_iterator"]["epoch"] = max(610                state["extra_state"]["train_iterator"].get("epoch", 1), 1611            )612        # --remove-bpe ==> --postprocess613        if hasattr(state["args"], "remove_bpe"):614            state["args"].post_process = state["args"].remove_bpe615        # --min-lr ==> --stop-min-lr616        if hasattr(state["args"], "min_lr"):617            state["args"].stop_min_lr = state["args"].min_lr618            del state["args"].min_lr619        # binary_cross_entropy / kd_binary_cross_entropy => wav2vec criterion620        if (621            hasattr(state["args"], "criterion")622            and state["args"].criterion in [623                "binary_cross_entropy",624                "kd_binary_cross_entropy",625            ]626        ):627            state["args"].criterion = "wav2vec"628        # remove log_keys if it's None (criteria will supply a default value of [])629        if hasattr(state["args"], "log_keys") and state["args"].log_keys is None:630            delattr(state["args"], "log_keys")631        # speech_pretraining => audio pretraining632        if (633            hasattr(state["args"], "task")634            and state["args"].task == "speech_pretraining"635        ):636            state["args"].task = "audio_pretraining"637        # audio_cpc => wav2vec638        if hasattr(state["args"], "arch") and state["args"].arch == "audio_cpc":639            state["args"].arch = "wav2vec"640        # convert legacy float learning rate to List[float]641        if hasattr(state["args"], "lr") and isinstance(state["args"].lr, float):642            state["args"].lr = [state["args"].lr]643        # convert task data arg to a string instead of List[string]644        if (645            hasattr(state["args"], "data")646            and isinstance(state["args"].data, list)647            and len(state["args"].data) > 0648        ):649            state["args"].data = state["args"].data[0]650        # remove keys in state["args"] related to teacher-student learning651        for key in [652            "static_teachers",653            "static_teacher_weights",654            "dynamic_teachers",655            "dynamic_teacher_weights",656        ]:657            if key in state["args"]:658                delattr(state["args"], key)659 660        state["cfg"] = convert_namespace_to_omegaconf(state["args"])661 662    if "cfg" in state and state["cfg"] is not None:663        cfg = state["cfg"]664        with open_dict(cfg):665            # any upgrades for Hydra-based configs666            if (667                "task" in cfg668                and "eval_wer_config" in cfg.task669                and isinstance(cfg.task.eval_wer_config.print_alignment, bool)670            ):671                cfg.task.eval_wer_config.print_alignment = "hard"672            if "generation" in cfg and isinstance(cfg.generation.print_alignment, bool):673                cfg.generation.print_alignment = "hard" if cfg.generation.print_alignment else None674            if (675                "model" in cfg676                and "w2v_args" in cfg.model677                and cfg.model.w2v_args is not None678                and (679                    hasattr(cfg.model.w2v_args, "task") or "task" in cfg.model.w2v_args680                )681                and hasattr(cfg.model.w2v_args.task, "eval_wer_config")682                and cfg.model.w2v_args.task.eval_wer_config is not None683                and isinstance(684                    cfg.model.w2v_args.task.eval_wer_config.print_alignment, bool685                )686            ):687                cfg.model.w2v_args.task.eval_wer_config.print_alignment = "hard"688 689    return state690 691 692def prune_state_dict(state_dict, model_cfg: Optional[DictConfig]):693    """Prune the given state_dict if desired for LayerDrop694    (https://arxiv.org/abs/1909.11556).695 696    Training with LayerDrop allows models to be robust to pruning at inference697    time. This function prunes state_dict to allow smaller models to be loaded698    from a larger model and re-maps the existing state_dict for this to occur.699 700    It's called by functions that load models from checkpoints and does not701    need to be called directly.702    """703    arch = None704    if model_cfg is not None:705        arch = (706            model_cfg._name707            if isinstance(model_cfg, DictConfig)708            else getattr(model_cfg, "arch", None)709        )710 711    if not model_cfg or arch is None or arch == "ptt_transformer":712        # args should not be none, but don't crash if it is.713        return state_dict714 715    encoder_layers_to_keep = getattr(model_cfg, "encoder_layers_to_keep", None)716    decoder_layers_to_keep = getattr(model_cfg, "decoder_layers_to_keep", None)717 718    if not encoder_layers_to_keep and not decoder_layers_to_keep:719        return state_dict720 721    # apply pruning722    logger.info(723        "Pruning model to specified layer configuration - this works best if the model was trained with LayerDrop"724    )725 726    def create_pruning_pass(layers_to_keep, layer_name):727        keep_layers = sorted(728            int(layer_string) for layer_string in layers_to_keep.split(",")729        )730        mapping_dict = {}731        for i in range(len(keep_layers)):732            mapping_dict[str(keep_layers[i])] = str(i)733 734        regex = re.compile(r"^{layer}.*\.layers\.(\d+)".format(layer=layer_name))735        return {"substitution_regex": regex, "mapping_dict": mapping_dict}736 737    pruning_passes = []738    if encoder_layers_to_keep:739        pruning_passes.append(create_pruning_pass(encoder_layers_to_keep, "encoder"))740    if decoder_layers_to_keep:741        pruning_passes.append(create_pruning_pass(decoder_layers_to_keep, "decoder"))742 743    new_state_dict = {}744    for layer_name in state_dict.keys():745        match = re.search(r"\.layers\.(\d+)\.", layer_name)746        # if layer has no number in it, it is a supporting layer, such as an747        # embedding748        if not match:749            new_state_dict[layer_name] = state_dict[layer_name]750            continue751 752        # otherwise, layer should be pruned.753        original_layer_number = match.group(1)754        # figure out which mapping dict to replace from755        for pruning_pass in pruning_passes:756            if original_layer_number in pruning_pass["mapping_dict"] and pruning_pass[757                "substitution_regex"758            ].search(layer_name):759                new_layer_number = pruning_pass["mapping_dict"][original_layer_number]760                substitution_match = pruning_pass["substitution_regex"].search(761                    layer_name762                )763                new_state_key = (764                    layer_name[: substitution_match.start(1)]765                    + new_layer_number766                    + layer_name[substitution_match.end(1) :]767                )768                new_state_dict[new_state_key] = state_dict[layer_name]769 770    # Since layers are now pruned, *_layers_to_keep are no longer needed.771    # This is more of "It would make it work fix" rather than a proper fix.772    if isinstance(model_cfg, DictConfig):773        context = open_dict(model_cfg)774    else:775        context = contextlib.ExitStack()776    with context:777        if hasattr(model_cfg, "encoder_layers_to_keep"):778            model_cfg.encoder_layers_to_keep = None779        if hasattr(model_cfg, "decoder_layers_to_keep"):780            model_cfg.decoder_layers_to_keep = None781 782    return new_state_dict783 784 785def load_pretrained_component_from_model(786    component: Union[FairseqEncoder, FairseqDecoder], checkpoint: str787):788    """789    Load a pretrained FairseqEncoder or FairseqDecoder from checkpoint into the790    provided `component` object. If state_dict fails to load, there may be a791    mismatch in the architecture of the corresponding `component` found in the792    `checkpoint` file.793    """794    if not PathManager.exists(checkpoint):795        raise IOError("Model file not found: {}".format(checkpoint))796    state = load_checkpoint_to_cpu(checkpoint)797    if isinstance(component, FairseqEncoder):798        component_type = "encoder"799    elif isinstance(component, FairseqDecoder):800        component_type = "decoder"801    else:802        raise ValueError(803            "component to load must be either a FairseqEncoder or "804            "FairseqDecoder. Loading other component types are not supported."805        )806    component_state_dict = OrderedDict()807    for key in state["model"].keys():808        if key.startswith(component_type):809            # encoder.input_layers.0.0.weight --> input_layers.0.0.weight810            component_subkey = key[len(component_type) + 1 :]811            component_state_dict[component_subkey] = state["model"][key]812    component.load_state_dict(component_state_dict, strict=True)813    return component814 815 816def verify_checkpoint_directory(save_dir: str) -> None:817    if not os.path.exists(save_dir):818        os.makedirs(save_dir, exist_ok=True)819    temp_file_path = os.path.join(save_dir, "dummy")820    try:821        with open(temp_file_path, "w"):822            pass823    except OSError as e:824        logger.warning(825            "Unable to access checkpoint save directory: {}".format(save_dir)826        )827        raise e828    else:829        os.remove(temp_file_path)830 831 832def load_ema_from_checkpoint(fpath):833    """Loads exponential moving averaged (EMA) checkpoint from input and834    returns a model with ema weights.835 836    Args:837      fpath: A string path of checkpoint to load from.838 839    Returns:840      A dict of string keys mapping to various values. The 'model' key841      from the returned dict should correspond to an OrderedDict mapping842      string parameter names to torch Tensors.843    """844    params_dict = collections.OrderedDict()845    new_state = None846 847    with PathManager.open(fpath, 'rb') as f:848        new_state = torch.load(849            f,850            map_location=(851                lambda s, _: torch.serialization.default_restore_location(s, 'cpu')852            ),853        )854 855        # EMA model is stored in a separate "extra state"856        model_params = new_state['extra_state']['ema']857 858        for key in list(model_params.keys()):859            p = model_params[key]860            if isinstance(p, torch.HalfTensor):861                p = p.float()862            if key not in params_dict:863                params_dict[key] = p.clone()864                # NOTE: clone() is needed in case of p is a shared parameter865            else:866                raise ValueError("Key {} is repeated in EMA model params.".format(key))867 868        if len(params_dict) == 0:869            raise ValueError(870                f"Input checkpoint path '{fpath}' does not contain "871                "ema model weights, is this model trained with EMA?"872            )873 874    new_state['model'] = params_dict875    return new_state876