| """Train a small MIRA latent world model on the CS2 codec, and test action-controllability. |
| |
| train: python scripts/train_cs2_wm.py |
| test : python scripts/train_cs2_wm.py --test (seed real context, feed fixed actions, generate) |
| |
| Frozen codec is loaded from the consolidated checkpoint; a tiny flow-matching diffusion transformer |
| learns to predict the next codec latent conditioned on CS2 keyboard+mouse actions. |
| """ |
| import os, sys, time, argparse |
| os.environ.setdefault("RS_DINO_HF", "facebook/dinov3-vitl16-pretrain-lvd1689m") |
| sys.path.insert(0, "src") |
| import torch |
|
|
| from mira.world_model.config import LatentWorldModelConfig |
| from mira.world_model.latent_world_model import LatentWorldModel |
| from mira.world_model.actions_config import ActionConfig |
| from mira.ml.image_config import ImageConfig |
| from mira.data.cs2_stream import create_cs2_loader, CS2_KEYS |
|
|
| dev = "cuda" if torch.cuda.is_available() else "cpu" |
| T = int(os.environ.get("CS2_T", "16")) |
| CODEC = os.environ.get("CS2_CODEC", "runs/cs2_codec_consolidated/codec.pth") |
| OUT = os.environ.get("CS2_WM_OUT", "runs/cs2_wm") |
|
|
|
|
| def build_model(): |
| actions = ActionConfig(valid_keys=list(CS2_KEYS), source_fps=24, target_fps=24) |
| video = ImageConfig(height=288, width=512, channels=3, timesteps=T, fps=24) |
| cfg = LatentWorldModelConfig( |
| actions=actions, video=video, codec_checkpoint=CODEC, |
| latent_mean_std=[0.0, 1.0], |
| causal=True, use_clean_past=True, learned_temporal_pool=True, |
| use_codec_posterior_mean=True, attention_gating=True, ada_attn_ln=True, |
| patch_size=1, n_register_tokens=0, n_context_frames=max(2, T // 2 - 2), |
| dropout_action_prob=0.1, |
| |
| hidden_dim=int(os.environ.get("WM_HIDDEN", "384")), |
| n_head=int(os.environ.get("WM_HEADS", "6")), |
| n_kv_head=int(os.environ.get("WM_KV", "2")), |
| n_layers=int(os.environ.get("WM_LAYERS", "4")), |
| time_attention_every=int(os.environ.get("WM_TATT", "2")), |
| activation_checkpointing=os.environ.get("WM_ACT_CKPT", "0") == "1", |
| ) |
| return LatentWorldModel(cfg).to(dev) |
|
|
|
|
| def train(): |
| from pathlib import Path |
| model = build_model().train() |
| n_tr = sum(p.numel() for p in model.parameters() if p.requires_grad) / 1e6 |
| print(f"[wm] trainable params: {n_tr:.1f}M (+ frozen codec)", flush=True) |
| opt = torch.optim.AdamW([p for p in model.parameters() if p.requires_grad], lr=1e-4) |
| loader = create_cs2_loader(subset=os.environ.get("CS2_SUBSET", "sample"), n_players=1, clip_len=T, |
| target_fps=24, frame_size=(288, 512), |
| batch_size=int(os.environ.get("WM_BATCH", "1")), |
| num_workers=int(os.environ.get("CS2_WORKERS", "0")), infinite=True) |
| ckpt_dir = Path(OUT); ckpt_dir.mkdir(parents=True, exist_ok=True) |
| STEPS = int(os.environ.get("CS2_WM_STEPS", "5000")) |
| step, ema, t0 = 0, None, time.time() |
| for batch, meta in loader: |
| losses = model(batch.to(dev)) |
| total = losses["loss_total"] |
| 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"[wm] step {step:5d} loss={float(total):.4f} ema={ema:.4f} " |
| f"({(time.time()-t0)/(step+1):.2f}s/step)", flush=True) |
| step += 1 |
| if step % 500 == 0: |
| |
| wm_only = {k: v for k, v in model.state_dict().items() if not k.startswith("codec.")} |
| torch.save({"step": step, "model": wm_only, "ema": ema}, ckpt_dir / f"wm_{step:06d}.pt") |
| for old in sorted(ckpt_dir.glob("wm_0*.pt"))[:-2]: |
| old.unlink(missing_ok=True) |
| print(f"[wm] saved wm_{step:06d}.pt (ema={ema:.4f})", flush=True) |
| if step >= STEPS: |
| break |
| print("[wm] done", flush=True) |
|
|
|
|
| @torch.no_grad() |
| def test(): |
| """Seed a real context clip, then generate under three fixed action regimes and save frames.""" |
| import glob, numpy as np |
| from PIL import Image |
| from mira.world_model.latent_world_model import WorldModelInferenceConfig |
| model = build_model().eval() |
| ck = os.environ.get("CS2_WM_CKPT") or (sorted(glob.glob(f"{OUT}/wm_0*.pt")) or [None])[-1] |
| if ck: |
| model.load_state_dict(torch.load(ck, map_location=dev, weights_only=False)["model"], strict=False) |
| print(f"[wm-test] loaded {ck}", flush=True) |
| loader = create_cs2_loader(subset="sample", n_players=1, clip_len=T, target_fps=24, |
| frame_size=(288, 512), batch_size=1, num_workers=0, infinite=True) |
| batch, _ = next(iter(loader)) |
|
|
| |
| kW, kMOUSE = CS2_KEYS.index("W"), None |
| regimes = { |
| "idle": lambda a: (a.key_presses.zero_(), a.mouse_movements.zero_()), |
| "forward": lambda a: (a.key_presses.zero_().__setitem__((slice(None), slice(None), kW), 1), |
| a.mouse_movements.zero_()), |
| "turn_right": lambda a: (a.key_presses.zero_(), |
| a.mouse_movements.zero_().__setitem__((slice(None), slice(None), 0), 25.0)), |
| } |
| icfg = WorldModelInferenceConfig() if False else None |
| outdir = os.path.join(OUT, "control"); os.makedirs(outdir, exist_ok=True) |
| for name, override in regimes.items(): |
| b = batch.clone() |
| override(b.actions) |
| out = model.inference(b.to(dev), progress_bar=False) |
| vid = out.output_video[0] |
| img = ((vid[-1] * 0.5 + 0.5).clamp(0, 1) * 255).byte().permute(1, 2, 0).cpu().numpy() |
| p = os.path.join(outdir, f"gen_{name}_lastframe.png") |
| Image.fromarray(img).save(p) |
| print(f"[wm-test] {name}: saved {p}", flush=True) |
| print("[wm-test] compare the three frames: forward/turn should differ from idle if controllable.", flush=True) |
|
|
|
|
| if __name__ == "__main__": |
| ap = argparse.ArgumentParser(); ap.add_argument("--test", action="store_true"); a = ap.parse_args() |
| (test if a.test else train)() |
|
|