File size: 6,887 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 63 64 65 66 67 68 69 70 71 72 73 74 75 76 77 78 79 80 81 82 83 84 85 86 87 88 89 90 91 92 93 94 95 96 97 98 99 100 101 102 103 104 105 106 107 108 109 110 111 112 113 114 115 116 117 118 119 120 121 122 123 124 125 126 127 128 129 | """Minimal codec training on STREAMED CS2-10k — proves the streamed data -> MIRA model -> loss
pipeline end to end, in the existing venv (no pixi, no hydra). Pure L1 reconstruction so it needs
NO external weights; `require_dino_weights=False` builds the DINOv3 backbone with random weights
(swap in real gated DINOv3 later for a real RAE codec). Small config to fit 16GB and run fast.
"""
import os, sys, time, traceback
sys.path.insert(0, "src")
import torch
from mira.codec import VideoCodec, CodecLoss
from mira.codec.config import (
VideoCodecConfig, RAEEncoderConfig, ViTDecoderConfig,
StridedConvBottleneckConfig, ImageConfig,
)
from mira.codec.loss import CodecLossWeights
from mira.data.cs2_stream import create_cs2_loader
dev = "cuda" if torch.cuda.is_available() else "cpu"
# MIRA codec specs (paper Sec 6.3 / Table 9): per-view 288x512, 40-frame clips, DINOv3-L backbone,
# 1 latent per 32x32 px at half the input rate. CS2 is 48fps -> 24fps (48 isn't divisible by 20).
# Env knobs for 16GB tuning: CS2_T (frames, default 40), CS2_BS (batch, default 1), CS2_SUBSET.
T = int(os.environ.get("CS2_T", "40"))
FPS = 24
H, W = 288, 512
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),
# decoder size from env (defaults small/16GB; paper ViT-XL ~ 1152/28/16 on big GPUs)
vit_width=int(os.environ.get("CODEC_VIT_WIDTH", "512")),
vit_depth=int(os.environ.get("CODEC_VIT_DEPTH", "6")),
vit_num_heads=int(os.environ.get("CODEC_VIT_HEADS", "8")),
mlp_dim_multiplier=4,
qk_norm="layernorm", patch_size=16, patch_size_t=2, video=img, # 16*bottleneck_stride(2)=32 == encoder /32
activation_checkpointing=True,
)
cfg = VideoCodecConfig(encoder=enc, decoder=dec)
# Real gated DINOv3-L weights are used iff RS_DINO_WEIGHTS_DIR points at them (download from Meta:
# dinov3_vitl16_pretrain_lvd1689m-*.pth). Otherwise fall back to a random backbone so it still runs.
want_dino = bool(os.environ.get("RS_DINO_WEIGHTS_DIR"))
print(f"[cs2-codec] device={dev} T={T} {H}x{W}@{FPS}fps real_dino_weights={want_dino}", flush=True)
try:
model = VideoCodec(cfg, require_dino_weights=want_dino).to(dev).train()
except FileNotFoundError:
print("[cs2-codec] RS_DINO_WEIGHTS_DIR set but weights not found; using RANDOM backbone. "
"For a real codec, put dinov3_vitl16_pretrain_lvd1689m-*.pth in that dir.", flush=True)
model = VideoCodec(cfg, require_dino_weights=False).to(dev).train()
except Exception:
traceback.print_exc(); sys.exit(3)
print(f"[cs2-codec] params: {sum(p.numel() for p in model.parameters())/1e6:.1f}M", flush=True)
loss_fn = CodecLoss(CodecLossWeights(
loss_mae=1.0, loss_lpips_perceptual=0.0, loss_dino_latent_consistency=0.0, auto_weight=False,
)).to(dev)
opt = torch.optim.AdamW(model.parameters(), lr=1e-4)
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=int(os.environ.get("CS2_BS", "1")),
num_workers=int(os.environ.get("CS2_WORKERS", "0")), # >0 on Linux (Modal) to parallelize decode
infinite=True,
)
print("[cs2-codec] streaming CS2 sample split; first batch downloads a few clips...", flush=True)
from pathlib import Path
STEPS = int(os.environ.get("CS2_STEPS", "20000"))
CKPT_EVERY = int(os.environ.get("CS2_CKPT_EVERY", "500"))
ckpt_dir = Path(os.environ.get("CS2_OUT", "runs/cs2_codec"))
ckpt_dir.mkdir(parents=True, exist_ok=True)
print(f"[cs2-codec] REAL RUN: {STEPS} steps, checkpoint every {CKPT_EVERY} -> {ckpt_dir}", flush=True)
# Resume from the newest checkpoint (or CS2_RESUME) so restarts don't lose progress.
import glob
_start_step, _ema0 = 0, None
_resume = os.environ.get("CS2_RESUME") or (sorted(glob.glob(str(ckpt_dir / "codec_0*.pt"))) or [None])[-1]
def _trainable_sd():
# exclude the frozen DINOv3 backbone (reloaded from HF/hub) -> small, portable checkpoints
return {k: v for k, v in model.state_dict().items() if "rae_dino.dino_model" not in k}
if _resume and os.path.exists(_resume):
_sd = torch.load(_resume, map_location=dev, weights_only=False)
model.load_state_dict(_sd["model"], strict=False) # backbone already loaded (HF/hub)
_start_step = int(_sd.get("step", 0)); _ema0 = _sd.get("ema_loss")
print(f"[cs2-codec] resumed from {_resume} at step {_start_step} (ema_loss={_ema0})", flush=True)
# M.2 thermal duty-cycle: every CS2_PAUSE_EVERY steps, sleep CS2_PAUSE_SECS to let the SSD cool.
PAUSE_EVERY = int(os.environ.get("CS2_PAUSE_EVERY", "1000"))
PAUSE_SECS = int(os.environ.get("CS2_PAUSE_SECS", "900"))
step, ema = _start_step, _ema0
_pace_t0, _pace_s0 = time.time(), _start_step
for batch, meta in loader:
batch = batch.to(dev)
out = model(batch)
losses = loss_fn(out, global_step=step)
total = losses["loss"] if "loss" in losses else sum(v for v in losses.values() if v.ndim == 0)
opt.zero_grad(); total.backward(); opt.step()
ema = float(total) if ema is None else 0.98 * ema + 0.02 * float(total)
if step % 20 == 0:
print(f"[cs2-codec] step {step:6d} loss={float(total):.4f} ema={ema:.4f} "
f"({(time.time()-_pace_t0)/(step-_pace_s0+1):.2f}s/step)", flush=True)
step += 1
if step % CKPT_EVERY == 0:
torch.save({"step": step, "model": _trainable_sd(), "config": cfg.model_dump(),
"ema_loss": ema}, ckpt_dir / f"codec_{step:06d}.pt")
print(f"[cs2-codec] saved {ckpt_dir / f'codec_{step:06d}.pt'} (ema_loss={ema:.4f})", flush=True)
keep = int(os.environ.get("CS2_KEEP", "2")) # cap disk: keep only the newest N
for old in sorted(ckpt_dir.glob("codec_0*.pt"))[:-keep]:
old.unlink(missing_ok=True)
if step >= STEPS:
break
if PAUSE_EVERY and step % PAUSE_EVERY == 0: # checkpoint just saved (1000%500==0)
import datetime as _dt
print(f"[cs2-codec] PAUSE {PAUSE_SECS}s at step {step} for M.2 cooldown "
f"(resume ~{(_dt.datetime.now()+_dt.timedelta(seconds=PAUSE_SECS)).strftime('%H:%M:%S')})",
flush=True)
time.sleep(PAUSE_SECS)
_pace_t0, _pace_s0 = time.time(), step # exclude the sleep from the pace estimate
print(f"[cs2-codec] RESUME at step {step}", flush=True)
torch.save({"step": step, "model": _trainable_sd(), "config": cfg.model_dump(), "ema_loss": ema},
ckpt_dir / "codec_final.pt")
print(f"[cs2-codec] DONE {step} steps, final ema_loss={ema:.4f} -> {ckpt_dir/'codec_final.pt'}", flush=True)
|