Compare commits

...
Author SHA1 Message Date
SolitaryThinker f0e727f5ee [misc]: point DreamX-World at published FastVideo Diffusers repos
Advertise FastVideo/DreamX-World-5B-Cam-Diffusers and
FastVideo/DreamX-World-5B-Diffusers as the loadable ids in the registry
(raw GD-ML path detection is kept via the pattern detectors), default the
example and local tests to the hub ids so they no longer skip when local
converted dirs are absent, and drop the now-redundant pipeline class
override in the example (model_index.json carries DreamXWorldPipeline).
2026-07-06 12:35:40 -07:00
SolitaryThinker de6a60ac90 [test]: randomize the zero-init head so DreamX AR tiny parity is not vacuous 2026-07-05 14:39:39 -07:00
SolitaryThinker 1b4fdc2c41 [refactor]: DreamX review wave 3 — AR DiT onto FastVideo layer primitives + real param_names_mapping 2026-07-05 14:39:39 -07:00
SolitaryThinker 385cd65abb [bugfix]: DreamX review wave 2 — AR pipeline actually encodes its conditioning image 2026-07-05 14:39:39 -07:00
SolitaryThinker 156d86b70c [bugfix]: DreamX review wave 1 — detector exclusivity, converter self-heal, SP hard-fail guard 2026-07-05 14:39:39 -07:00
Suckl bab79fb5f6 Add DreamX World 5B AR pipeline 2026-07-05 14:39:39 -07:00
Suckl 073a9e78f2 Add DreamX World 5B Cam pipeline 2026-07-05 14:39:39 -07:00
38 changed files with 5397 additions and 8 deletions
@@ -191,6 +191,9 @@ surfaces:
sources:
- fastvideo.configs.pipelines.gen3c.Gen3CConfig
- fastvideo.configs.pipelines.gen3c.Gen3CInferenceConfig
color_correction_strength:
sources:
- fastvideo.configs.pipelines.dreamx_world.DreamXWorld5BARPipelineConfig
default_camera_rotation:
sources:
- fastvideo.configs.pipelines.gen3c.Gen3CConfig
+2
View File
@@ -58,6 +58,8 @@ pipeline initialization and sampling.
| FastWan2.1 T2V 1.3B | `FastVideo/FastWan2.1-T2V-1.3B-Diffusers` | 480P | ⭕ | ⭕ | ⭕ | ✅ | ⭕ |
| FastWan2.2 TI2V 5B Full Attn* | `FastVideo/FastWan2.2-TI2V-5B-FullAttn-Diffusers` | 720P | ⭕ | ⭕ | ⭕ | ✅ | ⭕ |
| Wan2.2 TI2V 5B | `Wan-AI/Wan2.2-TI2V-5B-Diffusers` | 720P | ⭕ | ⭕ | ✅ | ⭕ | ⭕ |
| DreamX-World 5B Cam | `FastVideo/DreamX-World-5B-Cam-Diffusers` | 480P | ⭕ | ⭕ | ⭕ | ⭕ | ⭕ |
| DreamX-World 5B AR | `FastVideo/DreamX-World-5B-Diffusers` | 704px1280p | ⭕ | ⭕ | ⭕ | ⭕ | ⭕ |
| Lucy Edit Dev 5B*** | `decart-ai/Lucy-Edit-Dev` | 480P | ⭕ | ⭕ | ⭕ | ⭕ | ⭕ |
| Wan2.2 T2V A14B | `Wan-AI/Wan2.2-T2V-A14B-Diffusers` | 480P<br>720P | ❌ | ❌ | ✅ | ⭕ | ⭕ |
| Wan2.2 I2V A14B | `Wan-AI/Wan2.2-I2V-A14B-Diffusers` | 480P<br>720P | ❌ | ❌ | ✅ | ⭕ | ⭕ |
@@ -0,0 +1,70 @@
"""DreamX-World-5B-Cam camera-controlled video generation.
Uses the pre-converted Diffusers checkpoint FastVideo/DreamX-World-5B-Cam-Diffusers.
To convert the raw GD-ML/DreamX-World-5B-Cam checkpoint yourself, see
scripts/checkpoint_conversion/dreamx_world_to_diffusers.py.
"""
import os
from fastvideo import VideoGenerator
OUTPUT_PATH = os.getenv("DREAMX_WORLD_OUTPUT_PATH", "video_samples_dreamx_world")
def _env_int(name: str, default: int) -> int:
return int(os.getenv(name, str(default)))
def _env_float(name: str, default: float) -> float:
return float(os.getenv(name, str(default)))
def main():
model_name = os.getenv("DREAMX_WORLD_MODEL_DIR", "FastVideo/DreamX-World-5B-Cam-Diffusers")
generator = VideoGenerator.from_pretrained(
model_name,
num_gpus=1,
use_fsdp_inference=False,
dit_cpu_offload=False,
vae_cpu_offload=True,
text_encoder_cpu_offload=True,
pin_cpu_memory=False,
)
prompt = os.getenv(
"DREAMX_WORLD_PROMPT",
"A cinematic first-person drive through a futuristic coastal city at "
"sunrise, reflective glass towers, clean streets, soft volumetric light.",
)
image_path = os.getenv(
"DREAMX_WORLD_IMAGE_PATH",
"https://huggingface.co/datasets/YiYiXu/testing-images/resolve/main/wan_i2v_input.JPG",
)
kwargs = {
"output_path": OUTPUT_PATH,
"save_video": os.getenv("DREAMX_WORLD_SAVE_VIDEO", "1") != "0",
"height": _env_int("DREAMX_WORLD_HEIGHT", 480),
"width": _env_int("DREAMX_WORLD_WIDTH", 832),
"num_frames": _env_int("DREAMX_WORLD_NUM_FRAMES", 161),
"num_inference_steps": _env_int("DREAMX_WORLD_STEPS", 30),
"guidance_scale": _env_float("DREAMX_WORLD_GUIDANCE", 5.0),
"action_list": os.getenv("DREAMX_WORLD_ACTIONS", "w,d,w").split(","),
"action_speed_list": [
float(value)
for value in os.getenv("DREAMX_WORLD_ACTION_SPEEDS", "4.0,2.0,4.0").split(",")
],
}
if image_path:
kwargs["image_path"] = image_path
try:
generator.generate_video(prompt, **kwargs)
finally:
generator.shutdown()
if __name__ == "__main__":
main()
+4 -3
View File
@@ -1,5 +1,6 @@
from fastvideo.configs.models.dits.cosmos import CosmosVideoConfig
from fastvideo.configs.models.dits.cosmos2_5 import Cosmos25VideoConfig
from fastvideo.configs.models.dits.dreamx_world import DreamXWorldARConfig, DreamXWorldConfig
from fastvideo.configs.models.dits.flux_2 import Flux2Config
from fastvideo.configs.models.dits.hunyuangamecraft import HunyuanGameCraftConfig
from fastvideo.configs.models.dits.hunyuanvideo import HunyuanVideoConfig
@@ -13,7 +14,7 @@ from fastvideo.configs.models.dits.hyworld import HYWorldConfig
from fastvideo.configs.models.dits.kandinsky5 import Kandinsky5VideoConfig
__all__ = [
"HunyuanVideoConfig", "HunyuanVideo15Config", "HunyuanGameCraftConfig", "WanVideoConfig", "CosmosVideoConfig",
"Cosmos25VideoConfig", "LongCatVideoConfig", "LTX2VideoConfig", "HYWorldConfig", "Kandinsky5VideoConfig",
"MagiHumanVideoConfig", "StableAudioConfig", "Flux2Config"
"HunyuanVideoConfig", "HunyuanVideo15Config", "HunyuanGameCraftConfig", "WanVideoConfig", "DreamXWorldConfig",
"DreamXWorldARConfig", "CosmosVideoConfig", "Cosmos25VideoConfig", "LongCatVideoConfig", "LTX2VideoConfig",
"HYWorldConfig", "Kandinsky5VideoConfig", "MagiHumanVideoConfig", "StableAudioConfig", "Flux2Config"
]
@@ -0,0 +1,71 @@
# SPDX-License-Identifier: Apache-2.0
from dataclasses import dataclass, field
from fastvideo.configs.models.dits.base import DiTArchConfig, DiTConfig
from fastvideo.configs.models.dits.wanvideo import WanVideoArchConfig
@dataclass
class DreamXWorldArchConfig(WanVideoArchConfig):
"""DreamX-World DiT config with camera PRoPE control fields."""
add_control_adapter: bool = True
cam_method: str | None = "prope"
attn_compress: int = 1
cam_self_attn_layers: tuple[int, ...] | None = None
@dataclass
class DreamXWorldConfig(DiTConfig):
arch_config: DiTArchConfig = field(default_factory=DreamXWorldArchConfig)
prefix: str = "Wan"
@dataclass
class DreamXWorldARArchConfig(DreamXWorldArchConfig):
"""DreamX-World-5B autoregressive causal DiT config."""
model_type: str = "ti2v"
patch_size: tuple[int, int, int] = (1, 2, 2)
text_len: int = 512
text_dim: int = 4096
freq_dim: int = 256
attn_compress: int = 4
cam_self_attn_layers: tuple[int, ...] | None = tuple(range(30))
local_attn_size: int = 12
sink_size: int = 3
num_frames_per_block: int = 3
rope_cache_policy: str = "block_relativistic"
# The official AR checkpoint (AMAP-ML/DreamX-World ``model.safetensors``)
# already uses FastVideo's native key names and the converter copies the
# tensors verbatim, so every rule is an identity. The rules enumerate the
# full state-dict surface of ``DreamXWorldARTransformer3DModel`` (norm1 /
# norm2 / head.norm are affine-free and have no parameters).
param_names_mapping: dict = field(
default_factory=lambda: {
r"^patch_embedding\.(.*)$": r"patch_embedding.\1",
r"^text_embedding\.([02])\.(.*)$": r"text_embedding.\1.\2",
r"^time_embedding\.([02])\.(.*)$": r"time_embedding.\1.\2",
r"^time_projection\.1\.(.*)$": r"time_projection.1.\1",
r"^blocks\.(\d+)\.self_attn\.(q|k|v|o)\.(.*)$": r"blocks.\1.self_attn.\2.\3",
r"^blocks\.(\d+)\.self_attn\.norm_(q|k)\.weight$": r"blocks.\1.self_attn.norm_\2.weight",
r"^blocks\.(\d+)\.cross_attn\.(q|k|v|o)\.(.*)$": r"blocks.\1.cross_attn.\2.\3",
r"^blocks\.(\d+)\.cross_attn\.norm_(q|k)\.weight$": r"blocks.\1.cross_attn.norm_\2.weight",
r"^blocks\.(\d+)\.cam_self_attn\.(q_proj|k_proj|v_proj|out_proj)\.(.*)$": r"blocks.\1.cam_self_attn.\2.\3",
r"^blocks\.(\d+)\.cam_self_attn\.norm_(q|k)\.weight$": r"blocks.\1.cam_self_attn.norm_\2.weight",
r"^blocks\.(\d+)\.norm3\.(.*)$": r"blocks.\1.norm3.\2",
r"^blocks\.(\d+)\.ffn\.([02])\.(.*)$": r"blocks.\1.ffn.\2.\3",
r"^blocks\.(\d+)\.modulation$": r"blocks.\1.modulation",
r"^head\.head\.(.*)$": r"head.head.\1",
r"^head\.modulation$": r"head.modulation",
})
reverse_param_names_mapping: dict = field(default_factory=dict)
lora_param_names_mapping: dict = field(default_factory=dict)
@dataclass
class DreamXWorldARConfig(DreamXWorldConfig):
arch_config: DiTArchConfig = field(default_factory=DreamXWorldARArchConfig)
prefix: str = "Wan"
+3 -1
View File
@@ -1,6 +1,7 @@
from fastvideo.configs.pipelines.base import PipelineConfig
from fastvideo.configs.pipelines.cosmos import CosmosConfig
from fastvideo.configs.pipelines.cosmos2_5 import Cosmos25Config
from fastvideo.configs.pipelines.dreamx_world import DreamXWorld5BARPipelineConfig, DreamXWorld5BCamPipelineConfig
from fastvideo.configs.pipelines.hunyuan import FastHunyuanConfig, HunyuanConfig
from fastvideo.configs.pipelines.hunyuan15 import Hunyuan15T2V480PConfig, Hunyuan15T2V720PConfig
from fastvideo.configs.pipelines.hunyuangamecraft import HunyuanGameCraftPipelineConfig
@@ -16,5 +17,6 @@ __all__ = [
"HunyuanConfig", "FastHunyuanConfig", "HunyuanGameCraftPipelineConfig", "PipelineConfig", "Hunyuan15T2V480PConfig",
"Hunyuan15T2V720PConfig", "WanT2V480PConfig", "WanI2V480PConfig", "WanT2V720PConfig", "WanI2V720PConfig",
"SelfForcingWanT2V480PConfig", "LucyEditDevConfig", "CosmosConfig", "Cosmos25Config", "LTX2T2VConfig",
"HYWorldConfig", "MatrixGame2I2V480PConfig", "MatrixGame3I2V720PConfig", "get_pipeline_config_cls_from_name"
"DreamXWorld5BCamPipelineConfig", "DreamXWorld5BARPipelineConfig", "HYWorldConfig", "MatrixGame2I2V480PConfig",
"MatrixGame3I2V720PConfig", "get_pipeline_config_cls_from_name"
]
+127
View File
@@ -0,0 +1,127 @@
# SPDX-License-Identifier: Apache-2.0
"""DreamX-World-5B-Cam FastVideo model configuration helpers."""
from collections.abc import Callable
from dataclasses import dataclass, field
import torch
from fastvideo.configs.models import DiTConfig, EncoderConfig, VAEConfig
from fastvideo.configs.models.dits.dreamx_world import (DreamXWorldARArchConfig, DreamXWorldARConfig,
DreamXWorldArchConfig, DreamXWorldConfig)
from fastvideo.configs.models.encoders import BaseEncoderOutput, T5Config
from fastvideo.configs.models.encoders.t5 import T5ArchConfig
from fastvideo.configs.models.vaes import WanVAEConfig
from fastvideo.configs.pipelines.base import PipelineConfig
from fastvideo.configs.pipelines.wan import LucyEditDevConfig, t5_postprocess_text
def make_dreamx_world_5b_cam_dit_config() -> DreamXWorldConfig:
"""Return the DreamX-World DiT config matching DreamX-World-5B-Cam."""
return DreamXWorldConfig(arch_config=DreamXWorldArchConfig(
num_attention_heads=24,
attention_head_dim=128,
in_channels=48,
out_channels=48,
ffn_dim=14336,
num_layers=30,
cross_attn_norm=True,
qk_norm="rms_norm_across_heads",
add_control_adapter=True,
cam_method="prope",
attn_compress=1,
cam_self_attn_layers=None,
))
def make_dreamx_world_5b_ar_dit_config() -> DreamXWorldARConfig:
"""Return the DreamX-World-5B autoregressive causal DiT config."""
return DreamXWorldARConfig(arch_config=DreamXWorldARArchConfig(
model_type="ti2v",
num_attention_heads=24,
attention_head_dim=128,
in_channels=48,
out_channels=48,
ffn_dim=14336,
num_layers=30,
cross_attn_norm=True,
qk_norm=True,
add_control_adapter=True,
cam_method="prope",
attn_compress=4,
cam_self_attn_layers=tuple(range(30)),
local_attn_size=12,
sink_size=3,
num_frames_per_block=3,
))
def make_dreamx_world_5b_cam_vae_config() -> WanVAEConfig:
"""Return the Wan2.2 48-channel VAE config used by DreamX-World-5B-Cam."""
return LucyEditDevConfig().vae_config
def make_dreamx_world_5b_cam_text_encoder_config() -> T5Config:
"""Return the UMT5-XXL text encoder config used by DreamX-World-5B-Cam."""
return T5Config(
arch_config=T5ArchConfig(
vocab_size=256384,
d_model=4096,
d_kv=64,
d_ff=10240,
num_layers=24,
num_decoder_layers=None,
num_heads=64,
relative_attention_num_buckets=32,
dropout_rate=0.0,
text_len=512,
feed_forward_proj="gelu",
is_encoder_decoder=False,
),
prefix="umt5",
)
@dataclass
class DreamXWorld5BCamPipelineConfig(PipelineConfig):
"""Pipeline config for the first-scope DreamX-World-5B-Cam mode."""
dit_config: DiTConfig = field(default_factory=make_dreamx_world_5b_cam_dit_config)
vae_config: VAEConfig = field(default_factory=make_dreamx_world_5b_cam_vae_config)
text_encoder_configs: tuple[EncoderConfig,
...] = field(default_factory=lambda: (make_dreamx_world_5b_cam_text_encoder_config(), ))
postprocess_text_funcs: tuple[Callable[[BaseEncoderOutput], torch.Tensor],
...] = field(default_factory=lambda: (t5_postprocess_text, ))
text_encoder_precisions: tuple[str, ...] = field(default_factory=lambda: ("bf16", ))
flow_shift: float | None = 3.0
ti2v_task: bool = True
expand_timesteps: bool = True
vae_tiling: bool = False
vae_sp: bool = False
vae_precision: str = "fp32"
vae_decode_precision: str | None = "bf16"
dit_precision: str = "bf16"
def __post_init__(self) -> None:
self.vae_config.load_encoder = True
self.vae_config.load_decoder = True
self.dit_config.expand_timesteps = self.expand_timesteps
@dataclass
class DreamXWorld5BARPipelineConfig(DreamXWorld5BCamPipelineConfig):
"""Pipeline config for DreamX-World-5B autoregressive forcing."""
dit_config: DiTConfig = field(default_factory=make_dreamx_world_5b_ar_dit_config)
flow_shift: float | None = 5.0
ti2v_task: bool = True
is_causal: bool = True
dmd_denoising_steps: tuple[int, ...] = (1000, 750, 500, 250)
warp_denoising_step: bool = True
context_noise: float = 0.1
num_frames_per_block: int = 3
color_correction_strength: float = 1.0
def __post_init__(self) -> None:
super().__post_init__()
self.dit_config.expand_timesteps = True
+511
View File
@@ -0,0 +1,511 @@
# SPDX-License-Identifier: Apache-2.0
import math
from typing import Any
import torch
import torch.nn as nn
from fastvideo.configs.models.dits.dreamx_world import DreamXWorldConfig
from fastvideo.distributed.communication_op import (
sequence_model_parallel_all_gather_with_unpad,
sequence_model_parallel_shard)
from fastvideo.distributed.parallel_state import get_sp_world_size
from fastvideo.layers.layernorm import RMSNorm
from fastvideo.layers.linear import ReplicatedLinear
from fastvideo.layers.rotary_embedding import get_rotary_pos_embed
from fastvideo.models.dits.wanvideo import (LayerNormScaleShift,
PatchEmbed,
WanTimeTextImageEmbedding,
WanTransformer3DModel,
WanTransformerBlock)
from fastvideo.platforms import AttentionBackendEnum, current_platform
from fastvideo.attention import LocalAttention
from fastvideo.distributed.parallel_state import get_sp_world_size
from fastvideo.layers.quantization import QuantizationConfig
from fastvideo.models.dits.base import BaseDiT
def _dreamx_invert_se3(transforms: torch.Tensor) -> torch.Tensor:
assert transforms.shape[-2:] == (4, 4)
rot_inv = transforms[..., :3, :3].transpose(-1, -2)
out = torch.zeros_like(transforms)
out[..., :3, :3] = rot_inv
out[..., :3, 3] = -torch.einsum("...ij,...j->...i", rot_inv,
transforms[..., :3, 3])
out[..., 3, 3] = 1.0
return out.to(dtype=transforms.dtype)
def _dreamx_lift_k(intrinsics: torch.Tensor) -> torch.Tensor:
assert intrinsics.shape[-2:] == (3, 3)
out = torch.zeros(intrinsics.shape[:-2] + (4, 4),
device=intrinsics.device,
dtype=intrinsics.dtype)
out[..., :3, :3] = intrinsics
out[..., 3, 3] = 1.0
return out
def _dreamx_invert_k(intrinsics: torch.Tensor) -> torch.Tensor:
assert intrinsics.shape[-2:] == (3, 3)
out = torch.zeros_like(intrinsics)
out[..., 0, 0] = 1.0 / intrinsics[..., 0, 0]
out[..., 1, 1] = 1.0 / intrinsics[..., 1, 1]
out[..., 0, 2] = -intrinsics[..., 0, 2] / intrinsics[..., 0, 0]
out[..., 1, 2] = -intrinsics[..., 1, 2] / intrinsics[..., 1, 1]
out[..., 2, 2] = 1.0
return out.to(dtype=intrinsics.dtype)
def _dreamx_apply_tiled_projmat(feats: torch.Tensor,
matrix: torch.Tensor) -> torch.Tensor:
batch, num_heads, seq_len, feat_dim = feats.shape
proj_dim = matrix.shape[-1]
assert feat_dim % proj_dim == 0
if matrix.shape[1] == seq_len:
feats = feats.view(batch, num_heads, seq_len, feat_dim // proj_dim,
proj_dim)
out = torch.einsum("btij,bntpj->bntpi", matrix, feats)
return out.reshape(batch, num_heads, seq_len, feat_dim)
cameras = matrix.shape[1]
assert seq_len > cameras and seq_len % cameras == 0
feats = feats.reshape(batch, num_heads, cameras, -1,
feat_dim // proj_dim, proj_dim)
out = torch.einsum("bcij,bncpkj->bncpki", matrix, feats)
return out.reshape(batch, num_heads, seq_len, feat_dim)
def _dreamx_prope_qkv(q: torch.Tensor, k: torch.Tensor, v: torch.Tensor,
viewmats: torch.Tensor, intrinsics: torch.Tensor):
batch, num_heads, seq_len, head_dim = q.shape
cameras = viewmats.shape[1]
assert q.shape == k.shape == v.shape
assert viewmats.shape == (batch, cameras, 4, 4)
assert intrinsics.shape == (batch, cameras, 3, 3)
assert head_dim % 4 == 0
intrinsics_norm = torch.zeros_like(intrinsics)
intrinsics_norm[..., 0, 0] = intrinsics[..., 0, 0]
intrinsics_norm[..., 1, 1] = intrinsics[..., 1, 1]
intrinsics_norm[..., 2, 2] = 1.0
proj = torch.einsum("...ij,...jk->...ik",
_dreamx_lift_k(intrinsics_norm), viewmats)
proj_t = proj.transpose(-1, -2).to(dtype=viewmats.dtype)
proj_inv = torch.einsum(
"...ij,...jk->...ik",
_dreamx_invert_se3(viewmats),
_dreamx_lift_k(_dreamx_invert_k(intrinsics_norm)),
).to(dtype=viewmats.dtype)
q = _dreamx_apply_tiled_projmat(q, proj_t)
k = _dreamx_apply_tiled_projmat(k, proj_inv)
v = _dreamx_apply_tiled_projmat(v, proj_inv)
return q, k, v, proj
class DreamXPropeSelfAttention(nn.Module):
"""DreamX-World parallel PRoPE camera self-attention branch."""
def __init__(self,
dim: int,
attn_dim: int,
num_heads: int,
qk_norm: str | bool = True,
eps: float = 1e-6,
quant_config: QuantizationConfig | None = None,
prefix: str = ""):
super().__init__()
assert attn_dim % num_heads == 0
self.attn_dim = attn_dim
self.num_heads = num_heads
self.head_dim = attn_dim // num_heads
self.qk_norm = qk_norm
self.q_proj = ReplicatedLinear(dim,
attn_dim,
quant_config=quant_config,
prefix=f"{prefix}.q_proj")
self.k_proj = ReplicatedLinear(dim,
attn_dim,
quant_config=quant_config,
prefix=f"{prefix}.k_proj")
self.v_proj = ReplicatedLinear(dim,
attn_dim,
quant_config=quant_config,
prefix=f"{prefix}.v_proj")
self.out_proj = ReplicatedLinear(attn_dim,
dim,
quant_config=quant_config,
prefix=f"{prefix}.out_proj")
if qk_norm == "rms_norm":
self.norm_q = RMSNorm(self.head_dim, eps=eps)
self.norm_k = RMSNorm(self.head_dim, eps=eps)
elif qk_norm in (True, "rms_norm_across_heads"):
self.norm_q = RMSNorm(attn_dim, eps=eps)
self.norm_k = RMSNorm(attn_dim, eps=eps)
elif qk_norm is False:
self.norm_q = nn.Identity()
self.norm_k = nn.Identity()
else:
raise ValueError(f"Unsupported qk_norm for DreamX PRoPE: {qk_norm}")
nn.init.zeros_(self.out_proj.weight)
if self.out_proj.bias is not None:
nn.init.zeros_(self.out_proj.bias)
self.attn = LocalAttention(
num_heads=num_heads,
head_size=self.head_dim,
dropout_rate=0,
softmax_scale=None,
causal=False,
supported_attention_backends=(AttentionBackendEnum.FLASH_ATTN,
AttentionBackendEnum.TORCH_SDPA))
def forward(self, hidden_states: torch.Tensor,
y_camera: dict[str, torch.Tensor]) -> torch.Tensor:
if get_sp_world_size() > 1:
# The transformer shards the sequence before the block loop and
# this branch uses LocalAttention (no all-to-all): under
# sequence parallelism each rank would attend only within its
# own shard — silently wrong output. Fail loudly until this
# path is ported to DistributedAttention and validated.
raise NotImplementedError(
"DreamXPropeSelfAttention does not support sequence "
"parallelism yet (LocalAttention on a sharded sequence "
"corrupts output). Run with sp_size=1.")
batch_size, seq_len, _ = hidden_states.shape
query, _ = self.q_proj(hidden_states)
key, _ = self.k_proj(hidden_states)
value, _ = self.v_proj(hidden_states)
if self.qk_norm == "rms_norm":
query = query.view(batch_size, seq_len, self.num_heads,
self.head_dim)
key = key.view(batch_size, seq_len, self.num_heads, self.head_dim)
query = self.norm_q(query)
key = self.norm_k(key)
else:
query = self.norm_q(query).view(batch_size, seq_len,
self.num_heads, self.head_dim)
key = self.norm_k(key).view(batch_size, seq_len, self.num_heads,
self.head_dim)
value = value.view(batch_size, seq_len, self.num_heads, self.head_dim)
query = query.transpose(1, 2)
key = key.transpose(1, 2)
value = value.transpose(1, 2)
query, key, value, output_projection = _dreamx_prope_qkv(
query,
key,
value,
viewmats=y_camera["viewmats"],
intrinsics=y_camera["K"],
)
out = self.attn(query.transpose(1, 2), key.transpose(1, 2),
value.transpose(1, 2))
out = _dreamx_apply_tiled_projmat(out.transpose(1, 2),
output_projection).transpose(1, 2)
out = out.flatten(2)
out, _ = self.out_proj(out)
return out
class DreamXWorldTransformerBlock(WanTransformerBlock):
def __init__(self,
dim: int,
ffn_dim: int,
num_heads: int,
qk_norm: str = "rms_norm_across_heads",
cross_attn_norm: bool = False,
eps: float = 1e-6,
added_kv_proj_dim: int | None = None,
supported_attention_backends: tuple[AttentionBackendEnum, ...]
| None = None,
quant_config: QuantizationConfig | None = None,
prefix: str = "",
add_control_adapter: bool = True,
cam_method: str | None = "prope",
attn_compress: int = 1,
cam_self_attn_layers: tuple[int, ...] | None = None,
layer_idx: int | None = None):
super().__init__(dim, ffn_dim, num_heads, qk_norm, cross_attn_norm,
eps, added_kv_proj_dim,
supported_attention_backends, quant_config, prefix)
self.cam_self_attn = None
add_cam_attn = add_control_adapter and cam_method == "prope"
if add_cam_attn and cam_self_attn_layers is not None:
add_cam_attn = layer_idx in cam_self_attn_layers
if add_cam_attn:
if num_heads % attn_compress != 0 or dim % attn_compress != 0:
raise ValueError("DreamX attn_compress must divide dim and num_heads")
self.cam_self_attn = DreamXPropeSelfAttention(
dim,
dim // attn_compress,
num_heads // attn_compress,
qk_norm=qk_norm,
eps=eps,
quant_config=quant_config,
prefix=f"{prefix}.cam_self_attn")
def forward(
self,
hidden_states: torch.Tensor,
encoder_hidden_states: torch.Tensor,
temb: torch.Tensor,
freqs_cis: tuple[torch.Tensor, torch.Tensor],
original_seq_len: int,
y_camera: dict[str, torch.Tensor] | None = None,
) -> torch.Tensor:
if hidden_states.dim() == 4:
hidden_states = hidden_states.squeeze(1)
orig_dtype = hidden_states.dtype
if temb.dim() == 4:
shift_msa, scale_msa, gate_msa, c_shift_msa, c_scale_msa, c_gate_msa = (
self.scale_shift_table.unsqueeze(0) + temb.float()).chunk(
6, dim=2)
shift_msa = shift_msa.squeeze(2)
scale_msa = scale_msa.squeeze(2)
gate_msa = gate_msa.squeeze(2)
c_shift_msa = c_shift_msa.squeeze(2)
c_scale_msa = c_scale_msa.squeeze(2)
c_gate_msa = c_gate_msa.squeeze(2)
else:
e = self.scale_shift_table + temb.float()
shift_msa, scale_msa, gate_msa, c_shift_msa, c_scale_msa, c_gate_msa = e.chunk(
6, dim=1)
assert shift_msa.dtype == torch.float32
norm_hidden_states = (self.norm1(hidden_states.float()) *
(1 + scale_msa) + shift_msa).to(orig_dtype)
query, _ = self.to_q(norm_hidden_states)
key, _ = self.to_k(norm_hidden_states)
value, _ = self.to_v(norm_hidden_states)
if self.norm_q is not None:
query = self.norm_q(query)
if self.norm_k is not None:
key = self.norm_k(key)
query = query.squeeze(1).unflatten(2, (self.num_attention_heads, -1))
key = key.squeeze(1).unflatten(2, (self.num_attention_heads, -1))
value = value.squeeze(1).unflatten(2, (self.num_attention_heads, -1))
attn_output, _ = self.attn1(
query,
key,
value,
original_seq_len,
freqs_cis=freqs_cis,
)
attn_output = attn_output.flatten(2)
attn_output, _ = self.to_out(attn_output)
attn_output = attn_output.squeeze(1)
if self.cam_self_attn is not None and y_camera is not None:
attn_output = attn_output + self.cam_self_attn(
norm_hidden_states, y_camera)
null_shift = null_scale = torch.tensor([0], device=hidden_states.device)
norm_hidden_states, hidden_states = self.self_attn_residual_norm(
hidden_states, attn_output, gate_msa, null_shift, null_scale)
norm_hidden_states, hidden_states = norm_hidden_states.to(
orig_dtype), hidden_states.to(orig_dtype)
attn_output = self.attn2(norm_hidden_states,
context=encoder_hidden_states,
context_lens=None)
norm_hidden_states, hidden_states = self.cross_attn_residual_norm(
hidden_states, attn_output, 1, c_shift_msa, c_scale_msa)
norm_hidden_states, hidden_states = norm_hidden_states.to(
orig_dtype), hidden_states.to(orig_dtype)
ff_output = self.ffn(norm_hidden_states)
hidden_states = self.mlp_residual(hidden_states, ff_output, c_gate_msa)
hidden_states = hidden_states.to(orig_dtype)
return hidden_states
class DreamXWorldTransformer3DModel(WanTransformer3DModel):
_fsdp_shard_conditions = DreamXWorldConfig()._fsdp_shard_conditions
_compile_conditions = DreamXWorldConfig()._compile_conditions
_supported_attention_backends = DreamXWorldConfig(
)._supported_attention_backends
param_names_mapping = DreamXWorldConfig().param_names_mapping
reverse_param_names_mapping = DreamXWorldConfig().reverse_param_names_mapping
lora_param_names_mapping = DreamXWorldConfig().lora_param_names_mapping
def __init__(self, config: DreamXWorldConfig, hf_config: dict[str,
Any]) -> None:
BaseDiT.__init__(self, config=config, hf_config=hf_config)
self.quant_config = config.quant_config
inner_dim = config.num_attention_heads * config.attention_head_dim
self.hidden_size = config.hidden_size
self.num_attention_heads = config.num_attention_heads
self.in_channels = config.in_channels
self.out_channels = config.out_channels
self.num_channels_latents = config.num_channels_latents
self.patch_size = config.patch_size
self.text_len = config.text_len
assert config.num_attention_heads % get_sp_world_size() == 0, f"The number of attention heads ({config.num_attention_heads}) must be divisible by the sequence parallel size ({get_sp_world_size()})"
self.patch_embedding = PatchEmbed(in_chans=config.in_channels,
embed_dim=inner_dim,
patch_size=config.patch_size,
flatten=False)
self.condition_embedder = WanTimeTextImageEmbedding(
dim=inner_dim,
time_freq_dim=config.freq_dim,
text_embed_dim=config.text_dim,
image_embed_dim=config.image_dim,
)
self.blocks = nn.ModuleList([
DreamXWorldTransformerBlock(
inner_dim,
config.ffn_dim,
config.num_attention_heads,
config.qk_norm,
config.cross_attn_norm,
config.eps,
config.added_kv_proj_dim,
self._supported_attention_backends,
quant_config=config.quant_config,
prefix=f"{config.prefix}.blocks.{i}",
add_control_adapter=config.add_control_adapter,
cam_method=config.cam_method,
attn_compress=config.attn_compress,
cam_self_attn_layers=config.cam_self_attn_layers,
layer_idx=i)
for i in range(config.num_layers)
])
self.norm_out = LayerNormScaleShift(inner_dim,
norm_type="layer",
eps=config.eps,
elementwise_affine=False,
dtype=torch.float32,
compute_dtype=torch.float32)
self.proj_out = nn.Linear(
inner_dim, config.out_channels * math.prod(config.patch_size))
self.scale_shift_table = nn.Parameter(
torch.randn(1, 2, inner_dim) / inner_dim**0.5)
self.gradient_checkpointing = False
self.__post_init__()
def forward(self,
hidden_states: torch.Tensor,
encoder_hidden_states: torch.Tensor | list[torch.Tensor],
timestep: torch.LongTensor,
encoder_hidden_states_image: torch.Tensor | list[torch.Tensor]
| None = None,
guidance=None,
y_camera: dict[str, torch.Tensor] | None = None,
**kwargs) -> torch.Tensor:
orig_dtype = hidden_states.dtype
if encoder_hidden_states is not None and not isinstance(
encoder_hidden_states, torch.Tensor):
encoder_hidden_states = encoder_hidden_states[0]
if isinstance(encoder_hidden_states_image,
list) and len(encoder_hidden_states_image) > 0:
encoder_hidden_states_image = encoder_hidden_states_image[0]
else:
encoder_hidden_states_image = None
batch_size, _, num_frames, height, width = hidden_states.shape
p_t, p_h, p_w = self.patch_size
post_patch_num_frames = num_frames // p_t
post_patch_height = height // p_h
post_patch_width = width // p_w
d = self.hidden_size // self.num_attention_heads
rope_dim_list = [d - 4 * (d // 6), 2 * (d // 6), 2 * (d // 6)]
freqs_cos, freqs_sin = get_rotary_pos_embed(
(post_patch_num_frames, post_patch_height, post_patch_width),
self.hidden_size,
self.num_attention_heads,
rope_dim_list,
dtype=torch.float32 if current_platform.is_mps() else torch.float64,
rope_theta=10000)
freqs_cis = (freqs_cos.to(hidden_states.device).float(),
freqs_sin.to(hidden_states.device).float())
hidden_states = self.patch_embedding(hidden_states)
hidden_states = hidden_states.flatten(2).transpose(1, 2)
hidden_states, original_seq_len = sequence_model_parallel_shard(
hidden_states, dim=1)
if timestep.dim() == 2:
ts_seq_len = timestep.shape[1]
timestep = timestep.flatten()
else:
ts_seq_len = None
temb, timestep_proj, encoder_hidden_states, encoder_hidden_states_image = self.condition_embedder(
timestep,
encoder_hidden_states,
encoder_hidden_states_image,
timestep_seq_len=ts_seq_len)
if ts_seq_len is not None:
timestep_proj = timestep_proj.unflatten(2, (6, -1))
else:
timestep_proj = timestep_proj.unflatten(1, (6, -1))
if encoder_hidden_states_image is not None:
if encoder_hidden_states is not None:
encoder_hidden_states = torch.concat(
[encoder_hidden_states_image, encoder_hidden_states], dim=1)
else:
encoder_hidden_states = encoder_hidden_states_image
if current_platform.is_mps() or current_platform.is_npu():
encoder_hidden_states = encoder_hidden_states.to(orig_dtype)
assert encoder_hidden_states.dtype == orig_dtype
if torch.is_grad_enabled() and self.gradient_checkpointing:
for block in self.blocks:
hidden_states = self._gradient_checkpointing_func(
block, hidden_states, encoder_hidden_states,
timestep_proj, freqs_cis, original_seq_len, y_camera)
else:
for block in self.blocks:
hidden_states = block(hidden_states, encoder_hidden_states,
timestep_proj, freqs_cis, original_seq_len,
y_camera=y_camera)
if temb.dim() == 3:
shift, scale = (self.scale_shift_table.unsqueeze(0) +
temb.unsqueeze(2)).chunk(2, dim=2)
shift = shift.squeeze(2)
scale = scale.squeeze(2)
else:
shift, scale = (self.scale_shift_table +
temb.unsqueeze(1)).chunk(2, dim=1)
hidden_states = self.norm_out(hidden_states, shift, scale)
hidden_states = sequence_model_parallel_all_gather_with_unpad(
hidden_states, original_seq_len, dim=1)
hidden_states = self.proj_out(hidden_states)
hidden_states = hidden_states.reshape(batch_size, post_patch_num_frames,
post_patch_height,
post_patch_width, p_t, p_h, p_w,
-1)
hidden_states = hidden_states.permute(0, 7, 1, 4, 2, 5, 3, 6)
output = hidden_states.flatten(6, 7).flatten(4, 5).flatten(2, 3)
return output
EntryClass = DreamXWorldTransformer3DModel
+920
View File
@@ -0,0 +1,920 @@
# SPDX-License-Identifier: Apache-2.0
"""DreamX-World autoregressive causal DiT.
Adapted from DreamX-World's Apache-2.0
``wan/modules/causal_camera_model_2_2_prope_infinity.py``. The implementation is
kept native to FastVideo: no production import from DreamX, Diffusers, or
Transformers is required.
"""
from typing import Any
import math
import torch
import torch.nn as nn
import torch.nn.functional as F
from fastvideo.configs.models.dits.dreamx_world import DreamXWorldARConfig
from fastvideo.layers.linear import ReplicatedLinear
from fastvideo.models.dits.base import BaseDiT
from fastvideo.models.dits.dreamx_world import (_dreamx_apply_tiled_projmat,
_dreamx_prope_qkv)
def attention(q: torch.Tensor, k: torch.Tensor, v: torch.Tensor) -> torch.Tensor:
# Deliberately raw SDPA rather than fastvideo.attention.LocalAttention:
# (1) LocalAttention dispatches through the attention-backend registry, so
# FLASH_ATTN could be selected and its kernel is not bit-identical to
# torch SDPA — the AR KV-cache rollout must stay numerically frozen;
# (2) LocalAttention requires an active ForwardContext, which direct
# transformer invocations (parity tests) do not set;
# (3) the sibling causal model keeps raw SDPA in the same KV-cache window
# path (matrixgame2/causal_model.py).
# Sequence-parallel gap: this model never shards the sequence; run with
# sp_size=1 (see fastvideo/layers/AGENTS.md on documenting raw SDPA).
q_bhld = q.transpose(1, 2)
k_bhld = k.transpose(1, 2)
v_bhld = v.transpose(1, 2)
out = F.scaled_dot_product_attention(q_bhld, k_bhld, v_bhld, dropout_p=0.0)
return out.transpose(1, 2)
def prope_qkv(q: torch.Tensor, k: torch.Tensor, v: torch.Tensor,
viewmats: torch.Tensor, Ks: torch.Tensor):
q, k, v, output_projection = _dreamx_prope_qkv(q, k, v, viewmats, Ks)
def apply_fn_o(x: torch.Tensor) -> torch.Tensor:
return _dreamx_apply_tiled_projmat(x, output_projection)
return q, k, v, apply_fn_o
def sinusoidal_embedding_1d(dim, position):
assert dim % 2 == 0
half = dim // 2
position = position.type(torch.float64)
sinusoid = torch.outer(
position, torch.pow(10000, -torch.arange(half).to(position).div(half)))
return torch.cat([torch.cos(sinusoid), torch.sin(sinusoid)], dim=1)
def rope_params(max_seq_len, dim, theta=10000):
assert dim % 2 == 0
freqs = torch.outer(
torch.arange(max_seq_len),
1.0 / torch.pow(theta,
torch.arange(0, dim, 2).to(torch.float64).div(dim)))
return torch.polar(torch.ones_like(freqs), freqs)
class WanRMSNorm(nn.Module):
"""Kept private instead of fastvideo.layers.layernorm.RMSNorm.
The official DreamX-World ``model_2_2.py`` computes the RMS statistics in
the *input* dtype — the upstream code has the fp32 upcast explicitly
commented out (``# return self._norm(x.float())...``). FastVideo's RMSNorm
always normalizes in fp32, which is not bit-identical under bf16, so the
verbatim implementation stays.
"""
def __init__(self, dim, eps=1e-5):
super().__init__()
self.dim = dim
self.eps = eps
self.weight = nn.Parameter(torch.ones(dim))
def forward(self, x):
return self._norm(x).type_as(x) * self.weight
def _norm(self, x):
return x * torch.rsqrt(x.pow(2).mean(dim=-1, keepdim=True) + self.eps)
class WanLayerNorm(nn.LayerNorm):
"""Kept private instead of fastvideo.layers.layernorm.FP32LayerNorm.
The official DreamX-World ``model_2_2.py`` normalizes in the *input* dtype
(no ``x.float()`` upcast, unlike Wan2.1). FP32LayerNorm casts input and
affine params to fp32, which is not bit-identical under bf16, so the
verbatim implementation stays.
"""
def __init__(self, dim, eps=1e-6, elementwise_affine=False):
super().__init__(dim, elementwise_affine=elementwise_affine, eps=eps)
def forward(self, x):
return super().forward(x).type_as(x)
class WanCrossAttention(nn.Module):
def __init__(self, dim, num_heads, window_size=(-1, -1), qk_norm=True, eps=1e-6):
super().__init__()
self.num_heads = num_heads
self.head_dim = dim // num_heads
self.q = ReplicatedLinear(dim, dim)
self.k = ReplicatedLinear(dim, dim)
self.v = ReplicatedLinear(dim, dim)
self.o = ReplicatedLinear(dim, dim)
self.norm_q = WanRMSNorm(dim, eps=eps) if qk_norm else nn.Identity()
self.norm_k = WanRMSNorm(dim, eps=eps) if qk_norm else nn.Identity()
def _kv(self, context, b, n, d):
k, _ = self.k(context)
k = self.norm_k(k).view(b, -1, n, d)
v, _ = self.v(context)
v = v.view(b, -1, n, d)
return k, v
def forward(self, x, context, context_lens, crossattn_cache=None):
b, n, d = x.size(0), self.num_heads, self.head_dim
q, _ = self.q(x)
q = self.norm_q(q).view(b, -1, n, d)
if crossattn_cache is not None:
if not crossattn_cache["is_init"]:
crossattn_cache["is_init"] = True
k, v = self._kv(context, b, n, d)
crossattn_cache["k"] = k
crossattn_cache["v"] = v
else:
k = crossattn_cache["k"]
v = crossattn_cache["v"]
else:
k, v = self._kv(context, b, n, d)
x = attention(q, k, v)
x = x.flatten(2)
out, _ = self.o(x)
return out
def block_relativistic_rope(x, grid_sizes, freqs, start_frame=0, relative_frame_indices=None):
"""
Apply Block-Relativistic RoPE to input tensor.
Adapted from Infinity-RoPE (https://arxiv.org/abs/2511.20649).
Args:
x: Input tensor [B, L, num_heads, head_dim]
grid_sizes: Tensor [B, 3] containing (F, H, W)
freqs: RoPE frequencies
start_frame: Starting frame index for sequential RoPE
relative_frame_indices: Optional tensor [F] specifying explicit frame indices
for Block-Relativistic RoPE. Overrides start_frame if provided.
"""
n, c = x.size(2), x.size(3) // 2
freqs = freqs.split([c - 2 * (c // 3), c // 3, c // 3], dim=1)
output = []
for i, (f, h, w) in enumerate(grid_sizes.tolist()):
seq_len = f * h * w
x_i = torch.view_as_complex(x[i, :seq_len].to(torch.float64).reshape(
seq_len, n, -1, 2))
if relative_frame_indices is not None:
frame_indices = relative_frame_indices.long()
freqs_temporal = freqs[0][frame_indices].view(f, 1, 1, -1).expand(f, h, w, -1)
else:
freqs_temporal = freqs[0][start_frame:start_frame + f].view(f, 1, 1, -1).expand(f, h, w, -1)
freqs_i = torch.cat([
freqs_temporal,
freqs[1][:h].view(1, h, 1, -1).expand(f, h, w, -1),
freqs[2][:w].view(1, 1, w, -1).expand(f, h, w, -1)
], dim=-1).reshape(seq_len, 1, -1)
x_i = torch.view_as_real(x_i * freqs_i).flatten(2)
x_i = torch.cat([x_i, x[i, seq_len:]])
output.append(x_i)
return torch.stack(output).type_as(x)
class CausalWanSelfAttention(nn.Module):
"""Self-attention with KV cache and Block-Relativistic RoPE for causal inference."""
def __init__(self, dim, num_heads, local_attn_size=6, sink_size=1,
qk_norm=True, eps=1e-6):
assert dim % num_heads == 0
super().__init__()
self.dim = dim
self.num_heads = num_heads
self.head_dim = dim // num_heads
self.local_attn_size = local_attn_size
self.sink_size = sink_size
self.qk_norm = qk_norm
self.eps = eps
self.max_attention_size = 39600 if local_attn_size == -1 else local_attn_size * 880
self.q = ReplicatedLinear(dim, dim)
self.k = ReplicatedLinear(dim, dim)
self.v = ReplicatedLinear(dim, dim)
self.o = ReplicatedLinear(dim, dim)
self.norm_q = WanRMSNorm(dim, eps=eps) if qk_norm else nn.Identity()
self.norm_k = WanRMSNorm(dim, eps=eps) if qk_norm else nn.Identity()
def forward(self, x, seq_lens, grid_sizes, freqs, kv_cache,
current_start=0, cache_start=None, sink_recache_after_switch=False):
"""
Args:
x: Shape [B, L, C]
seq_lens: Shape [B]
grid_sizes: Shape [B, 3] containing (F, H, W)
freqs: RoPE frequencies [1024, head_dim / 2]
kv_cache: Dict with 'k', 'v', 'global_end_index', 'local_end_index'
current_start: Current position in the global token sequence
"""
b, s, n, d = *x.shape[:2], self.num_heads, self.head_dim
if cache_start is None:
cache_start = current_start
q, _ = self.q(x)
q = self.norm_q(q).view(b, s, n, d)
k, _ = self.k(x)
k = self.norm_k(k).view(b, s, n, d)
v, _ = self.v(x)
v = v.view(b, s, n, d)
frame_seqlen = math.prod(grid_sizes[0][1:]).item()
num_new_frames = grid_sizes[0][0].item()
current_end = current_start + q.shape[1]
sink_tokens = self.sink_size * frame_seqlen
kv_cache_size = kv_cache["k"].shape[1]
num_new_tokens = q.shape[1]
cache_update_info = None
is_recompute = current_end <= kv_cache["global_end_index"].item() and current_start > 0
if self.local_attn_size != -1 and (current_end > kv_cache["global_end_index"].item()) and (
num_new_tokens + kv_cache["local_end_index"].item() > kv_cache_size):
# === ROLLING MODE: cache full, evict oldest non-sink tokens ===
num_evicted_tokens = num_new_tokens + kv_cache["local_end_index"].item() - kv_cache_size
num_rolled_tokens = kv_cache["local_end_index"].item() - num_evicted_tokens - sink_tokens
local_end_index = kv_cache["local_end_index"].item() + current_end - \
kv_cache["global_end_index"].item() - num_evicted_tokens
local_start_index = local_end_index - num_new_tokens
temp_k = kv_cache["k"].detach().clone()
temp_v = kv_cache["v"].detach().clone()
temp_k[:, sink_tokens:sink_tokens + num_rolled_tokens] = \
temp_k[:, sink_tokens + num_evicted_tokens:sink_tokens + num_evicted_tokens + num_rolled_tokens].clone()
temp_v[:, sink_tokens:sink_tokens + num_rolled_tokens] = \
temp_v[:, sink_tokens + num_evicted_tokens:sink_tokens + num_evicted_tokens + num_rolled_tokens].clone()
write_start_index = max(local_start_index, sink_tokens) if is_recompute else local_start_index
roped_offset = max(0, write_start_index - local_start_index)
write_len = max(0, local_end_index - write_start_index)
if write_len > 0:
temp_k[:, write_start_index:local_end_index] = k[:, roped_offset:roped_offset + write_len]
temp_v[:, write_start_index:local_end_index] = v[:, roped_offset:roped_offset + write_len]
# Block-Relativistic RoPE: query uses window-relative indices
query_relative_indices = torch.arange(
self.local_attn_size - num_new_frames, self.local_attn_size, device=q.device)
roped_query = block_relativistic_rope(
q, grid_sizes, freqs, relative_frame_indices=query_relative_indices).type_as(v)
# Block-Relativistic RoPE: cached K uses position-in-window indices
num_cache_frames = local_end_index // frame_seqlen
cache_relative_indices = torch.arange(0, num_cache_frames, device=k.device)
cache_grid_sizes = grid_sizes.clone()
cache_grid_sizes[0, 0] = num_cache_frames
roped_temp_k = block_relativistic_rope(
temp_k[:, :local_end_index].view(b, num_cache_frames, frame_seqlen, n, d).flatten(1, 2),
cache_grid_sizes, freqs, relative_frame_indices=cache_relative_indices).type_as(v)
cache_update_info = {
"action": "roll_and_insert",
"sink_tokens": sink_tokens,
"num_rolled_tokens": num_rolled_tokens,
"num_evicted_tokens": num_evicted_tokens,
"local_start_index": local_start_index,
"local_end_index": local_end_index,
"write_start_index": write_start_index,
"write_end_index": local_end_index,
"new_k": k[:, roped_offset:roped_offset + write_len],
"new_v": v[:, roped_offset:roped_offset + write_len],
"current_end": current_end,
"is_recompute": is_recompute
}
else:
# === DIRECT INSERT MODE: cache not yet full ===
local_end_index = kv_cache["local_end_index"].item() + current_end - kv_cache["global_end_index"].item()
local_start_index = local_end_index - num_new_tokens
temp_k = kv_cache["k"].detach().clone()
temp_v = kv_cache["v"].detach().clone()
write_start_index = max(local_start_index, sink_tokens) if is_recompute else local_start_index
if sink_recache_after_switch:
write_start_index = local_start_index
roped_offset = max(0, write_start_index - local_start_index)
write_len = max(0, local_end_index - write_start_index)
if write_len > 0:
temp_k[:, write_start_index:local_end_index] = k[:, roped_offset:roped_offset + write_len]
temp_v[:, write_start_index:local_end_index] = v[:, roped_offset:roped_offset + write_len]
# RoPE with relative indices (growing sequentially before cache fills)
current_frame_in_window = local_start_index // frame_seqlen
query_relative_indices = torch.arange(
current_frame_in_window, current_frame_in_window + num_new_frames, device=q.device)
roped_query = block_relativistic_rope(
q, grid_sizes, freqs, relative_frame_indices=query_relative_indices).type_as(v)
num_cache_frames = local_end_index // frame_seqlen
cache_relative_indices = torch.arange(0, num_cache_frames, device=k.device)
cache_grid_sizes = grid_sizes.clone()
cache_grid_sizes[0, 0] = num_cache_frames
roped_temp_k = block_relativistic_rope(
temp_k[:, :local_end_index].view(b, num_cache_frames, frame_seqlen, n, d).flatten(1, 2),
cache_grid_sizes, freqs, relative_frame_indices=cache_relative_indices).type_as(v)
cache_update_info = {
"action": "direct_insert",
"local_start_index": local_start_index,
"local_end_index": local_end_index,
"write_start_index": write_start_index,
"write_end_index": local_end_index,
"new_k": k[:, roped_offset:roped_offset + write_len],
"new_v": v[:, roped_offset:roped_offset + write_len],
"current_end": current_end,
"is_recompute": is_recompute
}
# Attention: sink tokens + local window
if sink_tokens > 0:
local_budget = self.max_attention_size - sink_tokens
k_sink = roped_temp_k[:, :sink_tokens]
v_sink = temp_v[:, :sink_tokens]
if local_budget > 0:
local_start_for_window = max(sink_tokens, local_end_index - local_budget)
k_local = roped_temp_k[:, local_start_for_window:local_end_index]
v_local = temp_v[:, local_start_for_window:local_end_index]
k_cat = torch.cat([k_sink, k_local], dim=1)
v_cat = torch.cat([v_sink, v_local], dim=1)
else:
k_cat = k_sink
v_cat = v_sink
x = attention(roped_query, k_cat, v_cat)
else:
window_start = max(0, local_end_index - self.max_attention_size)
x = attention(
roped_query,
roped_temp_k[:, window_start:local_end_index],
temp_v[:, window_start:local_end_index])
x = x.flatten(2)
x, _ = self.o(x)
return x, (current_end, local_end_index, cache_update_info)
class CausalPropeSelfAttention(nn.Module):
"""PRoPE self-attention with optional KV cache for camera-controlled inference."""
def __init__(self, dim, attn_dim, num_heads, window_size=(-1, -1),
local_attn_size=-1, sink_size=0, qk_norm=True, eps=1e-6):
assert dim % num_heads == 0
assert attn_dim % num_heads == 0
super().__init__()
self.dim = dim
self.attn_dim = attn_dim
self.num_heads = num_heads
self.head_dim = attn_dim // num_heads
self.local_attn_size = local_attn_size
self.sink_size = sink_size
self.qk_norm = qk_norm
self.eps = eps
self.window_size = window_size
self.max_attention_size = 39600 if local_attn_size == -1 else local_attn_size * 880
self.q_proj = ReplicatedLinear(dim, attn_dim)
self.k_proj = ReplicatedLinear(dim, attn_dim)
self.v_proj = ReplicatedLinear(dim, attn_dim)
self.out_proj = ReplicatedLinear(attn_dim, dim)
self.norm_q = WanRMSNorm(attn_dim, eps=eps) if qk_norm else nn.Identity()
self.norm_k = WanRMSNorm(attn_dim, eps=eps) if qk_norm else nn.Identity()
nn.init.zeros_(self.out_proj.weight)
nn.init.zeros_(self.out_proj.bias)
def forward(self, x, cam_viewmats, cam_K, seq_lens, grid_sizes, freqs,
kv_cache=None, current_start=0, cache_start=None,
sink_recache_after_switch=False, cache_update_policy="commit_detached"):
"""
Args:
x: Shape [B, L, C]
cam_viewmats: Camera view matrices
cam_K: Camera intrinsics
kv_cache: Optional KV cache dict. When None, runs full attention over current chunk.
"""
b, s, n, d = *x.shape[:2], self.num_heads, self.head_dim
if cache_start is None:
cache_start = current_start
q, _ = self.q_proj(x)
q = self.norm_q(q).view(b, s, n, d)
k, _ = self.k_proj(x)
k = self.norm_k(k).view(b, s, n, d)
v, _ = self.v_proj(x)
v = v.view(b, s, n, d)
# Apply PRoPE (Positional Rotary Position Embedding from camera parameters)
q_t, k_t, v_t, apply_fn_o = prope_qkv(
q.transpose(1, 2), k.transpose(1, 2), v.transpose(1, 2),
viewmats=cam_viewmats, Ks=cam_K)
proped_q = q_t.transpose(1, 2)
proped_k = k_t.transpose(1, 2)
proped_v = v_t.transpose(1, 2)
if kv_cache is None:
# No cache: full attention over current chunk
x_out = attention(proped_q, proped_k, proped_v)
else:
# KV cache mode with rolling cache support
frame_seqlen = math.prod(grid_sizes[0][1:]).item()
num_new_tokens = s
current_end = current_start + num_new_tokens
sink_tokens = self.sink_size * frame_seqlen
kv_cache_size = kv_cache["k"].shape[1]
is_recompute = (current_end <= kv_cache["global_end_index"].item()) and (current_start > 0)
if self.local_attn_size != -1 and (current_end > kv_cache["global_end_index"].item()) and (
num_new_tokens + kv_cache["local_end_index"].item() > kv_cache_size):
# === ROLLING MODE ===
num_evicted_tokens = num_new_tokens + kv_cache["local_end_index"].item() - kv_cache_size
num_rolled_tokens = kv_cache["local_end_index"].item() - num_evicted_tokens - sink_tokens
local_end_index = kv_cache["local_end_index"].item() + current_end - \
kv_cache["global_end_index"].item() - num_evicted_tokens
local_start_index = local_end_index - num_new_tokens
if cache_update_policy != "none":
with torch.no_grad():
kv_cache["k"][:, sink_tokens:sink_tokens + num_rolled_tokens] = \
kv_cache["k"][:, sink_tokens + num_evicted_tokens:sink_tokens + num_evicted_tokens + num_rolled_tokens].clone()
kv_cache["v"][:, sink_tokens:sink_tokens + num_rolled_tokens] = \
kv_cache["v"][:, sink_tokens + num_evicted_tokens:sink_tokens + num_evicted_tokens + num_rolled_tokens].clone()
write_start_index = max(local_start_index, sink_tokens) if is_recompute else local_start_index
roped_offset = max(0, write_start_index - local_start_index)
write_len = max(0, local_end_index - write_start_index)
if write_len > 0:
with torch.no_grad():
kv_cache["k"][:, write_start_index:local_end_index] = proped_k[:, roped_offset:roped_offset + write_len].detach()
kv_cache["v"][:, write_start_index:local_end_index] = proped_v[:, roped_offset:roped_offset + write_len].detach()
else:
# === DIRECT INSERT MODE ===
local_end_index = kv_cache["local_end_index"].item() + current_end - kv_cache["global_end_index"].item()
local_start_index = local_end_index - num_new_tokens
if cache_update_policy != "none":
write_start_index = max(local_start_index, sink_tokens) if is_recompute else local_start_index
if sink_recache_after_switch:
write_start_index = local_start_index
roped_offset = max(0, write_start_index - local_start_index)
write_len = max(0, local_end_index - write_start_index)
if write_len > 0:
with torch.no_grad():
kv_cache["k"][:, write_start_index:local_end_index] = proped_k[:, roped_offset:roped_offset + write_len].detach()
kv_cache["v"][:, write_start_index:local_end_index] = proped_v[:, roped_offset:roped_offset + write_len].detach()
# Attention: sink tokens + local window
if sink_tokens > 0:
local_budget = self.max_attention_size - sink_tokens
k_sink = kv_cache["k"][:, :sink_tokens].detach()
v_sink = kv_cache["v"][:, :sink_tokens].detach()
if local_budget > 0:
local_start_for_window = max(sink_tokens, local_end_index - local_budget)
k_local = kv_cache["k"][:, local_start_for_window:local_end_index].detach()
v_local = kv_cache["v"][:, local_start_for_window:local_end_index].detach()
k_cat = torch.cat([k_sink, k_local], dim=1)
v_cat = torch.cat([v_sink, v_local], dim=1)
else:
k_cat = k_sink
v_cat = v_sink
x_out = attention(proped_q, k_cat, v_cat)
else:
window_start = max(0, local_end_index - self.max_attention_size)
x_out = attention(
proped_q,
kv_cache["k"][:, window_start:local_end_index].detach(),
kv_cache["v"][:, window_start:local_end_index].detach())
if not is_recompute and cache_update_policy != "none":
kv_cache["global_end_index"].fill_(current_end)
kv_cache["local_end_index"].fill_(local_end_index)
# Apply inverse PRoPE
x = apply_fn_o(x_out.transpose(1, 2)).transpose(1, 2)
x = x.flatten(2)
x, _ = self.out_proj(x)
return x
class CausalWanAttentionBlock(nn.Module):
def __init__(self, dim, ffn_dim, num_heads, local_attn_size=-1, sink_size=0,
qk_norm=True, cross_attn_norm=False, eps=1e-6, **kwargs):
super().__init__()
self.dim = dim
self.ffn_dim = ffn_dim
self.num_heads = num_heads
self.local_attn_size = local_attn_size
self.qk_norm = qk_norm
self.cross_attn_norm = cross_attn_norm
self.eps = eps
self.add_control_adapter = kwargs.get('add_control_adapter', False)
self.cam_method = kwargs.get('cam_method')
self.attn_compress = kwargs.get('attn_compress', 1)
self.layer_idx = kwargs.get('layer_idx')
cam_self_attn_layers = kwargs.get('cam_self_attn_layers')
# layers
self.norm1 = WanLayerNorm(dim, eps)
self.self_attn = CausalWanSelfAttention(
dim, num_heads, local_attn_size, sink_size, qk_norm, eps)
self.norm3 = WanLayerNorm(
dim, eps, elementwise_affine=True) if cross_attn_norm else nn.Identity()
self.cross_attn = WanCrossAttention(dim, num_heads, (-1, -1), qk_norm, eps)
self.norm2 = WanLayerNorm(dim, eps)
# nn.Linear (not ReplicatedLinear) on purpose: the official checkpoint
# stores these as positional Sequential keys (ffn.0 / ffn.2) that the
# copy-only converter and the strict-load tests require verbatim, and
# ReplicatedLinear's (out, bias) tuple return cannot compose inside
# nn.Sequential without renaming the state-dict surface.
self.ffn = nn.Sequential(
nn.Linear(dim, ffn_dim), nn.GELU(approximate='tanh'),
nn.Linear(ffn_dim, dim))
# PRoPE self-attention branch for camera control
add_cam_attn = self.add_control_adapter and self.cam_method == 'prope'
if add_cam_attn and cam_self_attn_layers is not None:
add_cam_attn = self.layer_idx in cam_self_attn_layers
if add_cam_attn:
self.cam_self_attn = CausalPropeSelfAttention(
dim, dim // self.attn_compress, num_heads,
local_attn_size=local_attn_size, sink_size=sink_size,
qk_norm=qk_norm, eps=eps)
# modulation
self.modulation = nn.Parameter(torch.randn(1, 6, dim) / dim**0.5)
def forward(self, x, e, seq_lens, grid_sizes, freqs, context, context_lens,
kv_cache, crossattn_cache=None, current_start=0, cache_start=None,
cam_viewmats=None, cam_K=None, sink_recache_after_switch=False,
cache_update_policy="commit_detached"):
num_frames, frame_seqlen = e.shape[1], x.shape[1] // e.shape[1]
e = (self.modulation.unsqueeze(1) + e).chunk(6, dim=2)
# self-attention
attn_input = (self.norm1(x).unflatten(
dim=1, sizes=(num_frames, frame_seqlen)) * (1 + e[1]) + e[0]).flatten(1, 2)
y, cache_update_info = self.self_attn(
attn_input, seq_lens, grid_sizes, freqs, kv_cache,
current_start, cache_start, sink_recache_after_switch)
# PRoPE camera attention (parallel branch)
if hasattr(self, 'cam_self_attn') and cam_viewmats is not None and cam_K is not None:
prope_kv_cache = None
if kv_cache is not None and "prope_k" in kv_cache:
prope_kv_cache = {
"k": kv_cache["prope_k"],
"v": kv_cache["prope_v"],
"global_end_index": kv_cache["prope_global_end_index"],
"local_end_index": kv_cache["prope_local_end_index"],
}
y = y + self.cam_self_attn(
attn_input, cam_viewmats, cam_K, seq_lens, grid_sizes, freqs,
kv_cache=prope_kv_cache, current_start=current_start,
cache_start=cache_start, cache_update_policy=cache_update_policy)
x = x + (y.unflatten(dim=1, sizes=(num_frames, frame_seqlen)) * e[2]).flatten(1, 2)
# cross-attention & FFN
x = x + self.cross_attn(self.norm3(x), context, context_lens,
crossattn_cache=crossattn_cache)
y = self.ffn(
(self.norm2(x).unflatten(dim=1, sizes=(num_frames, frame_seqlen))
* (1 + e[4]) + e[3]).flatten(1, 2))
x = x + (y.unflatten(dim=1, sizes=(num_frames, frame_seqlen)) * e[5]).flatten(1, 2)
return x, cache_update_info
class CausalHead(nn.Module):
def __init__(self, dim, out_dim, patch_size, eps=1e-6):
super().__init__()
self.dim = dim
self.out_dim = out_dim
self.patch_size = patch_size
self.eps = eps
out_dim = math.prod(patch_size) * out_dim
self.norm = WanLayerNorm(dim, eps)
self.head = ReplicatedLinear(dim, out_dim)
self.modulation = nn.Parameter(torch.randn(1, 2, dim) / dim**0.5)
def forward(self, x, e):
num_frames, frame_seqlen = e.shape[1], x.shape[1] // e.shape[1]
e = (self.modulation.unsqueeze(1) + e).chunk(2, dim=2)
x, _ = self.head(
self.norm(x).unflatten(dim=1, sizes=(num_frames, frame_seqlen))
* (1 + e[1]) + e[0])
return x
class DreamXWorldARTransformer3DModel(BaseDiT):
"""DreamX-World-5B autoregressive causal transformer."""
_fsdp_shard_conditions = DreamXWorldARConfig()._fsdp_shard_conditions
_compile_conditions = DreamXWorldARConfig()._compile_conditions
_supported_attention_backends = DreamXWorldARConfig()._supported_attention_backends
param_names_mapping = DreamXWorldARConfig().param_names_mapping
reverse_param_names_mapping = DreamXWorldARConfig().reverse_param_names_mapping
lora_param_names_mapping = DreamXWorldARConfig().lora_param_names_mapping
_no_split_modules = ["CausalWanAttentionBlock"]
def __init__(self, config: DreamXWorldARConfig, hf_config: dict[str, Any]) -> None:
super().__init__(config=config, hf_config=hf_config)
model_type = config.model_type
patch_size = config.patch_size
text_len = config.text_len
in_dim = config.in_channels
dim = config.hidden_size
ffn_dim = config.ffn_dim
freq_dim = config.freq_dim
text_dim = config.text_dim
out_dim = config.out_channels
num_heads = config.num_attention_heads
num_layers = config.num_layers
local_attn_size = config.local_attn_size
sink_size = config.sink_size
qk_norm = bool(config.qk_norm)
cross_attn_norm = config.cross_attn_norm
eps = config.eps
add_control_adapter = config.add_control_adapter
cam_method = config.cam_method
attn_compress = config.attn_compress
cam_self_attn_layers = config.cam_self_attn_layers
assert model_type in ['t2v', 'i2v', 'ti2v']
self.model_type = model_type
self.patch_size = patch_size
self.text_len = text_len
self.in_dim = in_dim
self.dim = dim
self.ffn_dim = ffn_dim
self.freq_dim = freq_dim
self.text_dim = text_dim
self.out_dim = out_dim
self.num_heads = num_heads
self.num_layers = num_layers
self.local_attn_size = local_attn_size
self.qk_norm = qk_norm
self.cross_attn_norm = cross_attn_norm
self.eps = eps
# embeddings — nn.Linear inside nn.Sequential on purpose: the official
# checkpoint keys are positional (text_embedding.0/.2, time_embedding.0/.2,
# time_projection.1) and must load verbatim (see ffn comment above).
self.patch_embedding = nn.Conv3d(
in_dim, dim, kernel_size=patch_size, stride=patch_size)
self.text_embedding = nn.Sequential(
nn.Linear(text_dim, dim), nn.GELU(approximate='tanh'),
nn.Linear(dim, dim))
self.time_embedding = nn.Sequential(
nn.Linear(freq_dim, dim), nn.SiLU(), nn.Linear(dim, dim))
self.time_projection = nn.Sequential(
nn.SiLU(), nn.Linear(dim, dim * 6))
# transformer blocks
self.blocks = nn.ModuleList([
CausalWanAttentionBlock(
dim, ffn_dim, num_heads, local_attn_size, sink_size,
qk_norm, cross_attn_norm, eps,
add_control_adapter=add_control_adapter,
cam_method=cam_method,
attn_compress=attn_compress,
layer_idx=layer_idx,
cam_self_attn_layers=cam_self_attn_layers)
for layer_idx in range(num_layers)
])
for layer_idx, block in enumerate(self.blocks):
block.self_attn.layer_idx = layer_idx
block.self_attn.num_layers = self.num_layers
# head
self.head = CausalHead(dim, out_dim, patch_size, eps)
# RoPE frequencies
assert (dim % num_heads) == 0 and (dim // num_heads) % 2 == 0
d = dim // num_heads
self.freqs = torch.cat([
rope_params(1024, d - 4 * (d // 6)),
rope_params(1024, 2 * (d // 6)),
rope_params(1024, 2 * (d // 6))
], dim=1)
self.num_attention_heads = num_heads
self.attention_head_dim = dim // num_heads
self.hidden_size = dim
self.in_channels = in_dim
self.out_channels = out_dim
self.num_channels_latents = out_dim
self.init_weights()
self.num_frame_per_block = config.arch_config.num_frames_per_block
self.__post_init__()
def forward(self, x=None, t=None, context=None, seq_len=None, y=None, y_camera=None,
kv_cache=None, crossattn_cache=None, current_start=0,
cache_start=0, cache_update_policy="commit_detached",
hidden_states=None, encoder_hidden_states=None, timestep=None, **kwargs):
"""
Causal inference with KV caching.
See Algorithm 2 of CausVid (https://arxiv.org/abs/2412.07772).
Args:
x: List of input video tensors [C_in, F, H, W]
t: Timestep tensor [B, L]
context: List of text embeddings [L, C]
seq_len: Maximum sequence length for positional encoding
y: Optional conditional video inputs (I2V mode)
y_camera: Camera parameters dict {'viewmats': ..., 'K': ...}
kv_cache: List of KV cache dicts per transformer block
crossattn_cache: List of cross-attention cache dicts
current_start: Current position in global token sequence
cache_start: Cache start position
cache_update_policy: Cache update strategy ('commit_detached' or 'none')
Returns:
Stacked output tensors [B, C_out, F, H/8, W/8]
"""
if x is None and hidden_states is not None:
x = [sample for sample in hidden_states]
if t is None and timestep is not None:
t = timestep
if context is None and encoder_hidden_states is not None:
if isinstance(encoder_hidden_states, torch.Tensor):
context = [sample for sample in encoder_hidden_states]
else:
context = encoder_hidden_states
if seq_len is None:
if torch.is_tensor(t):
seq_len = int(t.shape[1]) if t.dim() > 1 else int(t.numel())
elif x is not None:
sample = x[0]
seq_len = (sample.shape[1] // self.patch_size[0]) * (sample.shape[2] // self.patch_size[1]) * (sample.shape[3] // self.patch_size[2])
if x is None or t is None or context is None or seq_len is None:
raise ValueError("DreamXWorldARTransformer3DModel requires x/t/context/seq_len or FastVideo aliases")
device = self.patch_embedding.weight.device
if self.freqs.is_meta or self.freqs.device != device:
d = self.dim // self.num_heads
self.freqs = torch.cat([
rope_params(1024, d - 4 * (d // 6)),
rope_params(1024, 2 * (d // 6)),
rope_params(1024, 2 * (d // 6)),
], dim=1).to(device)
if y is not None:
x = [torch.cat([u, v], dim=0) for u, v in zip(x, y, strict=True)]
# patch embedding
x = [self.patch_embedding(u.unsqueeze(0)) for u in x]
grid_sizes = torch.stack(
[torch.tensor(u.shape[2:], dtype=torch.long) for u in x])
x = [u.flatten(2).transpose(1, 2) for u in x]
seq_lens = torch.tensor([u.size(1) for u in x], dtype=torch.long)
assert seq_lens.max() <= seq_len
x = torch.cat(x)
# time embedding
e = self.time_embedding(
sinusoidal_embedding_1d(self.freq_dim, t.flatten()).type_as(x))
e0 = self.time_projection(e).unflatten(
1, (6, self.dim)).unflatten(dim=0, sizes=t.shape)
# text embedding
context_lens = None
context = self.text_embedding(
torch.stack([
torch.cat(
[u, u.new_zeros(self.text_len - u.size(0), u.size(1))])
for u in context
]))
# camera parameters
if y_camera is not None and isinstance(y_camera, dict):
cam_viewmats = y_camera['viewmats']
cam_K = y_camera['K']
else:
cam_viewmats = None
cam_K = None
block_kwargs = dict(
e=e0, seq_lens=seq_lens, grid_sizes=grid_sizes, freqs=self.freqs,
context=context, context_lens=context_lens,
cam_viewmats=cam_viewmats, cam_K=cam_K,
cache_update_policy=cache_update_policy,
)
cache_update_infos = []
for block_index, block in enumerate(self.blocks):
block_kwargs.update({
"kv_cache": kv_cache[block_index] if kv_cache is not None else None,
"crossattn_cache": crossattn_cache[block_index] if crossattn_cache is not None else None,
"current_start": current_start,
"cache_start": cache_start,
})
x, block_cache_update_info = block(x, **block_kwargs)
if kv_cache is not None:
cache_update_infos.append((block_index, block_cache_update_info))
# Apply deferred cache updates
if kv_cache is not None and cache_update_infos and cache_update_policy != "none":
self._apply_cache_updates(kv_cache, cache_update_infos)
# head & unpatchify
x = self.head(x, e.unflatten(dim=0, sizes=t.shape).unsqueeze(2))
x = self.unpatchify(x, grid_sizes)
return torch.stack(x)
def _apply_cache_updates(self, kv_cache, cache_update_infos):
"""Apply deferred cache updates collected from all transformer blocks.
For Block-Relativistic RoPE, this stores un-roped K values in the cache.
RoPE is applied dynamically during attention based on each token's current
relative position in the sliding window.
"""
with torch.no_grad():
for block_index, (current_end, local_end_index, update_info) in cache_update_infos:
if update_info is not None:
cache = kv_cache[block_index]
if update_info["action"] == "roll_and_insert":
sink_tokens = update_info["sink_tokens"]
num_rolled_tokens = update_info["num_rolled_tokens"]
num_evicted_tokens = update_info["num_evicted_tokens"]
write_start_index = update_info.get("write_start_index", update_info["local_start_index"])
write_end_index = update_info.get("write_end_index", update_info["local_end_index"])
new_k = update_info["new_k"].detach()
new_v = update_info["new_v"].detach()
cache["k"][:, sink_tokens:sink_tokens + num_rolled_tokens] = \
cache["k"][:, sink_tokens + num_evicted_tokens:sink_tokens + num_evicted_tokens + num_rolled_tokens].clone()
cache["v"][:, sink_tokens:sink_tokens + num_rolled_tokens] = \
cache["v"][:, sink_tokens + num_evicted_tokens:sink_tokens + num_evicted_tokens + num_rolled_tokens].clone()
if write_end_index > write_start_index and new_k.shape[1] == (write_end_index - write_start_index):
cache["k"][:, write_start_index:write_end_index] = new_k
cache["v"][:, write_start_index:write_end_index] = new_v
elif update_info["action"] == "direct_insert":
write_start_index = update_info.get("write_start_index", update_info["local_start_index"])
write_end_index = update_info.get("write_end_index", update_info["local_end_index"])
new_k = update_info["new_k"].detach()
new_v = update_info["new_v"].detach()
if write_end_index > write_start_index and new_k.shape[1] == (write_end_index - write_start_index):
cache["k"][:, write_start_index:write_end_index] = new_k
cache["v"][:, write_start_index:write_end_index] = new_v
is_recompute = False if update_info is None else update_info.get("is_recompute", False)
if not is_recompute:
kv_cache[block_index]["global_end_index"].fill_(current_end)
kv_cache[block_index]["local_end_index"].fill_(local_end_index)
def unpatchify(self, x, grid_sizes):
"""Reconstruct video tensors from patch embeddings."""
c = self.out_dim
out = []
for u, v in zip(x, grid_sizes.tolist(), strict=True):
u = u[:math.prod(v)].view(*v, *self.patch_size, c)
u = torch.einsum('fhwpqrc->cfphqwr', u)
u = u.reshape(c, *[i * j for i, j in zip(v, self.patch_size, strict=True)])
out.append(u)
return out
def init_weights(self):
"""Initialize model parameters using Xavier initialization."""
for m in self.modules():
if isinstance(m, (nn.Linear, ReplicatedLinear)):
nn.init.xavier_uniform_(m.weight)
if m.bias is not None:
nn.init.zeros_(m.bias)
nn.init.xavier_uniform_(self.patch_embedding.weight.flatten(1))
for m in self.text_embedding.modules():
if isinstance(m, nn.Linear):
nn.init.normal_(m.weight, std=.02)
for m in self.time_embedding.modules():
if isinstance(m, nn.Linear):
nn.init.normal_(m.weight, std=.02)
nn.init.zeros_(self.head.head.weight)
EntryClass = DreamXWorldARTransformer3DModel
+7 -3
View File
@@ -11,7 +11,7 @@ import tempfile
from abc import ABC, abstractmethod
from collections.abc import Callable, Set
from dataclasses import dataclass, field
from functools import lru_cache
from functools import cache, lru_cache
from typing import NoReturn, TypeVar, cast
import cloudpickle
@@ -32,6 +32,8 @@ _TEXT_TO_VIDEO_DIT_MODELS = {
"HYWorldTransformer3DModel":
("dits", "hyworld", "HYWorldTransformer3DModel"),
"WanTransformer3DModel": ("dits", "wanvideo", "WanTransformer3DModel"),
"DreamXWorldTransformer3DModel": ("dits", "dreamx_world", "DreamXWorldTransformer3DModel"),
"DreamXWorldARTransformer3DModel": ("dits", "dreamx_world_ar", "DreamXWorldARTransformer3DModel"),
"CausalWanTransformer3DModel": ("dits", "causal_wanvideo", "CausalWanTransformer3DModel"),
"CosmosTransformer3DModel": ("dits", "cosmos", "CosmosTransformer3DModel"),
"Cosmos25Transformer3DModel": ("dits", "cosmos2_5", "Cosmos25Transformer3DModel"),
@@ -48,6 +50,8 @@ _TEXT_TO_VIDEO_DIT_MODELS = {
_IMAGE_TO_VIDEO_DIT_MODELS = {
# "HunyuanVideoTransformer3DModel": ("dits", "hunyuanvideo", "HunyuanVideoDiT"),
"WanTransformer3DModel": ("dits", "wanvideo", "WanTransformer3DModel"),
"DreamXWorldTransformer3DModel": ("dits", "dreamx_world", "DreamXWorldTransformer3DModel"),
"DreamXWorldARTransformer3DModel": ("dits", "dreamx_world_ar", "DreamXWorldARTransformer3DModel"),
"CausalWanTransformer3DModel": ("dits", "causal_wanvideo", "CausalWanTransformer3DModel"),
"MatrixGame2WanModel": ("dits", "matrixgame2", "MatrixGame2WanModel"),
"CausalMatrixGame2WanModel": ("dits", "matrixgame2", "CausalMatrixGame2WanModel"),
@@ -141,7 +145,7 @@ _LEGACY_FAST_VIDEO_MODELS = {
MODELS_PATH = os.path.dirname(__file__)
@lru_cache(maxsize=None)
@cache
def _discover_and_register_models() -> dict[str, tuple[str, str, str]]:
discovered_models: dict[str, tuple[str, str, str]] = {}
for root, dirs, files in os.walk(MODELS_PATH):
@@ -156,7 +160,7 @@ def _discover_and_register_models() -> dict[str, tuple[str, str, str]]:
filepath = os.path.join(root, filename)
try:
with open(filepath, "r", encoding="utf-8") as f:
with open(filepath, encoding="utf-8") as f:
source = f.read()
tree = ast.parse(source, filename=filename)
@@ -0,0 +1,28 @@
# SPDX-License-Identifier: Apache-2.0
from fastvideo.configs.pipelines.dreamx_world import (
DreamXWorld5BARPipelineConfig,
DreamXWorld5BCamPipelineConfig,
make_dreamx_world_5b_ar_dit_config,
make_dreamx_world_5b_cam_dit_config,
make_dreamx_world_5b_cam_text_encoder_config,
make_dreamx_world_5b_cam_vae_config,
)
from fastvideo.pipelines.basic.dreamx_world.dreamx_world_ar_pipeline import DreamXWorldARPipeline
from fastvideo.pipelines.basic.dreamx_world.dreamx_world_pipeline import DreamXWorldPipeline
from fastvideo.pipelines.basic.dreamx_world.stages import (
DREAMX_Y_CAMERA_KEY,
DreamXWorldCameraConditioningStage,
)
__all__ = [
"DREAMX_Y_CAMERA_KEY",
"DreamXWorld5BARPipelineConfig",
"DreamXWorld5BCamPipelineConfig",
"DreamXWorldCameraConditioningStage",
"DreamXWorldARPipeline",
"DreamXWorldPipeline",
"make_dreamx_world_5b_ar_dit_config",
"make_dreamx_world_5b_cam_dit_config",
"make_dreamx_world_5b_cam_text_encoder_config",
"make_dreamx_world_5b_cam_vae_config",
]
@@ -0,0 +1,219 @@
# SPDX-License-Identifier: Apache-2.0
"""DreamX-World autoregressive causal denoising stage."""
from __future__ import annotations
from typing import Any
import torch
from tqdm.auto import tqdm
from fastvideo.fastvideo_args import FastVideoArgs
from fastvideo.pipelines.basic.dreamx_world.stages import DREAMX_Y_CAMERA_KEY
from fastvideo.pipelines.pipeline_batch_info import ForwardBatch
from fastvideo.pipelines.stages.denoising import DenoisingStage
class DreamXWorldARCausalDenoisingStage(DenoisingStage):
"""Official DreamX AR-forcing denoising loop with KV cache."""
_AR_NOISE_SEED_OFFSET = 1_000_003
def __init__(self, transformer, scheduler, pipeline=None, vae=None) -> None:
super().__init__(transformer=transformer, scheduler=scheduler, pipeline=pipeline, vae=vae)
self.num_transformer_blocks = len(self.transformer.blocks)
self.num_frame_per_block = int(getattr(self.transformer, "num_frame_per_block", 3))
self.local_attn_size = int(getattr(self.transformer, "local_attn_size", 12))
def forward(self, batch: ForwardBatch, fastvideo_args: FastVideoArgs) -> ForwardBatch:
assert batch.latents is not None, "latents must be prepared before DreamX AR denoising"
assert batch.prompt_embeds, "prompt embeds must be prepared before DreamX AR denoising"
latents = batch.latents
device = latents.device
target_dtype = torch.bfloat16
autocast_enabled = device.type == "cuda" and not fastvideo_args.disable_autocast
frame_seq_length = (latents.shape[-2] // self.transformer.patch_size[1]) * (latents.shape[-1] //
self.transformer.patch_size[2])
timesteps = torch.tensor(
tuple(getattr(fastvideo_args.pipeline_config, "dmd_denoising_steps", (1000, 750, 500, 250))),
dtype=torch.long,
).cpu()
if getattr(fastvideo_args.pipeline_config, "warp_denoising_step", True):
self.scheduler.set_timesteps(1000)
scheduler_timesteps = torch.cat((self.scheduler.timesteps.cpu(), torch.tensor([0], dtype=torch.float32)))
timesteps = scheduler_timesteps[1000 - timesteps]
timesteps = timesteps.to(device)
if latents.shape[2] % self.num_frame_per_block != 0:
raise ValueError("DreamX AR latent frames must be divisible by num_frame_per_block")
y_camera = batch.extra.get(DREAMX_Y_CAMERA_KEY, batch.extra.get("y_camera"))
if isinstance(y_camera, dict):
y_camera = {
k: v.to(device=device, dtype=target_dtype) if torch.is_tensor(v) else v
for k, v in y_camera.items()
}
if batch.image_latent is not None and batch.image_latent.shape[1] == latents.shape[1]:
latents[:, :, :batch.image_latent.shape[2]] = batch.image_latent.to(device=device, dtype=latents.dtype)
kv_cache = self._initialize_kv_cache(latents.shape[0], target_dtype, device, frame_seq_length)
crossattn_cache = self._initialize_crossattn_cache(latents.shape[0], target_dtype, device)
prompt = batch.prompt_embeds[0]
if torch.is_tensor(prompt):
prompt = prompt.to(device=device, dtype=target_dtype)
context = [sample for sample in prompt]
else:
context = prompt
num_blocks = latents.shape[2] // self.num_frame_per_block
start = 0
first_frame_mask = torch.ones_like(latents)
first_frame_mask[:, :, 0] = 0
base_generator = batch.generator[0] if isinstance(batch.generator, list) else batch.generator
noise_generator = self._make_noise_generator(base_generator, device)
with tqdm(total=num_blocks * len(timesteps), desc="DreamX AR denoising", leave=False) as progress:
for _ in range(num_blocks):
current_num_frames = self.num_frame_per_block
block_latents = latents[:, :, start:start + current_num_frames]
noisy_input = block_latents.clone()
mask_block = first_frame_mask[:, :, start:start + current_num_frames]
camera_block = self._slice_camera(y_camera, start, current_num_frames)
for idx, current_timestep in enumerate(timesteps):
timestep = torch.full(
(latents.shape[0], current_num_frames * frame_seq_length),
int(current_timestep.item()),
device=device,
dtype=torch.long,
)
if start == 0:
timestep[:, :frame_seq_length] = 0
with torch.autocast(device_type="cuda", dtype=target_dtype, enabled=autocast_enabled):
denoised = self.transformer(
hidden_states=block_latents.to(target_dtype),
encoder_hidden_states=torch.stack(context).to(target_dtype),
timestep=timestep,
y_camera=camera_block,
kv_cache=kv_cache,
crossattn_cache=crossattn_cache,
current_start=start * frame_seq_length,
)
denoised = denoised.to(latents.dtype)
if idx < len(timesteps) - 1:
next_timestep = torch.full((latents.shape[0], current_num_frames),
int(timesteps[idx + 1].item()),
device=device,
dtype=torch.long)
noise_kwargs = {"device": device, "dtype": denoised.dtype}
if noise_generator is not None:
noise_kwargs["generator"] = noise_generator
noise = torch.randn(denoised.permute(0, 2, 1, 3, 4).shape, **noise_kwargs)
block_btchw = self.scheduler.add_noise(
denoised.permute(0, 2, 1, 3, 4).flatten(0, 1),
noise.flatten(0, 1),
next_timestep.flatten(),
).unflatten(0, (latents.shape[0], current_num_frames))
block_latents = block_btchw.permute(0, 2, 1, 3, 4)
block_latents = block_latents * mask_block + noisy_input * (1 - mask_block)
else:
block_latents = denoised * mask_block + noisy_input * (1 - mask_block)
progress.update()
latents[:, :, start:start + current_num_frames] = block_latents
self._update_context_cache(block_latents, context, camera_block, kv_cache, crossattn_cache, start,
frame_seq_length, target_dtype, autocast_enabled,
float(getattr(fastvideo_args.pipeline_config, "context_noise", 0.1)))
start += current_num_frames
batch.latents = latents
return batch
def _make_noise_generator(self, generator: torch.Generator | None, device: torch.device) -> torch.Generator | None:
if generator is None:
return None
if getattr(generator, "device", None) == device:
return generator
seed = int(generator.initial_seed()) + self._AR_NOISE_SEED_OFFSET
return torch.Generator(device=device).manual_seed(seed)
@staticmethod
def _context_noise_timestep(context_noise: float) -> int:
if 0.0 < context_noise <= 1.0:
return int(context_noise * 1000)
return int(context_noise)
def _slice_camera(self, y_camera: Any, start: int, num_frames: int):
if not isinstance(y_camera, dict):
return y_camera
return {
"viewmats": y_camera["viewmats"][:, start:start + num_frames],
"K": y_camera["K"][:, start:start + num_frames],
}
def _initialize_kv_cache(self, batch_size: int, dtype: torch.dtype, device: torch.device,
frame_seq_length: int) -> list[dict[str, Any]]:
size = self.local_attn_size * frame_seq_length if self.local_attn_size != -1 else 18480
heads = self.transformer.num_attention_heads
head_dim = self.transformer.attention_head_dim
cam_self_attn = next(
(getattr(block, "cam_self_attn", None)
for block in self.transformer.blocks if getattr(block, "cam_self_attn", None) is not None),
None,
)
caches = []
for _ in range(self.num_transformer_blocks):
cache = {
"k": torch.zeros(batch_size, size, heads, head_dim, dtype=dtype, device=device),
"v": torch.zeros(batch_size, size, heads, head_dim, dtype=dtype, device=device),
"global_end_index": torch.tensor([0], dtype=torch.long, device=device),
"local_end_index": torch.tensor([0], dtype=torch.long, device=device),
}
if cam_self_attn is not None:
cam_heads = int(cam_self_attn.num_heads)
cam_head_dim = int(cam_self_attn.head_dim)
cache.update({
"prope_k":
torch.zeros(batch_size, size, cam_heads, cam_head_dim, dtype=dtype, device=device),
"prope_v":
torch.zeros(batch_size, size, cam_heads, cam_head_dim, dtype=dtype, device=device),
"prope_global_end_index":
torch.tensor([0], dtype=torch.long, device=device),
"prope_local_end_index":
torch.tensor([0], dtype=torch.long, device=device),
})
caches.append(cache)
return caches
def _initialize_crossattn_cache(self, batch_size: int, dtype: torch.dtype, device: torch.device):
heads = self.transformer.num_attention_heads
head_dim = self.transformer.attention_head_dim
return [{
"k": torch.zeros(batch_size, 512, heads, head_dim, dtype=dtype, device=device),
"v": torch.zeros(batch_size, 512, heads, head_dim, dtype=dtype, device=device),
"is_init": False,
} for _ in range(self.num_transformer_blocks)]
def _update_context_cache(self, block_latents: torch.Tensor, context: Any, camera_block: Any,
kv_cache: list[dict[str, Any]], crossattn_cache: list[dict[str, Any]], start: int,
frame_seq_length: int, target_dtype: torch.dtype, autocast_enabled: bool,
context_noise: float) -> None:
timestep = torch.full(
(block_latents.shape[0], block_latents.shape[2] * frame_seq_length),
self._context_noise_timestep(context_noise),
device=block_latents.device,
dtype=torch.long,
)
with torch.autocast(device_type="cuda", dtype=target_dtype, enabled=autocast_enabled):
self.transformer(
hidden_states=block_latents.to(target_dtype),
encoder_hidden_states=torch.stack(context).to(target_dtype),
timestep=timestep,
y_camera=camera_block,
kv_cache=kv_cache,
crossattn_cache=crossattn_cache,
current_start=start * frame_seq_length,
)
@@ -0,0 +1,228 @@
# SPDX-License-Identifier: Apache-2.0
from __future__ import annotations
import math
from dataclasses import dataclass
import numpy as np
import torch
from scipy.interpolate import interp1d
from scipy.spatial.transform import Rotation, Slerp
_ACTION_TO_MOTION = {
"w": "forward",
"a": "left",
"d": "right",
"s": "backward",
"j": "left_rot",
"l": "right_rot",
"i": "up_rot",
"k": "down_rot",
}
_TRANSLATION_BASE_UNIT = 1.0
_ROTATION_BASE_UNIT = 10.0
_INTRINSIC_ROW = [0.8, 0.5, 0.5, 0.5]
@dataclass
class DreamXCamera:
fx: float
fy: float
cx: float
cy: float
w2c_mat: np.ndarray
@property
def c2w_mat(self) -> np.ndarray:
return np.linalg.inv(self.w2c_mat)
@classmethod
def from_pose_row(cls, row: list[float]) -> DreamXCamera:
w2c_mat = np.eye(4, dtype=np.float64)
w2c_mat[:3, :] = np.asarray(row[7:], dtype=np.float64).reshape(3, 4)
return cls(
fx=float(row[1]),
fy=float(row[2]),
cx=float(row[3]),
cy=float(row[4]),
w2c_mat=w2c_mat,
)
def _translation_step(motion_type: str, current_pose: dict[str, np.ndarray], value: float, duration: int) -> np.ndarray:
if motion_type in ("forward", "backward"):
yaw = np.radians(current_pose["rotation"][1])
pitch = np.radians(current_pose["rotation"][0])
forward = np.array([-math.sin(yaw) * math.cos(pitch), math.sin(pitch), math.cos(yaw) * math.cos(pitch)])
direction = 1 if motion_type == "forward" else -1
return forward * value * direction / duration
if motion_type in ("left", "right"):
yaw = np.radians(current_pose["rotation"][1])
right = np.array([math.cos(yaw), 0.0, math.sin(yaw)])
direction = -1 if motion_type == "left" else 1
return right * value * direction / duration
return np.zeros(3)
def _rotation_step(motion_type: str, value: float, duration: int) -> np.ndarray:
if not motion_type.endswith("rot"):
return np.zeros(3)
axis = motion_type.split("_")[0]
rotation = np.zeros(3)
if axis == "left":
rotation[1] = value
elif axis == "right":
rotation[1] = -value
elif axis == "up":
rotation[0] = -value
elif axis == "down":
rotation[0] = value
return rotation / duration
def _euler_to_quaternion(angles: np.ndarray) -> list[float]:
pitch, yaw, roll = np.radians(angles)
cy = math.cos(yaw * 0.5)
sy = math.sin(yaw * 0.5)
cp = math.cos(pitch * 0.5)
sp = math.sin(pitch * 0.5)
cr = math.cos(roll * 0.5)
sr = math.sin(roll * 0.5)
return [
cy * cp * cr + sy * sp * sr,
cy * sp * cr + sy * cp * sr,
sy * cp * cr - cy * sp * sr,
cy * cp * sr - sy * sp * cr,
]
def _quaternion_to_rotation_matrix(quaternion: list[float]) -> np.ndarray:
qw, qx, qy, qz = quaternion
return np.array([
[1 - 2 * (qy**2 + qz**2), 2 * (qx * qy - qw * qz), 2 * (qx * qz + qw * qy)],
[2 * (qx * qy + qw * qz), 1 - 2 * (qx**2 + qz**2), 2 * (qy * qz - qw * qx)],
[2 * (qx * qz - qw * qy), 2 * (qy * qz + qw * qx), 1 - 2 * (qx**2 + qy**2)],
])
def _pose_rows_from_actions(action_seq: list[str], action_speed_list: list[float], duration: int) -> list[list[float]]:
if len(action_seq) != len(action_speed_list):
raise ValueError("action_seq and action_speed_list must have the same length")
positions: list[np.ndarray] = []
rotations: list[np.ndarray] = []
current_pose = {
"position": np.array([0.0, 0.0, 0.0]),
"rotation": np.array([0.0, 0.0, 0.0]),
}
for action_id, speed in zip(action_seq, action_speed_list, strict=True):
motion_types = [_ACTION_TO_MOTION[key] for key in list(action_id)]
translation_step = np.zeros(3)
rotation_step = np.zeros(3)
for motion_type in motion_types:
translation_step += _translation_step(motion_type, current_pose,
float(speed) * _TRANSLATION_BASE_UNIT, duration)
rotation_step += _rotation_step(motion_type, float(speed) * _ROTATION_BASE_UNIT, duration)
segment_positions = []
segment_rotations = []
for index in range(1, duration + 1):
segment_positions.append(current_pose["position"] + translation_step * index)
segment_rotations.append(current_pose["rotation"] + rotation_step * index)
current_pose["position"] = segment_positions[-1].copy()
current_pose["rotation"] = segment_rotations[-1].copy()
positions.extend(segment_positions)
rotations.extend(segment_rotations)
rows: list[list[float]] = [[0.0] + _INTRINSIC_ROW + [0.0, 0.0] +
[1.0, 0.0, 0.0, 0.0, 0.0, 1.0, 0.0, 0.0, 0.0, 0.0, 1.0, 0.0]]
for index, (position, rotation) in enumerate(zip(positions, rotations, strict=False)):
rotation_matrix = _quaternion_to_rotation_matrix(_euler_to_quaternion(rotation))
translation = -rotation_matrix @ position
extrinsic = np.hstack([rotation_matrix, translation.reshape(3, 1)])
rows.append([float(index)] + _INTRINSIC_ROW + [0.0, 0.0] + extrinsic.flatten().tolist())
return rows
def _interpolate_camera_poses(
cameras: list[DreamXCamera],
src_indices: np.ndarray,
tgt_indices: np.ndarray,
) -> list[DreamXCamera]:
if len(cameras) <= 1:
return [cameras[0]] * len(tgt_indices) if cameras else []
src_rot_mat = np.array([camera.w2c_mat[:3, :3] for camera in cameras])
src_trans_vec = np.array([camera.w2c_mat[:3, 3] for camera in cameras])
dets = np.linalg.det(src_rot_mat)
flip_handedness = dets.size > 0 and np.median(dets) < 0.0
if flip_handedness:
flip_mat = np.diag([1.0, 1.0, -1.0]).astype(src_rot_mat.dtype)
src_rot_mat = src_rot_mat @ flip_mat
trans = interp1d(src_indices, src_trans_vec, axis=0, kind="linear", bounds_error=False,
fill_value="extrapolate")(tgt_indices)
quats = Rotation.from_matrix(src_rot_mat).as_quat().copy()
for index in range(1, len(quats)):
if np.dot(quats[index], quats[index - 1]) < 0:
quats[index] = -quats[index]
rot = Slerp(src_indices, Rotation.from_quat(quats))(tgt_indices).as_matrix()
if flip_handedness:
rot = rot @ flip_mat
ref = cameras[0]
result = []
for index in range(len(tgt_indices)):
w2c_mat = np.eye(4, dtype=np.float64)
w2c_mat[:3, :] = np.hstack([rot[index], trans[index].reshape(3, 1)])
result.append(DreamXCamera(ref.fx, ref.fy, ref.cx, ref.cy, w2c_mat))
return result
def _relative_c2w_poses(cameras: list[DreamXCamera]) -> np.ndarray:
abs_w2cs = [camera.w2c_mat for camera in cameras]
abs_c2ws = [camera.c2w_mat for camera in cameras]
target_cam_c2w = np.eye(4, dtype=np.float64)
abs2rel = target_cam_c2w @ abs_w2cs[0]
poses = [target_cam_c2w] + [abs2rel @ c2w for c2w in abs_c2ws[1:]]
return np.asarray(poses, dtype=np.float32)
def _invert_se3(transforms: torch.Tensor) -> torch.Tensor:
rotation_inv = transforms[..., :3, :3].transpose(-1, -2)
output = torch.zeros_like(transforms)
output[..., :3, :3] = rotation_inv
output[..., :3, 3] = -torch.einsum("...ij,...j->...i", rotation_inv, transforms[..., :3, 3])
output[..., 3, 3] = 1.0
return output
def build_dreamx_camera_condition(
action_seq: list[str],
action_speed_list: list[float],
*,
num_frames: int,
height: int,
width: int,
dtype: torch.dtype = torch.float32,
device: torch.device | str = "cpu",
) -> dict[str, torch.Tensor]:
del height, width # DreamX-World-5B-Cam uses fixed normalized intrinsics.
duration = math.ceil(num_frames / len(action_seq))
rows = _pose_rows_from_actions(action_seq, action_speed_list, duration)[:num_frames]
cameras = [DreamXCamera.from_pose_row(row) for row in rows]
latent_frame_count = 1 + (len(cameras) - 1) // 4
src_indices = np.arange(len(cameras), dtype=np.float64)
tgt_indices = np.linspace(0, len(cameras) - 1, latent_frame_count)
cameras = _interpolate_camera_poses(cameras, src_indices, tgt_indices)
c2ws = torch.as_tensor(_relative_c2w_poses(cameras), dtype=dtype, device=device)
viewmats = _invert_se3(c2ws)
intrinsics = torch.zeros((latent_frame_count, 3, 3), dtype=dtype, device=device)
intrinsics[:, 0, 0] = 969.6969696969696 / (960.0 * 2)
intrinsics[:, 1, 1] = 969.6969696969696 / (540.0 * 2)
intrinsics[:, 2, 2] = 1.0
return {"viewmats": viewmats, "K": intrinsics}
@@ -0,0 +1,20 @@
# SPDX-License-Identifier: Apache-2.0
"""Compatibility exports for DreamX-World pipeline configs."""
from fastvideo.configs.pipelines.dreamx_world import (
DreamXWorld5BARPipelineConfig,
DreamXWorld5BCamPipelineConfig,
make_dreamx_world_5b_ar_dit_config,
make_dreamx_world_5b_cam_dit_config,
make_dreamx_world_5b_cam_text_encoder_config,
make_dreamx_world_5b_cam_vae_config,
)
__all__ = [
"DreamXWorld5BARPipelineConfig",
"DreamXWorld5BCamPipelineConfig",
"make_dreamx_world_5b_ar_dit_config",
"make_dreamx_world_5b_cam_dit_config",
"make_dreamx_world_5b_cam_text_encoder_config",
"make_dreamx_world_5b_cam_vae_config",
]
@@ -0,0 +1,67 @@
# SPDX-License-Identifier: Apache-2.0
"""DreamX-World-5B autoregressive pipeline entrypoint."""
from fastvideo.api.sampling_param import SamplingParam
from fastvideo.configs.pipelines.dreamx_world import DreamXWorld5BARPipelineConfig
from fastvideo.fastvideo_args import FastVideoArgs
from fastvideo.logger import init_logger
from fastvideo.models.schedulers.scheduling_flow_match_euler_discrete import FlowMatchEulerDiscreteScheduler
from fastvideo.pipelines import ComposedPipelineBase, LoRAPipeline
from fastvideo.pipelines.basic.dreamx_world.ar_denoising import DreamXWorldARCausalDenoisingStage
from fastvideo.pipelines.basic.dreamx_world.stages import (
DreamXWorldCameraConditioningStage,
DreamXWorldImageVAEEncodingStage,
)
from fastvideo.pipelines.stages import (
ConditioningStage,
DecodingStage,
InputValidationStage,
LatentPreparationStage,
TextEncodingStage,
TimestepPreparationStage,
)
logger = init_logger(__name__)
class DreamXWorldARPipeline(LoRAPipeline, ComposedPipelineBase):
"""DreamX-World-5B autoregressive causal camera pipeline."""
_required_config_modules = ["text_encoder", "tokenizer", "vae", "transformer", "scheduler"]
pipeline_config_cls = DreamXWorld5BARPipelineConfig
sampling_params_cls = SamplingParam
def initialize_pipeline(self, fastvideo_args: FastVideoArgs):
self.modules["scheduler"] = FlowMatchEulerDiscreteScheduler(shift=fastvideo_args.pipeline_config.flow_shift)
self.modules["scheduler"].set_timesteps(1000)
def create_pipeline_stages(self, fastvideo_args: FastVideoArgs) -> None:
self.add_stage(stage_name="input_validation_stage", stage=InputValidationStage())
self.add_stage(stage_name="prompt_encoding_stage",
stage=TextEncodingStage(
text_encoders=[self.get_module("text_encoder")],
tokenizers=[self.get_module("tokenizer")],
))
self.add_stage(stage_name="conditioning_stage", stage=ConditioningStage())
self.add_stage(stage_name="timestep_preparation_stage",
stage=TimestepPreparationStage(scheduler=self.get_module("scheduler")))
self.add_stage(stage_name="latent_preparation_stage",
stage=LatentPreparationStage(
scheduler=self.get_module("scheduler"),
transformer=self.get_module("transformer", None),
))
self.add_stage(stage_name="image_latent_preparation_stage",
stage=DreamXWorldImageVAEEncodingStage(vae=self.get_module("vae")))
self.add_stage(stage_name="dreamx_camera_conditioning_stage", stage=DreamXWorldCameraConditioningStage())
self.add_stage(stage_name="denoising_stage",
stage=DreamXWorldARCausalDenoisingStage(
transformer=self.get_module("transformer"),
scheduler=self.get_module("scheduler"),
vae=self.get_module("vae"),
pipeline=self,
))
self.add_stage(stage_name="decoding_stage", stage=DecodingStage(vae=self.get_module("vae"), pipeline=self))
logger.info("DreamXWorldARPipeline initialized with autoregressive causal denoising")
EntryClass = DreamXWorldARPipeline
@@ -0,0 +1,78 @@
# SPDX-License-Identifier: Apache-2.0
"""DreamX-World video pipeline entrypoint."""
from fastvideo.api.sampling_param import SamplingParam
from fastvideo.fastvideo_args import FastVideoArgs
from fastvideo.logger import init_logger
from fastvideo.models.schedulers.scheduling_flow_match_euler_discrete import FlowMatchEulerDiscreteScheduler
from fastvideo.pipelines import ComposedPipelineBase, LoRAPipeline
from fastvideo.configs.pipelines.dreamx_world import DreamXWorld5BCamPipelineConfig
from fastvideo.pipelines.basic.dreamx_world.stages import DreamXWorldCameraConditioningStage
from fastvideo.pipelines.stages import (
ConditioningStage,
DecodingStage,
DenoisingStage,
InputValidationStage,
LatentPreparationStage,
TextEncodingStage,
TimestepPreparationStage,
)
logger = init_logger(__name__)
class DreamXWorldPipeline(LoRAPipeline, ComposedPipelineBase):
"""DreamX-World-5B-Cam pipeline with native FastVideo camera conditioning."""
_required_config_modules = ["text_encoder", "tokenizer", "vae", "transformer", "scheduler"]
pipeline_config_cls = DreamXWorld5BCamPipelineConfig
sampling_params_cls = SamplingParam
def initialize_pipeline(self, fastvideo_args: FastVideoArgs):
self.modules["scheduler"] = FlowMatchEulerDiscreteScheduler(shift=fastvideo_args.pipeline_config.flow_shift)
def create_pipeline_stages(self, fastvideo_args: FastVideoArgs) -> None:
self.add_stage(stage_name="input_validation_stage", stage=InputValidationStage())
self.add_stage(
stage_name="prompt_encoding_stage",
stage=TextEncodingStage(
text_encoders=[self.get_module("text_encoder")],
tokenizers=[self.get_module("tokenizer")],
),
)
self.add_stage(stage_name="conditioning_stage", stage=ConditioningStage())
self.add_stage(
stage_name="timestep_preparation_stage",
stage=TimestepPreparationStage(scheduler=self.get_module("scheduler")),
)
self.add_stage(
stage_name="latent_preparation_stage",
stage=LatentPreparationStage(
scheduler=self.get_module("scheduler"),
transformer=self.get_module("transformer", None),
),
)
self.add_stage(stage_name="dreamx_camera_conditioning_stage", stage=DreamXWorldCameraConditioningStage())
self.add_stage(
stage_name="denoising_stage",
stage=DenoisingStage(
transformer=self.get_module("transformer"),
transformer_2=self.get_module("transformer_2", None),
scheduler=self.get_module("scheduler"),
vae=self.get_module("vae"),
pipeline=self,
),
)
self.add_stage(stage_name="decoding_stage", stage=DecodingStage(vae=self.get_module("vae"), pipeline=self))
logger.info("DreamXWorldPipeline initialized with native camera conditioning")
EntryClass = DreamXWorldPipeline
@@ -0,0 +1,58 @@
# SPDX-License-Identifier: Apache-2.0
"""DreamX-World model family pipeline presets."""
from fastvideo.api.presets import InferencePreset, PresetStageSpec
_NEGATIVE_PROMPT_CN = ("色调艳丽,过曝,静态,细节模糊不清,字幕,风格,作品,画作,画面,"
"静止,整体发灰,最差质量,低质量,JPEG压缩残留,丑陋的,残缺的,"
"多余的手指,画得不好的手部,画得不好的脸部,畸形的,毁容的,"
"形态畸形的肢体,手指融合,静止不动的画面,杂乱的背景,三条腿,"
"背景人很多,倒着走")
_DENOISE_STAGE = PresetStageSpec(
name="denoise",
kind="denoising",
description="DreamX-World camera-conditioned denoising pass",
allowed_overrides=frozenset({
"num_inference_steps",
"guidance_scale",
}),
)
DREAMX_WORLD_5B_CAM = InferencePreset(
name="dreamx_world_5b_cam",
version=1,
model_family="dreamx_world",
description="DreamX-World 5B camera-control video generation",
workload_type="i2v",
stage_schemas=(_DENOISE_STAGE, ),
defaults={
"height": 480,
"width": 832,
"num_frames": 161,
"fps": 16,
"guidance_scale": 5.0,
"num_inference_steps": 30,
"negative_prompt": _NEGATIVE_PROMPT_CN,
},
)
DREAMX_WORLD_5B_AR = InferencePreset(
name="dreamx_world_5b_ar",
version=1,
model_family="dreamx_world",
description="DreamX-World 5B autoregressive camera-control generation",
workload_type="i2v",
stage_schemas=(_DENOISE_STAGE, ),
defaults={
"height": 704,
"width": 1280,
"num_frames": 1005,
"fps": 16,
"guidance_scale": 1.0,
"num_inference_steps": 4,
"negative_prompt": _NEGATIVE_PROMPT_CN,
},
)
ALL_PRESETS = (DREAMX_WORLD_5B_CAM, DREAMX_WORLD_5B_AR)
@@ -0,0 +1,156 @@
# SPDX-License-Identifier: Apache-2.0
"""DreamX-World pipeline stages."""
from __future__ import annotations
import torch
from fastvideo.fastvideo_args import FastVideoArgs
from fastvideo.logger import init_logger
from fastvideo.pipelines.pipeline_batch_info import ForwardBatch
from fastvideo.pipelines.stages.base import PipelineStage
from fastvideo.pipelines.stages.validators import VerificationResult
from fastvideo.pipelines.basic.dreamx_world.camera_conditioning import (
build_dreamx_camera_condition, )
DREAMX_Y_CAMERA_KEY = "dreamx_y_camera"
logger = init_logger(__name__)
class DreamXWorldCameraConditioningStage(PipelineStage):
"""Build PRoPE camera conditioning for DreamX-World-5B-Cam."""
def forward(
self,
batch: ForwardBatch,
fastvideo_args: FastVideoArgs,
) -> ForwardBatch:
del fastvideo_args
if DREAMX_Y_CAMERA_KEY in batch.extra:
return batch
action_seq = batch.extra.get("dreamx_action_seq", batch.action_list)
action_speed_list = batch.extra.get("dreamx_action_speed_list", batch.action_speed_list)
if action_seq is None:
action_seq = ["w"]
if action_speed_list is None:
action_speed_list = [4]
if isinstance(action_seq, str):
action_seq = [action_seq]
if isinstance(action_speed_list, int | float):
action_speed_list = [action_speed_list]
if len(action_speed_list) == 1 and len(action_seq) > 1:
action_speed_list = list(action_speed_list) * len(action_seq)
action_speed_list = [float(speed) for speed in action_speed_list]
height = int(batch.height) if batch.height is not None else 704
width = int(batch.width) if batch.width is not None else 1280
num_frames = int(batch.num_frames)
dtype = batch.latents.dtype if torch.is_tensor(batch.latents) else torch.float32
device = batch.latents.device if torch.is_tensor(batch.latents) else "cpu"
y_camera = build_dreamx_camera_condition(
list(action_seq),
action_speed_list,
num_frames=num_frames,
height=height,
width=width,
dtype=dtype,
device=device,
)
batch.extra[DREAMX_Y_CAMERA_KEY] = {key: value.unsqueeze(0) for key, value in y_camera.items()}
return batch
def verify_output(
self,
batch: ForwardBatch,
fastvideo_args: FastVideoArgs,
) -> VerificationResult:
del fastvideo_args
result = VerificationResult()
y_camera = batch.extra.get(DREAMX_Y_CAMERA_KEY)
result.add_check("dreamx_y_camera", y_camera, lambda value: isinstance(value, dict))
if isinstance(y_camera, dict):
result.add_check("dreamx_y_camera.viewmats", y_camera.get("viewmats"), torch.is_tensor)
result.add_check("dreamx_y_camera.K", y_camera.get("K"), torch.is_tensor)
return result
class DreamXWorldImageVAEEncodingStage(PipelineStage):
"""Encode the conditioning image into the first-frame latent.
Official AR-forcing flow (AMAP-ML/DreamX-World inference_ar_forcing.py):
the input image is resized, normalized to [-1, 1], VAE-encoded
deterministically, and written into frame 0 of the noise — the causal
denoiser then treats frame 0 as clean context. This stage produces
``batch.image_latent`` ([B, C, 1, H_lat, W_lat]); the injection into
the latents happens in DreamXWorldARCausalDenoisingStage.
"""
def __init__(self, vae) -> None:
super().__init__()
self.vae = vae
def forward(
self,
batch: ForwardBatch,
fastvideo_args: FastVideoArgs,
) -> ForwardBatch:
if batch.pil_image is None:
# No conditioning image: the causal denoiser falls back to
# running from pure noise (frame 0 uninitialized). Warn loudly —
# this pipeline is registered I2V and the official flow always
# forces from a frame.
logger.warning("DreamXWorldARPipeline called without an input image; "
"first-frame context will be noise (T2V-style). Pass an "
"image for the official AR-forcing behavior.")
return batch
from fastvideo.platforms import get_local_torch_device
from fastvideo.utils import PRECISION_TO_TYPE
device = get_local_torch_device()
image = batch.pil_image
if not isinstance(image, torch.Tensor):
import numpy as np
import PIL.Image
assert isinstance(image, PIL.Image.Image)
width = batch.width if isinstance(batch.width, int) else batch.width[0]
height = batch.height if isinstance(batch.height, int) else batch.height[0]
image = image.convert("RGB").resize((width, height), PIL.Image.Resampling.LANCZOS)
arr = torch.from_numpy(np.asarray(image)).float().permute(2, 0, 1) / 255.0
image = (arr - 0.5) / 0.5 # official Normalize([0.5], [0.5])
image = image.unsqueeze(0) # [1, C, H, W]
if image.dim() == 4:
image = image.unsqueeze(2) # [B, C, 1, H, W]
elif image.dim() == 5:
image = image[:, :, :1]
image = image.to(device=device, dtype=torch.float32)
vae_dtype = PRECISION_TO_TYPE[fastvideo_args.pipeline_config.vae_precision]
vae_autocast_enabled = (vae_dtype != torch.float32) and not fastvideo_args.disable_autocast
self.vae = self.vae.to(device)
with torch.autocast(device_type="cuda", dtype=vae_dtype, enabled=vae_autocast_enabled):
if not vae_autocast_enabled:
image = image.to(vae_dtype)
encoder_output = self.vae.encode(image)
# Official encode_to_latent is deterministic ((mean - mu) / sigma per
# channel); the posterior mean + shift/scale is the FastVideo
# equivalent of that normalization.
latent = encoder_output.mean
if getattr(self.vae, "shift_factor", None) is not None:
shift = self.vae.shift_factor
latent = latent - (shift.to(latent.device, latent.dtype) if isinstance(shift, torch.Tensor) else shift)
scale = self.vae.scaling_factor
latent = latent * (scale.to(latent.device, latent.dtype) if isinstance(scale, torch.Tensor) else scale)
batch.image_latent = latent
return batch
def verify_input(self, batch: ForwardBatch, fastvideo_args: FastVideoArgs) -> VerificationResult:
result = VerificationResult()
return result
+19
View File
@@ -191,6 +191,19 @@ class DenoisingStage(PipelineStage):
},
)
dreamx_y_camera = batch.extra.get("dreamx_y_camera", batch.extra.get("y_camera"))
if isinstance(dreamx_y_camera, dict):
dreamx_y_camera = {
key: value.to(device=local_device, dtype=target_dtype) if torch.is_tensor(value) else value
for key, value in dreamx_y_camera.items()
}
dreamx_camera_kwargs = self.prepare_extra_func_kwargs(
self.transformer.forward,
{
"y_camera": dreamx_y_camera,
},
)
for key in ("flux2_txt_ids", "flux2_img_ids"):
value = batch.extra.get(key)
if torch.is_tensor(value):
@@ -242,7 +255,11 @@ class DenoisingStage(PipelineStage):
# the image latent instead of appending along the channel dim
assert batch.image_latent is None, "TI2V task should not have image latents"
assert self.vae is not None, "VAE is not provided for TI2V task"
vae_device = next(self.vae.parameters()).device
self.vae = self.vae.to(local_device)
z = self.vae.encode(batch.pil_image).mean.float()
if getattr(fastvideo_args, "vae_cpu_offload", False):
self.vae = self.vae.to(vae_device)
if (hasattr(self.vae, "shift_factor") and self.vae.shift_factor is not None):
if isinstance(self.vae.shift_factor, torch.Tensor):
z -= self.vae.shift_factor.to(z.device, z.dtype)
@@ -495,6 +512,7 @@ class DenoisingStage(PipelineStage):
**pos_cond_kwargs,
**action_kwargs,
**camera_kwargs,
**dreamx_camera_kwargs,
**timesteps_r_kwarg,
**flux2_id_kwargs,
)
@@ -537,6 +555,7 @@ class DenoisingStage(PipelineStage):
**neg_cond_kwargs,
**action_kwargs,
**camera_kwargs,
**dreamx_camera_kwargs,
**timesteps_r_kwarg,
**flux2_id_kwargs,
)
+40
View File
@@ -20,6 +20,7 @@ from fastvideo.configs.pipelines.cosmos2_5 import (
Cosmos25Config,
Cosmos25_14BConfig,
)
from fastvideo.configs.pipelines.dreamx_world import DreamXWorld5BARPipelineConfig, DreamXWorld5BCamPipelineConfig
from fastvideo.configs.pipelines.hunyuan import FastHunyuanConfig, HunyuanConfig
from fastvideo.configs.pipelines.hunyuangamecraft import HunyuanGameCraftPipelineConfig
from fastvideo.configs.pipelines.gen3c import Gen3CConfig
@@ -773,6 +774,42 @@ def _register_configs() -> None:
model_family="wan",
default_preset="wan_2_2_ti2v_5b",
)
register_configs(
sampling_param_cls=None,
pipeline_config_cls=DreamXWorld5BCamPipelineConfig,
workload_types=(WorkloadType.I2V, ),
hf_model_paths=[
"FastVideo/DreamX-World-5B-Cam-Diffusers",
],
model_detectors=[
# Pattern also catches the raw GD-ML/DreamX-World-5B-Cam id and
# local converted dirs. Mutually exclusive with the AR detector
# below: Cam requires an explicit "cam" marker so hyphenated AR
# local paths (e.g. /ckpts/dreamx-world-5b-converted) don't
# first-match here — detector resolution is first-match in
# registration order.
lambda path:
("dreamx-world" in path.lower() and "cam" in path.lower()) or "dreamxworldpipeline" in path.lower()
],
model_family="dreamx_world",
default_preset="dreamx_world_5b_cam",
)
register_configs(
sampling_param_cls=None,
pipeline_config_cls=DreamXWorld5BARPipelineConfig,
workload_types=(WorkloadType.I2V, ),
hf_model_paths=[
"FastVideo/DreamX-World-5B-Diffusers",
],
model_detectors=[
lambda path:
("dreamx-world-5b" in path.lower() and "cam" not in path.lower()) or "dreamxworldarpipeline" in path.lower(
)
],
model_family="dreamx_world",
default_preset="dreamx_world_5b_ar",
)
register_configs(
sampling_param_cls=None,
pipeline_config_cls=FastWan2_2_TI2V_5B_Config,
@@ -951,6 +988,8 @@ def _register_presets() -> None:
from fastvideo.api.presets import register_preset
from fastvideo.pipelines.basic.cosmos.presets import (
ALL_PRESETS as COSMOS_PRESETS, )
from fastvideo.pipelines.basic.dreamx_world.presets import (
ALL_PRESETS as DREAMX_WORLD_PRESETS, )
from fastvideo.pipelines.basic.gamecraft.presets import (
ALL_PRESETS as GAMECRAFT_PRESETS, )
from fastvideo.pipelines.basic.gen3c.presets import (
@@ -984,6 +1023,7 @@ def _register_presets() -> None:
all_preset_groups = (
COSMOS_PRESETS,
DREAMX_WORLD_PRESETS,
FLUX2_PRESETS,
GAMECRAFT_PRESETS,
GEN3C_PRESETS,
@@ -3,9 +3,9 @@
from __future__ import annotations
import os
from collections.abc import Iterator
from contextlib import contextmanager
from logging import Logger
from typing import Iterator
from fastvideo import VideoGenerator
from fastvideo.tests.ssim.reference_utils import (
@@ -67,6 +67,12 @@ def _find_reference_video(reference_folder: str, prompt: str) -> str:
raise FileNotFoundError("Reference video missing")
def _remove_stale_generated_video(output_dir: str, output_video_name: str) -> None:
stale_path = os.path.join(output_dir, output_video_name)
if os.path.exists(stale_path):
os.remove(stale_path)
def _assert_similarity(
*,
logger: Logger,
@@ -214,6 +220,7 @@ def run_text_to_video_similarity_test(
)
output_video_name = f"{prompt[:100].strip()}.mp4"
os.makedirs(output_dir, exist_ok=True)
_remove_stale_generated_video(output_dir, output_video_name)
params_map = select_ssim_params(
default_params_map,
@@ -289,6 +296,7 @@ def run_image_to_video_similarity_test(
)
output_video_name = f"{prompt[:100].strip()}.mp4"
os.makedirs(output_dir, exist_ok=True)
_remove_stale_generated_video(output_dir, output_video_name)
params_map = select_ssim_params(
default_params_map,
@@ -0,0 +1,214 @@
# SPDX-License-Identifier: Apache-2.0
from __future__ import annotations
import os
from pathlib import Path
import pytest
from PIL import Image, ImageDraw
from fastvideo.logger import init_logger
from fastvideo.tests.ssim.inference_similarity_utils import (
resolve_inference_device_reference_folder,
run_image_to_video_similarity_test,
run_text_to_video_similarity_test,
)
logger = init_logger(__name__)
REQUIRED_GPUS = 1
device_reference_folder = resolve_inference_device_reference_folder(logger)
_LOCAL_CONVERTED_MODEL = Path("converted_weights/dreamx_world")
_MODEL_PATH = os.getenv(
"DREAMX_WORLD_SSIM_MODEL_PATH",
str(_LOCAL_CONVERTED_MODEL),
)
_LOCAL_AR_CANDIDATES = (
Path("/tmp/converted_dreamx_world_ar"),
Path("/root/data/dreamx_world_ar_converted"),
)
_DEFAULT_AR_MODEL_PATH = next(
(str(path) for path in _LOCAL_AR_CANDIDATES if path.exists()),
str(_LOCAL_AR_CANDIDATES[0]),
)
_AR_MODEL_PATH = os.getenv(
"DREAMX_WORLD_AR_SSIM_MODEL_PATH",
_DEFAULT_AR_MODEL_PATH,
)
DREAMX_WORLD_PARAMS = {
"num_gpus": 1,
"model_path": _MODEL_PATH,
"height": 64,
"width": 64,
"num_frames": 9,
"num_inference_steps": 1,
"guidance_scale": 1.0,
"seed": 1024,
"fps": 16,
}
DREAMX_WORLD_FULL_QUALITY_PARAMS = {
**DREAMX_WORLD_PARAMS,
"height": 480,
"width": 832,
"num_frames": 161,
"num_inference_steps": 30,
"guidance_scale": 5.0,
}
DREAMX_WORLD_AR_PARAMS = {
"num_gpus": 1,
"model_path": _AR_MODEL_PATH,
"height": 192,
"width": 192,
"num_frames": 81,
"num_inference_steps": 4,
"guidance_scale": 1.0,
"seed": 2048,
"fps": 16,
}
DREAMX_WORLD_AR_FULL_QUALITY_PARAMS = {
**DREAMX_WORLD_AR_PARAMS,
"height": 704,
"width": 1280,
"num_frames": 1005,
}
DREAMX_WORLD_MODEL_TO_PARAMS = {
"DreamX-World-5B-Cam": DREAMX_WORLD_PARAMS,
}
DREAMX_WORLD_AR_MODEL_TO_PARAMS = {
"DreamX-World-5B": DREAMX_WORLD_AR_PARAMS,
}
FULL_QUALITY_DREAMX_WORLD_MODEL_TO_PARAMS = {
"DreamX-World-5B-Cam": DREAMX_WORLD_FULL_QUALITY_PARAMS,
}
FULL_QUALITY_DREAMX_WORLD_AR_MODEL_TO_PARAMS = {
"DreamX-World-5B": DREAMX_WORLD_AR_FULL_QUALITY_PARAMS,
}
DREAMX_WORLD_TEST_CASES = [
(
"A cinematic first-person drive through a futuristic coastal city at sunrise, "
"reflective glass towers, clean streets, soft volumetric light.",
("w", "d", "w"),
(4.0, 2.0, 4.0),
),
]
DREAMX_WORLD_AR_TEST_CASES = [
(
"A long autonomous drive through a futuristic coastal city at sunrise, "
"smooth forward camera motion, reflective glass towers, clean streets.",
("w", "d", "w", "a"),
(2.0, 1.0, 2.0, 1.0),
),
]
def _write_deterministic_reference_image(path: Path) -> None:
path.parent.mkdir(parents=True, exist_ok=True)
image = Image.new("RGB", (96, 96), color=(42, 76, 112))
draw = ImageDraw.Draw(image)
draw.rectangle((0, 56, 96, 96), fill=(36, 44, 52))
draw.polygon([(0, 56), (48, 30), (96, 56)], fill=(120, 142, 158))
draw.rectangle((34, 42, 62, 72), fill=(178, 198, 212))
draw.line((0, 80, 96, 66), fill=(238, 209, 124), width=3)
image.save(path)
@pytest.mark.parametrize(("prompt", "action_list", "action_speed_list"), DREAMX_WORLD_TEST_CASES)
@pytest.mark.parametrize("attention_backend_name", ["TORCH_SDPA"])
@pytest.mark.parametrize("model_id", list(DREAMX_WORLD_MODEL_TO_PARAMS.keys()))
def test_dreamx_world_inference_similarity(
prompt: str,
action_list: tuple[str, ...],
action_speed_list: tuple[float, ...],
attention_backend_name: str,
model_id: str,
tmp_path: Path,
) -> None:
model_path = Path(str(DREAMX_WORLD_MODEL_TO_PARAMS[model_id]["model_path"]))
if not model_path.exists():
pytest.skip(
f"DreamX-World converted model path is missing: {model_path}. "
"Set DREAMX_WORLD_SSIM_MODEL_PATH to a FastVideo-loadable converted root."
)
image_path = tmp_path / "dreamx_world_ssim_input.png"
_write_deterministic_reference_image(image_path)
run_image_to_video_similarity_test(
logger=logger,
script_dir=os.path.dirname(os.path.abspath(__file__)),
device_reference_folder=device_reference_folder,
prompt=prompt,
image_path=str(image_path),
attention_backend_name=attention_backend_name,
model_id=model_id,
default_params_map=DREAMX_WORLD_MODEL_TO_PARAMS,
full_quality_params_map=FULL_QUALITY_DREAMX_WORLD_MODEL_TO_PARAMS,
min_acceptable_ssim=0.98,
init_kwargs_override={
"use_fsdp_inference": False,
"dit_cpu_offload": False,
"vae_cpu_offload": True,
"text_encoder_cpu_offload": True,
"pin_cpu_memory": False,
"override_pipeline_cls_name": "DreamXWorldPipeline",
},
generation_kwargs_override={
"action_list": list(action_list),
"action_speed_list": list(action_speed_list),
},
)
@pytest.mark.parametrize(("prompt", "action_list", "action_speed_list"), DREAMX_WORLD_AR_TEST_CASES)
@pytest.mark.parametrize("attention_backend_name", ["TORCH_SDPA"])
@pytest.mark.parametrize("model_id", list(DREAMX_WORLD_AR_MODEL_TO_PARAMS.keys()))
def test_dreamx_world_ar_inference_similarity(
prompt: str,
action_list: tuple[str, ...],
action_speed_list: tuple[float, ...],
attention_backend_name: str,
model_id: str,
) -> None:
model_path = Path(str(DREAMX_WORLD_AR_MODEL_TO_PARAMS[model_id]["model_path"]))
if not model_path.exists():
pytest.skip(
f"DreamX-World AR converted model path is missing: {model_path}. "
"Set DREAMX_WORLD_AR_SSIM_MODEL_PATH to a FastVideo-loadable converted root."
)
run_text_to_video_similarity_test(
logger=logger,
script_dir=os.path.dirname(os.path.abspath(__file__)),
device_reference_folder=device_reference_folder,
prompt=prompt,
attention_backend_name=attention_backend_name,
model_id=model_id,
default_params_map=DREAMX_WORLD_AR_MODEL_TO_PARAMS,
full_quality_params_map=FULL_QUALITY_DREAMX_WORLD_AR_MODEL_TO_PARAMS,
min_acceptable_ssim=0.98,
init_kwargs_override={
"use_fsdp_inference": False,
"dit_cpu_offload": False,
"vae_cpu_offload": True,
"text_encoder_cpu_offload": True,
"pin_cpu_memory": False,
"override_pipeline_cls_name": "DreamXWorldARPipeline",
},
generation_kwargs_override={
"action_list": list(action_list),
"action_speed_list": list(action_speed_list),
},
)
@@ -0,0 +1,124 @@
#!/usr/bin/env python3
# SPDX-License-Identifier: Apache-2.0
"""Convert DreamX-World-5B autoregressive weights to FastVideo layout.
The HF repository stores one raw official ``model.safetensors`` whose keys match
FastVideo's native ``DreamXWorldARTransformer3DModel``. The converter writes a
Diffusers-like root with ``transformer/config.json`` and reusable Wan2.2
components. Use ``--symlink-transformer`` locally to avoid duplicating the 21GB
AR tensor file.
"""
from __future__ import annotations
import argparse
import json
import shutil
from pathlib import Path
TRANSFORMER_CONFIG: dict[str, object] = {
"_class_name": "DreamXWorldARTransformer3DModel",
"model_type": "ti2v",
"patch_size": [1, 2, 2],
"text_len": 512,
"num_attention_heads": 24,
"attention_head_dim": 128,
"in_channels": 48,
"out_channels": 48,
"text_dim": 4096,
"freq_dim": 256,
"ffn_dim": 14336,
"num_layers": 30,
"local_attn_size": 12,
"sink_size": 3,
"cross_attn_norm": True,
"qk_norm": True,
"eps": 1e-6,
"add_control_adapter": True,
"cam_method": "prope",
"attn_compress": 4,
"cam_self_attn_layers": list(range(30)),
"num_frames_per_block": 3,
}
MODEL_INDEX: dict[str, object] = {
"_class_name": "DreamXWorldARPipeline",
"_diffusers_version": "0.31.0",
"scheduler": ["diffusers", "FlowMatchEulerDiscreteScheduler"],
"text_encoder": ["transformers", "UMT5EncoderModel"],
"tokenizer": ["transformers", "AutoTokenizer"],
"transformer": ["diffusers", "DreamXWorldARTransformer3DModel"],
"vae": ["diffusers", "AutoencoderKLWan"],
}
REUSED_COMPONENTS = ("scheduler", "text_encoder", "tokenizer", "vae")
def _source_safetensors(source: Path) -> Path:
if source.is_file():
return source
path = source / "model.safetensors"
if not path.exists():
raise FileNotFoundError(f"Missing AR model.safetensors under {source}")
return path
def convert_transformer(source: Path, output: Path, symlink_transformer: bool) -> None:
src = _source_safetensors(source)
transformer_dir = output / "transformer"
transformer_dir.mkdir(parents=True, exist_ok=True)
(transformer_dir / "config.json").write_text(json.dumps(TRANSFORMER_CONFIG, indent=2) + "\n")
dst = transformer_dir / "model.safetensors"
if dst.is_symlink() and not dst.exists():
dst.unlink() # dangling symlink: replace so re-conversion self-heals
if dst.exists():
return
if symlink_transformer:
dst.symlink_to(src.resolve())
else:
shutil.copy2(src, dst)
def _copy_or_link_component(component: str, component_source: Path, output: Path, symlink: bool) -> None:
src = component_source / component
dst = output / component
if not src.exists():
raise FileNotFoundError(f"Missing reused component source: {src}")
if dst.is_symlink() and not dst.exists():
dst.unlink() # dangling symlink: replace so re-conversion self-heals
if dst.exists():
return
if symlink:
dst.symlink_to(src.resolve(), target_is_directory=src.is_dir())
elif src.is_dir():
shutil.copytree(src, dst)
else:
shutil.copy2(src, dst)
def write_model_index(output: Path, component_source: Path | None, symlink_components: bool) -> None:
if component_source is not None:
for component in REUSED_COMPONENTS:
_copy_or_link_component(component, component_source, output, symlink_components)
missing = [component for component in REUSED_COMPONENTS if not (output / component).exists()]
if missing:
print("skipping model_index.json; missing reused components: " + ", ".join(missing))
return
(output / "model_index.json").write_text(json.dumps(MODEL_INDEX, indent=2) + "\n")
def main() -> None:
parser = argparse.ArgumentParser(description=__doc__)
parser.add_argument("--source", type=Path, required=True)
parser.add_argument("--output", type=Path, required=True)
parser.add_argument("--component-source", type=Path)
parser.add_argument("--symlink-components", action="store_true")
parser.add_argument("--symlink-transformer", action="store_true")
args = parser.parse_args()
convert_transformer(args.source, args.output, args.symlink_transformer)
write_model_index(args.output, args.component_source, args.symlink_components)
if __name__ == "__main__":
main()
@@ -0,0 +1,221 @@
#!/usr/bin/env python3
# SPDX-License-Identifier: Apache-2.0
"""Convert DreamX-World-5B-Cam raw transformer weights to FastVideo-loadable format.
The GD-ML/DreamX-World-5B-Cam repository stores the transformer as raw
DreamX/Wan official shards. FastVideo's TransformerLoader expects a Diffusers-like
transformer folder with a config.json and safetensors whose keys can be mapped by
WanVideoConfig.param_names_mapping. This script performs the raw official ->
Diffusers-like key rename and writes the DreamX 5B-Cam transformer config.
Example:
python scripts/checkpoint_conversion/dreamx_world_to_diffusers.py \
--source official_weights/dreamx_world \
--output converted_weights/dreamx_world
"""
from __future__ import annotations
import argparse
import json
import re
import shutil
from collections import OrderedDict
from pathlib import Path
import torch
from huggingface_hub import save_torch_state_dict
from safetensors import safe_open
from safetensors.torch import load_file
OFFICIAL_TO_DIFFUSERS_MAPPING: dict[str, str] = {
r"^text_embedding\.0\.(.*)$": r"condition_embedder.text_embedder.linear_1.\1",
r"^text_embedding\.2\.(.*)$": r"condition_embedder.text_embedder.linear_2.\1",
r"^time_embedding\.0\.(.*)$": r"condition_embedder.time_embedder.linear_1.\1",
r"^time_embedding\.2\.(.*)$": r"condition_embedder.time_embedder.linear_2.\1",
r"^time_projection\.1\.(.*)$": r"condition_embedder.time_proj.\1",
r"^img_emb\.proj\.0\.(.*)$": r"condition_embedder.image_embedder.norm1.\1",
r"^img_emb\.proj\.1\.(.*)$": r"condition_embedder.image_embedder.ff.net.0.proj.\1",
r"^img_emb\.proj\.3\.(.*)$": r"condition_embedder.image_embedder.ff.net.2.\1",
r"^img_emb\.proj\.4\.(.*)$": r"condition_embedder.image_embedder.norm2.\1",
r"^head\.modulation": r"scale_shift_table",
r"^head\.head\.(.*)$": r"proj_out.\1",
r"^blocks\.(\d+)\.self_attn\.q\.(.*)$": r"blocks.\1.attn1.to_q.\2",
r"^blocks\.(\d+)\.self_attn\.k\.(.*)$": r"blocks.\1.attn1.to_k.\2",
r"^blocks\.(\d+)\.self_attn\.v\.(.*)$": r"blocks.\1.attn1.to_v.\2",
r"^blocks\.(\d+)\.self_attn\.o\.(.*)$": r"blocks.\1.attn1.to_out.0.\2",
r"^blocks\.(\d+)\.self_attn\.norm_q\.(.*)$": r"blocks.\1.attn1.norm_q.\2",
r"^blocks\.(\d+)\.self_attn\.norm_k\.(.*)$": r"blocks.\1.attn1.norm_k.\2",
r"^blocks\.(\d+)\.cross_attn\.q\.(.*)$": r"blocks.\1.attn2.to_q.\2",
r"^blocks\.(\d+)\.cross_attn\.k\.(.*)$": r"blocks.\1.attn2.to_k.\2",
r"^blocks\.(\d+)\.cross_attn\.k_img\.(.*)$": r"blocks.\1.attn2.add_k_proj.\2",
r"^blocks\.(\d+)\.cross_attn\.v\.(.*)$": r"blocks.\1.attn2.to_v.\2",
r"^blocks\.(\d+)\.cross_attn\.v_img\.(.*)$": r"blocks.\1.attn2.add_v_proj.\2",
r"^blocks\.(\d+)\.cross_attn\.o\.(.*)$": r"blocks.\1.attn2.to_out.0.\2",
r"^blocks\.(\d+)\.cross_attn\.norm_q\.(.*)$": r"blocks.\1.attn2.norm_q.\2",
r"^blocks\.(\d+)\.cross_attn\.norm_k\.(.*)$": r"blocks.\1.attn2.norm_k.\2",
r"^blocks\.(\d+)\.cross_attn\.norm_k_img\.(.*)$": r"blocks.\1.attn2.norm_added_k.\2",
r"^blocks\.(\d+)\.ffn\.0\.(.*)$": r"blocks.\1.ffn.net.0.proj.\2",
r"^blocks\.(\d+)\.ffn\.2\.(.*)$": r"blocks.\1.ffn.net.2.\2",
r"^blocks\.(\d+)\.modulation": r"blocks.\1.scale_shift_table",
r"^blocks\.(\d+)\.norm3\.(.*)$": r"blocks.\1.norm2.\2",
}
TRANSFORMER_CONFIG: dict[str, object] = {
"_class_name": "DreamXWorldTransformer3DModel",
"patch_size": [1, 2, 2],
"text_len": 512,
"num_attention_heads": 24,
"attention_head_dim": 128,
"in_channels": 48,
"out_channels": 48,
"text_dim": 4096,
"freq_dim": 256,
"ffn_dim": 14336,
"num_layers": 30,
"cross_attn_norm": True,
"qk_norm": "rms_norm_across_heads",
"eps": 1e-6,
"image_dim": None,
"added_kv_proj_dim": None,
"rope_max_seq_len": 1024,
"add_control_adapter": True,
"cam_method": "prope",
"attn_compress": 1,
"cam_self_attn_layers": None,
}
MODEL_INDEX: dict[str, object] = {
"_class_name": "DreamXWorldPipeline",
"_diffusers_version": "0.31.0",
"scheduler": ["diffusers", "FlowMatchEulerDiscreteScheduler"],
"text_encoder": ["transformers", "UMT5EncoderModel"],
"tokenizer": ["transformers", "AutoTokenizer"],
"transformer": ["diffusers", "DreamXWorldTransformer3DModel"],
"vae": ["diffusers", "AutoencoderKLWan"],
}
REUSED_COMPONENTS = ("scheduler", "text_encoder", "tokenizer", "vae")
def map_transformer_key(key: str) -> str:
for pattern, replacement in OFFICIAL_TO_DIFFUSERS_MAPPING.items():
if re.match(pattern, key):
return re.sub(pattern, replacement, key)
return key
def _safetensor_files(source: Path) -> list[Path]:
if source.is_file():
if source.suffix != ".safetensors":
raise ValueError(f"Only .safetensors files are supported, got {source}")
return [source]
index_path = source / "diffusion_pytorch_model.safetensors.index.json"
if index_path.exists():
index = json.loads(index_path.read_text())
return sorted({source / shard for shard in index["weight_map"].values()})
files = sorted(source.glob("*.safetensors"))
if not files:
raise FileNotFoundError(f"No safetensors files found under {source}")
return files
def convert_transformer(source: Path, output: Path, max_shard_size: str) -> None:
transformer_dir = output / "transformer"
transformer_dir.mkdir(parents=True, exist_ok=True)
(transformer_dir / "config.json").write_text(json.dumps(TRANSFORMER_CONFIG, indent=2) + "\n")
converted: OrderedDict[str, torch.Tensor] = OrderedDict()
for shard in _safetensor_files(source):
print(f"loading {shard}")
for key, tensor in load_file(shard, device="cpu").items():
new_key = map_transformer_key(key)
if new_key in converted:
raise ValueError(f"Duplicate converted key: {new_key}")
converted[new_key] = tensor
print(f"saving {len(converted)} tensors to {transformer_dir}")
save_torch_state_dict(converted, transformer_dir, max_shard_size=max_shard_size)
def _copy_or_link_component(component: str, component_source: Path, output: Path, symlink: bool) -> None:
src = component_source / component
dst = output / component
if not src.exists():
raise FileNotFoundError(f"Missing reused component source: {src}")
if dst.is_symlink() and not dst.exists():
dst.unlink() # dangling symlink: replace so re-conversion self-heals
if dst.exists():
print(f"keeping existing {dst}")
return
if symlink:
dst.symlink_to(src.resolve(), target_is_directory=src.is_dir())
print(f"linked {dst} -> {src}")
elif src.is_dir():
shutil.copytree(src, dst)
print(f"copied {src} -> {dst}")
else:
shutil.copy2(src, dst)
print(f"copied {src} -> {dst}")
def write_model_index(output: Path, component_source: Path | None, symlink_components: bool) -> None:
if component_source is not None:
for component in REUSED_COMPONENTS:
_copy_or_link_component(component, component_source, output, symlink_components)
missing = [component for component in REUSED_COMPONENTS if not (output / component).exists()]
if missing:
print("skipping model_index.json; missing reused components: " + ", ".join(missing))
print("pass --component-source <Wan2.2 Diffusers root> to copy or link reused components")
return
(output / "model_index.json").write_text(json.dumps(MODEL_INDEX, indent=2) + "\n")
print(f"wrote {output / 'model_index.json'}")
def analyze(source: Path) -> None:
total = 0
unchanged = 0
examples: list[tuple[str, str]] = []
for shard in _safetensor_files(source):
with safe_open(shard, framework="pt", device="cpu") as tensors:
for key in tensors:
total += 1
new_key = map_transformer_key(key)
unchanged += int(new_key == key)
if len(examples) < 20 and new_key != key:
examples.append((key, new_key))
print(f"total_keys={total} unchanged_keys={unchanged}")
for old, new in examples:
print(f"{old} -> {new}")
def main() -> None:
parser = argparse.ArgumentParser(description=__doc__)
parser.add_argument("--source", type=Path, required=True, help="DreamX raw transformer directory or safetensors file")
parser.add_argument("--output", type=Path, required=True, help="Output model root; transformer/ is created inside it")
parser.add_argument("--max-shard-size", default="10GB")
parser.add_argument(
"--component-source",
type=Path,
help="Optional Wan2.2 Diffusers root whose scheduler/text_encoder/tokenizer/vae components are reused.",
)
parser.add_argument(
"--symlink-components",
action="store_true",
help="Symlink reused components from --component-source instead of copying them.",
)
parser.add_argument("--analyze", action="store_true", help="Only print key mapping summary")
args = parser.parse_args()
if args.analyze:
analyze(args.source)
else:
convert_transformer(args.source, args.output, args.max_shard_size)
write_model_index(args.output, args.component_source, args.symlink_components)
if __name__ == "__main__":
main()
@@ -0,0 +1,189 @@
# DreamX World Port Status
## Summary
- model_family: `dreamx_world`
- workload_types: `I2V camera-control compatibility shim`; `I2V autoregressive camera-control forcing`
- official_ref: `https://github.com/AMAP-ML/DreamX-World`
- official_ref_dir: `DreamX-World/`
- hf_weights_path: `GD-ML/DreamX-World-5B-Cam`
- local_weights_dir: `official_weights/dreamx_world`
- source_layout: `raw_official`
- local_tests_readme: `tests/local_tests/dreamx_world/README.md`
## Current Phase
- phase: `phase_11_post_parity_handoff`
- status: `complete`
- owner: `orchestrator`
- last_updated: `2026-07-02`
## Component Matrix
| Component | Type | Reuse/Port | Official Definition | Official Instantiation | FastVideo Target | Prototype | Conversion | Parity | Open Issues |
|---|---|---|---|---|---|---|---|---|---|
| transformer | dit | ported_dedicated | `DreamX-World/models/wan_transformer3d.py`; PRoPE helpers in `DreamX-World/models/prope_utils.py` | `DreamX-World/inference_dreamx5b.py::setup_models`, `Wan2_2Transformer3DModel.from_pretrained(... cam_method=prope, add_control_adapter=True)` | `fastvideo/models/dits/dreamx_world.py`; `fastvideo/configs/models/dits/dreamx_world.py`; DreamX pipeline config helper | native_prope_pass | real_conversion_pass | strict_load_and_forward_parity_pass | none |
| vae | vae | reuse_pending | `DreamX-World/models/wan_vae3_8.py` | `DreamX-World/inference_dreamx5b.py::setup_models`, `AutoencoderKLWan3_8.from_pretrained(Wan2.2_VAE.pth)` | `fastvideo/models/vaes/wanvae.py`; DreamX VAE config helper | config_smoke_pass | raw_key_mapping_pass | encode_parity_pass | none |
| text_encoder/tokenizer | encoder | reuse_pending | `DreamX-World/models/wan_text_encoder.py`; tokenizer via Wan2.2 base model | `DreamX-World/inference_dreamx5b.py::setup_models`, `WanT5EncoderModel` + tokenizer subpaths | `fastvideo/models/encoders/t5.py::UMT5EncoderModel`; DreamX UMT5 config helper | config_smoke_pass | staged_weight_load_pass | hidden_state_parity_pass | none |
| scheduler | generic | reuse_proven | Diffusers `FlowMatchEulerDiscreteScheduler` | `DreamX-World/inference_dreamx5b.py::setup_models`, default `sampler_name=Flow` | `fastvideo/models/schedulers/scheduling_flow_match_euler_discrete.py` | pass | not_required | non_skip_pass | Q003 |
| camera_conditioning | generic | port_pending | `DreamX-World/utils/inference_utils.py`, `DreamX-World/models/prope_utils.py`, `DreamX-World/wan/modules/camera_prope.py` | `DreamX-World/inference_dreamx5b.py::get_camera_sequence`, `pipeline(... control_camera_video=...)` | `fastvideo/pipelines/basic/dreamx_world/camera_conditioning.py` | pass | not_required | non_skip_pass | none |
| pipeline | pipeline | port_complete | `DreamX-World/pipeline/pipeline_dreamxworld.py` | `DreamX-World/inference_dreamx5b.py::process_inference_from_json` | `fastvideo/pipelines/basic/dreamx_world/` plus config/preset/registry | pipeline_load_generate_smoke_pass | model_index_and_config_consistency_smoke_pass | pipeline_api_vs_worker_forward_parity_pass | none |
| ar_transformer | dit | port_complete | `DreamX-World/wan/modules/causal_camera_model_2_2_prope_infinity.py::CausalWanModel` | `DreamX-World/inference_ar_forcing.py::load_pipeline` | `fastvideo/models/dits/dreamx_world_ar.py`; `fastvideo/configs/models/dits/dreamx_world.py::DreamXWorldARConfig` | tiny_official_forward_parity_pass | identity_conversion_pass | real_5b_strict_load_pass | none |
| ar_pipeline | pipeline | port_complete | `DreamX-World/pipeline/pipeline_causal_camera.py` | `DreamX-World/inference_ar_forcing.py::main` | `fastvideo/pipelines/basic/dreamx_world/dreamx_world_ar_pipeline.py`; `fastvideo/pipelines/basic/dreamx_world/ar_denoising.py`; registry/preset/config | config_registry_pass | symlink_model_index_pass | short_full_generation_pass | none |
## Conversion State
- conversion_script: `scripts/checkpoint_conversion/dreamx_world_to_diffusers.py`
- converted_weights_dir: `converted_weights/dreamx_world`
- source_layout: `raw_official`
- strict_load_status: `pass`
- conversion_script_status: `transformer_model_index_and_config_consistency_smoke_pass`
- model_index_status: `smoke_pass`
- passthrough_components: `Wan2.2 Diffusers scheduler, tokenizer, and text encoder are symlinked from official_weights/Wan2.2-TI2V-5B-Diffusers; VAE parity uses raw Wan2.2_VAE.pth with an explicit DreamX raw-to-FastVideo key mapper because official encode returns normalized latents.`
- retry_history: `none`
## Parity Commands
| Scope | Command | Last Result | Notes |
|---|---|---|---|
| transformer | `python -m pytest tests/local_tests/dreamx_world/test_dreamx_world_transformer_parity.py -v -s` | strict_load_and_forward_parity_pass | 2026-07-01: converted real 5B-Cam transformer shards strict-load into dedicated `DreamXWorldTransformer3DModel` with 0 shape mismatches; official-vs-FastVideo small-input fp32 forward parity passes on CUDA (`diff_max=0.072533`, `diff_mean=0.008014`). |
| vae | `python -m pytest tests/local_tests/dreamx_world/test_dreamx_world_vae_parity.py -v -s` | encode_parity_pass | 2026-06-30: official DreamX raw `Wan2.2_VAE.pth` maps 196/196 keys into FastVideo Wan VAE and encode parity passes after applying the same official latent normalization (`(mu - mean) / std`). |
| text_encoder | `python -m pytest tests/local_tests/dreamx_world/test_dreamx_world_text_encoder_parity.py -v -s` | hidden_state_parity_pass | 2026-06-30: official `WanT5EncoderModel` vs FastVideo `UMT5EncoderModel` hidden-state parity passes on CUDA using staged Wan2.2 text encoder/tokenizer weights and reference-only `xfuser` stubs. |
| scheduler | `python -m pytest tests/local_tests/dreamx_world/test_dreamx_world_scheduler_parity.py -v -s` | non_skip_pass | 2026-06-30: FastVideo FlowMatch scheduler matches official Diffusers timesteps and step output for DreamX default Flow sampler; `DreamXWorldPipeline` initializes FlowMatch with official `shift=3.0`. |
| camera_conditioning | `python -m pytest tests/local_tests/dreamx_world/test_dreamx_world_camera_conditioning_parity.py -v -s` | non_skip_pass | 2026-06-29: 3 parameterized cases passed against official reference on CPU. |
| pipeline_config | `python -m pytest tests/local_tests/dreamx_world/test_dreamx_world_pipeline_config.py -v -s` | pipeline_entry_preset_scheduler_modelinfo_and_camera_stage_smoke_pass | DreamX 5B-Cam PipelineConfig wires DiT/VAE/UMT5/Flow/TI2V settings and official `shift=3.0`; default preset is registered for `GD-ML/DreamX-World-5B-Cam`; local converted-style `model_index.json` resolves to `DreamXWorldPipeline`; the pipeline initializes FlowMatch, camera conditioning writes `batch.extra["dreamx_y_camera"]`, and generic denoising can pass it as `y_camera`. |
| ar_transformer | `PYTHONPATH=/workspace/FastVideo python -m pytest tests/local_tests/dreamx_world/test_dreamx_world_ar_conversion.py tests/local_tests/dreamx_world/test_dreamx_world_ar_transformer_parity.py -q -rs` | 6_passed_0_skipped | 2026-07-02: AR converter writes symlinked transformer layout/model_index; tiny official `CausalWanModel` vs FastVideo `DreamXWorldARTransformer3DModel` forward parity passes; real 5B `model.safetensors` strict-loads with zero missing/unexpected keys from `/tmp/converted_dreamx_world_ar`. |
| ar_pipeline_config | `PYTHONPATH=/workspace/FastVideo python -m pytest tests/local_tests/dreamx_world/test_dreamx_world_pipeline_config.py -q -rs` | 10_passed_0_skipped | 2026-07-02: `DreamXWorld5BARPipelineConfig`, `DreamXWorldARPipeline`, registry config selection for `GD-ML/DreamX-World-5B`, and `dreamx_world_5b_ar` preset pass. |
| ar_full_generation | `PYTHONPATH=/workspace/FastVideo python /tmp/run_dreamx_ar_video_smoke.py` | generated_video_pass | 2026-07-02: A40 short full-generation smoke passed from `/tmp/converted_dreamx_world_ar` with 64x64, 9 frames, 4 denoise steps, `output_type=pil`, `save_video=True`; MP4 saved at `outputs_video/dreamx_world_ar_video_smoke/a quiet road through a futuristic coastal city at sunrise.mp4` and decoded as 9 frames of `(64, 64, 3)` uint8. |
| ar_long_horizon | `PYTHONPATH=/workspace/FastVideo FASTVIDEO_ATTENTION_BACKEND=TORCH_SDPA python /tmp/run_dreamx_ar_long_horizon.py` | generated_video_pass | 2026-07-02: A40 long-horizon AR generation passed from `/tmp/converted_dreamx_world_ar` with 64x64, 1005 frames, 4 denoise steps, seed 4096; MP4 saved at `outputs_video/dreamx_world_ar_long_horizon/A long autonomous drive through a futuristic coastal city at sunrise, smooth forward camera motion,.mp4` and decoded as 1005 frames of `(64, 64, 3)`; end-to-end generation latency was 231.68s after load. |
| ar_ssim_default | `PYTHONPATH=/workspace/FastVideo FASTVIDEO_SSIM_MODEL_ID=DreamX-World-5B DREAMX_WORLD_AR_SSIM_MODEL_PATH=/tmp/converted_dreamx_world_ar python -m pytest fastvideo/tests/ssim/test_dreamx_world_similarity.py -q -rs --skip-ssim-reference-download` | 1_passed_0_skipped | 2026-07-02: A40 default AR SSIM reference seeded locally at `fastvideo/tests/ssim/reference_videos/default/A40_reference_videos/DreamX-World-5B/TORCH_SDPA/`; helper removes stale generated base MP4 before generation; default params are 192x192, 81 frames, 4 steps, seed 2048, min SSIM 0.98. |
| ar_ssim_modal_l40s | `modal run /tmp/modal_dreamx_ar_ssim_git.py` | 1_passed_0_skipped | 2026-07-02: Modal L40S run checked out `post-fix dreamx-world-5b-cam branch commit`, downloaded default `L40S_reference_videos` via `reference_videos_cli.py download`, used cached converted AR weights under `/root/data/dreamx_world_ar_converted`, seeded the missing AR L40S reference from generated output, reran a fresh generated-vs-reference compare successfully (`mean_ssim=1.0`), and exported the reference to Modal volume `hf-model-weights:dreamx_ar_ssim_l40s`. Downloaded local reference decodes as 81 frames of `(192, 192, 3)`. A40-vs-L40S reference spot check after deterministic AR noise fix: mean SSIM 0.7770, min 0.6139, max 0.9944. |
| pipeline_smoke | `python -m pytest tests/local_tests/pipelines/test_dreamx_world_pipeline_smoke.py tests/local_tests/pipelines/test_dreamx_world_pipeline_parity.py -q -rs` | 4_passed_0_skipped | 2026-06-30: combined smoke/parity passed. 2026-07-01: smoke alone passed with real `image_path` TI2V coverage (`3 passed`), validating image load, TI2V preprocessing, VAE first-frame encode under CPU offload, camera conditioning, and 1-step latent generation from `converted_weights/dreamx_world`. Tests force `FASTVIDEO_ATTENTION_BACKEND=TORCH_SDPA` to avoid the local FlashAttention-4 cute ABI mismatch. |
| basic_example | `PYTHONPATH=/workspace/FastVideo FASTVIDEO_ATTENTION_BACKEND=TORCH_SDPA DREAMX_WORLD_MODEL_DIR=converted_weights/dreamx_world DREAMX_WORLD_IMAGE_PATH= DREAMX_WORLD_HEIGHT=64 DREAMX_WORLD_WIDTH=64 DREAMX_WORLD_NUM_FRAMES=9 DREAMX_WORLD_STEPS=1 DREAMX_WORLD_GUIDANCE=1.0 DREAMX_WORLD_OUTPUT_PATH=outputs_video/dreamx_world_example_smoke python examples/inference/basic/basic_dreamx_world.py` | generated_video_pass | 2026-06-30: example saved an MP4 under `outputs_video/dreamx_world_example_smoke`; imageio/ffmpeg decoded frame 0 as `(64, 64, 3)` uint8, fps 16, duration 0.56s. |
## Open Questions
| ID | Question | Owner | Needed By Phase | Status | Resolution |
|---|---|---|---|---|---|
| Q001 | Should first PR expose only the `DreamX-World-5B-Cam` 5s camera-control mode and exclude AR long-horizon forcing? | user | prep | resolved | User approved starting with `DreamX-World-5B-Cam`; AR long-horizon is out of first-PR scope. |
| Q002 | Does FastVideo's existing Wan2.2 TI2V transformer support DreamX PRoPE/control adapter with a small extension, or is a DreamX-specific DiT required? | component:transformer | Phase 3 | resolved | Project guidance prefers a separate DreamX DiT for maintainability. DreamX PRoPE/control adapter now lives in `fastvideo/models/dits/dreamx_world.py`; Wan DiT/config have no DreamX-specific fields or `y_camera` signature. |
| Q003 | Which sampler is in first-PR scope: official default `Flow` only, or also `Flow_Unipc` and `Flow_DPM++`? | orchestrator | Phase 3 | resolved | First PR should support official default `Flow` only. FastVideo FlowMatch scheduler parity is non-skip PASS; optional `Flow_Unipc` and `Flow_DPM++` are out of first-PR scope. |
| Q004 | Which HF token env var should be used if rate limits or gated Wan2.2 base weights require auth? | user | Phase 5 | resolved | No auth was required for the completed local downloads; keep using env var names only if future gated repos require auth. |
| Q005 | Should native FastVideo production code depend on DreamX reference-only packages such as `xfuser` or OpenCV? | user | Phase 3 | resolved | No. These packages may be used only for official reference/local parity setup; native FastVideo integration must remove that runtime requirement. |
| Q006 | Should AR handoff require a full generated long-horizon video in this no-HF-token/no-GPU-budget pass? | user/runtime | pipeline | resolved | A40 long-horizon generation passes with 1005 frames at 64x64/4 steps. AR default SSIM passes locally on A40 and on Modal L40S with a 192x192/81-frame reference. HF upload/publication remains a separate operation if the reference dataset should be updated upstream. |
| Q007 | Can A40 references stand in for L40S CI references? | quality | release | resolved | No. After deterministic AR noise fix, A40-vs-L40S reference mean SSIM is 0.7770, below the 0.98 same-device threshold. Publish the L40S-specific reference for CI. |
## Issues And Blockers
| ID | Phase | Component | Severity | Issue | Evidence | Owner | Status | Resolution |
|---|---|---|---|---|---|---|---|---|
| I001 | prep | official_env | medium | Official import initially failed because `xfuser` was missing. | `ModuleNotFoundError: No module named 'xfuser'` from `python -c "import sys; sys.path.insert(0, 'DreamX-World'); import inference_dreamx5b"` | prep | resolved | Installed `xfuser==0.4.1`; import progressed. |
| I002 | prep | official_env | medium | Official import then failed because GUI OpenCV required missing system `libxcb.so.1`. | `ImportError: libxcb.so.1: cannot open shared object file` through `cv2` import in Diffusers ConsisID path. | prep | resolved | Installed `opencv-python-headless`; `import inference_dreamx5b` passed. |
| I003 | prep | weights | medium | HF repo has raw official transformer shards and no Diffusers `model_index.json`. | `inspect_hf_layout.py GD-ML/DreamX-World-5B-Cam --json` returned `source_layout=raw_official`, `needs_conversion=yes`, `model_index_class=null`. | conversion | resolved | Downloaded raw DreamX shards to `official_weights/dreamx_world`; converted transformer to `converted_weights/dreamx_world/transformer`; symlinked reusable Wan2.2 Diffusers components; real 5B transformer strict-load passes. |
| I004 | prep | dependencies | high | Official reference import required extra packages in the local environment, but FastVideo native runtime should not inherit those dependencies. | `xfuser==0.4.1` and `opencv-python-headless` were installed only to make `DreamX-World/inference_dreamx5b.py` import for reference/parity. | pipeline | resolved | Production DreamX FastVideo code uses native camera/image/video utilities and has no runtime `xfuser` or OpenCV import requirement; those packages remain reference-only local parity dependencies. |
| I005 | parity | transformer | medium | Transformer full forward parity initially failed in bf16 official harness. | Official CUDA bf16 LayerNorm path was unstable; fp32 small-input harness avoids that dtype issue and compares against FastVideo with single-process SP identity patches. | component:transformer | resolved | Official-vs-FastVideo forward parity now passes on CUDA with `diff_max=0.072533`, `diff_mean=0.008014`. |
| I006 | parity | vae/text_encoder | medium | VAE/text parity initially remained skipped after weights were staged. | Text official import needed a reference-only `xfuser` stub; VAE comparison initially used raw official normalized latents against FastVideo raw mu. | component:vae,component:text_encoder | resolved | Text hidden-state parity passes. VAE encode parity passes after raw key mapping and applying the official latent normalization to FastVideo output. |
| I007 | quality | pipeline_ti2v | medium | Real image-path TI2V smoke initially failed when `vae_cpu_offload=True` because DenoisingStage encoded the first frame while the VAE weights remained on CPU. | DreamX SSIM first run failed with `RuntimeError: Input type (torch.cuda.FloatTensor) and weight type (torch.FloatTensor) should be the same` at `fastvideo/pipelines/stages/denoising.py` VAE encode. | pipeline | resolved | DenoisingStage now moves the VAE to `local_device` before TI2V first-frame encode; image-path pipeline smoke and DreamX SSIM both pass. |
| I008 | quality | ssim_helper | high | SSIM helper could compare against a stale generated base MP4 when a rerun saved the new video as `_1.mp4`. | Existing generated outputs made AR reference seeding appear to pass before a fresh generated-vs-reference compare. | quality | resolved | `run_text_to_video_similarity_test` and `run_image_to_video_similarity_test` now remove the stale generated base MP4 before generation. A40 and Modal L40S AR SSIM were rerun after the fix. |
| I009 | quality | ar_denoising | high | AR denoising added CUDA noise without using the request seed when the original generator was CPU-backed. | Fresh reruns against old AR references produced mean SSIM near 0.05. | pipeline | resolved | `DreamXWorldARCausalDenoisingStage` now derives a device-local generator from the request seed for AR noise. Fresh same-device A40 and L40S reruns pass with mean SSIM 1.0 after reseeding references. |
## Escape Hatches
| ID | Phase | Decision Type | Question | Recommended Option | Status | Resolution |
|---|---|---|---|---|---|---|
## Decisions
| Date | Decision | Rationale | Impact |
|---|---|---|---|
| 2026-06-29 | First PR scope is `DreamX-World-5B-Cam` only. | Cam mode is closest to existing Wan2.2 TI2V support; AR forcing needs separate causal/KV pipeline work. | Component inventory and parity focus on `inference_dreamx5b.py` and `pipeline_dreamxworld.py`. |
| 2026-06-29 | Do not install full DreamX requirements during prep. | Full requirements pin core FastVideo stack packages. | Installed only `xfuser==0.4.1` and `opencv-python-headless` to make official imports work. |
| 2026-06-29 | Treat HF DreamX-World-5B-Cam weights as raw official transformer layout requiring conversion. | HF inspection found no `model_index.json`. | Phase 5 must create `scripts/checkpoint_conversion/dreamx_world_to_diffusers.py` after component prototype/key dumps. |
| 2026-06-29 | Do not add DreamX reference-only dependencies to FastVideo production requirements. | The current environment should remain the FastVideo environment; extra packages are only for official reference parity. | Native DreamX integration must avoid runtime `xfuser` and OpenCV requirements unless explicitly approved later. |
| 2026-06-29 | Implement DreamX camera conditioning as native FastVideo utility. | It is weightless and removes the need to import DreamX reference utilities at production runtime. | `fastvideo/pipelines/basic/dreamx_world/camera_conditioning.py` now has non-skip parity against official action-to-PRoPE tensors. |
| 2026-06-29 | First PR supports DreamX default `Flow` sampler only. | FastVideo FlowMatch Euler scheduler matches the official Diffusers scheduler for DreamX defaults. | Pipeline work can use FastVideo native FlowMatch scheduler; UniPC and DPM++ are out of first-PR scope. |
| 2026-07-01 | Keep DreamX PRoPE/control adapter in a dedicated DreamX DiT class. | Project guidance is that putting too much DreamX behavior into Wan makes the model hard to manage. | `fastvideo/models/dits/dreamx_world.py` defines `DreamXWorldTransformer3DModel`, `DreamXWorldTransformerBlock`, and `DreamXPropeSelfAttention`; `fastvideo/configs/models/dits/dreamx_world.py` owns DreamX adapter config fields; Wan DiT/config are unchanged from DreamX. |
| 2026-06-30 | Camera parity test loads official camera functions by file instead of importing the official `utils` package. | Official package initialization pulls unrelated dependencies that can require GUI OpenCV system libraries. | Camera parity remains non-skip without adding DreamX reference-only dependencies to FastVideo production requirements. |
| 2026-06-30 | Add DreamX-World-5B-Cam model and pipeline config helpers plus a conversion script. | Official HF DreamX 5B-Cam transformer config is 30 layers, hidden size 3072, 24 heads, 48 latent channels, plus Wan2.2 48-channel VAE and UMT5-XXL text encoder. | DreamX helpers wire DiT/VAE/UMT5/Flow/TI2V settings; `dreamx_world_to_diffusers.py` writes a FastVideo-loadable transformer config plus renamed safetensors; strict-load smoke passes on a tiny official DreamX transformer and the real 5B converted shards. |
| 2026-06-30 | Pass DreamX camera PRoPE condition through the FastVideo batch/denoising path. | DreamX transformer expects `y_camera={"viewmats", "K"}` at denoising time. | `DreamXWorldPipeline` is registered as a basic pipeline entry and initializes the official default FlowMatch scheduler; `dreamx_world_5b_cam` preset mirrors official 5B-Cam defaults; `DreamXWorldCameraConditioningStage` writes `batch.extra["dreamx_y_camera"]`; generic denoising filters and forwards it as `y_camera` only for compatible transformers. |
## Handoff Notes
- Prep, component parity, pipeline smoke/parity, and the basic example validation are complete for `DreamX-World-5B-Cam`.
- Official reference clone is staged at `DreamX-World/` and ignored by git.
- Workspace-local weights are staged: DreamX raw transformer shards under `official_weights/dreamx_world`, Wan2.2 raw base artifacts under `official_weights/Wan2.2-TI2V-5B`, and Wan2.2 Diffusers reusable components under `official_weights/Wan2.2-TI2V-5B-Diffusers`.
- Camera conditioning parity is active and passing without weights.
- Default Flow scheduler parity is active and passing without weights.
- Transformer has corrected official 5B-Cam architecture in dedicated DreamX DiT/config files, native PRoPE/control-adapter, conversion mapping, real converted 5B strict-load, and official-vs-FastVideo forward parity passing on CUDA. VAE encode parity and text hidden-state parity pass on CUDA. Pipeline entry/registry, local model_info resolution, preset, config, FlowMatch scheduler init, camera stage, denoising `y_camera` kwarg smokes, independent CUDA pipeline smoke/parity, and a small saved-video basic example pass.
- Full local DreamX component suite is non-skip PASS: `python -m pytest tests/local_tests/dreamx_world/ -q -rs` returned `26 passed` on 2026-07-01. Pipeline smoke/parity and SSIM quality regression are also non-skip PASS locally.
- Keep `xfuser` and OpenCV as reference-only parity dependencies. Do not add them
to FastVideo requirements or production imports.
## Quality Regression
- status: `added`
- test: `fastvideo/tests/ssim/test_dreamx_world_similarity.py`
- command: `PYTHONPATH=/workspace/FastVideo FASTVIDEO_SSIM_MODEL_ID=DreamX-World-5B-Cam python -m pytest fastvideo/tests/ssim/test_dreamx_world_similarity.py -q -rs`
- result: `1 passed, 0 skipped` on 2026-07-01
- reference: Local A40/TORCH_SDPA reference seeded from the generated candidate under `fastvideo/tests/ssim/reference_videos/default/A40_reference_videos/DreamX-World-5B-Cam/TORCH_SDPA/`. The test uses a deterministic generated input image, 64x64 request dimensions, 9 frames, 1 denoise step, seed 1024, and min SSIM 0.98. Full-quality params are present for 480x832/161 frames/30 steps.
- note: Modal L40S seeding passed using the configured Modal profile and unauthenticated HF public downloads. HF upload/publication still requires `HF_API_KEY`, `HUGGINGFACE_HUB_TOKEN`, or `HF_TOKEN` with write access; no token values were used or recorded.
## Final Handoff
```text
final_handoff:
prep_handoff_complete: yes
conversion_status: pass
components:
- name: transformer
reuse_or_port: ported_dedicated_dit
parity_test: tests/local_tests/dreamx_world/test_dreamx_world_transformer_parity.py
parity_status: non_skip_pass
concerns_or_unknowns: none
- name: vae
reuse_or_port: reused
parity_test: tests/local_tests/dreamx_world/test_dreamx_world_vae_parity.py
parity_status: non_skip_pass
concerns_or_unknowns: none
- name: text_encoder_tokenizer
reuse_or_port: reused
parity_test: tests/local_tests/dreamx_world/test_dreamx_world_text_encoder_parity.py
parity_status: non_skip_pass
concerns_or_unknowns: none
- name: scheduler
reuse_or_port: reused
parity_test: tests/local_tests/dreamx_world/test_dreamx_world_scheduler_parity.py
parity_status: non_skip_pass
concerns_or_unknowns: none
- name: camera_conditioning
reuse_or_port: ported
parity_test: tests/local_tests/dreamx_world/test_dreamx_world_camera_conditioning_parity.py
parity_status: non_skip_pass
concerns_or_unknowns: none
- name: ar_transformer
reuse_or_port: ported
parity_test: tests/local_tests/dreamx_world/test_dreamx_world_ar_transformer_parity.py
parity_status: non_skip_pass
concerns_or_unknowns: none
- name: ar_pipeline
reuse_or_port: ported
parity_test: tests/local_tests/dreamx_world/test_dreamx_world_pipeline_config.py plus AR smoke/SSIM commands listed above
parity_status: non_skip_pass
concerns_or_unknowns: none
pipeline_smoke: pass
pipeline_parity: pass
example_status: pass
quality_regression: added
local_tests_readme: tests/local_tests/dreamx_world/README.md
port_state_file: tests/local_tests/dreamx_world/PORT_STATUS.md
token_values_committed: no
runtime_third_party_model_imports: none
blockers: none
escape_hatch: none
```
| 2026-07-02 | Add DreamX-World-5B autoregressive support. | Official AR repo is raw single-safetensors layout and needs a dedicated causal/KV stage. | Added native AR DiT/config, identity converter, AR pipeline config/preset/registry, and targeted non-skip tests. Raw AR weights are staged at `/tmp/dreamx_world_ar_weights`; converted symlink layout at `/tmp/converted_dreamx_world_ar`. |
| 2026-07-02 | Validate DreamX-World-5B AR short full generation on A40. | Targeted parity/config tests prove components, but end-to-end runtime can still fail at scheduler, RoPE cache, device, or decode boundaries. | `DreamXWorldARPipeline` generated and saved a 64x64/9-frame/4-step MP4 from `/tmp/converted_dreamx_world_ar`; the saved MP4 decodes to 9 frames. |
| 2026-07-02 | Validate DreamX-World-5B AR long-horizon and default SSIM on A40. | AR needs coverage beyond the 9-frame smoke to exercise longer KV/cache progression and a quality regression path. | 1005-frame 64x64/4-step generation passes and decodes; default AR SSIM test uses 192x192/81 frames because MS-SSIM requires short side >160. |
| 2026-07-02 | Validate DreamX-World-5B AR default SSIM on Modal L40S. | CI references are device-specific; A40 alone is not enough for L40S reference coverage. | Modal checked out `post-fix dreamx-world-5b-cam branch commit`, downloaded existing default L40S references, seeded the missing AR reference, reran SSIM successfully (`mean_ssim=1.0`), and exported the L40S reference back to the workspace. |
+254
View File
@@ -0,0 +1,254 @@
# DreamX World Local Tests
Local-only parity and smoke tests for the `dreamx_world` FastVideo port. These
tests compare FastVideo against the official DreamX-World reference
implementation and are not expected to run in CI unless explicitly promoted
later.
Port progress, open questions, issues, and handoff notes live in
`tests/local_tests/dreamx_world/PORT_STATUS.md`.
## Reference Assets
| Field | Value |
|---|---|
| Model family | `dreamx_world` |
| First-PR scope | `DreamX-World-5B-Cam`; follow-up scope now includes `DreamX-World-5B` autoregressive forcing |
| Out-of-scope variants | none for the DreamX-World 5B/Cam paths currently ported |
| Workload types | I2V camera-control compatibility shim: image + prompt + action sequence to video |
| Official reference | `https://github.com/AMAP-ML/DreamX-World` |
| Local reference dir | `DreamX-World/` |
| Official commit/version | `221875811ba31f7eac6c3025b215c09ad2cefd1d` |
| HF weights | `GD-ML/DreamX-World-5B-Cam` |
| HF revision | default |
| Local weights dir | `official_weights/dreamx_world` |
| Source layout | `raw_official` |
| Needs conversion | `yes` |
Do not write token values in this file. Current token env var detected during
prep: `none`.
## Shared Environment Setup
Run from the FastVideo repo root in the same conda/env used for FastVideo. Do
not create a separate upstream environment for parity tests.
```bash
python ".agents/skills/add-model-01-prep/scripts/clone_reference_repo.py" \
"https://github.com/AMAP-ML/DreamX-World.git" \
"DreamX-World" \
--commit "221875811ba31f7eac6c3025b215c09ad2cefd1d" \
--update-gitignore
```
DreamX-World does not expose a packaging file for editable install. During prep
the official import check used `sys.path.insert(0, "DreamX-World")`.
Additional official deps installed into the current environment for imports:
```bash
uv pip install xfuser==0.4.1
uv pip install opencv-python-headless
```
These packages are for running the official DreamX reference during local
parity only. They must not become FastVideo production/runtime dependencies for
the native `dreamx_world` pipeline.
Do not install the full `DreamX-World/requirements.txt` without explicit
approval. It pins core FastVideo stack packages including `torch`, `torchvision`,
`triton`, `flash_attn`, and `diffusers`.
## Official Environment Status
```text
dependency_changes: installed official deps in current env
official_env_status: imports_ok
private_dep_stubs: none
blocked_on: none
```
Import check used during prep:
```bash
python -c "import sys; sys.path.insert(0, 'DreamX-World'); import inference_dreamx5b; print('imports_ok')"
```
## Weight Setup
HF layout inspection found no root `model_index.json`; the repo contains
`config.json`, a safetensors index, and three transformer safetensors shards.
This is a raw official transformer layout and requires conversion before
FastVideo can load it through `VideoGenerator.from_pretrained`.
```bash
python ".agents/skills/add-model-01-prep/scripts/inspect_hf_layout.py" \
"GD-ML/DreamX-World-5B-Cam" \
--json
```
Weights have been staged workspace-locally. The raw DreamX transformer repo lives at `official_weights/dreamx_world`; Wan2.2 raw base artifacts live at `official_weights/Wan2.2-TI2V-5B`; Wan2.2 Diffusers reusable components live at `official_weights/Wan2.2-TI2V-5B-Diffusers`. To reproduce the DreamX download:
```bash
python ".agents/skills/add-model-01-prep/scripts/download_hf_weights.py" \
"GD-ML/DreamX-World-5B-Cam" \
"official_weights/dreamx_world"
```
### DreamX-World-5B Autoregressive Setup
The AR repository `GD-ML/DreamX-World-5B` is also raw official layout: no
`model_index.json`, root `config.json`, and a single `model.safetensors`. The
current environment has the raw AR checkpoint staged outside the workspace at
`/tmp/dreamx_world_ar_weights` to avoid workspace quota pressure. The converted
FastVideo layout is staged at `/tmp/converted_dreamx_world_ar` with the 21GB
transformer safetensors symlinked instead of copied.
```bash
python ".agents/skills/add-model-01-prep/scripts/download_hf_weights.py" \
"GD-ML/DreamX-World-5B" \
"/tmp/dreamx_world_ar_weights"
python scripts/checkpoint_conversion/dreamx_world_ar_to_diffusers.py \
--source /tmp/dreamx_world_ar_weights \
--output /tmp/converted_dreamx_world_ar \
--component-source official_weights/Wan2.2-TI2V-5B-Diffusers \
--symlink-components \
--symlink-transformer
```
AR production code uses `DreamXWorldARTransformer3DModel`,
`DreamXWorld5BARPipelineConfig`, `DreamXWorldARPipeline`, and
`DreamXWorldARCausalDenoisingStage`. The AR DiT is a native FastVideo port of
the official Apache-2.0 `CausalWanModel`; it has no production DreamX, Diffusers
model-class, Transformers model-class, `xfuser`, or OpenCV import.
## Prototype And Conversion Artifacts
State-dict key/shape dumps are generated after FastVideo native prototypes exist
and are used to build the conversion mapping.
```text
official_key_dumps:
transformer: converted_weights/dreamx_world/_mapping/transformer_official_keys.json
fastvideo_key_dumps:
transformer: converted_weights/dreamx_world/_mapping/transformer_fastvideo_keys.json
conversion_script: scripts/checkpoint_conversion/dreamx_world_to_diffusers.py
conversion_script_status: transformer_model_index_and_config_consistency_smoke_pass
conversion_source_layout: raw_official
converted_weights_dir: converted_weights/dreamx_world
model_index_status: smoke_pass
strict_load_status: pass
```
The converter writes `transformer/` from raw DreamX shards. To create a full
FastVideo-loadable diffusers-style root, pass a Wan2.2 Diffusers directory as the
component source so reusable components are copied or symlinked before
`model_index.json` is emitted:
```bash
python scripts/checkpoint_conversion/dreamx_world_to_diffusers.py \
--source official_weights/dreamx_world \
--output converted_weights/dreamx_world \
--component-source /path/to/Wan2.2-TI2V-5B-Diffusers \
--symlink-components
```
## Expected Parity Tests
Planned local tests for this family:
| Component | Official files / args | Test | Concerns | Status |
|---|---|---|---|---|
| transformer | `DreamX-World/models/wan_transformer3d.py`; instantiated in `DreamX-World/inference_dreamx5b.py` with `cam_method=prope`, `add_control_adapter=True` | `tests/local_tests/dreamx_world/test_dreamx_world_transformer_parity.py` | FastVideo uses dedicated `DreamXWorldTransformer3DModel`/`DreamXWorldConfig` files for DreamX 5B-Cam config, PRoPE, conversion mapping, real converted 5B strict-load PASS, and official-vs-FastVideo small-input forward parity PASS on CUDA. Wan DiT/config have no DreamX-specific adapter fields. | strict_load_and_forward_parity_pass |
| vae | `DreamX-World/models/wan_vae3_8.py`; `vae_type=AutoencoderKLWan3_8`, `vae_subpath=Wan2.2_VAE.pth` | `tests/local_tests/dreamx_world/test_dreamx_world_vae_parity.py` | DreamX raw `Wan2.2_VAE.pth` maps 196/196 keys into FastVideo Wan VAE; encode parity passes after applying official latent normalization. | encode_parity_pass |
| text_encoder/tokenizer | `DreamX-World/models/wan_text_encoder.py`; T5 path from Wan2.2 base model | `tests/local_tests/dreamx_world/test_dreamx_world_text_encoder_parity.py` | Official `WanT5EncoderModel` vs FastVideo `UMT5EncoderModel` hidden-state parity passes on CUDA with staged Wan2.2 weights/tokenizer. | hidden_state_parity_pass |
| scheduler | Diffusers `FlowMatchEulerDiscreteScheduler`; selected by default `sampler_name=Flow` | `tests/local_tests/dreamx_world/test_dreamx_world_scheduler_parity.py` | First PR can support official default `Flow`; optional `Flow_Unipc` and `Flow_DPM++` are out of first-PR scope. | non_skip_pass |
| camera_conditioning | `DreamX-World/utils/inference_utils.py`, `models/prope_utils.py`, `wan/modules/camera_prope.py` | `tests/local_tests/dreamx_world/test_dreamx_world_camera_conditioning_parity.py` | Action sequence to PRoPE/control input must match official tensor shapes and values. | non_skip_pass |
| pipeline | `DreamX-World/pipeline/pipeline_dreamxworld.py`; call path in `DreamX-World/inference_dreamx5b.py` | `tests/local_tests/dreamx_world/test_dreamx_world_pipeline_config.py`; `tests/local_tests/pipelines/test_dreamx_world_pipeline_smoke.py`; `tests/local_tests/pipelines/test_dreamx_world_pipeline_parity.py` | DreamX PipelineConfig wires first-scope DiT/VAE/UMT5/Flow/TI2V settings, official `shift=3.0`, default preset values, FlowMatch scheduler initialization, and local `model_index.json` resolution; independent pipeline smoke covers real CUDA local load + latent generation, and parity compares public API output to worker-side explicit ForwardBatch execution. | pipeline_smoke_and_parity_pass |
| ar_transformer | `DreamX-World/wan/modules/causal_camera_model_2_2_prope_infinity.py`; instantiated by `DreamX-World/inference_ar_forcing.py` as `CausalWanModel` with `local_attn_size=12`, `sink_size=3`, `attn_compress=4` | `tests/local_tests/dreamx_world/test_dreamx_world_ar_transformer_parity.py` | Native `DreamXWorldARTransformer3DModel` keeps official identity key layout; tiny official-vs-FastVideo forward parity passes; real 5B AR safetensors strict-load passes from `/tmp/converted_dreamx_world_ar`. | tiny_forward_parity_and_real_strict_load_pass |
| ar_pipeline | `DreamX-World/pipeline/pipeline_causal_camera.py`; AR block/KV/context-noise loop | `tests/local_tests/dreamx_world/test_dreamx_world_pipeline_config.py`; A40 short full-generation smoke | Dedicated `DreamXWorldARCausalDenoisingStage` implements blockwise KV forcing; raw HF repo must be converted before `VideoGenerator.from_pretrained` because it has no `model_index.json`; converted AR layout generated a 64x64/9-frame/4-step MP4 on A40. | config_registry_and_short_full_generation_pass |
Include reused components in parity. Reuse is accepted only after the FastVideo
component definition and official instantiation arguments have both been checked
and the component parity test passes non-skip.
Run the relevant tests with:
```bash
python -m pytest tests/local_tests/dreamx_world/test_dreamx_world_camera_conditioning_parity.py -v -s
python -m pytest tests/local_tests/dreamx_world/test_dreamx_world_scheduler_parity.py -v -s
python -m pytest tests/local_tests/dreamx_world/test_dreamx_world_transformer_parity.py -v -s
python -m pytest tests/local_tests/dreamx_world/test_dreamx_world_vae_parity.py -v -s
python -m pytest tests/local_tests/dreamx_world/test_dreamx_world_text_encoder_parity.py -v -s
python -m pytest tests/local_tests/dreamx_world/test_dreamx_world_pipeline_config.py -v -s
python -m pytest tests/local_tests/dreamx_world/test_dreamx_world_conversion.py -v -s
python -m pytest tests/local_tests/pipelines/test_dreamx_world_pipeline_smoke.py tests/local_tests/pipelines/test_dreamx_world_pipeline_parity.py -q -rs
```
## Current Local Results
```bash
python -m pytest tests/local_tests/dreamx_world/ -v -s
# 2026-07-01: 26 passed, 0 skipped
# 2026-07-02: AR targeted suite passed: 16 passed, 0 skipped
PYTHONPATH=/workspace/FastVideo python /tmp/run_dreamx_ar_video_smoke.py
# 2026-07-02: DreamX-World-5B AR short full-generation smoke passed on A40
# output: outputs_video/dreamx_world_ar_video_smoke/a quiet road through a futuristic coastal city at sunrise.mp4
# decoded: 9 frames, (64, 64, 3), uint8
PYTHONPATH=/workspace/FastVideo FASTVIDEO_ATTENTION_BACKEND=TORCH_SDPA python /tmp/run_dreamx_ar_long_horizon.py
# 2026-07-02: DreamX-World-5B AR long-horizon generation passed on A40
# output: outputs_video/dreamx_world_ar_long_horizon/A long autonomous drive through a futuristic coastal city at sunrise, smooth forward camera motion,.mp4
# decoded: 1005 frames, (64, 64, 3)
PYTHONPATH=/workspace/FastVideo FASTVIDEO_SSIM_MODEL_ID=DreamX-World-5B DREAMX_WORLD_AR_SSIM_MODEL_PATH=/tmp/converted_dreamx_world_ar python -m pytest fastvideo/tests/ssim/test_dreamx_world_similarity.py -q -rs --skip-ssim-reference-download
# 2026-07-02: DreamX-World-5B AR default SSIM passed: 1 passed, 0 skipped
# local A40 reference: fastvideo/tests/ssim/reference_videos/default/A40_reference_videos/DreamX-World-5B/TORCH_SDPA/
modal run /tmp/modal_dreamx_ar_ssim_git.py
# 2026-07-02: Modal L40S default SSIM passed: 1 passed, 0 skipped
# checked out post-fix dreamx-world-5b-cam branch commit
# first downloaded default L40S references with reference_videos_cli.py download
# seeded missing AR L40S reference, then reran a fresh generated-vs-reference compare
# Modal JSON: mean_ssim=1.0, min_ssim=1.0, max_ssim=1.0
# local L40S reference: fastvideo/tests/ssim/reference_videos/default/L40S_reference_videos/DreamX-World-5B/TORCH_SDPA/
# A40-vs-L40S reference spot check after deterministic AR noise fix: mean SSIM 0.7770, min 0.6139, max 0.9944
python -m pytest tests/local_tests/pipelines/test_dreamx_world_pipeline_smoke.py tests/local_tests/pipelines/test_dreamx_world_pipeline_parity.py -q -rs
# 2026-06-30: 4 passed, 0 skipped
# 2026-07-01: smoke image-path TI2V coverage passed separately with 3 passed, 0 skipped
PYTHONPATH=/workspace/FastVideo FASTVIDEO_SSIM_MODEL_ID=DreamX-World-5B-Cam python -m pytest fastvideo/tests/ssim/test_dreamx_world_similarity.py -q -rs
# 2026-07-01: 1 passed, 0 skipped
```
`camera_conditioning`, default `Flow` scheduler, DreamX component/pipeline configs, default preset, pipeline entry/registry, FlowMatch scheduler initialization, DreamX camera stage, denoising `y_camera` pass-through, conversion model-index/config-consistency checks, real converted 5B dedicated DreamX transformer strict-load and forward parity on CUDA, VAE encode parity on CUDA, text hidden-state parity on CUDA, native DreamX PRoPE branch structure smoke, AR transformer parity/strict-load, AR config/registry, and AR short full-generation smoke are non-skip PASS results. The full local DreamX component suite, independent pipeline smoke/parity suite, image-path TI2V smoke, and SSIM quality regression currently have zero skips. The basic 5B-Cam example was run against `converted_weights/dreamx_world` with a 64x64/9-frame/1-step saved-video smoke, and imageio decoded the generated MP4 successfully; the AR pipeline separately saved and decoded a 64x64/9-frame/4-step MP4 from `/tmp/converted_dreamx_world_ar`.
## Review Notes
- Required before handoff: non-skip PASS for each required component parity
test, including reused components that own weights or numerical behavior.
- First PR scope originally targeted `DreamX-World-5B-Cam`; scope was later
expanded to include `DreamX-World-5B` AR support with a separate causal/KV
pipeline.
- AR support has targeted parity/config coverage, short full-generation smoke,
1005-frame A40 long-horizon generation, local A40 default SSIM coverage,
and Modal L40S default SSIM coverage. L40S validation was rerun after fixing
stale generated-output comparison in the SSIM helper and deterministic AR
noise seeding. HF reference publication remains a separate token-gated
release operation.
- FastVideo production code must not require `xfuser` or OpenCV just because the
official reference import needed them. Port camera/action preprocessing and
sequence-parallel behavior into existing FastVideo-native utilities or keep
reference-only imports inside local parity tests.
- Review agents should verify setup commands still match the PR, then run the
listed parity tests or report the exact blocker.
## Quality Regression
Quality regression is added in `fastvideo/tests/ssim/test_dreamx_world_similarity.py`. The 5B-Cam default test uses the workspace converted model root, `TORCH_SDPA`, a deterministic generated conditioning image, 9 frames, 1 denoise step, seed 1024, and min SSIM 0.98. A local A40 reference was seeded under `fastvideo/tests/ssim/reference_videos/default/A40_reference_videos/DreamX-World-5B-Cam/TORCH_SDPA/`, and the test passed non-skip locally on 2026-07-01. The AR default test uses `/tmp/converted_dreamx_world_ar` or `/root/data/dreamx_world_ar_converted`, `TORCH_SDPA`, 192x192, 81 frames, 4 steps, seed 2048, and min SSIM 0.98; local A40 and Modal L40S references were seeded under `fastvideo/tests/ssim/reference_videos/default/{A40,L40S}_reference_videos/DreamX-World-5B/TORCH_SDPA/`, and both tests passed non-skip on 2026-07-02 after fresh generated-output cleanup was added to the SSIM helper. Modal L40S validation checked out `post-fix dreamx-world-5b-cam branch commit`, seeded the missing AR L40S reference, reran the test, and wrote `mean_ssim=1.0`. A cross-device A40-vs-L40S reference spot check produced mean SSIM 0.7770, so CI should use the L40S-specific reference rather than the A40 artifact. Full-quality params are present for 5B-Cam 480x832/161 frames/30 steps and AR 704x1280/1005 frames/4 steps; publishing references to the HF dataset remains a release operation requiring a write-capable HF token env var, never a raw token value.
@@ -0,0 +1 @@
# SPDX-License-Identifier: Apache-2.0
@@ -0,0 +1,58 @@
# SPDX-License-Identifier: Apache-2.0
"""DreamX-World-5B autoregressive conversion smoke tests."""
import json
from fastvideo.configs.pipelines.dreamx_world import make_dreamx_world_5b_ar_dit_config
from fastvideo.models.registry import _LEGACY_FAST_VIDEO_MODELS
from scripts.checkpoint_conversion.dreamx_world_ar_to_diffusers import (
MODEL_INDEX,
REUSED_COMPONENTS,
TRANSFORMER_CONFIG,
convert_transformer,
write_model_index,
)
def test_dreamx_world_ar_converter_writes_symlinked_transformer_and_model_index(tmp_path):
source = tmp_path / "raw"
source.mkdir()
raw_tensor = source / "model.safetensors"
raw_tensor.write_bytes(b"placeholder")
component_source = tmp_path / "wan22"
output = tmp_path / "dreamx_ar"
for component in REUSED_COMPONENTS:
component_dir = component_source / component
component_dir.mkdir(parents=True)
(component_dir / "config.json").write_text("{}\n")
convert_transformer(source, output, symlink_transformer=True)
write_model_index(output, component_source, symlink_components=True)
assert (output / "transformer" / "model.safetensors").is_symlink()
model_index = json.loads((output / "model_index.json").read_text())
assert model_index == MODEL_INDEX
assert model_index["_class_name"] == "DreamXWorldARPipeline"
assert model_index["transformer"] == ["diffusers", "DreamXWorldARTransformer3DModel"]
def test_dreamx_world_ar_transformer_config_matches_pipeline_dit_config():
dit_config = make_dreamx_world_5b_ar_dit_config()
assert TRANSFORMER_CONFIG["_class_name"] == "DreamXWorldARTransformer3DModel"
assert TRANSFORMER_CONFIG["num_attention_heads"] == dit_config.num_attention_heads
assert TRANSFORMER_CONFIG["attention_head_dim"] == dit_config.attention_head_dim
assert TRANSFORMER_CONFIG["in_channels"] == dit_config.in_channels
assert TRANSFORMER_CONFIG["out_channels"] == dit_config.out_channels
assert TRANSFORMER_CONFIG["ffn_dim"] == dit_config.ffn_dim
assert TRANSFORMER_CONFIG["num_layers"] == dit_config.num_layers
assert TRANSFORMER_CONFIG["local_attn_size"] == dit_config.local_attn_size
assert TRANSFORMER_CONFIG["sink_size"] == dit_config.sink_size
assert TRANSFORMER_CONFIG["attn_compress"] == dit_config.attn_compress
assert tuple(TRANSFORMER_CONFIG["cam_self_attn_layers"]) == dit_config.cam_self_attn_layers
def test_dreamx_world_ar_model_index_component_classes_are_registered():
for component in ("scheduler", "text_encoder", "transformer", "vae"):
class_name = MODEL_INDEX[component][1]
assert class_name in _LEGACY_FAST_VIDEO_MODELS
@@ -0,0 +1,182 @@
# SPDX-License-Identifier: Apache-2.0
"""DreamX-World-5B autoregressive transformer parity.
Coverage scope: both. The tiny forward parity compares FastVideo's native AR DiT
against the official DreamX ``CausalWanModel`` implementation with identical
weights. The real 5B checkpoint gate strict-loads the downloaded safetensors
through the FastVideo model class.
"""
from __future__ import annotations
import os
from pathlib import Path
import sys
import pytest
import torch
from torch.testing import assert_close
from fastvideo.configs.models.dits.dreamx_world import DreamXWorldARArchConfig, DreamXWorldARConfig
from fastvideo.configs.pipelines.dreamx_world import make_dreamx_world_5b_ar_dit_config
from fastvideo.models.dits.dreamx_world_ar import DreamXWorldARTransformer3DModel
from fastvideo.models.loader.fsdp_load import load_model_from_full_model_state_dict
from fastvideo.models.loader.utils import get_param_names_mapping
from fastvideo.models.loader.weight_utils import resolve_safetensors_files, safetensors_weights_iterator
REPO_ROOT = Path(__file__).resolve().parents[3]
OFFICIAL_REF_DIR = Path(os.getenv("DREAMX_WORLD_OFFICIAL_REF_DIR", REPO_ROOT / "DreamX-World"))
CONVERTED_AR_DIR = Path(os.getenv("DREAMX_WORLD_AR_CONVERTED_DIR", "/tmp/converted_dreamx_world_ar"))
CONVERTED_AR_HF_REPO = "FastVideo/DreamX-World-5B-Diffusers"
PARITY_SCOPE = "both"
def _tiny_config() -> DreamXWorldARConfig:
return DreamXWorldARConfig(
arch_config=DreamXWorldARArchConfig(
num_attention_heads=1,
attention_head_dim=8,
in_channels=4,
out_channels=4,
ffn_dim=16,
num_layers=1,
text_dim=8,
freq_dim=8,
text_len=4,
local_attn_size=2,
sink_size=1,
attn_compress=1,
cam_self_attn_layers=(0,),
))
def _load_official_tiny():
if not OFFICIAL_REF_DIR.exists():
pytest.skip(f"Official DreamX reference missing: {OFFICIAL_REF_DIR}")
sys.path.insert(0, str(OFFICIAL_REF_DIR))
try:
from wan.modules import attention as official_attention
from wan.modules import causal_camera_model_2_2_prope_infinity as causal_module
from wan.modules import model_2_2 as official_model_2_2
from wan.modules.causal_camera_model_2_2_prope_infinity import CausalWanModel
except Exception as exc: # noqa: BLE001
pytest.skip(f"Cannot import official AR transformer: {exc}")
official_attention.FLASH_ATTN_2_AVAILABLE = False
official_attention.FLASH_ATTN_3_AVAILABLE = False
def _sdpa_same_dtype(q, k, v, **kwargs):
del kwargs
out = torch.nn.functional.scaled_dot_product_attention(
q.transpose(1, 2), k.transpose(1, 2), v.transpose(1, 2), dropout_p=0.0)
return out.transpose(1, 2).contiguous()
official_attention.attention = _sdpa_same_dtype
official_attention.flash_attention = _sdpa_same_dtype
official_model_2_2.flash_attention = _sdpa_same_dtype
causal_module.attention = _sdpa_same_dtype
return CausalWanModel(
model_type="ti2v",
patch_size=(1, 2, 2),
text_len=4,
in_dim=4,
dim=8,
ffn_dim=16,
freq_dim=8,
text_dim=8,
out_dim=4,
num_heads=1,
num_layers=1,
local_attn_size=2,
sink_size=1,
qk_norm=True,
cross_attn_norm=True,
add_control_adapter=True,
cam_method="prope",
attn_compress=1,
cam_self_attn_layers=(0,),
).eval()
def _make_inputs():
torch.manual_seed(123)
x = [torch.randn(4, 1, 4, 4)]
t = torch.zeros(1, 4, dtype=torch.long)
context = [torch.randn(2, 8)]
camera = {
"viewmats": torch.eye(4).reshape(1, 1, 4, 4).repeat(1, 4, 1, 1),
"K": torch.eye(3).reshape(1, 1, 3, 3).repeat(1, 4, 1, 1),
}
kv_cache = [{
"k": torch.zeros(1, 8, 1, 8),
"v": torch.zeros(1, 8, 1, 8),
"global_end_index": torch.tensor([0]),
"local_end_index": torch.tensor([0]),
"prope_k": torch.zeros(1, 8, 1, 8),
"prope_v": torch.zeros(1, 8, 1, 8),
"prope_global_end_index": torch.tensor([0]),
"prope_local_end_index": torch.tensor([0]),
}]
cross_cache = [{
"k": torch.zeros(1, 4, 1, 8),
"v": torch.zeros(1, 4, 1, 8),
"is_init": False,
}]
return x, t, context, camera, kv_cache, cross_cache
def test_dreamx_world_ar_tiny_forward_matches_official():
official = _load_official_tiny()
# The official init_weights zero-inits the output head (head.head.weight and
# biases), so both models would output exactly zero and the comparison would
# pass vacuously. Randomize the head (deterministically) before copying the
# state dict so the outputs reflect the internal computation.
generator = torch.Generator().manual_seed(7)
with torch.no_grad():
official.head.head.weight.normal_(std=0.5, generator=generator)
official.head.head.bias.normal_(std=0.5, generator=generator)
fastvideo = DreamXWorldARTransformer3DModel(_tiny_config(), {}).eval()
fastvideo.load_state_dict(official.state_dict(), strict=True)
x, t, context, camera, kv_cache, cross_cache = _make_inputs()
official_out = official(x=x, t=t, context=context, seq_len=4, y_camera=camera,
kv_cache=kv_cache, crossattn_cache=cross_cache).detach()
x, t, context, camera, kv_cache, cross_cache = _make_inputs()
fastvideo_out = fastvideo(x=x, t=t, context=context, seq_len=4, y_camera=camera,
kv_cache=kv_cache, crossattn_cache=cross_cache).detach()
assert official_out.abs().max() > 0, "official output is all-zero; parity comparison is vacuous"
assert_close(fastvideo_out, official_out, atol=1e-5, rtol=1e-5)
def test_dreamx_world_ar_5b_config_matches_official_shape():
config = make_dreamx_world_5b_ar_dit_config()
assert config.num_layers == 30
assert config.num_attention_heads == 24
assert config.attention_head_dim == 128
assert config.hidden_size == 3072
assert config.ffn_dim == 14336
assert config.local_attn_size == 12
assert config.sink_size == 3
assert config.attn_compress == 4
assert config.cam_self_attn_layers == tuple(range(30))
def test_dreamx_world_ar_converted_5b_transformer_strict_loads():
transformer_dir = CONVERTED_AR_DIR / "transformer"
if not transformer_dir.exists():
# No local conversion: pull the published Diffusers transformer from the hub.
from huggingface_hub import snapshot_download
transformer_dir = Path(snapshot_download(CONVERTED_AR_HF_REPO, allow_patterns=["transformer/*"])) / "transformer"
with torch.device("meta"):
model = DreamXWorldARTransformer3DModel(make_dreamx_world_5b_ar_dit_config(), {})
incompatible = load_model_from_full_model_state_dict(
model,
safetensors_weights_iterator(resolve_safetensors_files(str(transformer_dir)), to_cpu=True),
device=torch.device("cpu"),
param_dtype=torch.bfloat16,
strict=True,
param_names_mapping=get_param_names_mapping(model.param_names_mapping),
training_mode=False,
)
assert incompatible.missing_keys == []
assert incompatible.unexpected_keys == []
assert not any(param.is_meta for param in model.parameters())
@@ -0,0 +1,95 @@
# SPDX-License-Identifier: Apache-2.0
"""DreamX-World camera-conditioning parity against the official reference.
Coverage scope: implementation_subcomponent. This verifies the weightless
action-sequence to PRoPE camera tensor path used by DreamX-World-5B-Cam.
"""
from __future__ import annotations
import os
from pathlib import Path
import pytest
import torch
from torch.testing import assert_close
from fastvideo.pipelines.basic.dreamx_world.camera_conditioning import build_dreamx_camera_condition
REPO_ROOT = Path(__file__).resolve().parents[3]
OFFICIAL_REF_DIR = Path(os.getenv("DREAMX_WORLD_OFFICIAL_REF_DIR", REPO_ROOT / "DreamX-World"))
PARITY_SCOPE = "implementation_subcomponent"
def _load_official_functions():
if not OFFICIAL_REF_DIR.exists():
pytest.skip(f"Official reference missing: {OFFICIAL_REF_DIR}")
try:
import importlib.util
pose_path = OFFICIAL_REF_DIR / "utils" / "pose_utils.py"
pose_spec = importlib.util.spec_from_file_location(
"dreamx_world_pose_utils", pose_path)
if pose_spec is None or pose_spec.loader is None:
raise RuntimeError(f"Cannot load DreamX pose_utils: {pose_path}")
pose_module = importlib.util.module_from_spec(pose_spec)
pose_spec.loader.exec_module(pose_module)
source = (OFFICIAL_REF_DIR / "utils" / "inference_utils.py").read_text()
source = source.replace(
"from .pose_utils import interpolate_camera_poses\n", "")
namespace = {"interpolate_camera_poses": pose_module.interpolate_camera_poses}
exec(compile(source, str(OFFICIAL_REF_DIR / "utils" / "inference_utils.py"), "exec"), namespace)
except Exception as exc: # noqa: BLE001 - local parity should skip missing reference deps.
pytest.skip(f"Cannot load DreamX camera reference: {exc}")
return namespace["ActionToPoseFromID"], namespace["GetPoseEmbedsFromPosesPrope"]
def _official_camera_condition(
action_seq: list[str],
action_speed_list: list[float],
*,
num_frames: int,
height: int,
width: int,
dtype: torch.dtype,
):
action_to_pose, get_pose_embeds = _load_official_functions()
duration = -(-num_frames // len(action_seq))
poses = action_to_pose(action_seq, action_speed_list, duration=duration)[:num_frames]
condition, _ = get_pose_embeds(poses, height, width, len(poses), False, 0, dtype=dtype, device="cpu")
return condition
@pytest.mark.parametrize(
("action_seq", "action_speed_list", "num_frames"),
[
(["w"], [4], 81),
(["wj", "d"], [4, 6], 121),
(["i", "k", "l"], [3, 5, 2], 85),
],
)
def test_dreamx_world_camera_conditioning_matches_official(action_seq, action_speed_list, num_frames):
dtype = torch.float32
official = _official_camera_condition(
action_seq,
action_speed_list,
num_frames=num_frames,
height=704,
width=1280,
dtype=dtype,
)
fastvideo = build_dreamx_camera_condition(
action_seq,
action_speed_list,
num_frames=num_frames,
height=704,
width=1280,
dtype=dtype,
device="cpu",
)
assert official.keys() == fastvideo.keys() == {"viewmats", "K"}
for key in ("viewmats", "K"):
assert official[key].shape == fastvideo[key].shape
diff = (official[key] - fastvideo[key]).abs()
print(f"{key}: shape={tuple(fastvideo[key].shape)} diff_max={diff.max().item():.8f} diff_mean={diff.mean().item():.8f}")
assert_close(fastvideo[key], official[key], atol=1e-5, rtol=1e-5)
@@ -0,0 +1,79 @@
# SPDX-License-Identifier: Apache-2.0
"""DreamX-World conversion script smoke tests."""
import json
from fastvideo.configs.pipelines.dreamx_world import make_dreamx_world_5b_cam_dit_config
from fastvideo.models.registry import _LEGACY_FAST_VIDEO_MODELS
from scripts.checkpoint_conversion.dreamx_world_to_diffusers import (
MODEL_INDEX,
REUSED_COMPONENTS,
TRANSFORMER_CONFIG,
_copy_or_link_component,
write_model_index,
)
def test_dreamx_world_converter_writes_full_model_index_with_reused_components(tmp_path):
component_source = tmp_path / "wan22"
output = tmp_path / "dreamx"
output.mkdir()
for component in REUSED_COMPONENTS:
component_dir = component_source / component
component_dir.mkdir(parents=True)
(component_dir / "config.json").write_text("{}\n")
write_model_index(output, component_source, symlink_components=True)
model_index = json.loads((output / "model_index.json").read_text())
assert model_index == MODEL_INDEX
assert model_index["_class_name"] == "DreamXWorldPipeline"
assert model_index["transformer"] == ["diffusers", "DreamXWorldTransformer3DModel"]
for component in REUSED_COMPONENTS:
assert (output / component).is_symlink()
def test_dreamx_world_transformer_config_has_camera_adapter_enabled():
assert TRANSFORMER_CONFIG["_class_name"] == "DreamXWorldTransformer3DModel"
assert TRANSFORMER_CONFIG["add_control_adapter"] is True
assert TRANSFORMER_CONFIG["cam_method"] == "prope"
assert TRANSFORMER_CONFIG["num_layers"] == 30
def test_dreamx_world_converter_transformer_config_matches_pipeline_dit_config():
dit_config = make_dreamx_world_5b_cam_dit_config()
assert TRANSFORMER_CONFIG["num_attention_heads"] == dit_config.num_attention_heads
assert TRANSFORMER_CONFIG["attention_head_dim"] == dit_config.attention_head_dim
assert TRANSFORMER_CONFIG["in_channels"] == dit_config.in_channels
assert TRANSFORMER_CONFIG["out_channels"] == dit_config.out_channels
assert TRANSFORMER_CONFIG["ffn_dim"] == dit_config.ffn_dim
assert TRANSFORMER_CONFIG["num_layers"] == dit_config.num_layers
assert TRANSFORMER_CONFIG["cross_attn_norm"] == dit_config.cross_attn_norm
assert TRANSFORMER_CONFIG["qk_norm"] == dit_config.qk_norm
assert TRANSFORMER_CONFIG["add_control_adapter"] == dit_config.add_control_adapter
assert TRANSFORMER_CONFIG["cam_method"] == dit_config.cam_method
assert TRANSFORMER_CONFIG["attn_compress"] == dit_config.attn_compress
assert TRANSFORMER_CONFIG["cam_self_attn_layers"] == dit_config.cam_self_attn_layers
def test_dreamx_world_model_index_component_classes_are_registered():
for component in ("scheduler", "text_encoder", "transformer", "vae"):
class_name = MODEL_INDEX[component][1]
assert class_name in _LEGACY_FAST_VIDEO_MODELS
def test_dreamx_world_copy_or_link_component_keeps_broken_symlink(tmp_path):
component_source = tmp_path / "wan22"
src = component_source / "scheduler"
src.mkdir(parents=True)
output = tmp_path / "dreamx"
output.mkdir()
dst = output / "scheduler"
dst.symlink_to(tmp_path / "missing_scheduler", target_is_directory=True)
_copy_or_link_component("scheduler", component_source, output, symlink=True)
assert dst.is_symlink()
@@ -0,0 +1,306 @@
# SPDX-License-Identifier: Apache-2.0
"""DreamX-World pipeline config and conditioning smoke tests."""
from types import SimpleNamespace
import json
import numpy as np
import torch
from fastvideo.api.presets import get_preset
from fastvideo.fastvideo_args import WorkloadType
from fastvideo.models.schedulers.scheduling_flow_match_euler_discrete import FlowMatchEulerDiscreteScheduler
from fastvideo.pipelines.pipeline_batch_info import ForwardBatch
from fastvideo.pipelines.pipeline_registry import PipelineType, import_pipeline_classes
from fastvideo.pipelines.basic.dreamx_world.camera_conditioning import (
DreamXCamera,
_interpolate_camera_poses,
build_dreamx_camera_condition,
)
from fastvideo.configs.pipelines.dreamx_world import DreamXWorld5BARPipelineConfig, DreamXWorld5BCamPipelineConfig
from fastvideo.pipelines.basic.dreamx_world.ar_denoising import DreamXWorldARCausalDenoisingStage
from fastvideo.pipelines.basic.dreamx_world.dreamx_world_ar_pipeline import DreamXWorldARPipeline
from fastvideo.pipelines.basic.dreamx_world.dreamx_world_pipeline import DreamXWorldPipeline
from fastvideo.pipelines.basic.dreamx_world.stages import (
DREAMX_Y_CAMERA_KEY,
DreamXWorldCameraConditioningStage,
)
from fastvideo.pipelines.stages.denoising import DenoisingStage
from fastvideo.registry import get_default_preset, get_model_info, get_pipeline_config_cls_from_name
def test_dreamx_world_5b_cam_pipeline_config_wires_first_scope_components():
config = DreamXWorld5BCamPipelineConfig()
assert config.flow_shift == 3.0
assert config.ti2v_task is True
assert config.expand_timesteps is True
assert config.dit_config.expand_timesteps is True
assert config.dit_config.num_layers == 30
assert config.dit_config.add_control_adapter is True
assert config.dit_config.cam_method == "prope"
assert config.vae_config.load_encoder is True
assert config.vae_config.load_decoder is True
assert config.vae_config.z_dim == 48
assert config.vae_config.scale_factor_temporal == 4
assert config.vae_config.scale_factor_spatial == 16
assert len(config.text_encoder_configs) == 1
text_config = config.text_encoder_configs[0]
assert text_config.prefix == "umt5"
assert text_config.vocab_size == 256384
assert text_config.d_model == 4096
assert config.text_encoder_precisions == ("bf16",)
def test_dreamx_world_pipeline_registry_discovers_entrypoint():
pipelines = import_pipeline_classes(PipelineType.BASIC)
assert pipelines["basic"]["DreamXWorldPipeline"] is DreamXWorldPipeline
def test_dreamx_world_local_model_index_resolves_model_info(tmp_path):
model_dir = tmp_path / "DreamX-World-5B-Cam-converted"
model_dir.mkdir()
for component in ("scheduler", "text_encoder", "tokenizer", "transformer", "vae"):
(model_dir / component).mkdir()
(model_dir / "model_index.json").write_text(json.dumps({
"_class_name": "DreamXWorldPipeline",
"_diffusers_version": "0.31.0",
"scheduler": ["diffusers", "FlowMatchEulerDiscreteScheduler"],
"text_encoder": ["transformers", "UMT5EncoderModel"],
"tokenizer": ["transformers", "AutoTokenizer"],
"transformer": ["diffusers", "DreamXWorldTransformer3DModel"],
"vae": ["diffusers", "AutoencoderKLWan"],
}) + "\n")
info = get_model_info(str(model_dir), pipeline_type=PipelineType.BASIC, workload_type=WorkloadType.I2V)
assert info.pipeline_cls is DreamXWorldPipeline
assert info.pipeline_config_cls is DreamXWorld5BCamPipelineConfig
def test_dreamx_world_model_path_resolves_pipeline_config():
assert get_pipeline_config_cls_from_name("GD-ML/DreamX-World-5B-Cam") is DreamXWorld5BCamPipelineConfig
def test_dreamx_world_default_preset_is_registered():
preset_name = get_default_preset("GD-ML/DreamX-World-5B-Cam")
preset = get_preset(preset_name, "dreamx_world")
assert preset.name == "dreamx_world_5b_cam"
assert preset.workload_type == "i2v"
assert preset.defaults["height"] == 480
assert preset.defaults["width"] == 832
assert preset.defaults["num_frames"] == 161
assert preset.defaults["num_inference_steps"] == 30
assert preset.defaults["guidance_scale"] == 5.0
def test_dreamx_world_pipeline_initializes_official_flow_scheduler():
pipeline = DreamXWorldPipeline.__new__(DreamXWorldPipeline)
pipeline.modules = {}
fastvideo_args = SimpleNamespace(pipeline_config=DreamXWorld5BCamPipelineConfig())
pipeline.initialize_pipeline(fastvideo_args)
scheduler = pipeline.modules["scheduler"]
assert isinstance(scheduler, FlowMatchEulerDiscreteScheduler)
assert scheduler.config.shift == 3.0
def test_dreamx_world_camera_conditioning_stage_sets_y_camera_extra():
batch = ForwardBatch(
data_type="t2v",
action_list=["wj", "d"],
action_speed_list=[4, 6],
num_frames=17,
height=704,
width=1280,
latents=torch.zeros(1, 16, 5, 44, 80),
)
stage = DreamXWorldCameraConditioningStage()
out = stage.forward(batch, fastvideo_args=object())
y_camera = out.extra[DREAMX_Y_CAMERA_KEY]
expected = build_dreamx_camera_condition(
["wj", "d"],
[4, 6],
num_frames=17,
height=704,
width=1280,
dtype=torch.float32,
device="cpu",
)
assert set(y_camera) == {"viewmats", "K"}
for key, expected_value in expected.items():
assert y_camera[key].shape == (1, *expected_value.shape)
torch.testing.assert_close(y_camera[key][0], expected_value)
assert stage.verify_output(out, fastvideo_args=object()).is_valid()
def test_dreamx_world_denoising_kwargs_filter_for_y_camera():
y_camera = {"viewmats": torch.eye(4).reshape(1, 1, 4, 4), "K": torch.eye(3).reshape(1, 1, 3, 3)}
stage = DenoisingStage.__new__(DenoisingStage)
def accepts_y_camera(hidden_states, encoder_hidden_states, timestep, y_camera=None):
return y_camera
def no_y_camera(hidden_states, encoder_hidden_states, timestep):
return hidden_states
assert stage.prepare_extra_func_kwargs(accepts_y_camera, {"y_camera": y_camera}) == {"y_camera": y_camera}
assert stage.prepare_extra_func_kwargs(no_y_camera, {"y_camera": y_camera}) == {}
def test_dreamx_world_ar_pipeline_config_wires_components():
config = DreamXWorld5BARPipelineConfig()
assert config.is_causal is True
assert config.flow_shift == 5.0
assert config.dmd_denoising_steps == (1000, 750, 500, 250)
assert config.warp_denoising_step is True
assert config.context_noise == 0.1
assert config.dit_config.arch_config.local_attn_size == 12
assert config.dit_config.arch_config.sink_size == 3
assert config.dit_config.arch_config.attn_compress == 4
def test_dreamx_world_ar_pipeline_registry_and_preset():
from fastvideo.api.presets import get_preset
from fastvideo.registry import get_pipeline_config_cls_from_name, get_preset_selection
assert get_pipeline_config_cls_from_name("GD-ML/DreamX-World-5B") is DreamXWorld5BARPipelineConfig
preset_name, family = get_preset_selection("GD-ML/DreamX-World-5B")
assert (preset_name, family) == ("dreamx_world_5b_ar", "dreamx_world")
preset = get_preset("dreamx_world_5b_ar", "dreamx_world")
assert preset.defaults["num_inference_steps"] == 4
assert DreamXWorldARPipeline.pipeline_config_cls is DreamXWorld5BARPipelineConfig
def test_dreamx_world_camera_conditioning_stage_expands_scalar_speed():
batch = ForwardBatch(
data_type="t2v",
action_list=["w", "d"],
action_speed_list=2.0,
num_frames=17,
height=704,
width=1280,
latents=torch.zeros(1, 16, 5, 44, 80),
)
stage = DreamXWorldCameraConditioningStage()
out = stage.forward(batch, fastvideo_args=object())
assert set(out.extra[DREAMX_Y_CAMERA_KEY]) == {"viewmats", "K"}
def test_dreamx_world_camera_interpolation_handles_single_camera():
camera = DreamXCamera(
fx=0.8,
fy=0.8,
cx=0.5,
cy=0.5,
w2c_mat=np.eye(4, dtype=np.float64),
)
out = _interpolate_camera_poses(
[camera],
src_indices=np.array([0.0]),
tgt_indices=np.array([0.0, 1.0, 2.0]),
)
assert out == [camera, camera, camera]
def test_dreamx_world_ar_cache_initializes_camera_self_attention_entries():
cam_self_attn = SimpleNamespace(num_heads=3, head_dim=5)
transformer = SimpleNamespace(
blocks=[SimpleNamespace(cam_self_attn=cam_self_attn), SimpleNamespace(cam_self_attn=cam_self_attn)],
num_attention_heads=2,
attention_head_dim=4,
)
stage = DreamXWorldARCausalDenoisingStage.__new__(DreamXWorldARCausalDenoisingStage)
stage.transformer = transformer
stage.num_transformer_blocks = 2
stage.local_attn_size = 6
caches = stage._initialize_kv_cache(
batch_size=1,
dtype=torch.float32,
device=torch.device("cpu"),
frame_seq_length=7,
)
assert len(caches) == 2
assert caches[0]["k"].shape == (1, 42, 2, 4)
assert caches[0]["prope_k"].shape == (1, 42, 3, 5)
assert caches[0]["prope_v"].shape == (1, 42, 3, 5)
assert int(caches[0]["prope_global_end_index"].item()) == 0
assert int(caches[0]["prope_local_end_index"].item()) == 0
def test_dreamx_world_ar_context_noise_fraction_maps_to_scheduler_timestep():
assert DreamXWorldARCausalDenoisingStage._context_noise_timestep(0.1) == 100
assert DreamXWorldARCausalDenoisingStage._context_noise_timestep(100) == 100
def test_dreamx_world_ar_context_update_advances_camera_cache_indices():
class DummyTransformer:
def __call__(self, *, hidden_states, encoder_hidden_states, timestep, y_camera, kv_cache, crossattn_cache,
current_start):
del encoder_hidden_states, y_camera, crossattn_cache
assert current_start == 0
assert timestep.unique().tolist() == [100]
new_tokens = timestep.shape[1]
for cache in kv_cache:
cache["local_end_index"] += new_tokens
cache["global_end_index"] += new_tokens
cache["prope_local_end_index"] += new_tokens
cache["prope_global_end_index"] += new_tokens
cache["k"][:, :new_tokens] = 1
cache["prope_k"][:, :new_tokens] = 1
return hidden_states
cam_self_attn = SimpleNamespace(num_heads=3, head_dim=5)
cache_transformer = SimpleNamespace(
blocks=[SimpleNamespace(cam_self_attn=cam_self_attn)],
num_attention_heads=2,
attention_head_dim=4,
)
stage = DreamXWorldARCausalDenoisingStage.__new__(DreamXWorldARCausalDenoisingStage)
stage.transformer = cache_transformer
stage.num_transformer_blocks = 1
stage.local_attn_size = 6
caches = stage._initialize_kv_cache(
batch_size=1,
dtype=torch.float32,
device=torch.device("cpu"),
frame_seq_length=2,
)
# Keep the cache allocation source separate from the callable transformer used by _update_context_cache.
stage.transformer = DummyTransformer()
stage._update_context_cache(
block_latents=torch.zeros(1, 4, 3, 2, 2),
context=[torch.zeros(2, 4)],
camera_block={"viewmats": torch.eye(4).reshape(1, 1, 4, 4), "K": torch.eye(3).reshape(1, 1, 3, 3)},
kv_cache=caches,
crossattn_cache=[{}],
start=0,
frame_seq_length=2,
target_dtype=torch.float32,
autocast_enabled=False,
context_noise=0.1,
)
assert int(caches[0]["local_end_index"].item()) == 6
assert int(caches[0]["prope_local_end_index"].item()) == 6
assert caches[0]["k"][:, :6].sum().item() == 48
assert caches[0]["prope_k"][:, :6].sum().item() == 90
@@ -0,0 +1,53 @@
# SPDX-License-Identifier: Apache-2.0
"""DreamX-World default Flow scheduler parity.
Coverage scope: implementation_subcomponent. DreamX-World-5B-Cam defaults to
Diffusers FlowMatchEulerDiscreteScheduler for sampler_name=Flow. This test
checks that FastVideo's native FlowMatchEulerDiscreteScheduler matches the
timestep schedule and Euler step used by the official default path.
"""
from __future__ import annotations
from pathlib import Path
import inspect
import torch
from diffusers import FlowMatchEulerDiscreteScheduler as OfficialFlowMatchEulerDiscreteScheduler
from omegaconf import OmegaConf
from torch.testing import assert_close
from fastvideo.models.schedulers.scheduling_flow_match_euler_discrete import (
FlowMatchEulerDiscreteScheduler as FastVideoFlowMatchEulerDiscreteScheduler,
)
REPO_ROOT = Path(__file__).resolve().parents[3]
PARITY_SCOPE = "implementation_subcomponent"
def _scheduler_kwargs(cls):
config_path = REPO_ROOT / "DreamX-World" / "configs" / "wan2.2" / "wan_ti2v_5b.yaml"
config = OmegaConf.load(config_path)
raw_kwargs = OmegaConf.to_container(config["scheduler_kwargs"])
signature = inspect.signature(cls)
return {key: value for key, value in raw_kwargs.items() if key in signature.parameters}
def test_dreamx_world_flow_scheduler_timesteps_and_step_match():
official = OfficialFlowMatchEulerDiscreteScheduler(**_scheduler_kwargs(OfficialFlowMatchEulerDiscreteScheduler))
fastvideo = FastVideoFlowMatchEulerDiscreteScheduler(**_scheduler_kwargs(FastVideoFlowMatchEulerDiscreteScheduler))
official.set_timesteps(50, device="cpu", mu=1)
fastvideo.set_timesteps(50, device="cpu", mu=1)
assert_close(fastvideo.timesteps, official.timesteps, atol=0, rtol=0)
assert_close(fastvideo.sigmas, official.sigmas, atol=0, rtol=0)
torch.manual_seed(7)
sample = torch.randn(1, 4, 2, 8, 8)
model_output = torch.randn_like(sample)
timestep = official.timesteps[3]
official_prev = official.step(model_output, timestep, sample, return_dict=False)[0]
fastvideo_prev = fastvideo.step(model_output, fastvideo.timesteps[3], sample, return_dict=False)[0]
diff = (official_prev - fastvideo_prev).abs()
print(f"scheduler diff_max={diff.max().item():.8f} diff_mean={diff.mean().item():.8f}")
assert_close(fastvideo_prev, official_prev, atol=0, rtol=0)
@@ -0,0 +1,159 @@
# SPDX-License-Identifier: Apache-2.0
"""DreamX-World Wan T5 encoder reuse parity scaffold.
Coverage scope: implementation_subcomponent. It records the official
WanT5EncoderModel loading path and FastVideo T5 target for later activation
with staged Wan2.2 base text encoder/tokenizer weights.
"""
from __future__ import annotations
import os
from pathlib import Path
import sys
import types
import pytest
import torch
from omegaconf import OmegaConf
from torch.testing import assert_close
from transformers import AutoTokenizer
from fastvideo.fastvideo_args import FastVideoArgs
from fastvideo.configs.pipelines.dreamx_world import (
DreamXWorld5BCamPipelineConfig,
make_dreamx_world_5b_cam_text_encoder_config,
)
from fastvideo.models.loader.component_loader import TextEncoderLoader
REPO_ROOT = Path(__file__).resolve().parents[3]
OFFICIAL_REF_DIR = Path(os.getenv("DREAMX_WORLD_OFFICIAL_REF_DIR", REPO_ROOT / "DreamX-World"))
WAN_BASE_DIR = Path(os.getenv("DREAMX_WORLD_WAN_BASE_DIR", REPO_ROOT / "official_weights" / "Wan2.2-TI2V-5B"))
WAN_DIFFUSERS_DIR = Path(os.getenv("DREAMX_WORLD_WAN_DIFFUSERS_DIR", REPO_ROOT / "official_weights" / "Wan2.2-TI2V-5B-Diffusers"))
PARITY_SCOPE = "implementation_subcomponent"
def _add_official_to_path():
if not OFFICIAL_REF_DIR.exists():
pytest.skip(f"Official reference missing: {OFFICIAL_REF_DIR}")
if str(OFFICIAL_REF_DIR) not in sys.path:
sys.path.insert(0, str(OFFICIAL_REF_DIR))
def _install_xfuser_stub() -> None:
if "xfuser" in sys.modules:
return
xfuser = types.ModuleType("xfuser")
core = types.ModuleType("xfuser.core")
distributed = types.ModuleType("xfuser.core.distributed")
long_ctx = types.ModuleType("xfuser.core.long_ctx_attention")
distributed.get_sequence_parallel_rank = lambda: 0
distributed.get_sequence_parallel_world_size = lambda: 1
distributed.get_sp_group = lambda: None
distributed.get_world_group = lambda: types.SimpleNamespace(local_rank=0, rank=0)
distributed.init_distributed_environment = lambda *args, **kwargs: None
distributed.initialize_model_parallel = lambda *args, **kwargs: None
distributed.model_parallel_is_initialized = lambda: False
class XFuserLongContextAttention:
def __call__(self, *args, **kwargs):
raise RuntimeError("xfuser stub cannot execute attention")
long_ctx.xFuserLongContextAttention = XFuserLongContextAttention
sys.modules.update({
"xfuser": xfuser,
"xfuser.core": core,
"xfuser.core.distributed": distributed,
"xfuser.core.long_ctx_attention": long_ctx,
})
def _text_kwargs():
config = OmegaConf.load(OFFICIAL_REF_DIR / "configs" / "wan2.2" / "wan_ti2v_5b.yaml")
return OmegaConf.to_container(config["text_encoder_kwargs"])
def _patch_single_process_text_parallel(monkeypatch):
import fastvideo.layers.linear as fastvideo_linear
import fastvideo.layers.vocab_parallel_embedding as fastvideo_embedding
import fastvideo.models.encoders.t5 as fastvideo_t5
for module in (fastvideo_t5, fastvideo_embedding, fastvideo_linear):
if hasattr(module, "get_tp_rank"):
monkeypatch.setattr(module, "get_tp_rank", lambda: 0)
if hasattr(module, "get_tp_world_size"):
monkeypatch.setattr(module, "get_tp_world_size", lambda: 1)
monkeypatch.setattr(fastvideo_embedding, "tensor_model_parallel_all_reduce", lambda x: x)
def _load_official_text_encoder(device, dtype):
_add_official_to_path()
_install_xfuser_stub()
text_path = WAN_BASE_DIR / "models_t5_umt5-xxl-enc-bf16.pth"
if not text_path.exists():
pytest.skip(f"Wan2.2 text encoder weights missing: {text_path}")
try:
from models import WanT5EncoderModel
except Exception as exc: # noqa: BLE001
pytest.skip(f"Cannot import official DreamX text encoder: {exc}")
model = WanT5EncoderModel.from_pretrained(
str(text_path), additional_kwargs=_text_kwargs(), low_cpu_mem_usage=True, torch_dtype=dtype
)
return model.to(device=device, dtype=dtype).eval()
def _load_fastvideo_text_encoder(device, dtype, monkeypatch):
text_encoder_path = WAN_DIFFUSERS_DIR / "text_encoder"
if not text_encoder_path.exists():
pytest.skip(f"Wan2.2 Diffusers text encoder missing: {text_encoder_path}")
_patch_single_process_text_parallel(monkeypatch)
pipeline_config = DreamXWorld5BCamPipelineConfig()
pipeline_config.text_encoder_configs[0]._fsdp_shard_conditions = []
args = FastVideoArgs(
model_path=str(text_encoder_path),
pipeline_config=pipeline_config,
text_encoder_cpu_offload=(device.type == "cpu"),
)
args.model_paths = {}
return TextEncoderLoader().load(str(text_encoder_path), args).to(device=device, dtype=dtype).eval()
def test_dreamx_world_text_encoder_config_matches_umt5_xxl_shape():
config = make_dreamx_world_5b_cam_text_encoder_config()
assert config.vocab_size == 256384
assert config.d_model == 4096
assert config.d_kv == 64
assert config.d_ff == 10240
assert config.num_heads == 64
assert config.num_layers == 24
assert config.relative_attention_num_buckets == 32
assert config.dropout_rate == 0.0
assert config.text_len == 512
assert config.prefix == "umt5"
def test_dreamx_world_fastvideo_text_encoder_loads_staged_weights(monkeypatch):
device = torch.device("cuda:0" if torch.cuda.is_available() else "cpu")
model = _load_fastvideo_text_encoder(device, torch.bfloat16, monkeypatch)
assert model.__class__.__name__ == "UMT5EncoderModel"
assert next(model.parameters()).device.type == device.type
assert next(model.parameters()).dtype == torch.bfloat16
@pytest.mark.skipif(not torch.cuda.is_available(), reason="CUDA required for text encoder parity.")
def test_dreamx_world_text_encoder_parity_scaffold(monkeypatch):
device = torch.device("cuda:0")
dtype = torch.bfloat16
official = _load_official_text_encoder(device, dtype)
fastvideo = _load_fastvideo_text_encoder(device, dtype, monkeypatch)
tokenizer_path = WAN_BASE_DIR / "google" / "umt5-xxl"
if not tokenizer_path.exists():
pytest.skip(f"Wan2.2 tokenizer missing: {tokenizer_path}")
tokenizer = AutoTokenizer.from_pretrained(str(tokenizer_path))
batch = tokenizer(["A quiet forest trail at sunrise."], padding="max_length", max_length=512, return_tensors="pt")
input_ids = batch.input_ids.to(device)
attention_mask = batch.attention_mask.to(device)
with torch.inference_mode():
official_hidden = official(input_ids, attention_mask=attention_mask)[0].float().cpu()
fastvideo_hidden = fastvideo(input_ids, attention_mask=attention_mask).last_hidden_state.float().cpu()
assert official_hidden.shape == fastvideo_hidden.shape
assert_close(fastvideo_hidden, official_hidden, atol=1e-3, rtol=1e-3)
@@ -0,0 +1,335 @@
# SPDX-License-Identifier: Apache-2.0
"""DreamX-World transformer parity scaffold.
Coverage scope: both. The official side loads DreamX-World-5B-Cam through
Wan2_2Transformer3DModel.from_pretrained with PRoPE camera control enabled.
The FastVideo side strict-loads the converted DreamX transformer weights into
the native DreamX-World DiT implementation.
"""
from __future__ import annotations
import math
import os
from pathlib import Path
import sys
import types
import pytest
import torch
from omegaconf import OmegaConf
from torch.testing import assert_close
from fastvideo.forward_context import set_forward_context
from fastvideo.configs.models.dits.dreamx_world import (
DreamXWorldArchConfig, DreamXWorldConfig)
from fastvideo.pipelines.basic.dreamx_world.camera_conditioning import build_dreamx_camera_condition
from fastvideo.configs.pipelines.dreamx_world import make_dreamx_world_5b_cam_dit_config
from fastvideo.platforms import AttentionBackendEnum
from fastvideo.models.dits.dreamx_world import (
DreamXPropeSelfAttention, DreamXWorldTransformer3DModel,
DreamXWorldTransformerBlock)
from fastvideo.models.loader.fsdp_load import load_model_from_full_model_state_dict
from fastvideo.models.loader.utils import get_param_names_mapping
from fastvideo.models.loader.weight_utils import resolve_safetensors_files, safetensors_weights_iterator
from scripts.checkpoint_conversion.dreamx_world_to_diffusers import map_transformer_key
REPO_ROOT = Path(__file__).resolve().parents[3]
OFFICIAL_REF_DIR = Path(os.getenv("DREAMX_WORLD_OFFICIAL_REF_DIR", REPO_ROOT / "DreamX-World"))
LOCAL_WEIGHTS_DIR = Path(os.getenv("DREAMX_WORLD_LOCAL_WEIGHTS_DIR", REPO_ROOT / "official_weights" / "dreamx_world"))
WAN_BASE_DIR = Path(os.getenv("DREAMX_WORLD_WAN_BASE_DIR", REPO_ROOT / "official_weights" / "Wan2.2-TI2V-5B"))
CONVERTED_WEIGHTS_DIR = Path(os.getenv("DREAMX_WORLD_CONVERTED_WEIGHTS_DIR", REPO_ROOT / "converted_weights" / "dreamx_world"))
CONVERTED_HF_REPO = "FastVideo/DreamX-World-5B-Cam-Diffusers"
PARITY_SCOPE = "both"
def _install_xfuser_stub() -> None:
if "xfuser" in sys.modules:
return
xfuser = types.ModuleType("xfuser")
core = types.ModuleType("xfuser.core")
distributed = types.ModuleType("xfuser.core.distributed")
long_ctx = types.ModuleType("xfuser.core.long_ctx_attention")
distributed.get_sequence_parallel_rank = lambda: 0
distributed.get_sequence_parallel_world_size = lambda: 1
distributed.get_sp_group = lambda: None
distributed.get_world_group = lambda: types.SimpleNamespace(local_rank=0, rank=0)
distributed.init_distributed_environment = lambda *args, **kwargs: None
distributed.initialize_model_parallel = lambda *args, **kwargs: None
distributed.model_parallel_is_initialized = lambda: False
class XFuserLongContextAttention:
def __call__(self, *args, **kwargs):
raise RuntimeError("xfuser stub cannot execute attention")
long_ctx.xFuserLongContextAttention = XFuserLongContextAttention
sys.modules.update({
"xfuser": xfuser,
"xfuser.core": core,
"xfuser.core.distributed": distributed,
"xfuser.core.long_ctx_attention": long_ctx,
})
def _make_tiny_dreamx_config() -> DreamXWorldConfig:
return DreamXWorldConfig(
arch_config=DreamXWorldArchConfig(
num_attention_heads=1,
attention_head_dim=8,
in_channels=16,
out_channels=16,
ffn_dim=32,
num_layers=1,
cross_attn_norm=True,
qk_norm="rms_norm_across_heads",
add_control_adapter=True,
cam_method="prope",
attn_compress=1,
cam_self_attn_layers=None,
))
def _add_official_to_path() -> None:
if not OFFICIAL_REF_DIR.exists():
pytest.skip(f"Official reference missing: {OFFICIAL_REF_DIR}")
if str(OFFICIAL_REF_DIR) not in sys.path:
sys.path.insert(0, str(OFFICIAL_REF_DIR))
def _official_transformer_kwargs() -> dict:
config_path = OFFICIAL_REF_DIR / "configs" / "wan2.2" / "wan_ti2v_5b.yaml"
if not config_path.exists():
pytest.skip(f"DreamX Wan config missing: {config_path}")
config = OmegaConf.load(config_path)
kwargs = OmegaConf.to_container(config["transformer_additional_kwargs"])
kwargs["cam_method"] = "prope"
kwargs["add_control_adapter"] = True
return kwargs
def _load_official_transformer(device: torch.device, dtype: torch.dtype) -> torch.nn.Module:
_add_official_to_path()
_install_xfuser_stub()
if not LOCAL_WEIGHTS_DIR.exists():
pytest.skip(f"DreamX transformer weights missing: {LOCAL_WEIGHTS_DIR}")
try:
from models import Wan2_2Transformer3DModel
except Exception as exc: # noqa: BLE001 - local parity should skip missing refs.
pytest.skip(f"Cannot import official DreamX transformer: {exc}")
model = Wan2_2Transformer3DModel.from_pretrained(
str(LOCAL_WEIGHTS_DIR),
transformer_additional_kwargs=_official_transformer_kwargs(),
torch_dtype=dtype,
)
return model.to(device=device, dtype=dtype).eval()
def _load_fastvideo_transformer(device: torch.device, dtype: torch.dtype) -> torch.nn.Module:
model = _load_fastvideo_transformer_strict(torch.device("cpu"), dtype)
return model.to(device=device, dtype=dtype).eval()
def _load_fastvideo_transformer_strict(device: torch.device, dtype: torch.dtype) -> torch.nn.Module:
transformer_dir = CONVERTED_WEIGHTS_DIR / "transformer"
if not transformer_dir.exists():
# No local conversion: pull the published Diffusers transformer from the hub.
from huggingface_hub import snapshot_download
transformer_dir = Path(snapshot_download(CONVERTED_HF_REPO, allow_patterns=["transformer/*"])) / "transformer"
safetensors_files = resolve_safetensors_files(str(transformer_dir))
config = make_dreamx_world_5b_cam_dit_config()
import fastvideo.models.dits.dreamx_world as fastvideo_dreamx
original_get_sp_world_size = fastvideo_dreamx.get_sp_world_size
fastvideo_dreamx.get_sp_world_size = lambda: 1
try:
with torch.device("meta"):
model = DreamXWorldTransformer3DModel(config=config, hf_config={})
finally:
fastvideo_dreamx.get_sp_world_size = original_get_sp_world_size
incompatible = load_model_from_full_model_state_dict(
model,
safetensors_weights_iterator(safetensors_files, to_cpu=True),
device=device,
param_dtype=dtype,
strict=True,
param_names_mapping=get_param_names_mapping(model.param_names_mapping),
training_mode=False,
)
assert incompatible.missing_keys == []
assert incompatible.unexpected_keys == []
assert not any(param.is_meta for param in model.parameters())
return model.to(device=device, dtype=dtype).eval()
def _make_inputs(device: torch.device, dtype: torch.dtype):
torch.manual_seed(1234)
num_frames = 5
height = 64
width = 64
latent_frames = (num_frames - 1) // 4 + 1
latent_h = height // 16
latent_w = width // 16
x = torch.randn(1, 48, latent_frames, latent_h, latent_w, device=device, dtype=dtype)
context = [torch.randn(16, 4096, device=device, dtype=dtype)]
seq_len = math.ceil((latent_h * latent_w) / 4 * latent_frames)
timestep = torch.full((1, seq_len), 250, device=device, dtype=torch.long)
camera = build_dreamx_camera_condition(
["w"], [4], num_frames=num_frames, height=height, width=width, dtype=dtype, device=device
)
camera = {key: value.unsqueeze(0) for key, value in camera.items()}
return {"x": [x[0]], "context": context, "t": timestep, "seq_len": seq_len, "y_camera": camera}
def _run_official(model: torch.nn.Module, inputs: dict) -> torch.Tensor:
with torch.inference_mode():
output = model(**inputs)
if isinstance(output, list):
output = torch.stack(output, dim=0)
assert torch.is_tensor(output), f"official output is not a tensor: {type(output)}"
return output.detach().float().cpu()
def _run_fastvideo(model: torch.nn.Module, inputs: dict) -> torch.Tensor:
hidden_states = torch.stack(inputs["x"], dim=0)
encoder_hidden_states = torch.stack([
torch.cat([inputs["context"][0], inputs["context"][0].new_zeros(512 - inputs["context"][0].shape[0], 4096)])
])
with torch.inference_mode(), set_forward_context(current_timestep=0, attn_metadata=None):
output = model(
hidden_states=hidden_states,
encoder_hidden_states=encoder_hidden_states,
timestep=inputs["t"],
y_camera=inputs["y_camera"],
)
assert torch.is_tensor(output), f"FastVideo output is not a tensor: {type(output)}"
return output.detach().float().cpu()
def test_dreamx_world_conversion_mapping_strict_load_smoke(monkeypatch):
_add_official_to_path()
_install_xfuser_stub()
try:
from models import Wan2_2Transformer3DModel
except Exception as exc: # noqa: BLE001
pytest.skip(f"Cannot import official DreamX transformer: {exc}")
official = Wan2_2Transformer3DModel(
dim=8,
ffn_dim=32,
num_heads=1,
num_layers=1,
add_control_adapter=True,
cam_method="prope",
)
official_state = official.state_dict()
diffusers_like_state = {
map_transformer_key(key): value.detach().clone()
for key, value in official_state.items()
}
import fastvideo.models.dits.dreamx_world as fastvideo_dreamx
monkeypatch.setattr(fastvideo_dreamx, "get_sp_world_size", lambda: 1)
with torch.device("meta"):
fastvideo = DreamXWorldTransformer3DModel(config=_make_tiny_dreamx_config(), hf_config={})
incompatible = load_model_from_full_model_state_dict(
fastvideo,
iter(diffusers_like_state.items()),
device=torch.device("cpu"),
param_dtype=torch.float32,
strict=True,
param_names_mapping=get_param_names_mapping(fastvideo.param_names_mapping),
training_mode=False,
)
assert incompatible.missing_keys == []
assert incompatible.unexpected_keys == []
assert not any(param.is_meta for param in fastvideo.parameters())
def test_dreamx_world_5b_cam_dit_config_matches_official_shape():
config = make_dreamx_world_5b_cam_dit_config()
assert config.num_layers == 30
assert config.num_attention_heads == 24
assert config.attention_head_dim == 128
assert config.hidden_size == 3072
assert config.ffn_dim == 14336
assert config.add_control_adapter is True
assert config.cam_method == "prope"
assert config.attn_compress == 1
def test_dreamx_world_converted_5b_transformer_strict_loads():
device = torch.device("cuda:0" if torch.cuda.is_available() else "cpu")
_load_fastvideo_transformer_strict(device, torch.bfloat16)
def test_dreamx_world_fastvideo_prope_branch_smoke():
block = DreamXWorldTransformerBlock(
8,
32,
1,
cross_attn_norm=True,
add_control_adapter=True,
cam_method="prope",
attn_compress=1,
layer_idx=0,
supported_attention_backends=(AttentionBackendEnum.TORCH_SDPA,),
)
assert block.cam_self_attn is not None
assert [
name for name, _ in block.named_parameters()
if name.startswith("cam_self_attn.")
][:8] == [
"cam_self_attn.q_proj.weight",
"cam_self_attn.q_proj.bias",
"cam_self_attn.k_proj.weight",
"cam_self_attn.k_proj.bias",
"cam_self_attn.v_proj.weight",
"cam_self_attn.v_proj.bias",
"cam_self_attn.out_proj.weight",
"cam_self_attn.out_proj.bias",
]
module = DreamXPropeSelfAttention(
dim=8,
attn_dim=8,
num_heads=1,
qk_norm="rms_norm_across_heads",
).eval()
assert module.num_heads == 1
assert module.head_dim == 8
assert tuple(module.out_proj.weight.shape) == (8, 8)
assert torch.count_nonzero(module.out_proj.weight) == 0
@pytest.mark.skipif(not torch.cuda.is_available(), reason="CUDA required for transformer parity.")
def test_dreamx_world_transformer_parity_scaffold(monkeypatch):
device = torch.device("cuda:0")
dtype = torch.float32
inputs = _make_inputs(device, dtype)
official = _load_official_transformer(device, dtype)
official_out = _run_official(official, inputs)
del official
torch.cuda.empty_cache()
import fastvideo.attention.layer as attention_layer
import fastvideo.distributed.communication_op as communication_op
import fastvideo.models.dits.dreamx_world as fastvideo_dreamx
monkeypatch.setattr(attention_layer, "get_sp_parallel_rank", lambda: 0)
monkeypatch.setattr(attention_layer, "get_sp_world_size", lambda: 1)
monkeypatch.setattr(attention_layer, "sequence_model_parallel_all_to_all_4D", lambda tensor, scatter_dim=2, gather_dim=1: tensor)
monkeypatch.setattr(attention_layer, "sequence_model_parallel_all_gather", lambda tensor, dim=-1: tensor)
monkeypatch.setattr(fastvideo_dreamx, "get_sp_world_size", lambda: 1)
monkeypatch.setattr(communication_op, "get_sp_world_size", lambda: 1)
monkeypatch.setattr(fastvideo_dreamx, "sequence_model_parallel_shard", lambda tensor, dim=1: (tensor, tensor.shape[dim]))
monkeypatch.setattr(
fastvideo_dreamx,
"sequence_model_parallel_all_gather_with_unpad",
lambda tensor, original_seq_len, dim=1: tensor.narrow(dim, 0, original_seq_len),
)
fastvideo = _load_fastvideo_transformer(device, dtype)
fastvideo_out = _run_fastvideo(fastvideo, inputs)
assert official_out.shape == fastvideo_out.shape
diff = (official_out - fastvideo_out).abs()
print(f"diff_max={diff.max().item():.6f} diff_mean={diff.mean().item():.6f}")
assert_close(fastvideo_out, official_out, atol=1e-1, rtol=1e-1)
@@ -0,0 +1,220 @@
# SPDX-License-Identifier: Apache-2.0
"""DreamX-World Wan2.2 VAE reuse parity scaffold.
Coverage scope: implementation_subcomponent. The official side uses
AutoencoderKLWan3_8 from DreamX, while the FastVideo side targets the native
Wan VAE. This remains a scaffold until Wan2.2 base VAE weights are staged.
"""
from __future__ import annotations
import os
from pathlib import Path
import sys
import re
import types
import pytest
import torch
from omegaconf import OmegaConf
from torch.testing import assert_close
from fastvideo.configs.pipelines.dreamx_world import make_dreamx_world_5b_cam_vae_config
from fastvideo.models.vaes.wanvae import AutoencoderKLWan
REPO_ROOT = Path(__file__).resolve().parents[3]
OFFICIAL_REF_DIR = Path(os.getenv("DREAMX_WORLD_OFFICIAL_REF_DIR", REPO_ROOT / "DreamX-World"))
WAN_BASE_DIR = Path(os.getenv("DREAMX_WORLD_WAN_BASE_DIR", REPO_ROOT / "official_weights" / "Wan2.2-TI2V-5B"))
PARITY_SCOPE = "implementation_subcomponent"
def _add_official_to_path():
if not OFFICIAL_REF_DIR.exists():
pytest.skip(f"Official reference missing: {OFFICIAL_REF_DIR}")
if str(OFFICIAL_REF_DIR) not in sys.path:
sys.path.insert(0, str(OFFICIAL_REF_DIR))
def _install_xfuser_stub() -> None:
if "xfuser" in sys.modules:
return
xfuser = types.ModuleType("xfuser")
core = types.ModuleType("xfuser.core")
distributed = types.ModuleType("xfuser.core.distributed")
long_ctx = types.ModuleType("xfuser.core.long_ctx_attention")
distributed.get_sequence_parallel_rank = lambda: 0
distributed.get_sequence_parallel_world_size = lambda: 1
distributed.get_sp_group = lambda: None
distributed.get_world_group = lambda: types.SimpleNamespace(local_rank=0, rank=0)
distributed.init_distributed_environment = lambda *args, **kwargs: None
distributed.initialize_model_parallel = lambda *args, **kwargs: None
distributed.model_parallel_is_initialized = lambda: False
class XFuserLongContextAttention:
def __call__(self, *args, **kwargs):
raise RuntimeError("xfuser stub cannot execute attention")
long_ctx.xFuserLongContextAttention = XFuserLongContextAttention
sys.modules.update({
"xfuser": xfuser,
"xfuser.core": core,
"xfuser.core.distributed": distributed,
"xfuser.core.long_ctx_attention": long_ctx,
})
def _map_residual_subkey(prefix: str, sub: str) -> str | None:
if sub == "residual.0.gamma":
return f"{prefix}.norm1.gamma"
match = re.match(r"^residual\.2\.(weight|bias)$", sub)
if match:
return f"{prefix}.conv1.{match.group(1)}"
if sub == "residual.3.gamma":
return f"{prefix}.norm2.gamma"
match = re.match(r"^residual\.6\.(weight|bias)$", sub)
if match:
return f"{prefix}.conv2.{match.group(1)}"
match = re.match(r"^shortcut\.(weight|bias)$", sub)
if match:
return f"{prefix}.conv_shortcut.{match.group(1)}"
return None
def _map_attention_subkey(prefix: str, sub: str) -> str | None:
if sub == "norm.gamma":
return f"{prefix}.norm.gamma"
match = re.match(r"^(to_qkv|proj)\.(weight|bias)$", sub)
if match:
return f"{prefix}.{match.group(1)}.{match.group(2)}"
return None
def _map_resample_subkey(prefix: str, sub: str) -> str | None:
match = re.match(r"^resample\.1\.(weight|bias)$", sub)
if match:
return f"{prefix}.resample.1.{match.group(1)}"
match = re.match(r"^time_conv\.(weight|bias)$", sub)
if match:
return f"{prefix}.time_conv.{match.group(1)}"
return None
def _map_dreamx_raw_vae_key(key: str) -> str | None:
match = re.match(r"^(conv1|conv2)\.(weight|bias)$", key)
if match:
prefix = "quant_conv" if match.group(1) == "conv1" else "post_quant_conv"
return f"{prefix}.{match.group(2)}"
match = re.match(r"^(encoder|decoder)\.conv1\.(weight|bias)$", key)
if match:
return f"{match.group(1)}.conv_in.{match.group(2)}"
match = re.match(r"^(encoder|decoder)\.head\.0\.gamma$", key)
if match:
return f"{match.group(1)}.norm_out.gamma"
match = re.match(r"^(encoder|decoder)\.head\.2\.(weight|bias)$", key)
if match:
return f"{match.group(1)}.conv_out.{match.group(2)}"
match = re.match(r"^(encoder|decoder)\.middle\.0\.(.*)$", key)
if match:
return _map_residual_subkey(f"{match.group(1)}.mid_block.resnets.0", match.group(2))
match = re.match(r"^(encoder|decoder)\.middle\.1\.(.*)$", key)
if match:
return _map_attention_subkey(f"{match.group(1)}.mid_block.attentions.0", match.group(2))
match = re.match(r"^(encoder|decoder)\.middle\.2\.(.*)$", key)
if match:
return _map_residual_subkey(f"{match.group(1)}.mid_block.resnets.1", match.group(2))
match = re.match(r"^encoder\.downsamples\.(\d+)\.downsamples\.(\d+)\.(.*)$", key)
if match:
stage = int(match.group(1))
block = int(match.group(2))
sub = match.group(3)
if block in (0, 1):
return _map_residual_subkey(f"encoder.down_blocks.{stage}.resnets.{block}", sub)
if block == 2:
return _map_resample_subkey(f"encoder.down_blocks.{stage}.downsampler", sub)
return None
match = re.match(r"^decoder\.upsamples\.(\d+)\.upsamples\.(\d+)\.(.*)$", key)
if match:
stage = int(match.group(1))
block = int(match.group(2))
sub = match.group(3)
if block in (0, 1, 2):
return _map_residual_subkey(f"decoder.up_blocks.{stage}.resnets.{block}", sub)
if block == 3:
return _map_resample_subkey(f"decoder.up_blocks.{stage}.upsampler", sub)
return None
return None
def _vae_kwargs():
config = OmegaConf.load(OFFICIAL_REF_DIR / "configs" / "wan2.2" / "wan_ti2v_5b.yaml")
return OmegaConf.to_container(config["vae_kwargs"])
def _load_official_vae(device, dtype):
_add_official_to_path()
_install_xfuser_stub()
vae_path = WAN_BASE_DIR / "Wan2.2_VAE.pth"
if not vae_path.exists():
pytest.skip(f"Wan2.2 base VAE weights missing: {vae_path}")
try:
from models import AutoencoderKLWan3_8
except Exception as exc: # noqa: BLE001
pytest.skip(f"Cannot import official DreamX VAE: {exc}")
model = AutoencoderKLWan3_8.from_pretrained(str(vae_path), additional_kwargs=_vae_kwargs())
return model.to(device=device, dtype=dtype).eval()
def _load_fastvideo_vae(device, dtype):
vae_path = WAN_BASE_DIR / "Wan2.2_VAE.pth"
if not vae_path.exists():
pytest.skip(f"Wan2.2 raw VAE weights missing: {vae_path}")
config = make_dreamx_world_5b_cam_vae_config()
config.load_encoder = True
config.load_decoder = True
model = AutoencoderKLWan(config).to(device=device, dtype=dtype)
raw_state = torch.load(str(vae_path), map_location="cpu", weights_only=True)
mapped_state = {}
for key, value in raw_state.items():
mapped_key = _map_dreamx_raw_vae_key(key)
if mapped_key is None:
raise AssertionError(f"Unmapped DreamX raw VAE key: {key}")
mapped_state[mapped_key] = value
model.load_state_dict(mapped_state, strict=True)
return model.eval()
def _normalize_fastvideo_vae_latent(latent: torch.Tensor) -> torch.Tensor:
config = make_dreamx_world_5b_cam_vae_config()
mean = torch.tensor(config.latents_mean, device=latent.device, dtype=latent.dtype).view(1, -1, 1, 1, 1)
std = torch.tensor(config.latents_std, device=latent.device, dtype=latent.dtype).view(1, -1, 1, 1, 1)
return (latent - mean) / std
def test_dreamx_world_vae_config_matches_wan22_shape():
config = make_dreamx_world_5b_cam_vae_config()
assert config.z_dim == 48
assert config.in_channels == 12
assert config.out_channels == 12
assert config.base_dim == 160
assert config.decoder_base_dim == 256
assert config.scale_factor_temporal == 4
assert config.scale_factor_spatial == 16
assert config.patch_size == 2
assert config.is_residual is True
assert config.clip_output is False
assert len(config.latents_mean) == 48
assert len(config.latents_std) == 48
@pytest.mark.skipif(not torch.cuda.is_available(), reason="CUDA required for VAE parity.")
def test_dreamx_world_vae_encode_parity_scaffold():
device = torch.device("cuda:0")
dtype = torch.bfloat16
official = _load_official_vae(device, dtype)
fastvideo = _load_fastvideo_vae(device, dtype)
torch.manual_seed(123)
video = torch.randn(1, 3, 5, 64, 64, device=device, dtype=dtype).clamp(-1, 1)
with torch.inference_mode():
official_latent = official.encode(video).latent_dist.mean.float().cpu()
fastvideo_latent = _normalize_fastvideo_vae_latent(fastvideo.encode(video).mean).float().cpu()
assert official_latent.shape == fastvideo_latent.shape
assert_close(fastvideo_latent, official_latent, atol=5e-2, rtol=5e-2)
@@ -0,0 +1,110 @@
# SPDX-License-Identifier: Apache-2.0
"""DreamX-World pipeline parity checks.
This test compares the FastVideo pipeline's DreamX-specific conditioning and
single-step scheduler path against an explicit hand-rolled pass using the same
loaded modules. It is intentionally local and deterministic: component parity
against the official DreamX repository lives in ``tests/local_tests/dreamx_world``.
"""
from __future__ import annotations
import gc
import os
from pathlib import Path
from typing import Any, cast
import pytest
import torch
from torch.testing import assert_close
os.environ.setdefault("FASTVIDEO_ATTENTION_BACKEND", "TORCH_SDPA")
# Local converted dir or HF repo id; the loader downloads hub ids itself.
MODEL_DIR = os.getenv("DREAMX_WORLD_MODEL_DIR", "FastVideo/DreamX-World-5B-Cam-Diffusers")
pytestmark = pytest.mark.skipif(
not torch.cuda.is_available(),
reason="DreamX-World pipeline parity requires CUDA",
)
def _run_worker_forward_batch(worker_wrapper: Any, request_kwargs: dict[str, Any]) -> torch.Tensor:
from fastvideo.api.sampling_param import SamplingParam
from fastvideo.pipelines.pipeline_batch_info import ForwardBatch
from fastvideo.utils import shallow_asdict
fastvideo_args = worker_wrapper.worker.fastvideo_args
sampling_param = SamplingParam.from_pretrained(fastvideo_args.model_path)
sampling_param.update({
key: value
for key, value in request_kwargs.items()
if key not in {"prompt", "output_path"}
})
sampling_param.prompt = request_kwargs["prompt"]
latents_size = [
(sampling_param.num_frames - 1) // 4 + 1,
sampling_param.height // 8,
sampling_param.width // 8,
]
n_tokens = latents_size[0] * latents_size[1] * latents_size[2]
batch = ForwardBatch(
**shallow_asdict(sampling_param),
eta=0.0,
n_tokens=n_tokens,
VSA_sparsity=fastvideo_args.VSA_sparsity,
)
output_batch = worker_wrapper.worker.pipeline.forward(batch, fastvideo_args)
assert output_batch.output is not None
return output_batch.output.detach().cpu()
def _close_generator(generator: Any) -> None:
generator.shutdown()
gc.collect()
torch.cuda.empty_cache()
def test_dreamx_world_one_step_pipeline_latent_matches_manual_pass() -> None:
from fastvideo import VideoGenerator
common_kwargs = dict(
prompt="a quiet road through a futuristic city at sunrise",
output_path="outputs_video/dreamx_world_parity",
save_video=False,
return_frames=True,
height=64,
width=64,
num_frames=9,
num_inference_steps=1,
guidance_scale=1.0,
action_list=["w"],
action_speed_list=[2.0],
seed=123,
)
generator = VideoGenerator.from_pretrained(
MODEL_DIR,
num_gpus=1,
use_fsdp_inference=False,
dit_cpu_offload=False,
vae_cpu_offload=True,
text_encoder_cpu_offload=True,
pin_cpu_memory=False,
output_type="latent",
override_pipeline_cls_name="DreamXWorldPipeline",
)
try:
result = cast(dict[str, Any], generator.generate_video(**common_kwargs))
pipeline_latents = cast(torch.Tensor, result["samples"]).detach().cpu()
manual_latents = generator.executor.collective_rpc(
_run_worker_forward_batch,
kwargs={"request_kwargs": common_kwargs},
)[0]
finally:
_close_generator(generator)
assert_close(pipeline_latents, manual_latents, atol=0.0, rtol=0.0)
@@ -0,0 +1,157 @@
# SPDX-License-Identifier: Apache-2.0
"""Smoke tests for the DreamX-World-5B-Cam pipeline."""
from __future__ import annotations
import os
from pathlib import Path
from typing import Any, cast
import pytest
import torch
from PIL import Image, ImageDraw
os.environ.setdefault("FASTVIDEO_ATTENTION_BACKEND", "TORCH_SDPA")
# Local converted dir or HF repo id; the loader downloads hub ids itself.
MODEL_DIR = os.getenv("DREAMX_WORLD_MODEL_DIR", "FastVideo/DreamX-World-5B-Cam-Diffusers")
def _write_smoke_image(path: Path) -> None:
image = Image.new("RGB", (96, 96), color=(42, 76, 112))
draw = ImageDraw.Draw(image)
draw.rectangle((0, 56, 96, 96), fill=(36, 44, 52))
draw.polygon([(0, 56), (48, 30), (96, 56)], fill=(120, 142, 158))
draw.rectangle((34, 42, 62, 72), fill=(178, 198, 212))
image.save(path)
pytestmark = pytest.mark.skipif(
not torch.cuda.is_available(),
reason="DreamX-World pipeline smoke requires CUDA",
)
def test_dreamx_world_typed_surface_preflight() -> None:
import fastvideo.registry as registry
from fastvideo.api.presets import get_preset, get_presets_for_family
from fastvideo.configs.pipelines.dreamx_world import DreamXWorld5BCamPipelineConfig
from fastvideo.fastvideo_args import WorkloadType
from fastvideo.pipelines.basic.dreamx_world.dreamx_world_pipeline import (
DreamXWorldPipeline,
EntryClass,
)
assert DreamXWorldPipeline.__name__ == "DreamXWorldPipeline"
assert EntryClass is DreamXWorldPipeline
assert DreamXWorldPipeline._required_config_modules == [
"text_encoder",
"tokenizer",
"vae",
"transformer",
"scheduler",
]
default_preset, model_family = registry.get_preset_selection(
"GD-ML/DreamX-World-5B-Cam"
)
assert model_family == "dreamx_world"
assert default_preset == "dreamx_world_5b_cam"
info = registry.get_model_info(
"GD-ML/DreamX-World-5B-Cam",
workload_type=WorkloadType.I2V,
override_pipeline_cls_name="DreamXWorldPipeline",
)
assert info.pipeline_cls is DreamXWorldPipeline
assert info.pipeline_config_cls is DreamXWorld5BCamPipelineConfig
names = {p.name for p in get_presets_for_family("dreamx_world")}
assert "dreamx_world_5b_cam" in names
preset = get_preset("dreamx_world_5b_cam", "dreamx_world")
assert preset.defaults["num_inference_steps"] == 30
assert preset.defaults["height"] == 480
assert preset.defaults["width"] == 832
assert preset.defaults["num_frames"] == 161
assert preset.defaults["guidance_scale"] == 5.0
cfg = DreamXWorld5BCamPipelineConfig()
assert cfg.flow_shift == 3.0
assert cfg.ti2v_task is True
assert cfg.expand_timesteps is True
assert cfg.dit_config.arch_config.add_control_adapter is True
assert cfg.dit_config.arch_config.cam_method == "prope"
def test_dreamx_world_camera_stage_writes_y_camera() -> None:
from types import SimpleNamespace
from fastvideo.pipelines.basic.dreamx_world.stages import (
DREAMX_Y_CAMERA_KEY,
DreamXWorldCameraConditioningStage,
)
from fastvideo.pipelines.pipeline_batch_info import ForwardBatch
batch = ForwardBatch(
data_type="video",
prompt="camera smoke",
latents=torch.zeros(1, 48, 3, 8, 8, dtype=torch.bfloat16, device="cuda"),
num_frames=9,
height=64,
width=64,
action_list=["w", "d"],
action_speed_list=[2.0, 1.0],
)
out = DreamXWorldCameraConditioningStage().forward(batch, cast(Any, SimpleNamespace()))
y_camera = out.extra[DREAMX_Y_CAMERA_KEY]
assert set(y_camera) == {"viewmats", "K"}
assert y_camera["viewmats"].shape == (1, 3, 4, 4)
assert y_camera["K"].shape == (1, 3, 3, 3)
assert y_camera["viewmats"].device.type == "cuda"
assert y_camera["viewmats"].dtype == torch.bfloat16
def test_dreamx_world_pipeline_load_generate_latent_smoke(tmp_path: Path) -> None:
from fastvideo import VideoGenerator
image_path = tmp_path / "dreamx_world_smoke_input.png"
_write_smoke_image(image_path)
generator = VideoGenerator.from_pretrained(
MODEL_DIR,
num_gpus=1,
use_fsdp_inference=False,
dit_cpu_offload=False,
vae_cpu_offload=True,
text_encoder_cpu_offload=True,
pin_cpu_memory=False,
output_type="latent",
override_pipeline_cls_name="DreamXWorldPipeline",
)
try:
result = generator.generate_video(
prompt="a quiet road through a futuristic city at sunrise",
output_path="outputs_video/dreamx_world_smoke",
save_video=False,
return_frames=True,
height=64,
width=64,
num_frames=9,
num_inference_steps=1,
guidance_scale=1.0,
image_path=str(image_path),
action_list=["w"],
action_speed_list=[2.0],
seed=0,
)
finally:
generator.shutdown()
assert isinstance(result, dict)
samples = cast(dict[str, Any], result)["samples"]
assert torch.is_tensor(samples)
assert samples.ndim == 5
assert samples.shape[1] == 48
assert torch.isfinite(samples).all()