AD1TEYA/multiscale-keypoint-conditioned-LipReading
8
1import os2 3import torch4 5 6def average_checkpoints(last):7 avg = None8 for path in last:9 states = torch.load(path, map_location=lambda storage, loc: storage)["state_dict"]10 states = {k[6:]: v for k, v in states.items() if k.startswith("model.")}11 if avg is None:12 avg = states13 else:14 for k in avg.keys():15 avg[k] += states[k]16 # average17 for k in avg.keys():18 if avg[k] is not None:19 if avg[k].is_floating_point():20 avg[k] /= len(last)21 else:22 avg[k] //= len(last)23 return avg24 25 26def ensemble(args):27 last = [28 os.path.join(args.exp_dir, args.exp_name, f"epoch={n}.ckpt")29 for n in range(30 args.max_epochs - 10,31 args.max_epochs,32 )33 ]34 model_path = os.path.join(args.exp_dir, args.exp_name, f"model_avg_10.pth")35 torch.save(average_checkpoints(last), model_path)36 return model_path37 