--- library_name: pytorch tags: - video-generation - world-model - diffusion - diffusion-forcing - ablation datasets: - lerobot/fractal20220817_data --- # NanoWM-B/2 · RT-1 · Ablation: action_injection = adaln One of the action-injection ablation arms on RT-1 fractal. All other axes (pred_target=v, cosine + ZTSNR schedule, NanoWM-B/2 backbone, 50k steps) match the pred-v reference so the comparison isolates the injection mechanism. adaLN: condition-modulated LayerNorm (scale+shift) after each attention/MLP. ## Run identity - **wandb**: https://wandb.ai/better_guidance/nano-world-model-ablation/runs/d1wbbpsf - **launcher**: `src/scripts/ablation/inj_adaln.sh` - **collection**: https://huggingface.co/collections/knightnemo/nano-world-model ## Training setup | Key | Value | |---|---| | Architecture | NanoWM-B/2 (12 layers, d=768, patch=2, 158.6M params) | | Dataset | RT-1 fractal (`lerobot/fractal20220817_data`) | | Frames × resolution | 4 × 256² → 4 × 32² latents (SD-VAE) | | Context frames | 1 (sequential / self-forcing scheduling) | | Action injection | **adaln** (7-dim continuous) | | Steps | 50,000 | | Batch | 8/GPU × 8 × H20 = 64 effective | | Optimizer | AdamW, lr 1e-4, wd 0.01, warmup 1000, grad clip 0.1 after 20k | | Precision | bf16-mixed (params fp32), VAE fp32, `torch.compile` on | | Seed | 3407 | ## Diffusion setup | Key | Value | |---|---| | pred_name | v | | noise_schedule | `squaredcos_cap_v2` (cosine) | | zero_terminal_snr | true | | timestep_sampling | logit_normal (SD3-style, μ=0, σ=1) | | snr_gamma | 5.0 (Min-SNR loss weighting) | | diffusion_steps | 1000 train · 250 DDIM sample | | history_stabilization_level (inference) | 0.02 | ## Loading ```bash git clone git@github.com:knightnemo/nano-world-model.git cd nano-world-model huggingface-cli download knightnemo/nanowm-b2-rt1-abl-inj-adaln-50k --local-dir ./ckpt ``` ```python import sys from omegaconf import OmegaConf from safetensors.torch import load_file sys.path.insert(0, "src") from models import get_models cfg = OmegaConf.load("ckpt/config.yaml") cfg.experiment.infra.compile = False model = get_models(cfg).eval() state_dict = load_file("ckpt/model.safetensors") model.load_state_dict(state_dict, strict=True) ```