| import os, glob, sys |
| import torch |
| from torchvision.io import read_video |
| from torchvision.utils import save_image |
|
|
| ROOT = "workspace/outputs/gen_4chunk_motion_delta_vs_multires" |
| OUT = "workspace/outputs/gen_4chunk_motion_delta_vs_multires/_grids" |
| os.makedirs(OUT, exist_ok=True) |
| FRAMES = [0, 6, 12, 18] |
|
|
| def load(path): |
| v, _, _ = read_video(path, pts_unit="sec", output_format="TCHW") |
| return v.float() / 255.0 |
|
|
| |
| keys = ["skier", "golden retriever", "Mercedes", "flyby video", "Aerial drone"] |
| pairs = [] |
| for s in ["train", "ood"]: |
| for mp in sorted(glob.glob(os.path.join(ROOT, "motion_delta", s, "*.mp4"))): |
| name = os.path.basename(mp) |
| if any(k.lower() in name.lower() for k in keys): |
| rp = os.path.join(ROOT, "multires", s, name) |
| if os.path.exists(rp): |
| pairs.append((s, name, mp, rp)) |
|
|
| for s, name, mp, rp in pairs: |
| vm, vr = load(mp), load(rp) |
| n = min(vm.shape[0], vr.shape[0]) |
| idx = [i for i in FRAMES if i < n] |
| rows = torch.cat([vm[idx], vr[idx]], dim=0) |
| tag = (s + "_" + name[:32]).replace(" ", "_").replace("/", "_").replace(".", "") |
| outp = os.path.join(OUT, tag + ".png") |
| save_image(rows, outp, nrow=len(idx)) |
| print(f"{outp} frames={idx} (top=motion_delta bottom=multires)") |
|
|