doanh25032004's picture
Backup source tree of video_gen_physics (2026-07-31T14:21:08Z)
ec0a9aa verified
|
Raw
History Blame Contribute Delete
7.46 kB

DreamDojo Architecture

DreamDojo is an action-conditioned video world model built on the Diffusion Transformer (DiT) backbone with Rectified Flow formulation. It generates future video frames conditioned on an input image/video, text prompt, and robot action sequences.

Model Variants

Variant Blocks Heads Channels Head Dim Params
2B 28 16 2048 128 ~2B
7B 28 32 4096 128 ~7B
14B 36 40 5120 128 ~14B

Config definitions: cosmos_predict2/_src/predict2/action/configs/action_conditioned/net.py

Core Architecture

DiT Backbone: MiniTrainDIT

File: cosmos_predict2/_src/predict2/networks/minimal_v4_dit.py

The backbone is a standard Vision Transformer adapted for video:

Input (B, C_in, T, H, W)
  --> Patchify (patch_spatial=2, patch_temporal=1)
  --> Linear projection to model_channels
  --> + 3D RoPE positional embeddings
  --> N x TransformerBlock:
      |-- AdaLN (adaptive layer norm from timestep embedding)
      |-- Self-Attention (Q/K/V with RMSNorm + RoPE)
      |-- Cross-Attention (to text embeddings)
      |-- FFN (GPT-2 style: Linear -> GELU -> Linear)
  --> Unpatchify
  --> Output (B, C_out, T, H, W)

Key classes:

  • MiniTrainDIT (line 700+): Main model with forward() and forward_with_cfg()
  • Attention (line 389+): Multi-head attention with configurable backends
  • GPT2FeedForward (line 238+): MLP block
  • RMSNorm (line 220+): Root mean square normalization

Action Conditioning

File: cosmos_predict2/_src/predict2/action/networks/action_conditioned_minimal_v1_lvg_dit.py

Actions are injected into the model via the timestep embedding pathway:

action: (B, T-1, action_dim=384) 
  --> ActionEmbedder (Linear projection)
  --> Add to timestep embedding (t_emb)
  --> t_emb conditions AdaLN in each transformer block

Action dimension layout (384-dim):

Range Robot/Type
[0, 29) Fourier GR-1
[29, 58) Retargeted GR-1
[58, 101) Unitree G1
[101, 147) Bimanual YAM
[147, 169) AgiBot
[169, 220) Reserved
[220, 352) MANO hand actions
[352, 384) Latent actions

Each action frame contains: [delta_xyz(3), delta_rotation(3), gripper_state(1)] for the relevant robot, with unused dimensions zeroed.

Attention Mechanisms

File: cosmos_predict2/_src/predict2/networks/minimal_v4_dit.py (line 389+)

The Attention class supports multiple backends:

Backend Description Default For
torch F.scaled_dot_product_attention General
torch-flex FlexAttention with BlockMask Sparse patterns
minimal_a2a Custom A2A attention 14B model
i4 Imaginaire4 attention Alternative
transformer_engine NVIDIA TE attention Training

Attention flow:

Input x: (B, S, D)  where S = T*H*W (flattened video tokens)
  --> Q = q_proj(x): (B, S, n_heads*head_dim)
  --> K = k_proj(x): (B, S, n_heads*head_dim)  [self-attn]
       K = k_proj(context): (B, M, n_heads*head_dim)  [cross-attn]
  --> V = v_proj(...)
  --> Reshape: (B, S, n_heads, head_dim)
  --> Q, K = RMSNorm(Q), RMSNorm(K)
  --> Q, K = apply_RoPE_3D(Q, K)  [self-attn only]
  --> output = attn_op(Q, K, V)
  --> output = output_proj(output)

Rectified Flow Diffusion

File: cosmos_predict2/_src/predict2/models/text2world_model_rectified_flow.py

DreamDojo uses Rectified Flow (RF), a straight-path ODE formulation:

  • Forward process: x_t = (1 - t) * x_0 + t * noise, where t in [0, 1]
  • Training objective: Predict velocity v = noise - x_0
  • Inference: Solve ODE from t=1 (noise) to t=0 (clean) using Euler steps
  • Scheduler: FlowUniPCMultistepScheduler (2nd-order predictor-corrector)
  • Default steps: ~35 Euler steps with shift=5.0
  • CFG: Classifier-free guidance with guidance_scale parameter

VAE Tokenizer

File: cosmos_predict2/_src/predict2/tokenizers/cosmos.py

The Cosmos VAE compresses video to latent space:

Property Value
Spatial compression 8x (via patch_spatial=2 in model)
Temporal compression 4x
Latent channels 16
Input format (B, 3, T, H, W) uint8 [0, 255]
Latent format (B, 16, T/4, H/8, W/8) bf16

For 480x640 input with 13 frames: latent shape = (1, 16, 4, 60, 80)

Text Encoder

Uses T5-based text encoder for computing text embeddings from prompts. The embeddings condition the model through cross-attention in each transformer block.

For distilled models, pre-computed CR1 (empty-string) embeddings can be used for efficiency.

Config System

DreamDojo uses a layered config system:

  1. Base configs (cosmos_predict2/_src/predict2/action/configs/action_conditioned/config.py): Define model, net, optimizer, scheduler defaults
  2. Network configs (net.py): Register model architectures (2B, 7B, 14B)
  3. Experiment configs (cosmos_predict2/experiments/base/action.py): Auto-register from YAML files
  4. YAML overrides (configs/*.yaml): Per-experiment overrides (e.g., 14b_480_640_gr1.yaml)

Config loading: load_model_from_checkpoint() in cosmos_predict2/_src/predict2/utils/model_loader.py

Distilled Model (Student)

File: cosmos_predict2/_src/predict2/interactive/networks/dit_action_causal.py

The distilled model differs from the teacher:

Aspect Teacher Student (Distilled)
Denoising steps ~35 4 (DMD2)
Attention Bidirectional Temporal causal
KV cache No Yes (frame-indexed)
torch.compile Optional Recommended
Inference mode Batch Streaming (per-frame)

Distillation pipeline (3 stages):

  1. Teacher generation (launch_teacher_gen.sh): Generate training data with teacher
  2. Warmup (launch_warmup.sh): Train student to match teacher outputs
  3. Self-forcing (launch_self_forcing.sh): Finetune with autoregressive self-predictions

Key File Map

DreamDojo/
  cosmos_predict2/_src/predict2/
    networks/
      minimal_v4_dit.py          # Core DiT backbone + Attention class
      a2a_cp.py                  # A2A and NATTEN attention ops
    action/
      networks/                  # Action-conditioned model wrappers
      inference/
        inference.py             # Standard inference script
        inference_batch.py       # Batch inference (custom benchmark)
        inference_pipeline.py    # ActionVideo2WorldInference class
      configs/                   # Action-conditioned configs
    interactive/
      networks/dit_action_causal.py  # Distilled causal DiT
      inference/action_video2world.py # Distilled streaming inference
    models/
      text2world_model_rectified_flow.py  # Rectified flow model
    tokenizers/                  # VAE tokenizer
    schedulers/                  # ODE solvers
    utils/
      model_loader.py            # Checkpoint loading
      kv_cache.py                # KV cache utilities
  configs/                       # YAML experiment configs
  docs/                          # Documentation
  scripts/                       # Utilities (checkpoint conversion, etc.)