CoolFace
Apppublic

AD1TEYA/multiscale-keypoint-conditioned-LipReading

sourceHugging Facemitupdated 6mo agoView on Hugging Face
8likes
average_checkpoints.py37 linesDownload Raw Back to root
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