Skip to content

Latest commit

 

History

298 Commits

Folders and files

NameName
Last commit message
Last commit date
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 

Repository files navigation

mini_mira

A from-scratch reimplementation of the core architecture behind MIRA — an action-conditioned latent world model for Rocket League, built on a representation-autoencoder (RAEv2) codec and a flow-matching diffusion transformer. Reimplements mira-wm/mira (Apache 2.0) end to end at a smaller scale, verified against the real source rather than assumed.

Current status: the DINO-track codec is frozen at step 3,999; its world model has completed two real training runs, through step 5,500, on top of it. A second, parallel track — swapping the frozen DINOv3 backbone for V-JEPA 2.1, as a controlled benchmark rather than a replacement — has its own codec (bottleneck fix included) deliberately stopped at step 1,999, and its own world model has now completed a real 2,000-step run on top of it — both step counts cut well short of their original targets given a compressed timeline, not crashes. Training is done; real sampling and evaluation against DINO's own numbers is still outstanding before either track's results count as announced. Done: fine-tuned the V-JEPA world model with FD-loss (arXiv:2604.28190v1) for the full 500 steps on top of the step-1999 checkpoint — real, clean completion at step 2,499, every quality metric improved on the pre-fine-tune baseline. See Status below.

What this is

One full forward pass — codec encode → latent diffusion → codec decode — matching the real MIRA architecture. Built by tracing the real repo, verifying shape contracts, and checking that each mechanism actually does what it claims (see Verification). Simplifications are disclosed decisions, not gaps found later — see Scope.

Architecture

