File size: 3,320 Bytes
242cc21
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
22
23
24
25
26
27
28
29
30
31
32
33
34
35
36
37
38
39
40
41
42
43
44
45
46
47
48
49
50
51
52
53
54
55
56
57
58
59
60
61
62
"""Check codec rendering quality: load a checkpoint, reconstruct a streamed CS2 clip, and save
ground-truth | reconstruction frames side by side (+ PSNR). Point at a checkpoint via CS2_CKPT,
else the newest in runs/cs2_codec/.
"""
import os, sys, glob, math
os.environ.setdefault("RS_DINO_HF", "facebook/dinov3-vitl16-pretrain-lvd1689m")  # real backbone
sys.path.insert(0, "src")
import torch
from PIL import Image

from mira.codec import VideoCodec
from mira.codec.config import (
    VideoCodecConfig, RAEEncoderConfig, ViTDecoderConfig, StridedConvBottleneckConfig, ImageConfig,
)
from mira.codec.viz import visualize_side_by_side
from mira.data.cs2_stream import create_cs2_loader

dev = "cuda" if torch.cuda.is_available() else "cpu"
T = int(os.environ.get("CS2_T", "24")); H, W, FPS = 288, 512, 24

# same architecture as training (weights, incl. frozen DINOv3, come from the checkpoint)
img = ImageConfig(height=H, width=W, channels=3, timesteps=T, fps=FPS)
enc = RAEEncoderConfig(latent_dim=32, rae_model="dinov3_vitl16",
                       aggregation_layers=[11, 13, 15, 17, 19, 21, 23],
                       bottleneck=StridedConvBottleneckConfig(stride=2, temporal_stride=2, noise_tau=0.0),
                       compile_dino=False, video=img)
dec = ViTDecoderConfig(latent_dim=32, bottleneck=StridedConvBottleneckConfig(stride=2),
                       vit_width=512, vit_depth=6, vit_num_heads=8, mlp_dim_multiplier=4,
                       qk_norm="layernorm", patch_size=16, patch_size_t=2, video=img,
                       activation_checkpointing=False)
model = VideoCodec(VideoCodecConfig(encoder=enc, decoder=dec), require_dino_weights=False).to(dev).eval()

ckpt = os.environ.get("CS2_CKPT") or max(glob.glob("runs/cs2_codec/*.pt"), key=os.path.getmtime)
sd = torch.load(ckpt, map_location=dev, weights_only=False)
model.load_state_dict(sd["model"], strict=False)  # frozen backbone comes from HF, not the ckpt
print(f"[eval] loaded {ckpt} (step {sd.get('step')}, ema_loss {sd.get('ema_loss')})", flush=True)

loader = create_cs2_loader(subset=os.environ.get("CS2_SUBSET", "sample"), n_players=1, clip_len=T,
                           target_fps=FPS, frame_size=(H, W), batch_size=1, num_workers=0, infinite=True)
batch, meta = next(iter(loader))
with torch.no_grad():
    out = model(batch.to(dev))

# PSNR over the whole clip (input/output are ~[-1,1] -> [0,1])
gt = (out.input_video * 0.5 + 0.5).clamp(0, 1)
pr = (out.output_video * 0.5 + 0.5).clamp(0, 1)
mse = torch.mean((gt - pr) ** 2).item()
psnr = 10 * math.log10(1.0 / mse) if mse > 0 else float("inf")
print(f"[eval] reconstruction PSNR = {psnr:.2f} dB  (higher is better; ~20 rough, 28+ good)", flush=True)

viz = visualize_side_by_side(out)["viz_video"]        # (B,T,C,H,2W) uint8, GT|recon
viz = viz[0]                                          # (T,C,H,2W)
outdir = os.path.join(os.path.dirname(ckpt), "recon"); os.makedirs(outdir, exist_ok=True)
saved = []
for i in {0, T // 2, T - 1}:
    im = viz[i].permute(1, 2, 0).cpu().numpy()        # (H, 2W, C)
    p = os.path.join(outdir, f"step{sd.get('step')}_frame{i:02d}_gt-left_recon-right.png")
    Image.fromarray(im).save(p); saved.append(p)
print("[eval] saved (left = ground truth, right = reconstruction):", flush=True)
for p in saved:
    print("   ", p, flush=True)