HarryLee/eCommerceImageCaptioning
2
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 