Compare commits

...
38 changed files with 5385 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 | `GD-ML/DreamX-World-5B-Cam` | 480P | ⭕ | ⭕ | ⭕ | ⭕ | ⭕ |
| DreamX-World 5B AR | `GD-ML/DreamX-World-5B` | 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,64 @@
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", "GD-ML/DreamX-World-5B-Cam")
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,
override_pipeline_cls_name="DreamXWorldPipeline",
)
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
@@ -190,6 +190,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):
@@ -241,7 +254,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)
@@ -494,6 +511,7 @@ class DenoisingStage(PipelineStage):
**pos_cond_kwargs,
**action_kwargs,
**camera_kwargs,
**dreamx_camera_kwargs,
**timesteps_r_kwarg,
**flux2_id_kwargs,
)
@@ -536,6 +554,7 @@ class DenoisingStage(PipelineStage):
**neg_cond_kwargs,
**action_kwargs,
**camera_kwargs,
**dreamx_camera_kwargs,
**timesteps_r_kwarg,
**flux2_id_kwargs,
)
+38
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,40 @@ 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=[
"GD-ML/DreamX-World-5B-Cam",
],
model_detectors=[
# 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=[
"GD-ML/DreamX-World-5B",
],
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 +986,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 +1021,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,178 @@
# 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"))
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():
pytest.skip(f"Converted AR transformer missing: {transformer_dir}")
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,331 @@
# 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"))
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():
pytest.skip(f"Converted DreamX transformer missing: {transformer_dir}")
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,112 @@
# 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")
MODEL_DIR = Path(os.getenv("DREAMX_WORLD_MODEL_DIR", "converted_weights/dreamx_world"))
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:
if not MODEL_DIR.exists():
pytest.fail(f"DreamX-World converted model directory is missing: {MODEL_DIR}")
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(
str(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,159 @@
# 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")
MODEL_DIR = Path(os.getenv("DREAMX_WORLD_MODEL_DIR", "converted_weights/dreamx_world"))
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:
if not MODEL_DIR.exists():
pytest.fail(f"DreamX-World converted model directory is missing: {MODEL_DIR}")
from fastvideo import VideoGenerator
image_path = tmp_path / "dreamx_world_smoke_input.png"
_write_smoke_image(image_path)
generator = VideoGenerator.from_pretrained(
str(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()