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.

Downloads last month

-

Downloads are not tracked for this model. How to track
Inference Providers NEW
This model isn't deployed by any Inference Provider. 🙋 Ask for provider support