| """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) |
|
|
| |
| 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) |
|
|