Loading fairseq/utils.py +7 −4 Changes for fairseq/utils.py: 7 added lines, 4 removed lines. Original line number Diff line number Diff line Loading @@ -146,17 +146,20 @@ def load_ensemble_for_inference(filenames, task, model_arg_overrides=None): state = torch.load(filename, map_location=lambda s, l: default_restore_location(s, 'cpu')) state = _upgrade_state_dict(state) states.append(state) args = states[0]['args'] if model_arg_overrides is not None: args = _override_model_args(args, model_arg_overrides) # build ensemble ensemble = [] for state in states: args = state['args'] if model_arg_overrides is not None: args = _override_model_args(args, model_arg_overrides) # build model for ensemble model = task.build_model(args) model.upgrade_state_dict(state['model']) model.load_state_dict(state['model'], strict=True) ensemble.append(model) return ensemble, args Loading Loading
fairseq/utils.py +7 −4 Changes for fairseq/utils.py: 7 added lines, 4 removed lines. Original line number Diff line number Diff line Loading @@ -146,17 +146,20 @@ def load_ensemble_for_inference(filenames, task, model_arg_overrides=None): state = torch.load(filename, map_location=lambda s, l: default_restore_location(s, 'cpu')) state = _upgrade_state_dict(state) states.append(state) args = states[0]['args'] if model_arg_overrides is not None: args = _override_model_args(args, model_arg_overrides) # build ensemble ensemble = [] for state in states: args = state['args'] if model_arg_overrides is not None: args = _override_model_args(args, model_arg_overrides) # build model for ensemble model = task.build_model(args) model.upgrade_state_dict(state['model']) model.load_state_dict(state['model'], strict=True) ensemble.append(model) return ensemble, args Loading