DINO-shaped features  (B, T, dino_dim, H, W)
        │
        ▼
   StridedConvBottleneckConfig / MyBottleneck      strided-conv projection to a latent grid
        │
        ▼
   z  (B, T', latent_dim, H', W')
        │
        ▼
   LatentWorldModelConfig / DiffusionTransformer   AdaLN-conditioned space-time diffusion
        │        (multi-step Euler integration,     transformer; predicts flow-matching velocity
        │         flow matching: noise → data)
        ▼
   ViTDecoderConfig / ViTVideoDecoder               space-time ViT decoder, RoPE-based attention
        │
        ▼
   video  (B, T, 3, H_out, W_out)

Both the decoder and world model factorize attention into spatial (bidirectional, within a frame) and temporal (causal, across frames) sublayers, with RoPE for position instead of learned embeddings. The world model additionally conditions every block, via AdaLN, on the diffusion timestep tau and the player's key-press actions, plus a separate additive conditioning on the previous frame's clean latent content (clean_past).

Scale

Target ~300M params; configs/scaled_300m.yaml is the real attempt, measured from an instantiated model:

Component Parameters
Bottleneck 196,640
Decoder (width=1024, depth=6, 16 heads) 104,000,000
World model (hidden_dim=1024, depth=8, 16 heads) 188,110,880
Action encoder (9 keys) 149,792
bos 32
Total 292,457,344 (−2.51% vs. the 300M target)
DinoModel (frozen, real pretrained weights, not in the total) 85,669,632
VjepaModel (frozen, real pretrained weights, V-JEPA track only, not in the total) 86,833,152

Tuned so the world_model/decoder parameter ratio (1.81) roughly matches real mira's own shipped ratio (1.82), not picked by feel. Fast verification scripts instead use configs/small.yaml (~11.3M params, mirrors class defaults) for millisecond runs.

Status

Results summary

Codec — held-out reconstruction quality:

Track Real step PSNR (dB) SSIM LPIPS
DINOv3 3,999 of 8,000 target 19.56 0.552 0.486
V-JEPA 2.1 1,999 of 8,000 target 21.79 0.655 0.427

Both real, from evaluate_codec.py/evaluate_codec_vjepa.py on 20 genuinely held-out clips each. V-JEPA's codec beats DINO's on all three metrics — at roughly half the training steps (1,999 vs. 3,999). Worth being honest about why before reading too much into it: not a clean architecture comparison — different resolutions (288×512 vs. 448×768, V-JEPA's own native training scale), and V-JEPA's checkpoint also carries the shallow-texture-branch bottleneck fix that DINO's never got (a real, separate improvement, not inherent to the backbone swap). A genuinely controlled comparison would need to isolate that variable; this table reports what each track's own real, shipped checkpoint actually measures, not an ablation.

World model — same long-rollout benchmark, both tracks (scripts/sample_rollout.py/ sample_rollout_vjepa.py): 2 real context frames, 36 generated frames, same real held-out clips (matched --seed) before each track's own resize, same eval code path for both.

Metric DINOv3 (step 5,503) V-JEPA 2.1 (step 2,499)
PSNR 17.40 16.46
SSIM 0.724 0.697
LPIPS 0.397 0.383
Frechet DINO Distance 3.46 22.75
Frechet Inception Distance 64.31 107.57
drift_dino_l2_drift 0.052 0.517
Depth degradation (fdd_at_6→fdd_at_36) 2.69→5.81 (2.15x) 18.15→28.62 (1.58x)

Real, first genuinely matched-methodology comparison this project has had — same rollout length, same real clips, same code path, not two different training-time snapshots. DINO wins clearly on the Frechet/drift metrics (6.6x lower FDD, ~10x lower drift_dino_l2_drift); the gap on raw pixel metrics (PSNR/SSIM/LPIPS) is small, and V-JEPA is actually slightly ahead on LPIPS. Not evidence the backbone swap itself is behind — DINO's world model has more than double the training (5,503 vs. V-JEPA's 2,499 total steps) on top of a more mature codec (step 3,999/8,000 vs. 1,999/8,000, literally half). Matches an already-on-record finding from mira's own paper (Table 6): codec/perceptual quality disproportionately affects Frechet-distance-style generation metrics specifically, more than raw pixel ones — exactly the pattern here. The honest conclusion is "V-JEPA hasn't had a fair shot yet," not "V-JEPA is architecturally worse."

Separately, V-JEPA's own before/after FD-loss fine-tune (different eval settings — 2 context latents but full-training-cadence eval, not this same long-rollout script): frechet_dino_distance 32.71→28.40, psnr 15.82→18.88, ssim 0.678→0.747, lpips 0.4265→0.358 — a genuine, if smaller scale, internal improvement, not directly comparable to the table above (different rollout mechanics/eval cadence).

DINO's own PSNR here (17.40) is also a real correction worth recording: earlier text in this file described DINO's training-time PSNR as "stayed noisy ~12-15dB, never improved" — this is a cleaner, held-out, fixed-seed reading using the same matched methodology as V-JEPA's row, and reads meaningfully higher. Not a clean DINO-vs-V-JEPA comparison overall — different codec maturity, different step budgets, the benchmark-fairness questions below are still open — but this table is the first one built to actually be comparable in methodology, not just adjacent numbers.


Codec: step 3,999 of training, full resolution (288×512, 40-frame clips), real Rocket League data, real mira's cosine LR schedule and three-term loss (L1 + LPIPS + DINO latent-consistency). Real held-out evaluation on 20 clips it never trained on (scripts/evaluate_codec.py): PSNR 19.56dB, SSIM 0.552, LPIPS 0.486 — reads as "poor" by standard external benchmarks, though an earlier read based on training curves and preview videos alone (not these numbers) judged it "genuinely decent." Real mira's own recipe for this component runs 250,001 steps — we're at roughly 1.6% of that, which is the most likely explanation for the quality gap.

Decoder reconstruction preview

World model: scripts/train_world_model.py fully implements real mira's diagonal flow-matching loss (+ optional PSD self-distillation), real action conditioning from real streamed key-press data, checkpointing, and a full eval suite (drift metrics, Frechet DINO/Inception Distance, PSNR, LPIPS, SSIM, rendered rollout videos). First real, multi-hour training run completed: ~2,900 steps over ~11 hours on real data, real measured improvement (SSIM 0.48→0.65, LPIPS 0.59→0.40, Frechet DINO/Inception Distance both dropped by more than half). Also surfaced a real, concrete finding: quality degrades the deeper into a self-generated rollout it goes (e.g. Frechet DINO Distance 2.6 right after context vs. 26.6 by 28 frames deep — confirmed visually too, not just numerically) — the expected consequence of clean_past only ever being real during training but partly self-generated during rollout. Two things added in response: a timeout-safe SIGTERM handler (graceful checkpoint save + wandb sign-off instead of an abrupt kill), and opt-in scheduled sampling (--scheduled-sampling-prob, default off) — occasionally trains on a self-generated clean_past instead of always-real, directly targeting that gap. Verified via verify_world_model_training.py's CPU mechanism check, then run for real: a second training run, resuming the first run's checkpoint with --scheduled-sampling-prob 0.3, reached step 5,500. Real improvement across every headline metric versus the first run's own final numbers (SSIM 0.65→0.80, LPIPS 0.40→0.30, Frechet DINO Distance 10.3→3.86, Frechet Inception Distance 112.9→92.66, PSNR broke its earlier flat/noisy pattern to reach 20.17dB), and the rollout-depth degradation itself narrowed from a roughly 10x shallow-to-deep blowup to roughly 2.6x, now plateauing instead of continuing to climb — encouraging, though not an isolated before/after (more training steps and a changed eval window happened in the same run). A real, separate obstacle surfaced getting that run started: --resume replays every already-consumed training batch before continuing (by design, so a resumed run doesn't silently re-see old data) — at ~2,900 steps' worth, that meant ~11,600 batches to replay, a real measured 9-46 hour wait for zero new training. Fixed by patching the one relevant checkpoint field (dataloader_batches_consumed) directly, nothing else touched. Confirmed working end-to-end on real GPU hardware before that, including multi-session --resume — RNG state, dataloader position, and codec/latent-stats provenance all persist and were verified in a real two-phase smoke test (fresh run → resume). Two real bugs were found and fixed this way (a list-vs-tensor crash in drift-metric eval, a device-mismatch crash in rollout video rendering) — the kind reading-only verification can't catch. Real (full-scale) training config is confirmed and proven across two real runs — --batch-size 4 --grad-accum-steps 4 (effective batch 16), full resolution, --precision bf16 — each time-boxed rather than left running indefinitely, given the shared, rotating GPU access this project runs on. One open architectural divergence from real mira remains (an extra, unconditioned final LayerNorm) — low-risk, not urgent to fix.

V-JEPA track: a second, parallel pipeline, benchmarking V-JEPA 2.1 against the DINOv3 backbone above — not a replacement, and the DINO-track checkpoints stay untouched as the control. So far: VjepaModel (src/mini_mira/codec/vjepa.py), a frozen V-JEPA 2.1 ViT-B encoder shaped as a drop-in sibling to DinoModel (same dino_forward/.dino_dim contract), verified via scripts/verify_vjepa.py. Real facts confirmed live before writing it: V-JEPA 2.1 is ungated (unlike DINOv3), embed_dim=768 (matches DINOv3-B exactly), and it halves the frame count internally (tubelet_size=2 — dino_forward here returns t // 2 frames, a genuine, documented deviation from DinoModel's own contract). Multi-layer feature aggregation is supported too, matching DinoModel's own interface and layer indices.

Next goal: two design decisions for this track are now made, deliberately kept simple rather than general — this is a benchmark run by two interns, not infrastructure meant to outlive it. Scripts: no runtime backbone-selection flag; V-JEPA gets its own full copies of the training scripts (train_codec_vjepa.py, train_world_model_vjepa.py, mirroring train_codec.py/ train_world_model.py) rather than a shared core loop — the DINO scripts stay untouched, and each track's logic is fully readable on its own, at the accepted cost that a future bug fix has to be applied to both copies by hand. Checkpoints: separate directories (checkpoints_vjepa/, checkpoints_wm_vjepa/) and a _vjepa-tagged filename, no embedded backbone-metadata field — DINO's own checkpoint dirs and defaults are untouched. Config preset: done — configs/scaled_300m_vjepa.yaml, identical to scaled_300m.yaml except bottleneck.temporal_stride 2→1 (V-JEPA's own tubelet already halves frame count before the bottleneck sees it). Turned out no library code needs to change for any of this: MyBottleneck is already backbone-agnostic (works from dino_dim/temporal_stride alone, holds no reference to which encoder produced its input), and LatentWorldModel already accepts an injected dino module — both confirmed by reading the code, not assumed.

scripts/train_codec_vjepa.py is done — a full fork of train_codec.py, VjepaModel in place of DinoModel throughout, --checkpoint-dir defaulting to checkpoints_vjepa/ and its local/HF checkpoint filename tagged checkpoint_vjepa.pth so it can never collide with DINO's own upload to the same HF repo. One flag pair has no V-JEPA equivalent and was dropped rather than faked: --perceptual-dino-model/--perceptual-dino-multilayer swap in a second, differently-sized DINO variant for the consistency-loss term, and V-JEPA 2.1 only has the one variant.

Three real bugs found and fixed, only surfaced by actually running the training mechanism end to end (scripts/verify_codec_training_vjepa.py, new — same overfit-one-video proof verify_codec_training.py uses for the DINO track), not by reading the code:

  • VjepaModel.dino_forward crashed on fewer than tubelet_size (2) frames — CodecLoss's consistency term scores a random frame subset, sometimes chunked down to a single frame, which DinoModel always tolerated (no minimum) but V-JEPA's frame-pairing can't. Fixed by padding a too-short input up to 2 frames (repeating the last one) instead of rejecting it — contained entirely inside vjepa.py, no shared code touched.
  • CodecLoss assumed a pixel-space frame index always equals a feature-space frame index — true for DinoModel (frame count in = frame count out) but not V-JEPA (halves it). Fixed in loss.py: target-feature lookup now remaps indices by the encoder's own reduction ratio (inferred from tensor shapes, not a hardcoded per-backbone constant — a no-op for DinoModel, ratio 1 always).
  • The same term flattened batch and selected-frame together before chunking, so a chunk could straddle two different videos in the batch — invisible for DinoModel (no cross-frame interaction at all) but silently wrong for V-JEPA, which would pair the last selected frame of one video with the first frame of a completely different one and treat it as real motion. Silent, no crash — dormant at today's real 40-frame default settings by coincidence (the random selection size happens to be even), live the moment --frames or --perceptual-chunk-size changes to almost any other value. Fixed in loss.py by chunking within each video's own selected frames only, never across videos — confirmed directly: instrumented the real encoder call and proved, under the exact conditions that reproduced the bug (odd selection size, small chunk size), zero calls mix two videos, across 6 real calls.

Verified for real: a 100-step overfit run (batch=2, 16 frames — large enough to exercise all three bugs' trigger conditions, unlike a smaller test) drops loss_total 57.7%, matching the same >50% bar verify_codec_training.py itself uses. loss_dino_latent_consistency itself barely moves (expected — matches the DINO track's own already-documented finding that MAE dominates this term's gradient 4-8x; the auto-balancing mechanism visibly compensating, its weight climbing 36.2→47.0, is this working as designed, not a new problem).

Supervisor directive: train V-JEPA at native resolution (720×1280), not the 288×512 downscale — V-JEPA track only, DINO untouched. No code change needed: resize_to_canonical (shared with DINO, in video_prep.py) already no-ops once --height/--width match a clip's real shape, so passing --height 720 --width 1280 on the V-JEPA launch is the entire change. MyBottleneck (strided-conv, no absolute-position assumption) and the decoder (RoPE) are both already resolution-agnostic, confirmed by reading the code, and V-JEPA 2.1's own sincos position embeddings interpolate to arbitrary input shapes internally.

Native resolution was tried for real on the rented box and abandoned. --batch-size 2 OOM'd outright; --batch-size 1 OOM'd too and barely moved the memory number (44.21GB → 44.10GB, proving the cost is per-frame token count, not batch size); 720 also turns out to violate a real architectural requirement (height/width must divide evenly by patch_size(16) × bottleneck_stride(2) = 32 — 720/32=22.5, so the decoder silently reconstructed the wrong shape and crashed downstream instead of failing at startup). Cropping to 704 (22×32) plus --activation-checkpointing got it to fit — but at only ~1.4GB of headroom, judged too risky for an unattended multi-day run. Supervisor pivoted the target resolution twice more: 512×896 (architecturally clean, but OOM'd even with checkpointing) then settled at 448×768 (the supervisor's requested 448×784 also violates the same 32-divisibility rule — 784/32=24.5 — 768 was chosen as the nearest valid crop). Confirmed working with real margin: cuda_peak_reserved = 34.11GB (35.44GB with --compile, see below), out of a ~44.42GB usable pool.

Real, unrelated performance bug found and fixed while investigating why the GPU was reportedly only ~57% utilized: nvidia-smi dmon confirmed the pattern directly (sm% alternating 0/100 almost every second, ~62% average) and ruled out the dataloader first, with real evidence, before touching any training code — top showed the CPU ~60% idle and individual dataloader workers lightly loaded even at --num-workers 8 (all of ubuntu-gpu's real cores). Root cause: train_codec_vjepa.py's per-micro-step loss accumulation called .item() on every loss term, every micro-step — each call forces a full CUDA sync, serializing what should be an async pipeline. Fixed by accumulating losses as GPU tensors through the micro-step loop and converting to Python floats once per step instead of once per term per micro-step (same final numbers, far fewer sync points). Confirmed via nvidia-smi dmon again after the fix: sm% sustained at 100% across every sample. V-JEPA-track only, by decision — train_codec.py (DINO) has the identical pattern but was left untouched.

Supervisor separately asked to compile the model before the real training run — --compile added to train_codec_vjepa.py, wrapping just the trainable bottleneck/decoder in torch.compile() (the frozen V-JEPA encoder is deliberately excluded — external, git-cloned facebookresearch/vjepa2 code, real risk of graph breaks on unfamiliar ops). Confirmed working combined with --activation-checkpointing (a real, version-sensitive PyTorch interaction that had never been tested together before) — no crash, no torch._dynamo errors, across several real steps on GPU.

train_codec_vjepa.py now has the same SIGTERM/SIGINT handling train_world_model.py already had (graceful checkpoint save + forced HF upload + wandb sign-off instead of an abrupt kill) — train_codec.py, its un-forked DINO-track sibling, still lacks this, a pre-existing gap left as-is. Added ahead of the real training run on the newly rented GPU box, where an unattended timeout kill was judged worth guarding against rather than accepting the same risk DINO's own codec run already ran with successfully.

Real bug found and fixed in codec/checkpoint.py (shared, so it also protects train_codec.py even though DINO never triggers it): train_codec_vjepa.py's new --compile flag wraps bottleneck/decoder in torch.compile(), whose OptimizedModule wrapper adds a real "_orig_mod." prefix to every state_dict() key (confirmed both locally, with a CPU reproduction, and on real GPU: saved a checkpoint under --compile, then failed to load it into evaluate_codec_vjepa.py's plain, uncompiled modules — Missing key(s) ... Unexpected key(s): "_orig_mod.projection.weight"). Fixed by unwrapping to the real underlying module (via OptimizedModule's own ._orig_mod attribute) on both save and load, so every checkpoint is always stored in one plain, portable format regardless of whether --compile produced or is loading it — a no-op whenever --compile was never used, which is every existing DINO-track and pre-fix V-JEPA-track checkpoint.

scripts/evaluate_codec_vjepa.py is done — a full fork of evaluate_codec.py, VjepaModel in place of DinoModel throughout, --config defaulting to configs/scaled_300m_vjepa.yaml. Built ahead of there being a real trained V-JEPA codec checkpoint to point it at (the real 4,000-step run hasn't launched yet) — syntax/import-checked, not yet run against real output, since there's nothing real to evaluate until that checkpoint exists. Also served as the real, independent proof that the checkpoint.py fix above actually works: saved a checkpoint under --compile on GPU, loaded it through this (non-compiled) script, got real eval numbers back instead of the Missing/Unexpected key(s) crash.

Real per-step timing measured at the settled config (448×768, --compile --activation-checkpointing, batch=2/accum=16): a clean, post-compile-warmup read via the real filesystem timestamps of two saved preview videos, 4 steps apart — ≈94.6 sec/step, so ≈105 hours (≈4.4 days) for the full 4,000-step run. This is the number the GPU rental decision below is built on.

GPU rental decision in progress: real requirements (≥40GB VRAM with margin, bf16/Ampere-or- newer, CUDA-only — rules out non-NVIDIA accelerators regardless of specs) plus the timing number above were used to evaluate a real cloud GPU price list against the current A40 rental. vCPU count and system RAM both turned out not to be differentiators (confirmed: --num-workers 8 left the CPU ~60% idle once the sync-stall bug above was fixed; every viable VRAM-qualifying option already ships far more RAM than the ~18GB actually observed in use). Recommendation: A100 80GB — cheaper per hour and faster and more VRAM than the A40 already tested, beating it on every axis with no tradeoff. Decision handed to the supervisor; not yet acted on.

What's still ahead: a deferred overfit-one-clip convergence check at 448×768 specifically (the existing ~83%-loss-drop proof was at 288×512, a different resolution) — queued for just before the real launch, not done yet. A GPU still needs to be chosen and the real 4,000-step run needs to produce a real checkpoint before any of the world-model scripts below can run for real. Three real methodology questions still need a decision before any final numbers count as comparable: whether both tracks get scored by the same fixed judge rather than each by its own backbone, what step budget each track gets, and whether hyperparameters stay identical across both.

World-model scaffolding, built ahead of the codec checkpoint existing: scripts/compute_latent_stats_vjepa.py and scripts/train_world_model_vjepa.py are done — full forks of compute_latent_stats.py/train_world_model.py, VjepaModel injected via LatentWorldModel's existing dino: seam. Same "build and verify now, can't run for real yet" shape as evaluate_codec_vjepa.py — CPU-only checks (syntax, real module import, --help) all pass, plus a real end-to-end run of compute_latent_stats_vjepa.py on synthetic data (random-init VjepaModel, loaded from the already-cached facebookresearch/vjepa2 clone, no network needed).

Along the way, found and fixed a real bug in shared code (latent_world_model.py, used by both tracks): self.temporal_downsampling was computed from bottleneck_config.temporal_stride alone, blind to an injected encoder's own temporal reduction. Correct for DINO (dino_forward never touches time); silently wrong for V-JEPA, whose dino_forward halves time internally (tubelet_size=2) before the bottleneck ever sees it — every latent frame's action-conditioning window would have covered half the raw frames it actually represents. This is exactly the risk already flagged, unresolved, in notes/deviations.md 1.21. Fixed: bottleneck_config.temporal_stride * getattr(self.dino, "tubelet_size", 1), computed once self.dino is known — a provable no-op for DINO (no such attribute, falls back to 1), confirmed empirically by re-running verify_world_model_training.py (DINO's own 5-check suite, unchanged results). New scripts/verify_world_model_training_vjepa.py mechanically proves the fix itself, with a time-halving stand-in encoder — confirms temporal_downsampling comes out as stride * 2, not stride, plus a real forward/backward/overfit/checkpoint round-trip.

Pre-launch silent-bug audit, done ahead of the next real GPU window (an A40, for testing/ probing, not the box that hosts the eventual real run). Same hypothesis as the temporal_downsampling fix above — code that implicitly assumes DINO's "never touches time" behavior — hunted systematically instead of waiting to hit each instance for real. Found four more, all fixed:

  • loss.py's DINO-consistency term (most serious — active by default at the real codec launch config): sampled k independent, arbitrarily-scattered frames and fed them to dino_forward as if they were one contiguous clip. V-JEPA's tubelet pairing then paired consecutive positions in that fake sequence — real frames 3-5 apart, not temporal neighbors — and MSE'd the resulting fabricated-pair embedding against the real encoder's genuine adjacent-pair target features, with zero frame correspondence between the two sides. auto_weight (on by default) rescaled this corrupted term to match loss_mae's gradient magnitude rather than letting it fade into noise. Fixed: a new _sample_frame_indices helper samples whole tubelet-sized adjacent groups instead of independent frames when the bound encoder has its own temporal reduction (getattr(dino, "tubelet_size", 1)), and _select_target_features deduplicates to one representative feature per real group — same compute-cost budget as before (~10 frames/step), genuinely correct pairing instead of fabricated pairs. Byte-identical no-op for DINO (reduction=1), confirmed by a direct fake-DINO check (this dev machine's real DinoModel loader hits a separate, pre-existing torch.hub bug, unrelated). Two new regression checks in verify_codec_training_vjepa.py, run against the real VjepaModel: every dino_forward call now provably pairs genuine temporal neighbors, and _select_target_features's dedup output matches a hand-derived expected result exactly.
  • train_world_model_vjepa.py's pre-flight assertions used the raw bottleneck stride instead of the true total downsampling, computed before the model (and its V-JEPA tubelet factor) even existed — crashed with an AssertionError at the script's own documented defaults, before any GPU work. Fixed via a new sanity-checked constant, VJEPA_TUBELET_SIZE_EXPECTED in vjepa.py (mirrors the existing VJEPA_DIM_EXPECTED pattern, with its own defensive assert in VjepaModel.__init__), so the pre-flight math is correct without constructing the model early.
  • eval_metrics.py/full_eval_metrics.py assumed DINO-re-encoded features (from re-running dino_forward on the decoded video) were in video-frame units — true for DINO, false for V-JEPA, whose re-encoding halves time again and lands back in latent-frame units. Silently dropped the wrong window from drift metrics every eval, and crashed OnlineGaussian.compute() on the first full eval by starving later Frechet-distance slices to zero samples. Fixed: a new optional dino_temporal_scale parameter (defaults to the old formula — zero changes needed in train_world_model.py), correctly computed in the V-JEPA script as temporal_downsampling // getattr(model.dino, "tubelet_size", 1). New regression check in verify_world_model_training_vjepa.py proves the correct slice length (3) differs from what the old formula would silently produce (2).
  • The height/width-divisible-by-32 requirement (already discovered the hard way once, per the resolution journey above) was never asserted anywhere — a bad --height/--width produced a confusing shape-mismatch crash elsewhere instead of a clear error. Now asserted right after config loading in all six scripts that exercise the decoder (both codec tracks' train/eval scripts, both world-model scripts) — purely additive, a no-op for every already-proven-valid real launch.
  • Checkpoint provenance (both codec and world-model checkpoint.py) didn't cover temporal_downsampling — a --resume against a mismatched config that happens to keep every parameter shape identical would load with zero warning. Not reversing the project's earlier deliberate "keep checkpoints simple" decision — added one narrow, optional, backward-compatible field plus one more warn-don't-block check in train_world_model_vjepa.py's resume block, mirroring the codec_checkpoint/latent_mean checks already there.

Real GPU validation of the V-JEPA world-model mechanism-only path, done on an A40 window explicitly reserved for world-model debugging (not the codec) while a new box is awaited for the real codec run. Ran compute_latent_stats_vjepa.py + train_world_model_vjepa.py for real — real pretrained V-JEPA encoder, random-init bottleneck/decoder (no --codec-checkpoint, per its now-optional mechanism-only mode) — and confirmed the entire pipeline works end to end: real training loop, validation, the full eval suite (drift metrics, Frechet DINO/Inception distance, PSNR/LPIPS/SSIM, rendered rollout preview videos), and checkpointing. Real numbers from that run: ~18.5–19.2GiB peak memory and roughly 50–60 sec/step at 448×768×40/batch=2/accum=16 — notably lighter than the codec's own measured 94.6 sec/step at the same resolution, as expected (no decoder backward pass, no LPIPS/DINO-consistency compute in the world model's own loss).

Two more real bugs were found and fixed live during this validation, same class as the pre-launch audit above:

  • The .item()-per-micro-step sync-stall pattern — identical root cause to the one already found and fixed for train_codec_vjepa.py (calling .item() on every loss term inside the grad-accum micro-step loop forces a full CUDA sync each time), but never ported to the world-model scripts. Confirmed live via nvidia-smi dmon showing the same alternating 0%/100% pattern on a real run. Fixed in both train_world_model.py and train_world_model_vjepa.py this time (backbone-agnostic bug — leaving it only in the V-JEPA script would have made V-JEPA look artificially faster than DINO for reasons unrelated to real architectural cost, a real confound for the eventual benchmark).
  • full_eval_metrics.py crashed on the real full eval (AssertionError: Need at least 2 samples to compute statistics.) — a gap in the earlier dino_temporal_scale fix: it corrected where the generated region starts in DINO-feature space, but not that its total length is also proportionally shorter than video-frame space for a time-halving encoder. The same fdd_slice_frames-wide windows, sized for the longer video tensor, ran past the end of the shorter DINO-feature one for later slices, leaving them with zero samples. Fixed with a new dino_fdd_slice_frames parameter, scaled proportionally (fdd_slice_frames * dino_temporal_scale // temporal_downsampling) — reproduced the exact real crash in a CPU regression test first, confirmed the fix resolves it, confirmed the un-fixed formula still crashes the same way, then confirmed on the real GPU run itself.

V-JEPA world-model real training recipe — run for real, completed (configs/runs/ wm_first_run_vjepa.yaml). --steps 2000, not DINO's own ~5,500-step second-run precedent — a deliberate, explicit cut given a compressed timeline (internship ending in days), not a fallback default. --scheduled-sampling-prob 0.3 from step 0 (DINO discovered this helps only on its second run; starting with it here skips re-discovering an already-proven finding). --height 448 --width 768 --frames 40 (matches the real V-JEPA codec's own trained resolution, not DINO's 288×512). --batch-size 4 --grad-accum-steps 4 --compile --num-workers 6 — real, measured cuda_peak_reserved=20.46GiB of ~46GB, flat for the whole run, no memory creep; notably lower than DINO's own compiled number (19.79GiB) at the same resolution/batch, despite V-JEPA's codec needing ~2x DINO's VRAM — the frozen V-JEPA+bottleneck encode step runs once per training step under plain torch.no_grad(), so that multiplier (measured with V-JEPA backpropped through the codec's own decoder+perceptual losses) doesn't carry over to the world-model's memory profile. bf16, no PSD — matching DINO. --num-workers 6 under --resume and the non-finite-gradient guard (ported from the codec's own real fix, see below) were both validated/added before this launch, not left as risk.

Real final numbers, step 1999: loss_total=0.3263, psnr=15.82, ssim=0.678, lpips=0.427, frechet_dino_distance=32.71, frechet_inception_distance=101.90. Rollout-depth degradation (fdd_at_6=28.56 → fdd_at_36=41.01, ~1.4x) is comparatively mild against DINO's own first (no-scheduled-sampling) run's ~10x degradation (fdd_at_7=2.55 → fdd_at_28=26.56) — consistent with scheduled sampling doing its job from step 0, though not a clean apples-to-apples number (different slice points, different codec maturity, the benchmark-fairness questions below still open). No NaN/non-finite-gradient skips occurred during this run.

FD-loss fine-tune (arXiv:2604.28190v1, "Representation Fréchet Loss for Visual Generation") — implemented, CPU-verified, not yet run on the real GPU. Supervisor-directed: 500 more steps on top of the step-1999 checkpoint, targeting quality (not the paper's other headline use case, distilling multi-step generators into one-step ones). The paper makes Fréchet Distance directly optimizable as a training loss by decoupling the population size used for the distance estimate (an EMA over many steps) from the batch size used for gradients. Two real findings shaped the design, not assumptions:

  • This project's existing Fréchet-distance code (full_eval_metrics.py) round-trips through scipy.linalg.sqrtm/numpy — non-differentiable by construction, not reusable as a loss. New pure-PyTorch implementation (src/mini_mira/world_model/fd_loss.py): the standard Tr((cov_r@cov_g)^0.5) cross term is computed via M = cov_sqrt_r @ cov_g @ cov_sqrt_r, Tr(M^0.5) = sum(sqrt(eigvalsh(M))) — cov_sqrt_r (the real, fixed side) computed once; eigvalsh on M is autograd-differentiable, unlike scipy.sqrtm.
  • LatentWorldModel.rollout() (full multi-step autoregressive sampling) is far too expensive to differentiate through in a 500-step budget, and is no_grad()-wrapped by design. Instead, the "generated" side reuses the one-step estimate already implemented for scheduled sampling (z1 ≈ z_t + (1-tau)*pred_v), computed from the exact z_t/pred_v/tau the normal training step already produces — no extra forward pass for generation, only the subsequent decode + V-JEPA feature extraction is new per-step cost.

Additive, not paper-literal: loss_total = loss_diffusion + fd_loss_weight * loss_fd, matching this codebase's existing PSD-loss pattern rather than the paper's own FD-loss-only post-training, to keep the base flow-matching signal alive over a short window. EMA decay defaults to 0.97 (deliberately far below the paper's own 0.999, tuned for thousands of post-training steps — at 500 steps it would barely move from its seed). New offline script, scripts/compute_real_frechet_stats_vjepa.py, precomputes real mean/covariance in V-JEPA's own 768-dim feature space from held-out clips (OnlineGaussian, reused as-is — no differentiability needed for a fixed, precomputed statistic). New recipe: configs/runs/wm_finetune_fd_loss_vjepa.yaml (--steps 2500, --lr 1e-5, --fd-loss-weight 1.0 as a starting point — genuinely untuned, the one real empirical unknown here). scripts/verify_fd_loss.py (new, 4 checks, all passing): differentiability + sane values, the EMA update rule matches the paper's formula exactly, fd_loss_weight=0 is a provable no-op (no state built at all), and a full synthetic forward/backward reaches every trainable parameter.

Real-GPU debugging done, config validated, the real 500-step run is complete. FD-loss is the first time this script backprops through the full decoder+V-JEPA encoder (eval/rollout always run no_grad) — real OOM at --batch-size 4/--grad-accum-steps 4 even with a first (too-coarse, whole-function) checkpointing attempt; fixed by reusing ViTVideoDecoder's own real per-block use_checkpointing (the same lever that got the codec itself fitting at this resolution — gated on self.training, toggled just for this call since LatentWorldModel normally keeps the decoder permanently .eval()). --batch-size 2 --grad-accum-steps 8 (same effective batch) is the real, confirmed-fitting config, cuda_peak_reserved=29.64GiB. --fd-loss-weight 1.0 was also real-world validated as drastically oversized (loss_fd completely dominated loss_diffusion, frechet_dino_distance/frechet_inception_distance measurably got worse); rescaled to 0.002. Re-validated with a real 10-step smoke test at the corrected config: frechet_dino_distance improved 32.71→23.42 (−28%), psnr/ssim/lpips all improved too, loss_diffusion stayed in a tight band instead of trending up. Two new CLI flags added along the way: --skip-dataloader-fastforward (a --resume replaying 8,000 real batches via next() before training even starts is real, not hypothetical — this is a new training phase, not same-phase crash recovery, so exact dataloader-position continuity was never actually needed) and --wandb-new-run (resume model/optimizer state normally, start a fresh wandb run).

One real interruption along the way, cause not fully confirmed: the launch died silently around step 2,299 (no Python traceback anywhere in the log — dmesg was inaccessible without root to confirm an OOM-kill, best guess given hours of --num-workers 6 runtime, not proven). --resume picked it back up cleanly with no data loss (--checkpoint-every 100 had already saved past that point) and it ran the remaining ~200 steps to completion without recurring.

Real final numbers, step 2,499, vs. the pre-fine-tune baseline (step 1,999) — every metric improved: frechet_dino_distance 32.71→28.40 (−13%), frechet_inception_distance 101.90→94.45 (−7%), psnr 15.82→18.88 (+19%), ssim 0.678→0.747 (+10%), lpips 0.427→0.358 (+16%, lower is better). loss_fd converged from ~225 at step 2000 to a stable ~70 plateau; loss_diffusion stayed in a healthy 0.25–0.32 band the whole run, never trended toward degradation. Worth noting honestly: frechet_dino_distance wasn't perfectly monotonic along the way (23.42 at step 2009 → 29.78 at step 2249 → 28.40 at step 2499) — small-eval-batch noise, expected — but the net trend from baseline held up across the full run, not just an early lucky reading. Real checkpoint: checkpoints_wm_vjepa_fd_finetune_v2/checkpoint_wm_vjepa.pth, backed up to BahaXD/mini-mira-codec-checkpoints.

Two of the four original benchmark-fairness questions are eval/comparison-time decisions, not training-launch blockers — corrected framing from an earlier pass: whether both tracks get scored by the same fixed judge, and which DINO checkpoint counts as the control, only matter once someone is actually writing up the DINO-vs-V-JEPA comparison. The code has no way to do "same fixed judge" during training anyway — decode_and_dino always reuses model.dino (whichever backbone that model trained under) for its own periodic eval, to avoid loading a second backbone; a shared-judge comparison would need a separate script built later. Still genuinely open: full hyperparameter parity beyond what's settled above, and formally which DINO checkpoint is the control (the settled recipe above already implies DINO's second/5,500-step run, but that's an implication, not yet an explicit statement for a real write-up).

Reconstruction-quality investigation, resolved with a real fix now training for real. The DINO-track codec's PSNR (19.56dB) reads as "poor" by standard benchmarks — reconstructions look blocky. Researched real, published techniques rather than guessing further, and found the actual cause by reading this project's own code: MyBottleneck blends 4 encoder-feature layers (shallow to deep) into one average before compression, diluting the shallow, texture-rich layer into the more abstract deep ones. A very recent paper on an almost identical setup (DINO-Tok, arXiv:2511.20565) hit this project's exact PSNR number scaling decoder capacity alone, and fixed it the same way this project independently arrived at: keep the shallow layer as its own channel group instead of averaging it away (StridedConvBottleneckConfig.use_shallow_texture_branch, off by default — the decoder needs no changes, since it only ever consumes the bottleneck's unchanged output shape). Validated on a one-clip diagnostic: wins on loss, pixel accuracy, and perceptual detail quality at every checkpoint measured, no instability. Adopted — a real 8,000-step V-JEPA codec run with this enabled is training now.

Two other quality levers were investigated and explicitly not adopted, kept in the codebase as tried-and-documented rather than removed: a small decoder-side "refinement head" (ViTDecoderConfig.use_refinement_head) crashed with a real NaN on its first design, and still showed persistent instability in perceptual scoring after a standard fix (LayerScale) removed the crash itself; and three tested learning-rate pacing variants all underperformed the pacing already in use.

The real 8,000-step run hit its own NaN crash, deterministically, at the same step twice — traced to a real gap where bf16 training (this project's default) had none of the automatic non-finite- gradient protection fp16-hybrid's GradScaler provides for free. Fixed with two additions to train_codec_vjepa.py: the optimizer step is skipped (not applied) on a non-finite gradient, and every checkpoint save — local, HF, and a new --local-checkpoint-history sliding window of independent local snapshots — checks weights for NaN/Inf first and skips the save entirely rather than overwriting a known-good state with a corrupted one. A further, not-yet-conclusively- validated mitigation (--max-auto-weight, capping how far the adaptive loss-balancing mechanism can rescale a term) is in progress. Full research trail, real numbers, and the real production- launch findings (A40-specific memory/speed tuning, a real wandb permission bug and its fix) are in notes/vjepa_codec_quality_research.md (git-ignored, local only).

The same non-finite-gradient-skip fix was ported to train_world_model_vjepa.py ahead of its own first real sustained launch — that script also defaults to --precision bf16 and had none of this protection at all, not yet having hit a real crash, only found by inspection before one happened.

Full bug-by-bug history and the evidence trail behind every claim above: notes/deviations.md and notes/session_handoff.md (both git-ignored, local only).

Project layout

File Contents
src/mini_mira/codec/bottleneck.py Encoder-side strided-conv projection into the latent, with an opt-in shallow-texture-branch variant
src/mini_mira/codec/decoder.py Space-time ViT decoder
src/mini_mira/world_model/diffusion_transformer.py AdaLN-conditioned diffusion transformer
src/mini_mira/ml/blocks.py Shared attention/MLP/AdaLN blocks (decoder + world model)
src/mini_mira/ml/init.py Mira-matching weight initialization
src/mini_mira/ml/rope.py Shared RoPE implementation (temporal + spatial)
src/mini_mira/world_model/timestep_encoder.py Sinusoidal embedding of diffusion timestep tau
src/mini_mira/world_model/action_encoder.py Encodes key-press actions into conditioning vectors
src/mini_mira/pipeline.py Architecture-demo pipeline, no real checkpoint loading — see LatentWorldModel for the real trainer
src/mini_mira/world_model/latent_world_model.py Real training wrapper: frozen codec + trainable world model, real flow-matching + PSD loss
src/mini_mira/world_model/checkpoint.py Save/resume for world-model training
src/mini_mira/world_model/eval_metrics.py Cheap, always-on drift-metric eval
src/mini_mira/world_model/full_eval_metrics.py Frechet DINO/Inception Distance, PSNR, LPIPS, SSIM
src/mini_mira/world_model/rollout_visualization.py Renders rollout videos with an action HUD overlay
src/mini_mira/codec/dino.py Real, frozen DINOv3 backbone
src/mini_mira/codec/vjepa.py Real, frozen V-JEPA 2.1 backbone — DinoModel-shaped sibling
src/mini_mira/codec/loss.py Codec training loss: L1 + LPIPS + DINO latent-consistency
src/mini_mira/codec/video_prep.py Resizes/pads real clips to canonical shape
src/mini_mira/codec/checkpoint.py Save/resume for codec training
src/mini_mira/codec/logging_utils.py Optional wandb logging, shared by both codec tracks — preview videos encode at crf=18 explicitly (found and fixed for real: no explicit value meant falling back to ffmpeg's own default, visibly compressed — confirmed by comparing a lossless raw frame against the same frame through this encoder)
scripts/test_shapes.py Shape-correctness checks
scripts/verify_rope.py Behavioral checks for RoPE
scripts/verify_conditioning.py Behavioral checks for AdaLN/clean-past/actions
scripts/verify_dino.py Behavioral checks for the real DINOv3 backbone (needs gated weights)
scripts/verify_vjepa.py Behavioral checks for the real V-JEPA 2.1 backbone (ungated, no weights needed to set up first)
scripts/test_dino.py Raw DINOv3 sanity check, bypassing DinoModel
src/mini_mira/ml/config_loading.py Builds a PipelineConfig (architecture) or a run-config (hyperparameters) from YAML
src/mini_mira/ml/run_config.py WorldModelRunConfig/CodecRunConfig — the hyperparameter axis, loaded via --run-config
configs/small.yaml Fast-test preset, mirrors class defaults
configs/scaled_300m.yaml ~300M-param target preset
configs/runs/ Real --run-config examples (hyperparameters — batch size, steps, eval cadence, ...)
scripts/verify_codec_training.py Mechanism proof the codec trains (synthetic data, no GPU needed)
scripts/download_shards.py Downloads real Rocket League shards from kyutai/rocket-science
scripts/train_codec.py Real GPU codec training
scripts/train_codec_vjepa.py Same, V-JEPA track — full fork, VjepaModel in place of DinoModel
scripts/verify_codec_training_vjepa.py Mechanism proof the V-JEPA-track codec trains (synthetic data, no GPU needed)
scripts/overfit_one_clip_vjepa.py Real-data, real-GPU diagnostic: overfit one real clip, wandb only, no checkpoints
scripts/reconstruct.py Mechanism smoke test: runs a video through the codec (random-init weights)
scripts/evaluate_codec.py Real quantitative eval of a trained codec checkpoint on held-out data
scripts/evaluate_codec_vjepa.py Same, V-JEPA track — full fork, VjepaModel in place of DinoModel
scripts/compute_latent_stats.py One-shot latent mean/std computation, feeds train_world_model.py
scripts/compute_latent_stats_vjepa.py Same, V-JEPA track — full fork, VjepaModel in place of DinoModel
scripts/train_world_model.py Real GPU world-model training
scripts/train_world_model_vjepa.py Same, V-JEPA track — full fork, VjepaModel injected via LatentWorldModel's dino: seam
scripts/sample_rollout.py Standalone: load a trained world-model checkpoint, roll out on real held-out clips, save videos + Frechet/PSNR/LPIPS/SSIM/drift metrics — no training loop
scripts/sample_rollout_vjepa.py Same, V-JEPA track — full fork
scripts/serve_interactive_vjepa.py Small local web server: browser keypresses drive live, turn-based rollout generation (no kv-cache, so not real-time — see its own docstring)
scripts/verify_interactive_rollout.py Proves the server's incremental generation is bit-identical to LatentWorldModel.rollout() (CPU-only, no GPU needed)
scripts/verify_world_model_training.py CPU mechanism proof for train_world_model.py
scripts/verify_world_model_training_vjepa.py CPU mechanism proof that temporal_downsampling correctly accounts for a time-halving encoder (V-JEPA-specific, not a fork of the above)
scripts/verify_full_eval_metrics.py CPU mechanism proof for the full eval suite
scripts/verify_run_config.py CPU mechanism proof for the --run-config system

Scope

Implemented: strided-conv bottleneck + ViT space-time decoder matching the real codec's shape contract; RoPE (temporal + axial spatial); QK-norm and mira-matching weight init; AdaLN conditioning on tau, clean-past, and actions; flow-matching sampling; the real, frozen DINOv3 backbone with real pretrained weights; real codec training on real data with the real loss, adaptive loss balancing, checkpoint save/resume, and an opt-in bottleneck variant (use_shallow_texture_branch) that preserves shallow-layer texture instead of averaging it away; real world-model training mechanism with real flow-matching + PSD loss, real action-conditioned data, a full eval suite, and opt-in scheduled sampling for rollout-depth robustness.

Deliberately simplified / not yet implemented (disclosed decisions, not gaps found later):

  • Actions are keyboard keys only, no mouse — matches the real released data, which has no real mouse signal either. Also simplified vs. mira's ActionEncoder: no dropout, mean-pooling instead of a learned temporal pool, plain Linear instead of mira's per-key dimension split.
  • clean_past is real encoded input by default (never the model's own previous output) — matches mira's own default too. --scheduled-sampling-prob opts into training on a self-generated estimate instead, some fraction of the time; off unless explicitly set.
  • No streaming inference / KV-cache — every diffusion step recomputes the whole sequence.
  • No grouped-query attention — always as many KV heads as query heads.
  • One shared implementation where the real repo has two separate ones (codec vs. world model) for identical logic — consolidated into blocks.py/rope.py instead.

Full audit trail for every intentional or since-corrected difference from real mira: notes/deviations.md.

Verification

Shape correctness alone doesn't prove a mechanism works, so each is checked behaviorally too:

  • scripts/test_shapes.py — every stage's output shape against the real codec's config.
  • scripts/verify_rope.py — RoPE causality and position sensitivity.
  • scripts/verify_conditioning.py — tau/clean_past/action sensitivity and determinism, plus end-to-end regression checks that output actually depends on input and on actions.
  • scripts/verify_dino.py — the real DINOv3 backbone is frozen, correctly shaped, handles non-multiple-of-16 resolutions, and produces non-degenerate output.
  • scripts/verify_codec_training.py — overfitting one synthetic video with real optimizer steps must substantially reduce the loss (catches dead gradients, detached graphs, wrong losses).

Getting started

Requirements: Python ≥ 3.10, plus requirements.txt. wandb/huggingface_hub are optional, lazily-imported (--wandb-project/--hf-backup-repo only); torchcodec is needed only for real streamed data (--index-path). No packaging config — scripts add src/ to sys.path directly.

pip install -r requirements.txt
python scripts/test_shapes.py
python scripts/verify_rope.py
python scripts/verify_conditioning.py
python scripts/verify_codec_training.py

The first run of anything using CodecLoss downloads pretrained VGG16 weights for LPIPS (~528MB, one-time). scripts/verify_dino.py additionally needs real, gated DINOv3 weights on disk (request access at ai.meta.com/resources/models-and-libraries/dinov3-downloads), pointed to via RS_DINO_WEIGHTS_DIR.

Real training needs a CUDA GPU and real data:

python scripts/download_shards.py --shards 50   # prints the local index path to pass below
python scripts/train_codec.py --config configs/scaled_300m.yaml \
  --index-path <path printed above> --require-pretrained-dino \
  --height 288 --width 512 --frames 40 --batch-size 4 --grad-accum-steps 8 \
  --activation-checkpointing --steps 1600
python scripts/compute_latent_stats.py --codec-checkpoint <checkpoint.pth> --index-path <path> \
  --output latent_stats.json
python scripts/train_world_model.py --config configs/scaled_300m.yaml \
  --codec-checkpoint <checkpoint.pth> --latent-stats latent_stats.json \
  --index-path <train data> --test-index-path <held-out data> --require-pretrained-dino

--help on either script lists the full flag set (LR schedule, precision, PSD weights, eval cadence, checkpoint/resume, wandb/HF Hub backup) — or skip retyping it every launch with --run-config <path> (see configs/runs/ for real examples): any flag passed explicitly still overrides the file, and omitting --run-config entirely reproduces the exact defaults above.

Resume gotchas (each has bitten this project for real at least once — see notes/session_handoff.md for the full writeup of each):

  • The resumed step count is checkpoint's saved step + 1, not the saved step itself.
  • --resume alone continues the checkpoint's original LR curve; add --reset-lr-schedule to start a fresh warmup/decay shaped for a new phase, and pass this run's own --lr/--lr-min/ --lr-warmup-steps/--lr-decay-steps explicitly either way.
  • Continuing the same phase across multiple sessions needs --lr-warmup-steps/ --lr-decay-steps passed identically on every session, sized to the whole multi-session arc — the script recomputes them from --steps unconditionally, so omitting them silently corrupts the curve instead of erroring.
  • train_codec.py: resuming under a different --precision than the checkpoint was saved with crashes (a known, unfixed gap) — keep --precision consistent across a checkpoint's resumes.
  • train_world_model.py's --resume fast-forwards past already-consumed batches by calling next() on a freshly-built loader — only correct because that loader is single-process, unseeded (deterministic by default). The codec scripts (train_codec.py/train_codec_vjepa.py) have no such mechanism, so raising --num-workers there is safe (--num-workers 6 is now the default, matching this project's real GPU box's 8 real CPU cores — recheck against nproc on different hardware); doing the same for train_world_model.py would need that resume mechanism addressed first, not done unprompted.
  • Both world-model scripts (train_world_model.py/train_world_model_vjepa.py) now refuse a --resume whose checkpoint step is already >= --steps — previously this silently trained nothing (range(start_step, steps) is just empty) and exited 0 with no warning, indistinguishable from a real completed run.

Configs

Named presets for PipelineConfig live in configs/ at the repo root (data, not package code) — one YAML file mirrors PipelineConfig's whole nested shape, no Hydra/config-group system.

File Purpose
configs/scaled_300m.yaml The intended architecture scale (~300M params)
configs/scaled_300m_vjepa.yaml Same, temporal_stride 2→1 for the V-JEPA track, use_shallow_texture_branch: true (matches the real trained codec checkpoint — compute_latent_stats_vjepa.py/train_world_model_vjepa.py have no CLI override for this field, only --config sets it)
configs/small.yaml Mirrors class defaults — what fast verification scripts use
import sys
from pathlib import Path
sys.path.insert(0, str(Path(__file__).resolve().parent / "src"))

from mini_mira.ml.config_loading import load_pipeline_config
from mini_mira.pipeline import MyPipeline

config = load_pipeline_config("configs/scaled_300m.yaml")
pipeline = MyPipeline(config)

Relationship to mira and attribution

Built by tracing the official mira-wm/mira release (Apache License 2.0). Class/config names match the real repository wherever there's a genuine one-to-one correspondence (ViTVideoDecoder, SelfAttention, DiffusionTransformer); MyBottleneck and MyPipeline keep their own names since they have no direct equivalent in the real repo. A few small, self-contained pieces (AdaptiveLayerNorm, the sinusoidal timestep embedding, the RoPE frequency computation) are adapted directly from the original source under its license terms.

About

No description, website, or topics provided.

Resources

Stars

0 stars

Watchers

0 watching

Forks

Releases

Packages

Contributors

Languages