Video DiT: Stage A, autoregressive k=4
Three approximately 95M-parameter Video DiT shapes trained to generate fixed-camera, three-face 2x2 Rubik's Cube videos. This release contains all seven saved training milestones per shape and the final full-recovery checkpoint: 24 checkpoint files. The milestones count training videos, not optimizer steps.
Architectures
| Shape | Model width | Layers | Heads | FFN width | Generator parameters |
|---|---|---|---|---|---|
| Deep–narrow | 512 | 22 | 8 | 1408 | 95,033,360 |
| Balanced | 640 | 14 | 10 | 1728 | 94,200,336 |
| Wide–shallow | 768 | 10 | 12 | 2048 | 96,909,328 |
All use 3D RoPE, head dimension 64, SwiGLU, 16 latent channels, and latent patches of (1,2,2). The model conditions on an initial frame, the complete nine-action language prompt, and visual history. No vector-state conditioning or auxiliary state loss is enabled. Auxiliary state-related config fields are inactive defaults.
Frozen representations: Wan2.1 VAE latents and 768-dimensional ModernBERT language features. These frozen backbones, their tokenizer/action-token adaptations, the training corpus, and executable inference/training code are not bundled here. This is a native research checkpoint archive, not a Transformers/Diffusers pipeline.
Training and milestones
Each run sees the same 1M distinct videos in the same order, with training seed 20260727 and global batch size 16. One H200 per run; 62,500 optimizer updates. Each video supplies 5,120 future target patch tokens (5.12B at the final milestone). The 3M-video source pool is not the exposure of these Stage-A runs. This 1M cap does not define a Stage-B training budget.
AR k=4 predicts five chunks of four future latent frames. Training uses parallel teacher-forced chunk losses averaged before one optimizer update; inference uses generated history. AdamW: betas (0.9,0.95), epsilon 1e-8, weight decay 0.05, gradient clipping 1.0. LR warms up over 16,384 videos to 4e-4 and decays by cosine to 4e-5 at 1M. Flow times follow sigmoid(N(0,1)); EMA decay is 0.9999.
| Videos seen | Optimizer updates | Target patch tokens |
|---|---|---|
| 32,768 | 2,048 | 167,772,160 |
| 65,536 | 4,096 | 335,544,320 |
| 131,072 | 8,192 | 671,088,640 |
| 262,144 | 16,384 | 1,342,177,280 |
| 524,288 | 32,768 | 2,684,354,560 |
| 786,432 | 49,152 | 4,026,531,840 |
| 1,000,000 | 62,500 | 5,120,000,000 |
Files and loading
Each of deep-narrow/, balanced/, and wide-shallow/ contains:
weights/videos-NNNNNNN.pt: raw generator state, EMA parameter state, model config and training metadata. The filename is zero-padded video exposure.recovery/videos-1000000.pt: the above plus optimizer state and RNG states.config.json,protocol.json, and per-checkpoint metadata JSON files.
manifest.json records SHA-256 hashes and sizes of all checkpoint files.
latent_normalization.json preserves the training channel normalization statistics.
evaluation/ contains the existing 100-episode development rollout summaries at
524,288 and 1M videos. These are not final locked-test or multiple-seed results.
import torch
from huggingface_hub import hf_hub_download
path = hf_hub_download(
"weihang44/video-dit-rubik-stage-a-ar-k4",
"wide-shallow/weights/videos-0786432.pt",
)
checkpoint = torch.load(path, map_location="cpu", weights_only=True)
config = checkpoint["model_config"]
raw_weights = checkpoint["model_state_dict"]
ema_weights = checkpoint["ema_parameter_state"]
# Instantiate the matching research generator, then load raw_weights strictly.
The exact matching implementation is required; these are custom PyTorch checkpoints.
Full recovery files also contain NumPy RNG state and may require trusted
weights_only=False loading. Only unpickle files whose source you trust.
Public copies preserve tensor values, optimizer state and RNG state. Cluster-local
metadata paths are replaced with local-assets/<basename> and W&B resume identities
are removed. Original checkpoints remain unchanged. To resume on another machine,
restore the dataset/representation assets and map the metadata paths to local paths;
the original trainer checks the data and protocol contracts exactly. Removing W&B
identity avoids resuming the original tracking run. No license grant is specified
by this checkpoint release.
Development rollout results at 1M videos
Raw weights, identical 100 held-out development episodes and sampling noise, 16 midpoint sampling steps per chunk, generated visual history. Means average the nine post-action boundaries, excluding the initial frame. Exact frame requires all 12 visible stickers to be correct.
| Shape | Mean sticker accuracy | Mean exact-frame accuracy | Action-9 sticker accuracy |
|---|---|---|---|
| Deep–narrow | 22.22% | 1.56% | 15.83% |
| Balanced | 21.69% | 0.78% | 16.17% |
| Wide–shallow | 69.74% | 28.44% | 34.17% |
These are single-seed, matched-data results, not matched-FLOP results or proof of convergence. Wide–shallow's advantage is concentrated at short horizons; its exact-frame accuracy is zero for actions 6–9 in this evaluation. Do not interpret this release as a general-purpose video model or a reliable Rubik's Cube solver.