File size: 2,735 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
"""Consolidate our decoder-only CS2 codec checkpoint into the format MIRA's world model expects:
a full VideoCodec state_dict (real HF DINOv3 backbone + our trained decoder) + a codec_config.yaml
alongside, so LatentWorldModel can load it via VideoCodec.load_from_checkpoint.
"""
import os, sys
os.environ.setdefault("RS_DINO_HF", "facebook/dinov3-vitl16-pretrain-lvd1689m")
sys.path.insert(0, "src")
import torch
from omegaconf import OmegaConf
from mira.codec import VideoCodec
from mira.codec.config import (
    VideoCodecConfig, RAEEncoderConfig, ViTDecoderConfig, StridedConvBottleneckConfig, ImageConfig,
)

SRC = os.environ.get("CS2_CODEC_CKPT", "runs/cs2_codec_dino/codec_036000.pt")
OUT = os.environ.get("CS2_CODEC_OUT", "runs/cs2_codec_consolidated")
T, H, W = int(os.environ.get("CS2_T", "16")), 288, 512

img = ImageConfig(height=H, width=W, channels=3, timesteps=T, fps=24)
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=True)
cfg = VideoCodecConfig(encoder=enc, decoder=dec)

print("[consolidate] building codec (real HF DINOv3 backbone) + loading trained decoder...", flush=True)
model = VideoCodec(cfg, require_dino_weights=False).eval()
sd = torch.load(SRC, map_location="cpu", weights_only=False)
missing, unexpected = model.load_state_dict(sd["model"], strict=False)
print(f"[consolidate] loaded decoder (missing={len(missing)} keys are the frozen backbone; "
      f"unexpected={len(unexpected)})", flush=True)

os.makedirs(OUT, exist_ok=True)
OmegaConf.save(OmegaConf.create({"model": {"architecture": {"config": cfg.model_dump()}}}),
               os.path.join(OUT, VideoCodec.CONFIG_FILENAME))
ckpt_path = os.path.join(OUT, "codec.pth")
torch.save({"state_dict": model.state_dict()}, ckpt_path)
print(f"[consolidate] wrote {ckpt_path} + {VideoCodec.CONFIG_FILENAME}", flush=True)

# verify round-trip through MIRA's loader
loaded = VideoCodec.load_from_checkpoint(ckpt_path, device="cpu")
print(f"[ok] round-trips via load_from_checkpoint: "
      f"{sum(p.numel() for p in loaded.parameters())/1e6:.0f}M params, "
      f"latent_dim={loaded.latent_dim}, td={loaded.temporal_downsampling}, "
      f"sd={loaded.spatial_downsampling}", flush=True)