Compare commits
7
Commits
| Author | SHA1 | Date | |
|---|---|---|---|
|
|
f0e727f5ee | ||
|
|
de6a60ac90 | ||
|
|
1b4fdc2c41 | ||
|
|
385cd65abb | ||
|
|
156d86b70c | ||
|
|
bab79fb5f6 | ||
|
|
073a9e78f2 |
@@ -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
|
||||
|
||||
@@ -58,6 +58,8 @@ pipeline initialization and sampling.
|
||||
| FastWan2.1 T2V 1.3B | `FastVideo/FastWan2.1-T2V-1.3B-Diffusers` | 480P | ⭕ | ⭕ | ⭕ | ✅ | ⭕ |
|
||||
| FastWan2.2 TI2V 5B Full Attn* | `FastVideo/FastWan2.2-TI2V-5B-FullAttn-Diffusers` | 720P | ⭕ | ⭕ | ⭕ | ✅ | ⭕ |
|
||||
| Wan2.2 TI2V 5B | `Wan-AI/Wan2.2-TI2V-5B-Diffusers` | 720P | ⭕ | ⭕ | ✅ | ⭕ | ⭕ |
|
||||
| DreamX-World 5B Cam | `FastVideo/DreamX-World-5B-Cam-Diffusers` | 480P | ⭕ | ⭕ | ⭕ | ⭕ | ⭕ |
|
||||
| DreamX-World 5B AR | `FastVideo/DreamX-World-5B-Diffusers` | 704px1280p | ⭕ | ⭕ | ⭕ | ⭕ | ⭕ |
|
||||
| Lucy Edit Dev 5B*** | `decart-ai/Lucy-Edit-Dev` | 480P | ⭕ | ⭕ | ⭕ | ⭕ | ⭕ |
|
||||
| Wan2.2 T2V A14B | `Wan-AI/Wan2.2-T2V-A14B-Diffusers` | 480P<br>720P | ❌ | ❌ | ✅ | ⭕ | ⭕ |
|
||||
| Wan2.2 I2V A14B | `Wan-AI/Wan2.2-I2V-A14B-Diffusers` | 480P<br>720P | ❌ | ❌ | ✅ | ⭕ | ⭕ |
|
||||
|
||||
@@ -0,0 +1,70 @@
|
||||
"""DreamX-World-5B-Cam camera-controlled video generation.
|
||||
|
||||
Uses the pre-converted Diffusers checkpoint FastVideo/DreamX-World-5B-Cam-Diffusers.
|
||||
To convert the raw GD-ML/DreamX-World-5B-Cam checkpoint yourself, see
|
||||
scripts/checkpoint_conversion/dreamx_world_to_diffusers.py.
|
||||
"""
|
||||
|
||||
import os
|
||||
|
||||
from fastvideo import VideoGenerator
|
||||
|
||||
|
||||
OUTPUT_PATH = os.getenv("DREAMX_WORLD_OUTPUT_PATH", "video_samples_dreamx_world")
|
||||
|
||||
|
||||
def _env_int(name: str, default: int) -> int:
|
||||
return int(os.getenv(name, str(default)))
|
||||
|
||||
|
||||
def _env_float(name: str, default: float) -> float:
|
||||
return float(os.getenv(name, str(default)))
|
||||
|
||||
|
||||
def main():
|
||||
model_name = os.getenv("DREAMX_WORLD_MODEL_DIR", "FastVideo/DreamX-World-5B-Cam-Diffusers")
|
||||
generator = VideoGenerator.from_pretrained(
|
||||
model_name,
|
||||
num_gpus=1,
|
||||
use_fsdp_inference=False,
|
||||
dit_cpu_offload=False,
|
||||
vae_cpu_offload=True,
|
||||
text_encoder_cpu_offload=True,
|
||||
pin_cpu_memory=False,
|
||||
)
|
||||
|
||||
prompt = os.getenv(
|
||||
"DREAMX_WORLD_PROMPT",
|
||||
"A cinematic first-person drive through a futuristic coastal city at "
|
||||
"sunrise, reflective glass towers, clean streets, soft volumetric light.",
|
||||
)
|
||||
image_path = os.getenv(
|
||||
"DREAMX_WORLD_IMAGE_PATH",
|
||||
"https://huggingface.co/datasets/YiYiXu/testing-images/resolve/main/wan_i2v_input.JPG",
|
||||
)
|
||||
|
||||
kwargs = {
|
||||
"output_path": OUTPUT_PATH,
|
||||
"save_video": os.getenv("DREAMX_WORLD_SAVE_VIDEO", "1") != "0",
|
||||
"height": _env_int("DREAMX_WORLD_HEIGHT", 480),
|
||||
"width": _env_int("DREAMX_WORLD_WIDTH", 832),
|
||||
"num_frames": _env_int("DREAMX_WORLD_NUM_FRAMES", 161),
|
||||
"num_inference_steps": _env_int("DREAMX_WORLD_STEPS", 30),
|
||||
"guidance_scale": _env_float("DREAMX_WORLD_GUIDANCE", 5.0),
|
||||
"action_list": os.getenv("DREAMX_WORLD_ACTIONS", "w,d,w").split(","),
|
||||
"action_speed_list": [
|
||||
float(value)
|
||||
for value in os.getenv("DREAMX_WORLD_ACTION_SPEEDS", "4.0,2.0,4.0").split(",")
|
||||
],
|
||||
}
|
||||
if image_path:
|
||||
kwargs["image_path"] = image_path
|
||||
|
||||
try:
|
||||
generator.generate_video(prompt, **kwargs)
|
||||
finally:
|
||||
generator.shutdown()
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
main()
|
||||
@@ -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"
|
||||
@@ -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"
|
||||
]
|
||||
|
||||
@@ -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
|
||||
@@ -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
|
||||
@@ -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
|
||||
@@ -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
|
||||
@@ -191,6 +191,19 @@ class DenoisingStage(PipelineStage):
|
||||
},
|
||||
)
|
||||
|
||||
dreamx_y_camera = batch.extra.get("dreamx_y_camera", batch.extra.get("y_camera"))
|
||||
if isinstance(dreamx_y_camera, dict):
|
||||
dreamx_y_camera = {
|
||||
key: value.to(device=local_device, dtype=target_dtype) if torch.is_tensor(value) else value
|
||||
for key, value in dreamx_y_camera.items()
|
||||
}
|
||||
dreamx_camera_kwargs = self.prepare_extra_func_kwargs(
|
||||
self.transformer.forward,
|
||||
{
|
||||
"y_camera": dreamx_y_camera,
|
||||
},
|
||||
)
|
||||
|
||||
for key in ("flux2_txt_ids", "flux2_img_ids"):
|
||||
value = batch.extra.get(key)
|
||||
if torch.is_tensor(value):
|
||||
@@ -242,7 +255,11 @@ class DenoisingStage(PipelineStage):
|
||||
# the image latent instead of appending along the channel dim
|
||||
assert batch.image_latent is None, "TI2V task should not have image latents"
|
||||
assert self.vae is not None, "VAE is not provided for TI2V task"
|
||||
vae_device = next(self.vae.parameters()).device
|
||||
self.vae = self.vae.to(local_device)
|
||||
z = self.vae.encode(batch.pil_image).mean.float()
|
||||
if getattr(fastvideo_args, "vae_cpu_offload", False):
|
||||
self.vae = self.vae.to(vae_device)
|
||||
if (hasattr(self.vae, "shift_factor") and self.vae.shift_factor is not None):
|
||||
if isinstance(self.vae.shift_factor, torch.Tensor):
|
||||
z -= self.vae.shift_factor.to(z.device, z.dtype)
|
||||
@@ -495,6 +512,7 @@ class DenoisingStage(PipelineStage):
|
||||
**pos_cond_kwargs,
|
||||
**action_kwargs,
|
||||
**camera_kwargs,
|
||||
**dreamx_camera_kwargs,
|
||||
**timesteps_r_kwarg,
|
||||
**flux2_id_kwargs,
|
||||
)
|
||||
@@ -537,6 +555,7 @@ class DenoisingStage(PipelineStage):
|
||||
**neg_cond_kwargs,
|
||||
**action_kwargs,
|
||||
**camera_kwargs,
|
||||
**dreamx_camera_kwargs,
|
||||
**timesteps_r_kwarg,
|
||||
**flux2_id_kwargs,
|
||||
)
|
||||
|
||||
@@ -20,6 +20,7 @@ from fastvideo.configs.pipelines.cosmos2_5 import (
|
||||
Cosmos25Config,
|
||||
Cosmos25_14BConfig,
|
||||
)
|
||||
from fastvideo.configs.pipelines.dreamx_world import DreamXWorld5BARPipelineConfig, DreamXWorld5BCamPipelineConfig
|
||||
from fastvideo.configs.pipelines.hunyuan import FastHunyuanConfig, HunyuanConfig
|
||||
from fastvideo.configs.pipelines.hunyuangamecraft import HunyuanGameCraftPipelineConfig
|
||||
from fastvideo.configs.pipelines.gen3c import Gen3CConfig
|
||||
@@ -773,6 +774,42 @@ def _register_configs() -> None:
|
||||
model_family="wan",
|
||||
default_preset="wan_2_2_ti2v_5b",
|
||||
)
|
||||
|
||||
register_configs(
|
||||
sampling_param_cls=None,
|
||||
pipeline_config_cls=DreamXWorld5BCamPipelineConfig,
|
||||
workload_types=(WorkloadType.I2V, ),
|
||||
hf_model_paths=[
|
||||
"FastVideo/DreamX-World-5B-Cam-Diffusers",
|
||||
],
|
||||
model_detectors=[
|
||||
# Pattern also catches the raw GD-ML/DreamX-World-5B-Cam id and
|
||||
# local converted dirs. Mutually exclusive with the AR detector
|
||||
# below: Cam requires an explicit "cam" marker so hyphenated AR
|
||||
# local paths (e.g. /ckpts/dreamx-world-5b-converted) don't
|
||||
# first-match here — detector resolution is first-match in
|
||||
# registration order.
|
||||
lambda path:
|
||||
("dreamx-world" in path.lower() and "cam" in path.lower()) or "dreamxworldpipeline" in path.lower()
|
||||
],
|
||||
model_family="dreamx_world",
|
||||
default_preset="dreamx_world_5b_cam",
|
||||
)
|
||||
register_configs(
|
||||
sampling_param_cls=None,
|
||||
pipeline_config_cls=DreamXWorld5BARPipelineConfig,
|
||||
workload_types=(WorkloadType.I2V, ),
|
||||
hf_model_paths=[
|
||||
"FastVideo/DreamX-World-5B-Diffusers",
|
||||
],
|
||||
model_detectors=[
|
||||
lambda path:
|
||||
("dreamx-world-5b" in path.lower() and "cam" not in path.lower()) or "dreamxworldarpipeline" in path.lower(
|
||||
)
|
||||
],
|
||||
model_family="dreamx_world",
|
||||
default_preset="dreamx_world_5b_ar",
|
||||
)
|
||||
register_configs(
|
||||
sampling_param_cls=None,
|
||||
pipeline_config_cls=FastWan2_2_TI2V_5B_Config,
|
||||
@@ -951,6 +988,8 @@ def _register_presets() -> None:
|
||||
from fastvideo.api.presets import register_preset
|
||||
from fastvideo.pipelines.basic.cosmos.presets import (
|
||||
ALL_PRESETS as COSMOS_PRESETS, )
|
||||
from fastvideo.pipelines.basic.dreamx_world.presets import (
|
||||
ALL_PRESETS as DREAMX_WORLD_PRESETS, )
|
||||
from fastvideo.pipelines.basic.gamecraft.presets import (
|
||||
ALL_PRESETS as GAMECRAFT_PRESETS, )
|
||||
from fastvideo.pipelines.basic.gen3c.presets import (
|
||||
@@ -984,6 +1023,7 @@ def _register_presets() -> None:
|
||||
|
||||
all_preset_groups = (
|
||||
COSMOS_PRESETS,
|
||||
DREAMX_WORLD_PRESETS,
|
||||
FLUX2_PRESETS,
|
||||
GAMECRAFT_PRESETS,
|
||||
GEN3C_PRESETS,
|
||||
|
||||
@@ -3,9 +3,9 @@
|
||||
from __future__ import annotations
|
||||
|
||||
import os
|
||||
from collections.abc import Iterator
|
||||
from contextlib import contextmanager
|
||||
from logging import Logger
|
||||
from typing import Iterator
|
||||
|
||||
from fastvideo import VideoGenerator
|
||||
from fastvideo.tests.ssim.reference_utils import (
|
||||
@@ -67,6 +67,12 @@ def _find_reference_video(reference_folder: str, prompt: str) -> str:
|
||||
raise FileNotFoundError("Reference video missing")
|
||||
|
||||
|
||||
def _remove_stale_generated_video(output_dir: str, output_video_name: str) -> None:
|
||||
stale_path = os.path.join(output_dir, output_video_name)
|
||||
if os.path.exists(stale_path):
|
||||
os.remove(stale_path)
|
||||
|
||||
|
||||
def _assert_similarity(
|
||||
*,
|
||||
logger: Logger,
|
||||
@@ -214,6 +220,7 @@ def run_text_to_video_similarity_test(
|
||||
)
|
||||
output_video_name = f"{prompt[:100].strip()}.mp4"
|
||||
os.makedirs(output_dir, exist_ok=True)
|
||||
_remove_stale_generated_video(output_dir, output_video_name)
|
||||
|
||||
params_map = select_ssim_params(
|
||||
default_params_map,
|
||||
@@ -289,6 +296,7 @@ def run_image_to_video_similarity_test(
|
||||
)
|
||||
output_video_name = f"{prompt[:100].strip()}.mp4"
|
||||
os.makedirs(output_dir, exist_ok=True)
|
||||
_remove_stale_generated_video(output_dir, output_video_name)
|
||||
|
||||
params_map = select_ssim_params(
|
||||
default_params_map,
|
||||
|
||||
@@ -0,0 +1,214 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
from __future__ import annotations
|
||||
|
||||
import os
|
||||
from pathlib import Path
|
||||
|
||||
import pytest
|
||||
from PIL import Image, ImageDraw
|
||||
|
||||
from fastvideo.logger import init_logger
|
||||
from fastvideo.tests.ssim.inference_similarity_utils import (
|
||||
resolve_inference_device_reference_folder,
|
||||
run_image_to_video_similarity_test,
|
||||
run_text_to_video_similarity_test,
|
||||
)
|
||||
|
||||
logger = init_logger(__name__)
|
||||
|
||||
REQUIRED_GPUS = 1
|
||||
|
||||
device_reference_folder = resolve_inference_device_reference_folder(logger)
|
||||
|
||||
_LOCAL_CONVERTED_MODEL = Path("converted_weights/dreamx_world")
|
||||
_MODEL_PATH = os.getenv(
|
||||
"DREAMX_WORLD_SSIM_MODEL_PATH",
|
||||
str(_LOCAL_CONVERTED_MODEL),
|
||||
)
|
||||
_LOCAL_AR_CANDIDATES = (
|
||||
Path("/tmp/converted_dreamx_world_ar"),
|
||||
Path("/root/data/dreamx_world_ar_converted"),
|
||||
)
|
||||
_DEFAULT_AR_MODEL_PATH = next(
|
||||
(str(path) for path in _LOCAL_AR_CANDIDATES if path.exists()),
|
||||
str(_LOCAL_AR_CANDIDATES[0]),
|
||||
)
|
||||
_AR_MODEL_PATH = os.getenv(
|
||||
"DREAMX_WORLD_AR_SSIM_MODEL_PATH",
|
||||
_DEFAULT_AR_MODEL_PATH,
|
||||
)
|
||||
|
||||
DREAMX_WORLD_PARAMS = {
|
||||
"num_gpus": 1,
|
||||
"model_path": _MODEL_PATH,
|
||||
"height": 64,
|
||||
"width": 64,
|
||||
"num_frames": 9,
|
||||
"num_inference_steps": 1,
|
||||
"guidance_scale": 1.0,
|
||||
"seed": 1024,
|
||||
"fps": 16,
|
||||
}
|
||||
|
||||
DREAMX_WORLD_FULL_QUALITY_PARAMS = {
|
||||
**DREAMX_WORLD_PARAMS,
|
||||
"height": 480,
|
||||
"width": 832,
|
||||
"num_frames": 161,
|
||||
"num_inference_steps": 30,
|
||||
"guidance_scale": 5.0,
|
||||
}
|
||||
|
||||
|
||||
DREAMX_WORLD_AR_PARAMS = {
|
||||
"num_gpus": 1,
|
||||
"model_path": _AR_MODEL_PATH,
|
||||
"height": 192,
|
||||
"width": 192,
|
||||
"num_frames": 81,
|
||||
"num_inference_steps": 4,
|
||||
"guidance_scale": 1.0,
|
||||
"seed": 2048,
|
||||
"fps": 16,
|
||||
}
|
||||
|
||||
DREAMX_WORLD_AR_FULL_QUALITY_PARAMS = {
|
||||
**DREAMX_WORLD_AR_PARAMS,
|
||||
"height": 704,
|
||||
"width": 1280,
|
||||
"num_frames": 1005,
|
||||
}
|
||||
|
||||
DREAMX_WORLD_MODEL_TO_PARAMS = {
|
||||
"DreamX-World-5B-Cam": DREAMX_WORLD_PARAMS,
|
||||
}
|
||||
|
||||
DREAMX_WORLD_AR_MODEL_TO_PARAMS = {
|
||||
"DreamX-World-5B": DREAMX_WORLD_AR_PARAMS,
|
||||
}
|
||||
|
||||
FULL_QUALITY_DREAMX_WORLD_MODEL_TO_PARAMS = {
|
||||
"DreamX-World-5B-Cam": DREAMX_WORLD_FULL_QUALITY_PARAMS,
|
||||
}
|
||||
|
||||
FULL_QUALITY_DREAMX_WORLD_AR_MODEL_TO_PARAMS = {
|
||||
"DreamX-World-5B": DREAMX_WORLD_AR_FULL_QUALITY_PARAMS,
|
||||
}
|
||||
|
||||
DREAMX_WORLD_TEST_CASES = [
|
||||
(
|
||||
"A cinematic first-person drive through a futuristic coastal city at sunrise, "
|
||||
"reflective glass towers, clean streets, soft volumetric light.",
|
||||
("w", "d", "w"),
|
||||
(4.0, 2.0, 4.0),
|
||||
),
|
||||
]
|
||||
|
||||
|
||||
DREAMX_WORLD_AR_TEST_CASES = [
|
||||
(
|
||||
"A long autonomous drive through a futuristic coastal city at sunrise, "
|
||||
"smooth forward camera motion, reflective glass towers, clean streets.",
|
||||
("w", "d", "w", "a"),
|
||||
(2.0, 1.0, 2.0, 1.0),
|
||||
),
|
||||
]
|
||||
|
||||
|
||||
def _write_deterministic_reference_image(path: Path) -> None:
|
||||
path.parent.mkdir(parents=True, exist_ok=True)
|
||||
image = Image.new("RGB", (96, 96), color=(42, 76, 112))
|
||||
draw = ImageDraw.Draw(image)
|
||||
draw.rectangle((0, 56, 96, 96), fill=(36, 44, 52))
|
||||
draw.polygon([(0, 56), (48, 30), (96, 56)], fill=(120, 142, 158))
|
||||
draw.rectangle((34, 42, 62, 72), fill=(178, 198, 212))
|
||||
draw.line((0, 80, 96, 66), fill=(238, 209, 124), width=3)
|
||||
image.save(path)
|
||||
|
||||
|
||||
@pytest.mark.parametrize(("prompt", "action_list", "action_speed_list"), DREAMX_WORLD_TEST_CASES)
|
||||
@pytest.mark.parametrize("attention_backend_name", ["TORCH_SDPA"])
|
||||
@pytest.mark.parametrize("model_id", list(DREAMX_WORLD_MODEL_TO_PARAMS.keys()))
|
||||
def test_dreamx_world_inference_similarity(
|
||||
prompt: str,
|
||||
action_list: tuple[str, ...],
|
||||
action_speed_list: tuple[float, ...],
|
||||
attention_backend_name: str,
|
||||
model_id: str,
|
||||
tmp_path: Path,
|
||||
) -> None:
|
||||
model_path = Path(str(DREAMX_WORLD_MODEL_TO_PARAMS[model_id]["model_path"]))
|
||||
if not model_path.exists():
|
||||
pytest.skip(
|
||||
f"DreamX-World converted model path is missing: {model_path}. "
|
||||
"Set DREAMX_WORLD_SSIM_MODEL_PATH to a FastVideo-loadable converted root."
|
||||
)
|
||||
|
||||
image_path = tmp_path / "dreamx_world_ssim_input.png"
|
||||
_write_deterministic_reference_image(image_path)
|
||||
|
||||
run_image_to_video_similarity_test(
|
||||
logger=logger,
|
||||
script_dir=os.path.dirname(os.path.abspath(__file__)),
|
||||
device_reference_folder=device_reference_folder,
|
||||
prompt=prompt,
|
||||
image_path=str(image_path),
|
||||
attention_backend_name=attention_backend_name,
|
||||
model_id=model_id,
|
||||
default_params_map=DREAMX_WORLD_MODEL_TO_PARAMS,
|
||||
full_quality_params_map=FULL_QUALITY_DREAMX_WORLD_MODEL_TO_PARAMS,
|
||||
min_acceptable_ssim=0.98,
|
||||
init_kwargs_override={
|
||||
"use_fsdp_inference": False,
|
||||
"dit_cpu_offload": False,
|
||||
"vae_cpu_offload": True,
|
||||
"text_encoder_cpu_offload": True,
|
||||
"pin_cpu_memory": False,
|
||||
"override_pipeline_cls_name": "DreamXWorldPipeline",
|
||||
},
|
||||
generation_kwargs_override={
|
||||
"action_list": list(action_list),
|
||||
"action_speed_list": list(action_speed_list),
|
||||
},
|
||||
)
|
||||
|
||||
@pytest.mark.parametrize(("prompt", "action_list", "action_speed_list"), DREAMX_WORLD_AR_TEST_CASES)
|
||||
@pytest.mark.parametrize("attention_backend_name", ["TORCH_SDPA"])
|
||||
@pytest.mark.parametrize("model_id", list(DREAMX_WORLD_AR_MODEL_TO_PARAMS.keys()))
|
||||
def test_dreamx_world_ar_inference_similarity(
|
||||
prompt: str,
|
||||
action_list: tuple[str, ...],
|
||||
action_speed_list: tuple[float, ...],
|
||||
attention_backend_name: str,
|
||||
model_id: str,
|
||||
) -> None:
|
||||
model_path = Path(str(DREAMX_WORLD_AR_MODEL_TO_PARAMS[model_id]["model_path"]))
|
||||
if not model_path.exists():
|
||||
pytest.skip(
|
||||
f"DreamX-World AR converted model path is missing: {model_path}. "
|
||||
"Set DREAMX_WORLD_AR_SSIM_MODEL_PATH to a FastVideo-loadable converted root."
|
||||
)
|
||||
|
||||
run_text_to_video_similarity_test(
|
||||
logger=logger,
|
||||
script_dir=os.path.dirname(os.path.abspath(__file__)),
|
||||
device_reference_folder=device_reference_folder,
|
||||
prompt=prompt,
|
||||
attention_backend_name=attention_backend_name,
|
||||
model_id=model_id,
|
||||
default_params_map=DREAMX_WORLD_AR_MODEL_TO_PARAMS,
|
||||
full_quality_params_map=FULL_QUALITY_DREAMX_WORLD_AR_MODEL_TO_PARAMS,
|
||||
min_acceptable_ssim=0.98,
|
||||
init_kwargs_override={
|
||||
"use_fsdp_inference": False,
|
||||
"dit_cpu_offload": False,
|
||||
"vae_cpu_offload": True,
|
||||
"text_encoder_cpu_offload": True,
|
||||
"pin_cpu_memory": False,
|
||||
"override_pipeline_cls_name": "DreamXWorldARPipeline",
|
||||
},
|
||||
generation_kwargs_override={
|
||||
"action_list": list(action_list),
|
||||
"action_speed_list": list(action_speed_list),
|
||||
},
|
||||
)
|
||||
@@ -0,0 +1,124 @@
|
||||
#!/usr/bin/env python3
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
"""Convert DreamX-World-5B autoregressive weights to FastVideo layout.
|
||||
|
||||
The HF repository stores one raw official ``model.safetensors`` whose keys match
|
||||
FastVideo's native ``DreamXWorldARTransformer3DModel``. The converter writes a
|
||||
Diffusers-like root with ``transformer/config.json`` and reusable Wan2.2
|
||||
components. Use ``--symlink-transformer`` locally to avoid duplicating the 21GB
|
||||
AR tensor file.
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import argparse
|
||||
import json
|
||||
import shutil
|
||||
from pathlib import Path
|
||||
|
||||
TRANSFORMER_CONFIG: dict[str, object] = {
|
||||
"_class_name": "DreamXWorldARTransformer3DModel",
|
||||
"model_type": "ti2v",
|
||||
"patch_size": [1, 2, 2],
|
||||
"text_len": 512,
|
||||
"num_attention_heads": 24,
|
||||
"attention_head_dim": 128,
|
||||
"in_channels": 48,
|
||||
"out_channels": 48,
|
||||
"text_dim": 4096,
|
||||
"freq_dim": 256,
|
||||
"ffn_dim": 14336,
|
||||
"num_layers": 30,
|
||||
"local_attn_size": 12,
|
||||
"sink_size": 3,
|
||||
"cross_attn_norm": True,
|
||||
"qk_norm": True,
|
||||
"eps": 1e-6,
|
||||
"add_control_adapter": True,
|
||||
"cam_method": "prope",
|
||||
"attn_compress": 4,
|
||||
"cam_self_attn_layers": list(range(30)),
|
||||
"num_frames_per_block": 3,
|
||||
}
|
||||
|
||||
MODEL_INDEX: dict[str, object] = {
|
||||
"_class_name": "DreamXWorldARPipeline",
|
||||
"_diffusers_version": "0.31.0",
|
||||
"scheduler": ["diffusers", "FlowMatchEulerDiscreteScheduler"],
|
||||
"text_encoder": ["transformers", "UMT5EncoderModel"],
|
||||
"tokenizer": ["transformers", "AutoTokenizer"],
|
||||
"transformer": ["diffusers", "DreamXWorldARTransformer3DModel"],
|
||||
"vae": ["diffusers", "AutoencoderKLWan"],
|
||||
}
|
||||
|
||||
REUSED_COMPONENTS = ("scheduler", "text_encoder", "tokenizer", "vae")
|
||||
|
||||
|
||||
def _source_safetensors(source: Path) -> Path:
|
||||
if source.is_file():
|
||||
return source
|
||||
path = source / "model.safetensors"
|
||||
if not path.exists():
|
||||
raise FileNotFoundError(f"Missing AR model.safetensors under {source}")
|
||||
return path
|
||||
|
||||
|
||||
def convert_transformer(source: Path, output: Path, symlink_transformer: bool) -> None:
|
||||
src = _source_safetensors(source)
|
||||
transformer_dir = output / "transformer"
|
||||
transformer_dir.mkdir(parents=True, exist_ok=True)
|
||||
(transformer_dir / "config.json").write_text(json.dumps(TRANSFORMER_CONFIG, indent=2) + "\n")
|
||||
dst = transformer_dir / "model.safetensors"
|
||||
if dst.is_symlink() and not dst.exists():
|
||||
dst.unlink() # dangling symlink: replace so re-conversion self-heals
|
||||
if dst.exists():
|
||||
return
|
||||
if symlink_transformer:
|
||||
dst.symlink_to(src.resolve())
|
||||
else:
|
||||
shutil.copy2(src, dst)
|
||||
|
||||
|
||||
def _copy_or_link_component(component: str, component_source: Path, output: Path, symlink: bool) -> None:
|
||||
src = component_source / component
|
||||
dst = output / component
|
||||
if not src.exists():
|
||||
raise FileNotFoundError(f"Missing reused component source: {src}")
|
||||
if dst.is_symlink() and not dst.exists():
|
||||
dst.unlink() # dangling symlink: replace so re-conversion self-heals
|
||||
if dst.exists():
|
||||
return
|
||||
if symlink:
|
||||
dst.symlink_to(src.resolve(), target_is_directory=src.is_dir())
|
||||
elif src.is_dir():
|
||||
shutil.copytree(src, dst)
|
||||
else:
|
||||
shutil.copy2(src, dst)
|
||||
|
||||
|
||||
def write_model_index(output: Path, component_source: Path | None, symlink_components: bool) -> None:
|
||||
if component_source is not None:
|
||||
for component in REUSED_COMPONENTS:
|
||||
_copy_or_link_component(component, component_source, output, symlink_components)
|
||||
missing = [component for component in REUSED_COMPONENTS if not (output / component).exists()]
|
||||
if missing:
|
||||
print("skipping model_index.json; missing reused components: " + ", ".join(missing))
|
||||
return
|
||||
(output / "model_index.json").write_text(json.dumps(MODEL_INDEX, indent=2) + "\n")
|
||||
|
||||
|
||||
def main() -> None:
|
||||
parser = argparse.ArgumentParser(description=__doc__)
|
||||
parser.add_argument("--source", type=Path, required=True)
|
||||
parser.add_argument("--output", type=Path, required=True)
|
||||
parser.add_argument("--component-source", type=Path)
|
||||
parser.add_argument("--symlink-components", action="store_true")
|
||||
parser.add_argument("--symlink-transformer", action="store_true")
|
||||
args = parser.parse_args()
|
||||
|
||||
convert_transformer(args.source, args.output, args.symlink_transformer)
|
||||
write_model_index(args.output, args.component_source, args.symlink_components)
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
main()
|
||||
@@ -0,0 +1,221 @@
|
||||
#!/usr/bin/env python3
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
"""Convert DreamX-World-5B-Cam raw transformer weights to FastVideo-loadable format.
|
||||
|
||||
The GD-ML/DreamX-World-5B-Cam repository stores the transformer as raw
|
||||
DreamX/Wan official shards. FastVideo's TransformerLoader expects a Diffusers-like
|
||||
transformer folder with a config.json and safetensors whose keys can be mapped by
|
||||
WanVideoConfig.param_names_mapping. This script performs the raw official ->
|
||||
Diffusers-like key rename and writes the DreamX 5B-Cam transformer config.
|
||||
|
||||
Example:
|
||||
python scripts/checkpoint_conversion/dreamx_world_to_diffusers.py \
|
||||
--source official_weights/dreamx_world \
|
||||
--output converted_weights/dreamx_world
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import argparse
|
||||
import json
|
||||
import re
|
||||
import shutil
|
||||
from collections import OrderedDict
|
||||
from pathlib import Path
|
||||
|
||||
import torch
|
||||
from huggingface_hub import save_torch_state_dict
|
||||
from safetensors import safe_open
|
||||
from safetensors.torch import load_file
|
||||
|
||||
|
||||
OFFICIAL_TO_DIFFUSERS_MAPPING: dict[str, str] = {
|
||||
r"^text_embedding\.0\.(.*)$": r"condition_embedder.text_embedder.linear_1.\1",
|
||||
r"^text_embedding\.2\.(.*)$": r"condition_embedder.text_embedder.linear_2.\1",
|
||||
r"^time_embedding\.0\.(.*)$": r"condition_embedder.time_embedder.linear_1.\1",
|
||||
r"^time_embedding\.2\.(.*)$": r"condition_embedder.time_embedder.linear_2.\1",
|
||||
r"^time_projection\.1\.(.*)$": r"condition_embedder.time_proj.\1",
|
||||
r"^img_emb\.proj\.0\.(.*)$": r"condition_embedder.image_embedder.norm1.\1",
|
||||
r"^img_emb\.proj\.1\.(.*)$": r"condition_embedder.image_embedder.ff.net.0.proj.\1",
|
||||
r"^img_emb\.proj\.3\.(.*)$": r"condition_embedder.image_embedder.ff.net.2.\1",
|
||||
r"^img_emb\.proj\.4\.(.*)$": r"condition_embedder.image_embedder.norm2.\1",
|
||||
r"^head\.modulation": r"scale_shift_table",
|
||||
r"^head\.head\.(.*)$": r"proj_out.\1",
|
||||
r"^blocks\.(\d+)\.self_attn\.q\.(.*)$": r"blocks.\1.attn1.to_q.\2",
|
||||
r"^blocks\.(\d+)\.self_attn\.k\.(.*)$": r"blocks.\1.attn1.to_k.\2",
|
||||
r"^blocks\.(\d+)\.self_attn\.v\.(.*)$": r"blocks.\1.attn1.to_v.\2",
|
||||
r"^blocks\.(\d+)\.self_attn\.o\.(.*)$": r"blocks.\1.attn1.to_out.0.\2",
|
||||
r"^blocks\.(\d+)\.self_attn\.norm_q\.(.*)$": r"blocks.\1.attn1.norm_q.\2",
|
||||
r"^blocks\.(\d+)\.self_attn\.norm_k\.(.*)$": r"blocks.\1.attn1.norm_k.\2",
|
||||
r"^blocks\.(\d+)\.cross_attn\.q\.(.*)$": r"blocks.\1.attn2.to_q.\2",
|
||||
r"^blocks\.(\d+)\.cross_attn\.k\.(.*)$": r"blocks.\1.attn2.to_k.\2",
|
||||
r"^blocks\.(\d+)\.cross_attn\.k_img\.(.*)$": r"blocks.\1.attn2.add_k_proj.\2",
|
||||
r"^blocks\.(\d+)\.cross_attn\.v\.(.*)$": r"blocks.\1.attn2.to_v.\2",
|
||||
r"^blocks\.(\d+)\.cross_attn\.v_img\.(.*)$": r"blocks.\1.attn2.add_v_proj.\2",
|
||||
r"^blocks\.(\d+)\.cross_attn\.o\.(.*)$": r"blocks.\1.attn2.to_out.0.\2",
|
||||
r"^blocks\.(\d+)\.cross_attn\.norm_q\.(.*)$": r"blocks.\1.attn2.norm_q.\2",
|
||||
r"^blocks\.(\d+)\.cross_attn\.norm_k\.(.*)$": r"blocks.\1.attn2.norm_k.\2",
|
||||
r"^blocks\.(\d+)\.cross_attn\.norm_k_img\.(.*)$": r"blocks.\1.attn2.norm_added_k.\2",
|
||||
r"^blocks\.(\d+)\.ffn\.0\.(.*)$": r"blocks.\1.ffn.net.0.proj.\2",
|
||||
r"^blocks\.(\d+)\.ffn\.2\.(.*)$": r"blocks.\1.ffn.net.2.\2",
|
||||
r"^blocks\.(\d+)\.modulation": r"blocks.\1.scale_shift_table",
|
||||
r"^blocks\.(\d+)\.norm3\.(.*)$": r"blocks.\1.norm2.\2",
|
||||
}
|
||||
|
||||
TRANSFORMER_CONFIG: dict[str, object] = {
|
||||
"_class_name": "DreamXWorldTransformer3DModel",
|
||||
"patch_size": [1, 2, 2],
|
||||
"text_len": 512,
|
||||
"num_attention_heads": 24,
|
||||
"attention_head_dim": 128,
|
||||
"in_channels": 48,
|
||||
"out_channels": 48,
|
||||
"text_dim": 4096,
|
||||
"freq_dim": 256,
|
||||
"ffn_dim": 14336,
|
||||
"num_layers": 30,
|
||||
"cross_attn_norm": True,
|
||||
"qk_norm": "rms_norm_across_heads",
|
||||
"eps": 1e-6,
|
||||
"image_dim": None,
|
||||
"added_kv_proj_dim": None,
|
||||
"rope_max_seq_len": 1024,
|
||||
"add_control_adapter": True,
|
||||
"cam_method": "prope",
|
||||
"attn_compress": 1,
|
||||
"cam_self_attn_layers": None,
|
||||
}
|
||||
|
||||
MODEL_INDEX: dict[str, object] = {
|
||||
"_class_name": "DreamXWorldPipeline",
|
||||
"_diffusers_version": "0.31.0",
|
||||
"scheduler": ["diffusers", "FlowMatchEulerDiscreteScheduler"],
|
||||
"text_encoder": ["transformers", "UMT5EncoderModel"],
|
||||
"tokenizer": ["transformers", "AutoTokenizer"],
|
||||
"transformer": ["diffusers", "DreamXWorldTransformer3DModel"],
|
||||
"vae": ["diffusers", "AutoencoderKLWan"],
|
||||
}
|
||||
|
||||
REUSED_COMPONENTS = ("scheduler", "text_encoder", "tokenizer", "vae")
|
||||
|
||||
|
||||
def map_transformer_key(key: str) -> str:
|
||||
for pattern, replacement in OFFICIAL_TO_DIFFUSERS_MAPPING.items():
|
||||
if re.match(pattern, key):
|
||||
return re.sub(pattern, replacement, key)
|
||||
return key
|
||||
|
||||
|
||||
def _safetensor_files(source: Path) -> list[Path]:
|
||||
if source.is_file():
|
||||
if source.suffix != ".safetensors":
|
||||
raise ValueError(f"Only .safetensors files are supported, got {source}")
|
||||
return [source]
|
||||
|
||||
index_path = source / "diffusion_pytorch_model.safetensors.index.json"
|
||||
if index_path.exists():
|
||||
index = json.loads(index_path.read_text())
|
||||
return sorted({source / shard for shard in index["weight_map"].values()})
|
||||
|
||||
files = sorted(source.glob("*.safetensors"))
|
||||
if not files:
|
||||
raise FileNotFoundError(f"No safetensors files found under {source}")
|
||||
return files
|
||||
|
||||
|
||||
def convert_transformer(source: Path, output: Path, max_shard_size: str) -> None:
|
||||
transformer_dir = output / "transformer"
|
||||
transformer_dir.mkdir(parents=True, exist_ok=True)
|
||||
(transformer_dir / "config.json").write_text(json.dumps(TRANSFORMER_CONFIG, indent=2) + "\n")
|
||||
|
||||
converted: OrderedDict[str, torch.Tensor] = OrderedDict()
|
||||
for shard in _safetensor_files(source):
|
||||
print(f"loading {shard}")
|
||||
for key, tensor in load_file(shard, device="cpu").items():
|
||||
new_key = map_transformer_key(key)
|
||||
if new_key in converted:
|
||||
raise ValueError(f"Duplicate converted key: {new_key}")
|
||||
converted[new_key] = tensor
|
||||
|
||||
print(f"saving {len(converted)} tensors to {transformer_dir}")
|
||||
save_torch_state_dict(converted, transformer_dir, max_shard_size=max_shard_size)
|
||||
|
||||
|
||||
def _copy_or_link_component(component: str, component_source: Path, output: Path, symlink: bool) -> None:
|
||||
src = component_source / component
|
||||
dst = output / component
|
||||
if not src.exists():
|
||||
raise FileNotFoundError(f"Missing reused component source: {src}")
|
||||
if dst.is_symlink() and not dst.exists():
|
||||
dst.unlink() # dangling symlink: replace so re-conversion self-heals
|
||||
if dst.exists():
|
||||
print(f"keeping existing {dst}")
|
||||
return
|
||||
if symlink:
|
||||
dst.symlink_to(src.resolve(), target_is_directory=src.is_dir())
|
||||
print(f"linked {dst} -> {src}")
|
||||
elif src.is_dir():
|
||||
shutil.copytree(src, dst)
|
||||
print(f"copied {src} -> {dst}")
|
||||
else:
|
||||
shutil.copy2(src, dst)
|
||||
print(f"copied {src} -> {dst}")
|
||||
|
||||
|
||||
def write_model_index(output: Path, component_source: Path | None, symlink_components: bool) -> None:
|
||||
if component_source is not None:
|
||||
for component in REUSED_COMPONENTS:
|
||||
_copy_or_link_component(component, component_source, output, symlink_components)
|
||||
missing = [component for component in REUSED_COMPONENTS if not (output / component).exists()]
|
||||
if missing:
|
||||
print("skipping model_index.json; missing reused components: " + ", ".join(missing))
|
||||
print("pass --component-source <Wan2.2 Diffusers root> to copy or link reused components")
|
||||
return
|
||||
(output / "model_index.json").write_text(json.dumps(MODEL_INDEX, indent=2) + "\n")
|
||||
print(f"wrote {output / 'model_index.json'}")
|
||||
|
||||
|
||||
def analyze(source: Path) -> None:
|
||||
total = 0
|
||||
unchanged = 0
|
||||
examples: list[tuple[str, str]] = []
|
||||
for shard in _safetensor_files(source):
|
||||
with safe_open(shard, framework="pt", device="cpu") as tensors:
|
||||
for key in tensors:
|
||||
total += 1
|
||||
new_key = map_transformer_key(key)
|
||||
unchanged += int(new_key == key)
|
||||
if len(examples) < 20 and new_key != key:
|
||||
examples.append((key, new_key))
|
||||
print(f"total_keys={total} unchanged_keys={unchanged}")
|
||||
for old, new in examples:
|
||||
print(f"{old} -> {new}")
|
||||
|
||||
|
||||
def main() -> None:
|
||||
parser = argparse.ArgumentParser(description=__doc__)
|
||||
parser.add_argument("--source", type=Path, required=True, help="DreamX raw transformer directory or safetensors file")
|
||||
parser.add_argument("--output", type=Path, required=True, help="Output model root; transformer/ is created inside it")
|
||||
parser.add_argument("--max-shard-size", default="10GB")
|
||||
parser.add_argument(
|
||||
"--component-source",
|
||||
type=Path,
|
||||
help="Optional Wan2.2 Diffusers root whose scheduler/text_encoder/tokenizer/vae components are reused.",
|
||||
)
|
||||
parser.add_argument(
|
||||
"--symlink-components",
|
||||
action="store_true",
|
||||
help="Symlink reused components from --component-source instead of copying them.",
|
||||
)
|
||||
parser.add_argument("--analyze", action="store_true", help="Only print key mapping summary")
|
||||
args = parser.parse_args()
|
||||
|
||||
if args.analyze:
|
||||
analyze(args.source)
|
||||
else:
|
||||
convert_transformer(args.source, args.output, args.max_shard_size)
|
||||
write_model_index(args.output, args.component_source, args.symlink_components)
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
main()
|
||||
@@ -0,0 +1,189 @@
|
||||
# DreamX World Port Status
|
||||
|
||||
## Summary
|
||||
|
||||
- model_family: `dreamx_world`
|
||||
- workload_types: `I2V camera-control compatibility shim`; `I2V autoregressive camera-control forcing`
|
||||
- official_ref: `https://github.com/AMAP-ML/DreamX-World`
|
||||
- official_ref_dir: `DreamX-World/`
|
||||
- hf_weights_path: `GD-ML/DreamX-World-5B-Cam`
|
||||
- local_weights_dir: `official_weights/dreamx_world`
|
||||
- source_layout: `raw_official`
|
||||
- local_tests_readme: `tests/local_tests/dreamx_world/README.md`
|
||||
|
||||
## Current Phase
|
||||
|
||||
- phase: `phase_11_post_parity_handoff`
|
||||
- status: `complete`
|
||||
- owner: `orchestrator`
|
||||
- last_updated: `2026-07-02`
|
||||
|
||||
## Component Matrix
|
||||
|
||||
| Component | Type | Reuse/Port | Official Definition | Official Instantiation | FastVideo Target | Prototype | Conversion | Parity | Open Issues |
|
||||
|---|---|---|---|---|---|---|---|---|---|
|
||||
| transformer | dit | ported_dedicated | `DreamX-World/models/wan_transformer3d.py`; PRoPE helpers in `DreamX-World/models/prope_utils.py` | `DreamX-World/inference_dreamx5b.py::setup_models`, `Wan2_2Transformer3DModel.from_pretrained(... cam_method=prope, add_control_adapter=True)` | `fastvideo/models/dits/dreamx_world.py`; `fastvideo/configs/models/dits/dreamx_world.py`; DreamX pipeline config helper | native_prope_pass | real_conversion_pass | strict_load_and_forward_parity_pass | none |
|
||||
| vae | vae | reuse_pending | `DreamX-World/models/wan_vae3_8.py` | `DreamX-World/inference_dreamx5b.py::setup_models`, `AutoencoderKLWan3_8.from_pretrained(Wan2.2_VAE.pth)` | `fastvideo/models/vaes/wanvae.py`; DreamX VAE config helper | config_smoke_pass | raw_key_mapping_pass | encode_parity_pass | none |
|
||||
| text_encoder/tokenizer | encoder | reuse_pending | `DreamX-World/models/wan_text_encoder.py`; tokenizer via Wan2.2 base model | `DreamX-World/inference_dreamx5b.py::setup_models`, `WanT5EncoderModel` + tokenizer subpaths | `fastvideo/models/encoders/t5.py::UMT5EncoderModel`; DreamX UMT5 config helper | config_smoke_pass | staged_weight_load_pass | hidden_state_parity_pass | none |
|
||||
| scheduler | generic | reuse_proven | Diffusers `FlowMatchEulerDiscreteScheduler` | `DreamX-World/inference_dreamx5b.py::setup_models`, default `sampler_name=Flow` | `fastvideo/models/schedulers/scheduling_flow_match_euler_discrete.py` | pass | not_required | non_skip_pass | Q003 |
|
||||
| camera_conditioning | generic | port_pending | `DreamX-World/utils/inference_utils.py`, `DreamX-World/models/prope_utils.py`, `DreamX-World/wan/modules/camera_prope.py` | `DreamX-World/inference_dreamx5b.py::get_camera_sequence`, `pipeline(... control_camera_video=...)` | `fastvideo/pipelines/basic/dreamx_world/camera_conditioning.py` | pass | not_required | non_skip_pass | none |
|
||||
| pipeline | pipeline | port_complete | `DreamX-World/pipeline/pipeline_dreamxworld.py` | `DreamX-World/inference_dreamx5b.py::process_inference_from_json` | `fastvideo/pipelines/basic/dreamx_world/` plus config/preset/registry | pipeline_load_generate_smoke_pass | model_index_and_config_consistency_smoke_pass | pipeline_api_vs_worker_forward_parity_pass | none |
|
||||
| ar_transformer | dit | port_complete | `DreamX-World/wan/modules/causal_camera_model_2_2_prope_infinity.py::CausalWanModel` | `DreamX-World/inference_ar_forcing.py::load_pipeline` | `fastvideo/models/dits/dreamx_world_ar.py`; `fastvideo/configs/models/dits/dreamx_world.py::DreamXWorldARConfig` | tiny_official_forward_parity_pass | identity_conversion_pass | real_5b_strict_load_pass | none |
|
||||
| ar_pipeline | pipeline | port_complete | `DreamX-World/pipeline/pipeline_causal_camera.py` | `DreamX-World/inference_ar_forcing.py::main` | `fastvideo/pipelines/basic/dreamx_world/dreamx_world_ar_pipeline.py`; `fastvideo/pipelines/basic/dreamx_world/ar_denoising.py`; registry/preset/config | config_registry_pass | symlink_model_index_pass | short_full_generation_pass | none |
|
||||
|
||||
## Conversion State
|
||||
|
||||
- conversion_script: `scripts/checkpoint_conversion/dreamx_world_to_diffusers.py`
|
||||
- converted_weights_dir: `converted_weights/dreamx_world`
|
||||
- source_layout: `raw_official`
|
||||
- strict_load_status: `pass`
|
||||
- conversion_script_status: `transformer_model_index_and_config_consistency_smoke_pass`
|
||||
- model_index_status: `smoke_pass`
|
||||
- passthrough_components: `Wan2.2 Diffusers scheduler, tokenizer, and text encoder are symlinked from official_weights/Wan2.2-TI2V-5B-Diffusers; VAE parity uses raw Wan2.2_VAE.pth with an explicit DreamX raw-to-FastVideo key mapper because official encode returns normalized latents.`
|
||||
- retry_history: `none`
|
||||
|
||||
## Parity Commands
|
||||
|
||||
| Scope | Command | Last Result | Notes |
|
||||
|---|---|---|---|
|
||||
| transformer | `python -m pytest tests/local_tests/dreamx_world/test_dreamx_world_transformer_parity.py -v -s` | strict_load_and_forward_parity_pass | 2026-07-01: converted real 5B-Cam transformer shards strict-load into dedicated `DreamXWorldTransformer3DModel` with 0 shape mismatches; official-vs-FastVideo small-input fp32 forward parity passes on CUDA (`diff_max=0.072533`, `diff_mean=0.008014`). |
|
||||
| vae | `python -m pytest tests/local_tests/dreamx_world/test_dreamx_world_vae_parity.py -v -s` | encode_parity_pass | 2026-06-30: official DreamX raw `Wan2.2_VAE.pth` maps 196/196 keys into FastVideo Wan VAE and encode parity passes after applying the same official latent normalization (`(mu - mean) / std`). |
|
||||
| text_encoder | `python -m pytest tests/local_tests/dreamx_world/test_dreamx_world_text_encoder_parity.py -v -s` | hidden_state_parity_pass | 2026-06-30: official `WanT5EncoderModel` vs FastVideo `UMT5EncoderModel` hidden-state parity passes on CUDA using staged Wan2.2 text encoder/tokenizer weights and reference-only `xfuser` stubs. |
|
||||
| scheduler | `python -m pytest tests/local_tests/dreamx_world/test_dreamx_world_scheduler_parity.py -v -s` | non_skip_pass | 2026-06-30: FastVideo FlowMatch scheduler matches official Diffusers timesteps and step output for DreamX default Flow sampler; `DreamXWorldPipeline` initializes FlowMatch with official `shift=3.0`. |
|
||||
| camera_conditioning | `python -m pytest tests/local_tests/dreamx_world/test_dreamx_world_camera_conditioning_parity.py -v -s` | non_skip_pass | 2026-06-29: 3 parameterized cases passed against official reference on CPU. |
|
||||
| pipeline_config | `python -m pytest tests/local_tests/dreamx_world/test_dreamx_world_pipeline_config.py -v -s` | pipeline_entry_preset_scheduler_modelinfo_and_camera_stage_smoke_pass | DreamX 5B-Cam PipelineConfig wires DiT/VAE/UMT5/Flow/TI2V settings and official `shift=3.0`; default preset is registered for `GD-ML/DreamX-World-5B-Cam`; local converted-style `model_index.json` resolves to `DreamXWorldPipeline`; the pipeline initializes FlowMatch, camera conditioning writes `batch.extra["dreamx_y_camera"]`, and generic denoising can pass it as `y_camera`. |
|
||||
| ar_transformer | `PYTHONPATH=/workspace/FastVideo python -m pytest tests/local_tests/dreamx_world/test_dreamx_world_ar_conversion.py tests/local_tests/dreamx_world/test_dreamx_world_ar_transformer_parity.py -q -rs` | 6_passed_0_skipped | 2026-07-02: AR converter writes symlinked transformer layout/model_index; tiny official `CausalWanModel` vs FastVideo `DreamXWorldARTransformer3DModel` forward parity passes; real 5B `model.safetensors` strict-loads with zero missing/unexpected keys from `/tmp/converted_dreamx_world_ar`. |
|
||||
| ar_pipeline_config | `PYTHONPATH=/workspace/FastVideo python -m pytest tests/local_tests/dreamx_world/test_dreamx_world_pipeline_config.py -q -rs` | 10_passed_0_skipped | 2026-07-02: `DreamXWorld5BARPipelineConfig`, `DreamXWorldARPipeline`, registry config selection for `GD-ML/DreamX-World-5B`, and `dreamx_world_5b_ar` preset pass. |
|
||||
| ar_full_generation | `PYTHONPATH=/workspace/FastVideo python /tmp/run_dreamx_ar_video_smoke.py` | generated_video_pass | 2026-07-02: A40 short full-generation smoke passed from `/tmp/converted_dreamx_world_ar` with 64x64, 9 frames, 4 denoise steps, `output_type=pil`, `save_video=True`; MP4 saved at `outputs_video/dreamx_world_ar_video_smoke/a quiet road through a futuristic coastal city at sunrise.mp4` and decoded as 9 frames of `(64, 64, 3)` uint8. |
|
||||
| ar_long_horizon | `PYTHONPATH=/workspace/FastVideo FASTVIDEO_ATTENTION_BACKEND=TORCH_SDPA python /tmp/run_dreamx_ar_long_horizon.py` | generated_video_pass | 2026-07-02: A40 long-horizon AR generation passed from `/tmp/converted_dreamx_world_ar` with 64x64, 1005 frames, 4 denoise steps, seed 4096; MP4 saved at `outputs_video/dreamx_world_ar_long_horizon/A long autonomous drive through a futuristic coastal city at sunrise, smooth forward camera motion,.mp4` and decoded as 1005 frames of `(64, 64, 3)`; end-to-end generation latency was 231.68s after load. |
|
||||
| ar_ssim_default | `PYTHONPATH=/workspace/FastVideo FASTVIDEO_SSIM_MODEL_ID=DreamX-World-5B DREAMX_WORLD_AR_SSIM_MODEL_PATH=/tmp/converted_dreamx_world_ar python -m pytest fastvideo/tests/ssim/test_dreamx_world_similarity.py -q -rs --skip-ssim-reference-download` | 1_passed_0_skipped | 2026-07-02: A40 default AR SSIM reference seeded locally at `fastvideo/tests/ssim/reference_videos/default/A40_reference_videos/DreamX-World-5B/TORCH_SDPA/`; helper removes stale generated base MP4 before generation; default params are 192x192, 81 frames, 4 steps, seed 2048, min SSIM 0.98. |
|
||||
| ar_ssim_modal_l40s | `modal run /tmp/modal_dreamx_ar_ssim_git.py` | 1_passed_0_skipped | 2026-07-02: Modal L40S run checked out `post-fix dreamx-world-5b-cam branch commit`, downloaded default `L40S_reference_videos` via `reference_videos_cli.py download`, used cached converted AR weights under `/root/data/dreamx_world_ar_converted`, seeded the missing AR L40S reference from generated output, reran a fresh generated-vs-reference compare successfully (`mean_ssim=1.0`), and exported the reference to Modal volume `hf-model-weights:dreamx_ar_ssim_l40s`. Downloaded local reference decodes as 81 frames of `(192, 192, 3)`. A40-vs-L40S reference spot check after deterministic AR noise fix: mean SSIM 0.7770, min 0.6139, max 0.9944. |
|
||||
| pipeline_smoke | `python -m pytest tests/local_tests/pipelines/test_dreamx_world_pipeline_smoke.py tests/local_tests/pipelines/test_dreamx_world_pipeline_parity.py -q -rs` | 4_passed_0_skipped | 2026-06-30: combined smoke/parity passed. 2026-07-01: smoke alone passed with real `image_path` TI2V coverage (`3 passed`), validating image load, TI2V preprocessing, VAE first-frame encode under CPU offload, camera conditioning, and 1-step latent generation from `converted_weights/dreamx_world`. Tests force `FASTVIDEO_ATTENTION_BACKEND=TORCH_SDPA` to avoid the local FlashAttention-4 cute ABI mismatch. |
|
||||
| basic_example | `PYTHONPATH=/workspace/FastVideo FASTVIDEO_ATTENTION_BACKEND=TORCH_SDPA DREAMX_WORLD_MODEL_DIR=converted_weights/dreamx_world DREAMX_WORLD_IMAGE_PATH= DREAMX_WORLD_HEIGHT=64 DREAMX_WORLD_WIDTH=64 DREAMX_WORLD_NUM_FRAMES=9 DREAMX_WORLD_STEPS=1 DREAMX_WORLD_GUIDANCE=1.0 DREAMX_WORLD_OUTPUT_PATH=outputs_video/dreamx_world_example_smoke python examples/inference/basic/basic_dreamx_world.py` | generated_video_pass | 2026-06-30: example saved an MP4 under `outputs_video/dreamx_world_example_smoke`; imageio/ffmpeg decoded frame 0 as `(64, 64, 3)` uint8, fps 16, duration 0.56s. |
|
||||
|
||||
## Open Questions
|
||||
|
||||
| ID | Question | Owner | Needed By Phase | Status | Resolution |
|
||||
|---|---|---|---|---|---|
|
||||
| Q001 | Should first PR expose only the `DreamX-World-5B-Cam` 5s camera-control mode and exclude AR long-horizon forcing? | user | prep | resolved | User approved starting with `DreamX-World-5B-Cam`; AR long-horizon is out of first-PR scope. |
|
||||
| Q002 | Does FastVideo's existing Wan2.2 TI2V transformer support DreamX PRoPE/control adapter with a small extension, or is a DreamX-specific DiT required? | component:transformer | Phase 3 | resolved | Project guidance prefers a separate DreamX DiT for maintainability. DreamX PRoPE/control adapter now lives in `fastvideo/models/dits/dreamx_world.py`; Wan DiT/config have no DreamX-specific fields or `y_camera` signature. |
|
||||
| Q003 | Which sampler is in first-PR scope: official default `Flow` only, or also `Flow_Unipc` and `Flow_DPM++`? | orchestrator | Phase 3 | resolved | First PR should support official default `Flow` only. FastVideo FlowMatch scheduler parity is non-skip PASS; optional `Flow_Unipc` and `Flow_DPM++` are out of first-PR scope. |
|
||||
| Q004 | Which HF token env var should be used if rate limits or gated Wan2.2 base weights require auth? | user | Phase 5 | resolved | No auth was required for the completed local downloads; keep using env var names only if future gated repos require auth. |
|
||||
| Q005 | Should native FastVideo production code depend on DreamX reference-only packages such as `xfuser` or OpenCV? | user | Phase 3 | resolved | No. These packages may be used only for official reference/local parity setup; native FastVideo integration must remove that runtime requirement. |
|
||||
| Q006 | Should AR handoff require a full generated long-horizon video in this no-HF-token/no-GPU-budget pass? | user/runtime | pipeline | resolved | A40 long-horizon generation passes with 1005 frames at 64x64/4 steps. AR default SSIM passes locally on A40 and on Modal L40S with a 192x192/81-frame reference. HF upload/publication remains a separate operation if the reference dataset should be updated upstream. |
|
||||
| Q007 | Can A40 references stand in for L40S CI references? | quality | release | resolved | No. After deterministic AR noise fix, A40-vs-L40S reference mean SSIM is 0.7770, below the 0.98 same-device threshold. Publish the L40S-specific reference for CI. |
|
||||
|
||||
## Issues And Blockers
|
||||
|
||||
| ID | Phase | Component | Severity | Issue | Evidence | Owner | Status | Resolution |
|
||||
|---|---|---|---|---|---|---|---|---|
|
||||
| I001 | prep | official_env | medium | Official import initially failed because `xfuser` was missing. | `ModuleNotFoundError: No module named 'xfuser'` from `python -c "import sys; sys.path.insert(0, 'DreamX-World'); import inference_dreamx5b"` | prep | resolved | Installed `xfuser==0.4.1`; import progressed. |
|
||||
| I002 | prep | official_env | medium | Official import then failed because GUI OpenCV required missing system `libxcb.so.1`. | `ImportError: libxcb.so.1: cannot open shared object file` through `cv2` import in Diffusers ConsisID path. | prep | resolved | Installed `opencv-python-headless`; `import inference_dreamx5b` passed. |
|
||||
| I003 | prep | weights | medium | HF repo has raw official transformer shards and no Diffusers `model_index.json`. | `inspect_hf_layout.py GD-ML/DreamX-World-5B-Cam --json` returned `source_layout=raw_official`, `needs_conversion=yes`, `model_index_class=null`. | conversion | resolved | Downloaded raw DreamX shards to `official_weights/dreamx_world`; converted transformer to `converted_weights/dreamx_world/transformer`; symlinked reusable Wan2.2 Diffusers components; real 5B transformer strict-load passes. |
|
||||
| I004 | prep | dependencies | high | Official reference import required extra packages in the local environment, but FastVideo native runtime should not inherit those dependencies. | `xfuser==0.4.1` and `opencv-python-headless` were installed only to make `DreamX-World/inference_dreamx5b.py` import for reference/parity. | pipeline | resolved | Production DreamX FastVideo code uses native camera/image/video utilities and has no runtime `xfuser` or OpenCV import requirement; those packages remain reference-only local parity dependencies. |
|
||||
| I005 | parity | transformer | medium | Transformer full forward parity initially failed in bf16 official harness. | Official CUDA bf16 LayerNorm path was unstable; fp32 small-input harness avoids that dtype issue and compares against FastVideo with single-process SP identity patches. | component:transformer | resolved | Official-vs-FastVideo forward parity now passes on CUDA with `diff_max=0.072533`, `diff_mean=0.008014`. |
|
||||
| I006 | parity | vae/text_encoder | medium | VAE/text parity initially remained skipped after weights were staged. | Text official import needed a reference-only `xfuser` stub; VAE comparison initially used raw official normalized latents against FastVideo raw mu. | component:vae,component:text_encoder | resolved | Text hidden-state parity passes. VAE encode parity passes after raw key mapping and applying the official latent normalization to FastVideo output. |
|
||||
| I007 | quality | pipeline_ti2v | medium | Real image-path TI2V smoke initially failed when `vae_cpu_offload=True` because DenoisingStage encoded the first frame while the VAE weights remained on CPU. | DreamX SSIM first run failed with `RuntimeError: Input type (torch.cuda.FloatTensor) and weight type (torch.FloatTensor) should be the same` at `fastvideo/pipelines/stages/denoising.py` VAE encode. | pipeline | resolved | DenoisingStage now moves the VAE to `local_device` before TI2V first-frame encode; image-path pipeline smoke and DreamX SSIM both pass. |
|
||||
| I008 | quality | ssim_helper | high | SSIM helper could compare against a stale generated base MP4 when a rerun saved the new video as `_1.mp4`. | Existing generated outputs made AR reference seeding appear to pass before a fresh generated-vs-reference compare. | quality | resolved | `run_text_to_video_similarity_test` and `run_image_to_video_similarity_test` now remove the stale generated base MP4 before generation. A40 and Modal L40S AR SSIM were rerun after the fix. |
|
||||
| I009 | quality | ar_denoising | high | AR denoising added CUDA noise without using the request seed when the original generator was CPU-backed. | Fresh reruns against old AR references produced mean SSIM near 0.05. | pipeline | resolved | `DreamXWorldARCausalDenoisingStage` now derives a device-local generator from the request seed for AR noise. Fresh same-device A40 and L40S reruns pass with mean SSIM 1.0 after reseeding references. |
|
||||
|
||||
## Escape Hatches
|
||||
|
||||
| ID | Phase | Decision Type | Question | Recommended Option | Status | Resolution |
|
||||
|---|---|---|---|---|---|---|
|
||||
|
||||
## Decisions
|
||||
|
||||
| Date | Decision | Rationale | Impact |
|
||||
|---|---|---|---|
|
||||
| 2026-06-29 | First PR scope is `DreamX-World-5B-Cam` only. | Cam mode is closest to existing Wan2.2 TI2V support; AR forcing needs separate causal/KV pipeline work. | Component inventory and parity focus on `inference_dreamx5b.py` and `pipeline_dreamxworld.py`. |
|
||||
| 2026-06-29 | Do not install full DreamX requirements during prep. | Full requirements pin core FastVideo stack packages. | Installed only `xfuser==0.4.1` and `opencv-python-headless` to make official imports work. |
|
||||
| 2026-06-29 | Treat HF DreamX-World-5B-Cam weights as raw official transformer layout requiring conversion. | HF inspection found no `model_index.json`. | Phase 5 must create `scripts/checkpoint_conversion/dreamx_world_to_diffusers.py` after component prototype/key dumps. |
|
||||
| 2026-06-29 | Do not add DreamX reference-only dependencies to FastVideo production requirements. | The current environment should remain the FastVideo environment; extra packages are only for official reference parity. | Native DreamX integration must avoid runtime `xfuser` and OpenCV requirements unless explicitly approved later. |
|
||||
| 2026-06-29 | Implement DreamX camera conditioning as native FastVideo utility. | It is weightless and removes the need to import DreamX reference utilities at production runtime. | `fastvideo/pipelines/basic/dreamx_world/camera_conditioning.py` now has non-skip parity against official action-to-PRoPE tensors. |
|
||||
| 2026-06-29 | First PR supports DreamX default `Flow` sampler only. | FastVideo FlowMatch Euler scheduler matches the official Diffusers scheduler for DreamX defaults. | Pipeline work can use FastVideo native FlowMatch scheduler; UniPC and DPM++ are out of first-PR scope. |
|
||||
| 2026-07-01 | Keep DreamX PRoPE/control adapter in a dedicated DreamX DiT class. | Project guidance is that putting too much DreamX behavior into Wan makes the model hard to manage. | `fastvideo/models/dits/dreamx_world.py` defines `DreamXWorldTransformer3DModel`, `DreamXWorldTransformerBlock`, and `DreamXPropeSelfAttention`; `fastvideo/configs/models/dits/dreamx_world.py` owns DreamX adapter config fields; Wan DiT/config are unchanged from DreamX. |
|
||||
| 2026-06-30 | Camera parity test loads official camera functions by file instead of importing the official `utils` package. | Official package initialization pulls unrelated dependencies that can require GUI OpenCV system libraries. | Camera parity remains non-skip without adding DreamX reference-only dependencies to FastVideo production requirements. |
|
||||
| 2026-06-30 | Add DreamX-World-5B-Cam model and pipeline config helpers plus a conversion script. | Official HF DreamX 5B-Cam transformer config is 30 layers, hidden size 3072, 24 heads, 48 latent channels, plus Wan2.2 48-channel VAE and UMT5-XXL text encoder. | DreamX helpers wire DiT/VAE/UMT5/Flow/TI2V settings; `dreamx_world_to_diffusers.py` writes a FastVideo-loadable transformer config plus renamed safetensors; strict-load smoke passes on a tiny official DreamX transformer and the real 5B converted shards. |
|
||||
| 2026-06-30 | Pass DreamX camera PRoPE condition through the FastVideo batch/denoising path. | DreamX transformer expects `y_camera={"viewmats", "K"}` at denoising time. | `DreamXWorldPipeline` is registered as a basic pipeline entry and initializes the official default FlowMatch scheduler; `dreamx_world_5b_cam` preset mirrors official 5B-Cam defaults; `DreamXWorldCameraConditioningStage` writes `batch.extra["dreamx_y_camera"]`; generic denoising filters and forwards it as `y_camera` only for compatible transformers. |
|
||||
|
||||
## Handoff Notes
|
||||
|
||||
- Prep, component parity, pipeline smoke/parity, and the basic example validation are complete for `DreamX-World-5B-Cam`.
|
||||
- Official reference clone is staged at `DreamX-World/` and ignored by git.
|
||||
- Workspace-local weights are staged: DreamX raw transformer shards under `official_weights/dreamx_world`, Wan2.2 raw base artifacts under `official_weights/Wan2.2-TI2V-5B`, and Wan2.2 Diffusers reusable components under `official_weights/Wan2.2-TI2V-5B-Diffusers`.
|
||||
- Camera conditioning parity is active and passing without weights.
|
||||
- Default Flow scheduler parity is active and passing without weights.
|
||||
- Transformer has corrected official 5B-Cam architecture in dedicated DreamX DiT/config files, native PRoPE/control-adapter, conversion mapping, real converted 5B strict-load, and official-vs-FastVideo forward parity passing on CUDA. VAE encode parity and text hidden-state parity pass on CUDA. Pipeline entry/registry, local model_info resolution, preset, config, FlowMatch scheduler init, camera stage, denoising `y_camera` kwarg smokes, independent CUDA pipeline smoke/parity, and a small saved-video basic example pass.
|
||||
- Full local DreamX component suite is non-skip PASS: `python -m pytest tests/local_tests/dreamx_world/ -q -rs` returned `26 passed` on 2026-07-01. Pipeline smoke/parity and SSIM quality regression are also non-skip PASS locally.
|
||||
- Keep `xfuser` and OpenCV as reference-only parity dependencies. Do not add them
|
||||
to FastVideo requirements or production imports.
|
||||
|
||||
|
||||
## Quality Regression
|
||||
|
||||
- status: `added`
|
||||
- test: `fastvideo/tests/ssim/test_dreamx_world_similarity.py`
|
||||
- command: `PYTHONPATH=/workspace/FastVideo FASTVIDEO_SSIM_MODEL_ID=DreamX-World-5B-Cam python -m pytest fastvideo/tests/ssim/test_dreamx_world_similarity.py -q -rs`
|
||||
- result: `1 passed, 0 skipped` on 2026-07-01
|
||||
- reference: Local A40/TORCH_SDPA reference seeded from the generated candidate under `fastvideo/tests/ssim/reference_videos/default/A40_reference_videos/DreamX-World-5B-Cam/TORCH_SDPA/`. The test uses a deterministic generated input image, 64x64 request dimensions, 9 frames, 1 denoise step, seed 1024, and min SSIM 0.98. Full-quality params are present for 480x832/161 frames/30 steps.
|
||||
- note: Modal L40S seeding passed using the configured Modal profile and unauthenticated HF public downloads. HF upload/publication still requires `HF_API_KEY`, `HUGGINGFACE_HUB_TOKEN`, or `HF_TOKEN` with write access; no token values were used or recorded.
|
||||
|
||||
## Final Handoff
|
||||
|
||||
```text
|
||||
final_handoff:
|
||||
prep_handoff_complete: yes
|
||||
conversion_status: pass
|
||||
components:
|
||||
- name: transformer
|
||||
reuse_or_port: ported_dedicated_dit
|
||||
parity_test: tests/local_tests/dreamx_world/test_dreamx_world_transformer_parity.py
|
||||
parity_status: non_skip_pass
|
||||
concerns_or_unknowns: none
|
||||
- name: vae
|
||||
reuse_or_port: reused
|
||||
parity_test: tests/local_tests/dreamx_world/test_dreamx_world_vae_parity.py
|
||||
parity_status: non_skip_pass
|
||||
concerns_or_unknowns: none
|
||||
- name: text_encoder_tokenizer
|
||||
reuse_or_port: reused
|
||||
parity_test: tests/local_tests/dreamx_world/test_dreamx_world_text_encoder_parity.py
|
||||
parity_status: non_skip_pass
|
||||
concerns_or_unknowns: none
|
||||
- name: scheduler
|
||||
reuse_or_port: reused
|
||||
parity_test: tests/local_tests/dreamx_world/test_dreamx_world_scheduler_parity.py
|
||||
parity_status: non_skip_pass
|
||||
concerns_or_unknowns: none
|
||||
- name: camera_conditioning
|
||||
reuse_or_port: ported
|
||||
parity_test: tests/local_tests/dreamx_world/test_dreamx_world_camera_conditioning_parity.py
|
||||
parity_status: non_skip_pass
|
||||
concerns_or_unknowns: none
|
||||
- name: ar_transformer
|
||||
reuse_or_port: ported
|
||||
parity_test: tests/local_tests/dreamx_world/test_dreamx_world_ar_transformer_parity.py
|
||||
parity_status: non_skip_pass
|
||||
concerns_or_unknowns: none
|
||||
- name: ar_pipeline
|
||||
reuse_or_port: ported
|
||||
parity_test: tests/local_tests/dreamx_world/test_dreamx_world_pipeline_config.py plus AR smoke/SSIM commands listed above
|
||||
parity_status: non_skip_pass
|
||||
concerns_or_unknowns: none
|
||||
pipeline_smoke: pass
|
||||
pipeline_parity: pass
|
||||
example_status: pass
|
||||
quality_regression: added
|
||||
local_tests_readme: tests/local_tests/dreamx_world/README.md
|
||||
port_state_file: tests/local_tests/dreamx_world/PORT_STATUS.md
|
||||
token_values_committed: no
|
||||
runtime_third_party_model_imports: none
|
||||
blockers: none
|
||||
escape_hatch: none
|
||||
```
|
||||
|
||||
| 2026-07-02 | Add DreamX-World-5B autoregressive support. | Official AR repo is raw single-safetensors layout and needs a dedicated causal/KV stage. | Added native AR DiT/config, identity converter, AR pipeline config/preset/registry, and targeted non-skip tests. Raw AR weights are staged at `/tmp/dreamx_world_ar_weights`; converted symlink layout at `/tmp/converted_dreamx_world_ar`. |
|
||||
| 2026-07-02 | Validate DreamX-World-5B AR short full generation on A40. | Targeted parity/config tests prove components, but end-to-end runtime can still fail at scheduler, RoPE cache, device, or decode boundaries. | `DreamXWorldARPipeline` generated and saved a 64x64/9-frame/4-step MP4 from `/tmp/converted_dreamx_world_ar`; the saved MP4 decodes to 9 frames. |
|
||||
| 2026-07-02 | Validate DreamX-World-5B AR long-horizon and default SSIM on A40. | AR needs coverage beyond the 9-frame smoke to exercise longer KV/cache progression and a quality regression path. | 1005-frame 64x64/4-step generation passes and decodes; default AR SSIM test uses 192x192/81 frames because MS-SSIM requires short side >160. |
|
||||
| 2026-07-02 | Validate DreamX-World-5B AR default SSIM on Modal L40S. | CI references are device-specific; A40 alone is not enough for L40S reference coverage. | Modal checked out `post-fix dreamx-world-5b-cam branch commit`, downloaded existing default L40S references, seeded the missing AR reference, reran SSIM successfully (`mean_ssim=1.0`), and exported the L40S reference back to the workspace. |
|
||||
@@ -0,0 +1,254 @@
|
||||
# DreamX World Local Tests
|
||||
|
||||
Local-only parity and smoke tests for the `dreamx_world` FastVideo port. These
|
||||
tests compare FastVideo against the official DreamX-World reference
|
||||
implementation and are not expected to run in CI unless explicitly promoted
|
||||
later.
|
||||
|
||||
Port progress, open questions, issues, and handoff notes live in
|
||||
`tests/local_tests/dreamx_world/PORT_STATUS.md`.
|
||||
|
||||
## Reference Assets
|
||||
|
||||
| Field | Value |
|
||||
|---|---|
|
||||
| Model family | `dreamx_world` |
|
||||
| First-PR scope | `DreamX-World-5B-Cam`; follow-up scope now includes `DreamX-World-5B` autoregressive forcing |
|
||||
| Out-of-scope variants | none for the DreamX-World 5B/Cam paths currently ported |
|
||||
| Workload types | I2V camera-control compatibility shim: image + prompt + action sequence to video |
|
||||
| Official reference | `https://github.com/AMAP-ML/DreamX-World` |
|
||||
| Local reference dir | `DreamX-World/` |
|
||||
| Official commit/version | `221875811ba31f7eac6c3025b215c09ad2cefd1d` |
|
||||
| HF weights | `GD-ML/DreamX-World-5B-Cam` |
|
||||
| HF revision | default |
|
||||
| Local weights dir | `official_weights/dreamx_world` |
|
||||
| Source layout | `raw_official` |
|
||||
| Needs conversion | `yes` |
|
||||
|
||||
Do not write token values in this file. Current token env var detected during
|
||||
prep: `none`.
|
||||
|
||||
## Shared Environment Setup
|
||||
|
||||
Run from the FastVideo repo root in the same conda/env used for FastVideo. Do
|
||||
not create a separate upstream environment for parity tests.
|
||||
|
||||
```bash
|
||||
python ".agents/skills/add-model-01-prep/scripts/clone_reference_repo.py" \
|
||||
"https://github.com/AMAP-ML/DreamX-World.git" \
|
||||
"DreamX-World" \
|
||||
--commit "221875811ba31f7eac6c3025b215c09ad2cefd1d" \
|
||||
--update-gitignore
|
||||
```
|
||||
|
||||
DreamX-World does not expose a packaging file for editable install. During prep
|
||||
the official import check used `sys.path.insert(0, "DreamX-World")`.
|
||||
|
||||
Additional official deps installed into the current environment for imports:
|
||||
|
||||
```bash
|
||||
uv pip install xfuser==0.4.1
|
||||
uv pip install opencv-python-headless
|
||||
```
|
||||
|
||||
These packages are for running the official DreamX reference during local
|
||||
parity only. They must not become FastVideo production/runtime dependencies for
|
||||
the native `dreamx_world` pipeline.
|
||||
|
||||
Do not install the full `DreamX-World/requirements.txt` without explicit
|
||||
approval. It pins core FastVideo stack packages including `torch`, `torchvision`,
|
||||
`triton`, `flash_attn`, and `diffusers`.
|
||||
|
||||
## Official Environment Status
|
||||
|
||||
```text
|
||||
dependency_changes: installed official deps in current env
|
||||
official_env_status: imports_ok
|
||||
private_dep_stubs: none
|
||||
blocked_on: none
|
||||
```
|
||||
|
||||
Import check used during prep:
|
||||
|
||||
```bash
|
||||
python -c "import sys; sys.path.insert(0, 'DreamX-World'); import inference_dreamx5b; print('imports_ok')"
|
||||
```
|
||||
|
||||
## Weight Setup
|
||||
|
||||
HF layout inspection found no root `model_index.json`; the repo contains
|
||||
`config.json`, a safetensors index, and three transformer safetensors shards.
|
||||
This is a raw official transformer layout and requires conversion before
|
||||
FastVideo can load it through `VideoGenerator.from_pretrained`.
|
||||
|
||||
```bash
|
||||
python ".agents/skills/add-model-01-prep/scripts/inspect_hf_layout.py" \
|
||||
"GD-ML/DreamX-World-5B-Cam" \
|
||||
--json
|
||||
```
|
||||
|
||||
Weights have been staged workspace-locally. The raw DreamX transformer repo lives at `official_weights/dreamx_world`; Wan2.2 raw base artifacts live at `official_weights/Wan2.2-TI2V-5B`; Wan2.2 Diffusers reusable components live at `official_weights/Wan2.2-TI2V-5B-Diffusers`. To reproduce the DreamX download:
|
||||
|
||||
```bash
|
||||
python ".agents/skills/add-model-01-prep/scripts/download_hf_weights.py" \
|
||||
"GD-ML/DreamX-World-5B-Cam" \
|
||||
"official_weights/dreamx_world"
|
||||
```
|
||||
|
||||
|
||||
|
||||
### DreamX-World-5B Autoregressive Setup
|
||||
|
||||
The AR repository `GD-ML/DreamX-World-5B` is also raw official layout: no
|
||||
`model_index.json`, root `config.json`, and a single `model.safetensors`. The
|
||||
current environment has the raw AR checkpoint staged outside the workspace at
|
||||
`/tmp/dreamx_world_ar_weights` to avoid workspace quota pressure. The converted
|
||||
FastVideo layout is staged at `/tmp/converted_dreamx_world_ar` with the 21GB
|
||||
transformer safetensors symlinked instead of copied.
|
||||
|
||||
```bash
|
||||
python ".agents/skills/add-model-01-prep/scripts/download_hf_weights.py" \
|
||||
"GD-ML/DreamX-World-5B" \
|
||||
"/tmp/dreamx_world_ar_weights"
|
||||
|
||||
python scripts/checkpoint_conversion/dreamx_world_ar_to_diffusers.py \
|
||||
--source /tmp/dreamx_world_ar_weights \
|
||||
--output /tmp/converted_dreamx_world_ar \
|
||||
--component-source official_weights/Wan2.2-TI2V-5B-Diffusers \
|
||||
--symlink-components \
|
||||
--symlink-transformer
|
||||
```
|
||||
|
||||
AR production code uses `DreamXWorldARTransformer3DModel`,
|
||||
`DreamXWorld5BARPipelineConfig`, `DreamXWorldARPipeline`, and
|
||||
`DreamXWorldARCausalDenoisingStage`. The AR DiT is a native FastVideo port of
|
||||
the official Apache-2.0 `CausalWanModel`; it has no production DreamX, Diffusers
|
||||
model-class, Transformers model-class, `xfuser`, or OpenCV import.
|
||||
|
||||
## Prototype And Conversion Artifacts
|
||||
|
||||
State-dict key/shape dumps are generated after FastVideo native prototypes exist
|
||||
and are used to build the conversion mapping.
|
||||
|
||||
```text
|
||||
official_key_dumps:
|
||||
transformer: converted_weights/dreamx_world/_mapping/transformer_official_keys.json
|
||||
fastvideo_key_dumps:
|
||||
transformer: converted_weights/dreamx_world/_mapping/transformer_fastvideo_keys.json
|
||||
conversion_script: scripts/checkpoint_conversion/dreamx_world_to_diffusers.py
|
||||
conversion_script_status: transformer_model_index_and_config_consistency_smoke_pass
|
||||
conversion_source_layout: raw_official
|
||||
converted_weights_dir: converted_weights/dreamx_world
|
||||
model_index_status: smoke_pass
|
||||
strict_load_status: pass
|
||||
```
|
||||
|
||||
The converter writes `transformer/` from raw DreamX shards. To create a full
|
||||
FastVideo-loadable diffusers-style root, pass a Wan2.2 Diffusers directory as the
|
||||
component source so reusable components are copied or symlinked before
|
||||
`model_index.json` is emitted:
|
||||
|
||||
```bash
|
||||
python scripts/checkpoint_conversion/dreamx_world_to_diffusers.py \
|
||||
--source official_weights/dreamx_world \
|
||||
--output converted_weights/dreamx_world \
|
||||
--component-source /path/to/Wan2.2-TI2V-5B-Diffusers \
|
||||
--symlink-components
|
||||
```
|
||||
|
||||
## Expected Parity Tests
|
||||
|
||||
Planned local tests for this family:
|
||||
|
||||
| Component | Official files / args | Test | Concerns | Status |
|
||||
|---|---|---|---|---|
|
||||
| transformer | `DreamX-World/models/wan_transformer3d.py`; instantiated in `DreamX-World/inference_dreamx5b.py` with `cam_method=prope`, `add_control_adapter=True` | `tests/local_tests/dreamx_world/test_dreamx_world_transformer_parity.py` | FastVideo uses dedicated `DreamXWorldTransformer3DModel`/`DreamXWorldConfig` files for DreamX 5B-Cam config, PRoPE, conversion mapping, real converted 5B strict-load PASS, and official-vs-FastVideo small-input forward parity PASS on CUDA. Wan DiT/config have no DreamX-specific adapter fields. | strict_load_and_forward_parity_pass |
|
||||
| vae | `DreamX-World/models/wan_vae3_8.py`; `vae_type=AutoencoderKLWan3_8`, `vae_subpath=Wan2.2_VAE.pth` | `tests/local_tests/dreamx_world/test_dreamx_world_vae_parity.py` | DreamX raw `Wan2.2_VAE.pth` maps 196/196 keys into FastVideo Wan VAE; encode parity passes after applying official latent normalization. | encode_parity_pass |
|
||||
| text_encoder/tokenizer | `DreamX-World/models/wan_text_encoder.py`; T5 path from Wan2.2 base model | `tests/local_tests/dreamx_world/test_dreamx_world_text_encoder_parity.py` | Official `WanT5EncoderModel` vs FastVideo `UMT5EncoderModel` hidden-state parity passes on CUDA with staged Wan2.2 weights/tokenizer. | hidden_state_parity_pass |
|
||||
| scheduler | Diffusers `FlowMatchEulerDiscreteScheduler`; selected by default `sampler_name=Flow` | `tests/local_tests/dreamx_world/test_dreamx_world_scheduler_parity.py` | First PR can support official default `Flow`; optional `Flow_Unipc` and `Flow_DPM++` are out of first-PR scope. | non_skip_pass |
|
||||
| camera_conditioning | `DreamX-World/utils/inference_utils.py`, `models/prope_utils.py`, `wan/modules/camera_prope.py` | `tests/local_tests/dreamx_world/test_dreamx_world_camera_conditioning_parity.py` | Action sequence to PRoPE/control input must match official tensor shapes and values. | non_skip_pass |
|
||||
| pipeline | `DreamX-World/pipeline/pipeline_dreamxworld.py`; call path in `DreamX-World/inference_dreamx5b.py` | `tests/local_tests/dreamx_world/test_dreamx_world_pipeline_config.py`; `tests/local_tests/pipelines/test_dreamx_world_pipeline_smoke.py`; `tests/local_tests/pipelines/test_dreamx_world_pipeline_parity.py` | DreamX PipelineConfig wires first-scope DiT/VAE/UMT5/Flow/TI2V settings, official `shift=3.0`, default preset values, FlowMatch scheduler initialization, and local `model_index.json` resolution; independent pipeline smoke covers real CUDA local load + latent generation, and parity compares public API output to worker-side explicit ForwardBatch execution. | pipeline_smoke_and_parity_pass |
|
||||
| ar_transformer | `DreamX-World/wan/modules/causal_camera_model_2_2_prope_infinity.py`; instantiated by `DreamX-World/inference_ar_forcing.py` as `CausalWanModel` with `local_attn_size=12`, `sink_size=3`, `attn_compress=4` | `tests/local_tests/dreamx_world/test_dreamx_world_ar_transformer_parity.py` | Native `DreamXWorldARTransformer3DModel` keeps official identity key layout; tiny official-vs-FastVideo forward parity passes; real 5B AR safetensors strict-load passes from `/tmp/converted_dreamx_world_ar`. | tiny_forward_parity_and_real_strict_load_pass |
|
||||
| ar_pipeline | `DreamX-World/pipeline/pipeline_causal_camera.py`; AR block/KV/context-noise loop | `tests/local_tests/dreamx_world/test_dreamx_world_pipeline_config.py`; A40 short full-generation smoke | Dedicated `DreamXWorldARCausalDenoisingStage` implements blockwise KV forcing; raw HF repo must be converted before `VideoGenerator.from_pretrained` because it has no `model_index.json`; converted AR layout generated a 64x64/9-frame/4-step MP4 on A40. | config_registry_and_short_full_generation_pass |
|
||||
|
||||
Include reused components in parity. Reuse is accepted only after the FastVideo
|
||||
component definition and official instantiation arguments have both been checked
|
||||
and the component parity test passes non-skip.
|
||||
|
||||
Run the relevant tests with:
|
||||
|
||||
```bash
|
||||
python -m pytest tests/local_tests/dreamx_world/test_dreamx_world_camera_conditioning_parity.py -v -s
|
||||
python -m pytest tests/local_tests/dreamx_world/test_dreamx_world_scheduler_parity.py -v -s
|
||||
python -m pytest tests/local_tests/dreamx_world/test_dreamx_world_transformer_parity.py -v -s
|
||||
python -m pytest tests/local_tests/dreamx_world/test_dreamx_world_vae_parity.py -v -s
|
||||
python -m pytest tests/local_tests/dreamx_world/test_dreamx_world_text_encoder_parity.py -v -s
|
||||
python -m pytest tests/local_tests/dreamx_world/test_dreamx_world_pipeline_config.py -v -s
|
||||
python -m pytest tests/local_tests/dreamx_world/test_dreamx_world_conversion.py -v -s
|
||||
python -m pytest tests/local_tests/pipelines/test_dreamx_world_pipeline_smoke.py tests/local_tests/pipelines/test_dreamx_world_pipeline_parity.py -q -rs
|
||||
```
|
||||
|
||||
## Current Local Results
|
||||
|
||||
```bash
|
||||
python -m pytest tests/local_tests/dreamx_world/ -v -s
|
||||
# 2026-07-01: 26 passed, 0 skipped
|
||||
# 2026-07-02: AR targeted suite passed: 16 passed, 0 skipped
|
||||
|
||||
PYTHONPATH=/workspace/FastVideo python /tmp/run_dreamx_ar_video_smoke.py
|
||||
# 2026-07-02: DreamX-World-5B AR short full-generation smoke passed on A40
|
||||
# output: outputs_video/dreamx_world_ar_video_smoke/a quiet road through a futuristic coastal city at sunrise.mp4
|
||||
# decoded: 9 frames, (64, 64, 3), uint8
|
||||
|
||||
PYTHONPATH=/workspace/FastVideo FASTVIDEO_ATTENTION_BACKEND=TORCH_SDPA python /tmp/run_dreamx_ar_long_horizon.py
|
||||
# 2026-07-02: DreamX-World-5B AR long-horizon generation passed on A40
|
||||
# output: outputs_video/dreamx_world_ar_long_horizon/A long autonomous drive through a futuristic coastal city at sunrise, smooth forward camera motion,.mp4
|
||||
# decoded: 1005 frames, (64, 64, 3)
|
||||
|
||||
PYTHONPATH=/workspace/FastVideo FASTVIDEO_SSIM_MODEL_ID=DreamX-World-5B DREAMX_WORLD_AR_SSIM_MODEL_PATH=/tmp/converted_dreamx_world_ar python -m pytest fastvideo/tests/ssim/test_dreamx_world_similarity.py -q -rs --skip-ssim-reference-download
|
||||
# 2026-07-02: DreamX-World-5B AR default SSIM passed: 1 passed, 0 skipped
|
||||
# local A40 reference: fastvideo/tests/ssim/reference_videos/default/A40_reference_videos/DreamX-World-5B/TORCH_SDPA/
|
||||
|
||||
modal run /tmp/modal_dreamx_ar_ssim_git.py
|
||||
# 2026-07-02: Modal L40S default SSIM passed: 1 passed, 0 skipped
|
||||
# checked out post-fix dreamx-world-5b-cam branch commit
|
||||
# first downloaded default L40S references with reference_videos_cli.py download
|
||||
# seeded missing AR L40S reference, then reran a fresh generated-vs-reference compare
|
||||
# Modal JSON: mean_ssim=1.0, min_ssim=1.0, max_ssim=1.0
|
||||
# local L40S reference: fastvideo/tests/ssim/reference_videos/default/L40S_reference_videos/DreamX-World-5B/TORCH_SDPA/
|
||||
# A40-vs-L40S reference spot check after deterministic AR noise fix: mean SSIM 0.7770, min 0.6139, max 0.9944
|
||||
|
||||
python -m pytest tests/local_tests/pipelines/test_dreamx_world_pipeline_smoke.py tests/local_tests/pipelines/test_dreamx_world_pipeline_parity.py -q -rs
|
||||
# 2026-06-30: 4 passed, 0 skipped
|
||||
# 2026-07-01: smoke image-path TI2V coverage passed separately with 3 passed, 0 skipped
|
||||
|
||||
PYTHONPATH=/workspace/FastVideo FASTVIDEO_SSIM_MODEL_ID=DreamX-World-5B-Cam python -m pytest fastvideo/tests/ssim/test_dreamx_world_similarity.py -q -rs
|
||||
# 2026-07-01: 1 passed, 0 skipped
|
||||
```
|
||||
|
||||
`camera_conditioning`, default `Flow` scheduler, DreamX component/pipeline configs, default preset, pipeline entry/registry, FlowMatch scheduler initialization, DreamX camera stage, denoising `y_camera` pass-through, conversion model-index/config-consistency checks, real converted 5B dedicated DreamX transformer strict-load and forward parity on CUDA, VAE encode parity on CUDA, text hidden-state parity on CUDA, native DreamX PRoPE branch structure smoke, AR transformer parity/strict-load, AR config/registry, and AR short full-generation smoke are non-skip PASS results. The full local DreamX component suite, independent pipeline smoke/parity suite, image-path TI2V smoke, and SSIM quality regression currently have zero skips. The basic 5B-Cam example was run against `converted_weights/dreamx_world` with a 64x64/9-frame/1-step saved-video smoke, and imageio decoded the generated MP4 successfully; the AR pipeline separately saved and decoded a 64x64/9-frame/4-step MP4 from `/tmp/converted_dreamx_world_ar`.
|
||||
|
||||
## Review Notes
|
||||
|
||||
- Required before handoff: non-skip PASS for each required component parity
|
||||
test, including reused components that own weights or numerical behavior.
|
||||
- First PR scope originally targeted `DreamX-World-5B-Cam`; scope was later
|
||||
expanded to include `DreamX-World-5B` AR support with a separate causal/KV
|
||||
pipeline.
|
||||
- AR support has targeted parity/config coverage, short full-generation smoke,
|
||||
1005-frame A40 long-horizon generation, local A40 default SSIM coverage,
|
||||
and Modal L40S default SSIM coverage. L40S validation was rerun after fixing
|
||||
stale generated-output comparison in the SSIM helper and deterministic AR
|
||||
noise seeding. HF reference publication remains a separate token-gated
|
||||
release operation.
|
||||
- FastVideo production code must not require `xfuser` or OpenCV just because the
|
||||
official reference import needed them. Port camera/action preprocessing and
|
||||
sequence-parallel behavior into existing FastVideo-native utilities or keep
|
||||
reference-only imports inside local parity tests.
|
||||
- Review agents should verify setup commands still match the PR, then run the
|
||||
listed parity tests or report the exact blocker.
|
||||
|
||||
|
||||
## Quality Regression
|
||||
|
||||
Quality regression is added in `fastvideo/tests/ssim/test_dreamx_world_similarity.py`. The 5B-Cam default test uses the workspace converted model root, `TORCH_SDPA`, a deterministic generated conditioning image, 9 frames, 1 denoise step, seed 1024, and min SSIM 0.98. A local A40 reference was seeded under `fastvideo/tests/ssim/reference_videos/default/A40_reference_videos/DreamX-World-5B-Cam/TORCH_SDPA/`, and the test passed non-skip locally on 2026-07-01. The AR default test uses `/tmp/converted_dreamx_world_ar` or `/root/data/dreamx_world_ar_converted`, `TORCH_SDPA`, 192x192, 81 frames, 4 steps, seed 2048, and min SSIM 0.98; local A40 and Modal L40S references were seeded under `fastvideo/tests/ssim/reference_videos/default/{A40,L40S}_reference_videos/DreamX-World-5B/TORCH_SDPA/`, and both tests passed non-skip on 2026-07-02 after fresh generated-output cleanup was added to the SSIM helper. Modal L40S validation checked out `post-fix dreamx-world-5b-cam branch commit`, seeded the missing AR L40S reference, reran the test, and wrote `mean_ssim=1.0`. A cross-device A40-vs-L40S reference spot check produced mean SSIM 0.7770, so CI should use the L40S-specific reference rather than the A40 artifact. Full-quality params are present for 5B-Cam 480x832/161 frames/30 steps and AR 704x1280/1005 frames/4 steps; publishing references to the HF dataset remains a release operation requiring a write-capable HF token env var, never a raw token value.
|
||||
@@ -0,0 +1 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
@@ -0,0 +1,58 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
"""DreamX-World-5B autoregressive conversion smoke tests."""
|
||||
|
||||
import json
|
||||
|
||||
from fastvideo.configs.pipelines.dreamx_world import make_dreamx_world_5b_ar_dit_config
|
||||
from fastvideo.models.registry import _LEGACY_FAST_VIDEO_MODELS
|
||||
from scripts.checkpoint_conversion.dreamx_world_ar_to_diffusers import (
|
||||
MODEL_INDEX,
|
||||
REUSED_COMPONENTS,
|
||||
TRANSFORMER_CONFIG,
|
||||
convert_transformer,
|
||||
write_model_index,
|
||||
)
|
||||
|
||||
|
||||
def test_dreamx_world_ar_converter_writes_symlinked_transformer_and_model_index(tmp_path):
|
||||
source = tmp_path / "raw"
|
||||
source.mkdir()
|
||||
raw_tensor = source / "model.safetensors"
|
||||
raw_tensor.write_bytes(b"placeholder")
|
||||
component_source = tmp_path / "wan22"
|
||||
output = tmp_path / "dreamx_ar"
|
||||
|
||||
for component in REUSED_COMPONENTS:
|
||||
component_dir = component_source / component
|
||||
component_dir.mkdir(parents=True)
|
||||
(component_dir / "config.json").write_text("{}\n")
|
||||
|
||||
convert_transformer(source, output, symlink_transformer=True)
|
||||
write_model_index(output, component_source, symlink_components=True)
|
||||
|
||||
assert (output / "transformer" / "model.safetensors").is_symlink()
|
||||
model_index = json.loads((output / "model_index.json").read_text())
|
||||
assert model_index == MODEL_INDEX
|
||||
assert model_index["_class_name"] == "DreamXWorldARPipeline"
|
||||
assert model_index["transformer"] == ["diffusers", "DreamXWorldARTransformer3DModel"]
|
||||
|
||||
|
||||
def test_dreamx_world_ar_transformer_config_matches_pipeline_dit_config():
|
||||
dit_config = make_dreamx_world_5b_ar_dit_config()
|
||||
assert TRANSFORMER_CONFIG["_class_name"] == "DreamXWorldARTransformer3DModel"
|
||||
assert TRANSFORMER_CONFIG["num_attention_heads"] == dit_config.num_attention_heads
|
||||
assert TRANSFORMER_CONFIG["attention_head_dim"] == dit_config.attention_head_dim
|
||||
assert TRANSFORMER_CONFIG["in_channels"] == dit_config.in_channels
|
||||
assert TRANSFORMER_CONFIG["out_channels"] == dit_config.out_channels
|
||||
assert TRANSFORMER_CONFIG["ffn_dim"] == dit_config.ffn_dim
|
||||
assert TRANSFORMER_CONFIG["num_layers"] == dit_config.num_layers
|
||||
assert TRANSFORMER_CONFIG["local_attn_size"] == dit_config.local_attn_size
|
||||
assert TRANSFORMER_CONFIG["sink_size"] == dit_config.sink_size
|
||||
assert TRANSFORMER_CONFIG["attn_compress"] == dit_config.attn_compress
|
||||
assert tuple(TRANSFORMER_CONFIG["cam_self_attn_layers"]) == dit_config.cam_self_attn_layers
|
||||
|
||||
|
||||
def test_dreamx_world_ar_model_index_component_classes_are_registered():
|
||||
for component in ("scheduler", "text_encoder", "transformer", "vae"):
|
||||
class_name = MODEL_INDEX[component][1]
|
||||
assert class_name in _LEGACY_FAST_VIDEO_MODELS
|
||||
@@ -0,0 +1,182 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
"""DreamX-World-5B autoregressive transformer parity.
|
||||
|
||||
Coverage scope: both. The tiny forward parity compares FastVideo's native AR DiT
|
||||
against the official DreamX ``CausalWanModel`` implementation with identical
|
||||
weights. The real 5B checkpoint gate strict-loads the downloaded safetensors
|
||||
through the FastVideo model class.
|
||||
"""
|
||||
from __future__ import annotations
|
||||
|
||||
import os
|
||||
from pathlib import Path
|
||||
import sys
|
||||
|
||||
import pytest
|
||||
import torch
|
||||
from torch.testing import assert_close
|
||||
|
||||
from fastvideo.configs.models.dits.dreamx_world import DreamXWorldARArchConfig, DreamXWorldARConfig
|
||||
from fastvideo.configs.pipelines.dreamx_world import make_dreamx_world_5b_ar_dit_config
|
||||
from fastvideo.models.dits.dreamx_world_ar import DreamXWorldARTransformer3DModel
|
||||
from fastvideo.models.loader.fsdp_load import load_model_from_full_model_state_dict
|
||||
from fastvideo.models.loader.utils import get_param_names_mapping
|
||||
from fastvideo.models.loader.weight_utils import resolve_safetensors_files, safetensors_weights_iterator
|
||||
|
||||
REPO_ROOT = Path(__file__).resolve().parents[3]
|
||||
OFFICIAL_REF_DIR = Path(os.getenv("DREAMX_WORLD_OFFICIAL_REF_DIR", REPO_ROOT / "DreamX-World"))
|
||||
CONVERTED_AR_DIR = Path(os.getenv("DREAMX_WORLD_AR_CONVERTED_DIR", "/tmp/converted_dreamx_world_ar"))
|
||||
CONVERTED_AR_HF_REPO = "FastVideo/DreamX-World-5B-Diffusers"
|
||||
PARITY_SCOPE = "both"
|
||||
|
||||
|
||||
def _tiny_config() -> DreamXWorldARConfig:
|
||||
return DreamXWorldARConfig(
|
||||
arch_config=DreamXWorldARArchConfig(
|
||||
num_attention_heads=1,
|
||||
attention_head_dim=8,
|
||||
in_channels=4,
|
||||
out_channels=4,
|
||||
ffn_dim=16,
|
||||
num_layers=1,
|
||||
text_dim=8,
|
||||
freq_dim=8,
|
||||
text_len=4,
|
||||
local_attn_size=2,
|
||||
sink_size=1,
|
||||
attn_compress=1,
|
||||
cam_self_attn_layers=(0,),
|
||||
))
|
||||
|
||||
|
||||
def _load_official_tiny():
|
||||
if not OFFICIAL_REF_DIR.exists():
|
||||
pytest.skip(f"Official DreamX reference missing: {OFFICIAL_REF_DIR}")
|
||||
sys.path.insert(0, str(OFFICIAL_REF_DIR))
|
||||
try:
|
||||
from wan.modules import attention as official_attention
|
||||
from wan.modules import causal_camera_model_2_2_prope_infinity as causal_module
|
||||
from wan.modules import model_2_2 as official_model_2_2
|
||||
from wan.modules.causal_camera_model_2_2_prope_infinity import CausalWanModel
|
||||
except Exception as exc: # noqa: BLE001
|
||||
pytest.skip(f"Cannot import official AR transformer: {exc}")
|
||||
official_attention.FLASH_ATTN_2_AVAILABLE = False
|
||||
official_attention.FLASH_ATTN_3_AVAILABLE = False
|
||||
|
||||
def _sdpa_same_dtype(q, k, v, **kwargs):
|
||||
del kwargs
|
||||
out = torch.nn.functional.scaled_dot_product_attention(
|
||||
q.transpose(1, 2), k.transpose(1, 2), v.transpose(1, 2), dropout_p=0.0)
|
||||
return out.transpose(1, 2).contiguous()
|
||||
|
||||
official_attention.attention = _sdpa_same_dtype
|
||||
official_attention.flash_attention = _sdpa_same_dtype
|
||||
official_model_2_2.flash_attention = _sdpa_same_dtype
|
||||
causal_module.attention = _sdpa_same_dtype
|
||||
return CausalWanModel(
|
||||
model_type="ti2v",
|
||||
patch_size=(1, 2, 2),
|
||||
text_len=4,
|
||||
in_dim=4,
|
||||
dim=8,
|
||||
ffn_dim=16,
|
||||
freq_dim=8,
|
||||
text_dim=8,
|
||||
out_dim=4,
|
||||
num_heads=1,
|
||||
num_layers=1,
|
||||
local_attn_size=2,
|
||||
sink_size=1,
|
||||
qk_norm=True,
|
||||
cross_attn_norm=True,
|
||||
add_control_adapter=True,
|
||||
cam_method="prope",
|
||||
attn_compress=1,
|
||||
cam_self_attn_layers=(0,),
|
||||
).eval()
|
||||
|
||||
|
||||
def _make_inputs():
|
||||
torch.manual_seed(123)
|
||||
x = [torch.randn(4, 1, 4, 4)]
|
||||
t = torch.zeros(1, 4, dtype=torch.long)
|
||||
context = [torch.randn(2, 8)]
|
||||
camera = {
|
||||
"viewmats": torch.eye(4).reshape(1, 1, 4, 4).repeat(1, 4, 1, 1),
|
||||
"K": torch.eye(3).reshape(1, 1, 3, 3).repeat(1, 4, 1, 1),
|
||||
}
|
||||
kv_cache = [{
|
||||
"k": torch.zeros(1, 8, 1, 8),
|
||||
"v": torch.zeros(1, 8, 1, 8),
|
||||
"global_end_index": torch.tensor([0]),
|
||||
"local_end_index": torch.tensor([0]),
|
||||
"prope_k": torch.zeros(1, 8, 1, 8),
|
||||
"prope_v": torch.zeros(1, 8, 1, 8),
|
||||
"prope_global_end_index": torch.tensor([0]),
|
||||
"prope_local_end_index": torch.tensor([0]),
|
||||
}]
|
||||
cross_cache = [{
|
||||
"k": torch.zeros(1, 4, 1, 8),
|
||||
"v": torch.zeros(1, 4, 1, 8),
|
||||
"is_init": False,
|
||||
}]
|
||||
return x, t, context, camera, kv_cache, cross_cache
|
||||
|
||||
|
||||
def test_dreamx_world_ar_tiny_forward_matches_official():
|
||||
official = _load_official_tiny()
|
||||
# The official init_weights zero-inits the output head (head.head.weight and
|
||||
# biases), so both models would output exactly zero and the comparison would
|
||||
# pass vacuously. Randomize the head (deterministically) before copying the
|
||||
# state dict so the outputs reflect the internal computation.
|
||||
generator = torch.Generator().manual_seed(7)
|
||||
with torch.no_grad():
|
||||
official.head.head.weight.normal_(std=0.5, generator=generator)
|
||||
official.head.head.bias.normal_(std=0.5, generator=generator)
|
||||
fastvideo = DreamXWorldARTransformer3DModel(_tiny_config(), {}).eval()
|
||||
fastvideo.load_state_dict(official.state_dict(), strict=True)
|
||||
|
||||
x, t, context, camera, kv_cache, cross_cache = _make_inputs()
|
||||
official_out = official(x=x, t=t, context=context, seq_len=4, y_camera=camera,
|
||||
kv_cache=kv_cache, crossattn_cache=cross_cache).detach()
|
||||
x, t, context, camera, kv_cache, cross_cache = _make_inputs()
|
||||
fastvideo_out = fastvideo(x=x, t=t, context=context, seq_len=4, y_camera=camera,
|
||||
kv_cache=kv_cache, crossattn_cache=cross_cache).detach()
|
||||
assert official_out.abs().max() > 0, "official output is all-zero; parity comparison is vacuous"
|
||||
assert_close(fastvideo_out, official_out, atol=1e-5, rtol=1e-5)
|
||||
|
||||
|
||||
def test_dreamx_world_ar_5b_config_matches_official_shape():
|
||||
config = make_dreamx_world_5b_ar_dit_config()
|
||||
assert config.num_layers == 30
|
||||
assert config.num_attention_heads == 24
|
||||
assert config.attention_head_dim == 128
|
||||
assert config.hidden_size == 3072
|
||||
assert config.ffn_dim == 14336
|
||||
assert config.local_attn_size == 12
|
||||
assert config.sink_size == 3
|
||||
assert config.attn_compress == 4
|
||||
assert config.cam_self_attn_layers == tuple(range(30))
|
||||
|
||||
|
||||
def test_dreamx_world_ar_converted_5b_transformer_strict_loads():
|
||||
transformer_dir = CONVERTED_AR_DIR / "transformer"
|
||||
if not transformer_dir.exists():
|
||||
# No local conversion: pull the published Diffusers transformer from the hub.
|
||||
from huggingface_hub import snapshot_download
|
||||
|
||||
transformer_dir = Path(snapshot_download(CONVERTED_AR_HF_REPO, allow_patterns=["transformer/*"])) / "transformer"
|
||||
with torch.device("meta"):
|
||||
model = DreamXWorldARTransformer3DModel(make_dreamx_world_5b_ar_dit_config(), {})
|
||||
incompatible = load_model_from_full_model_state_dict(
|
||||
model,
|
||||
safetensors_weights_iterator(resolve_safetensors_files(str(transformer_dir)), to_cpu=True),
|
||||
device=torch.device("cpu"),
|
||||
param_dtype=torch.bfloat16,
|
||||
strict=True,
|
||||
param_names_mapping=get_param_names_mapping(model.param_names_mapping),
|
||||
training_mode=False,
|
||||
)
|
||||
assert incompatible.missing_keys == []
|
||||
assert incompatible.unexpected_keys == []
|
||||
assert not any(param.is_meta for param in model.parameters())
|
||||
@@ -0,0 +1,95 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
"""DreamX-World camera-conditioning parity against the official reference.
|
||||
|
||||
Coverage scope: implementation_subcomponent. This verifies the weightless
|
||||
action-sequence to PRoPE camera tensor path used by DreamX-World-5B-Cam.
|
||||
"""
|
||||
from __future__ import annotations
|
||||
|
||||
import os
|
||||
from pathlib import Path
|
||||
import pytest
|
||||
import torch
|
||||
from torch.testing import assert_close
|
||||
|
||||
from fastvideo.pipelines.basic.dreamx_world.camera_conditioning import build_dreamx_camera_condition
|
||||
|
||||
REPO_ROOT = Path(__file__).resolve().parents[3]
|
||||
OFFICIAL_REF_DIR = Path(os.getenv("DREAMX_WORLD_OFFICIAL_REF_DIR", REPO_ROOT / "DreamX-World"))
|
||||
PARITY_SCOPE = "implementation_subcomponent"
|
||||
|
||||
|
||||
def _load_official_functions():
|
||||
if not OFFICIAL_REF_DIR.exists():
|
||||
pytest.skip(f"Official reference missing: {OFFICIAL_REF_DIR}")
|
||||
try:
|
||||
import importlib.util
|
||||
|
||||
pose_path = OFFICIAL_REF_DIR / "utils" / "pose_utils.py"
|
||||
pose_spec = importlib.util.spec_from_file_location(
|
||||
"dreamx_world_pose_utils", pose_path)
|
||||
if pose_spec is None or pose_spec.loader is None:
|
||||
raise RuntimeError(f"Cannot load DreamX pose_utils: {pose_path}")
|
||||
pose_module = importlib.util.module_from_spec(pose_spec)
|
||||
pose_spec.loader.exec_module(pose_module)
|
||||
|
||||
source = (OFFICIAL_REF_DIR / "utils" / "inference_utils.py").read_text()
|
||||
source = source.replace(
|
||||
"from .pose_utils import interpolate_camera_poses\n", "")
|
||||
namespace = {"interpolate_camera_poses": pose_module.interpolate_camera_poses}
|
||||
exec(compile(source, str(OFFICIAL_REF_DIR / "utils" / "inference_utils.py"), "exec"), namespace)
|
||||
except Exception as exc: # noqa: BLE001 - local parity should skip missing reference deps.
|
||||
pytest.skip(f"Cannot load DreamX camera reference: {exc}")
|
||||
return namespace["ActionToPoseFromID"], namespace["GetPoseEmbedsFromPosesPrope"]
|
||||
|
||||
|
||||
def _official_camera_condition(
|
||||
action_seq: list[str],
|
||||
action_speed_list: list[float],
|
||||
*,
|
||||
num_frames: int,
|
||||
height: int,
|
||||
width: int,
|
||||
dtype: torch.dtype,
|
||||
):
|
||||
action_to_pose, get_pose_embeds = _load_official_functions()
|
||||
duration = -(-num_frames // len(action_seq))
|
||||
poses = action_to_pose(action_seq, action_speed_list, duration=duration)[:num_frames]
|
||||
condition, _ = get_pose_embeds(poses, height, width, len(poses), False, 0, dtype=dtype, device="cpu")
|
||||
return condition
|
||||
|
||||
|
||||
@pytest.mark.parametrize(
|
||||
("action_seq", "action_speed_list", "num_frames"),
|
||||
[
|
||||
(["w"], [4], 81),
|
||||
(["wj", "d"], [4, 6], 121),
|
||||
(["i", "k", "l"], [3, 5, 2], 85),
|
||||
],
|
||||
)
|
||||
def test_dreamx_world_camera_conditioning_matches_official(action_seq, action_speed_list, num_frames):
|
||||
dtype = torch.float32
|
||||
official = _official_camera_condition(
|
||||
action_seq,
|
||||
action_speed_list,
|
||||
num_frames=num_frames,
|
||||
height=704,
|
||||
width=1280,
|
||||
dtype=dtype,
|
||||
)
|
||||
fastvideo = build_dreamx_camera_condition(
|
||||
action_seq,
|
||||
action_speed_list,
|
||||
num_frames=num_frames,
|
||||
height=704,
|
||||
width=1280,
|
||||
dtype=dtype,
|
||||
device="cpu",
|
||||
)
|
||||
|
||||
assert official.keys() == fastvideo.keys() == {"viewmats", "K"}
|
||||
for key in ("viewmats", "K"):
|
||||
assert official[key].shape == fastvideo[key].shape
|
||||
diff = (official[key] - fastvideo[key]).abs()
|
||||
print(f"{key}: shape={tuple(fastvideo[key].shape)} diff_max={diff.max().item():.8f} diff_mean={diff.mean().item():.8f}")
|
||||
assert_close(fastvideo[key], official[key], atol=1e-5, rtol=1e-5)
|
||||
@@ -0,0 +1,79 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
"""DreamX-World conversion script smoke tests."""
|
||||
|
||||
import json
|
||||
|
||||
from fastvideo.configs.pipelines.dreamx_world import make_dreamx_world_5b_cam_dit_config
|
||||
from fastvideo.models.registry import _LEGACY_FAST_VIDEO_MODELS
|
||||
|
||||
from scripts.checkpoint_conversion.dreamx_world_to_diffusers import (
|
||||
MODEL_INDEX,
|
||||
REUSED_COMPONENTS,
|
||||
TRANSFORMER_CONFIG,
|
||||
_copy_or_link_component,
|
||||
write_model_index,
|
||||
)
|
||||
|
||||
|
||||
def test_dreamx_world_converter_writes_full_model_index_with_reused_components(tmp_path):
|
||||
component_source = tmp_path / "wan22"
|
||||
output = tmp_path / "dreamx"
|
||||
output.mkdir()
|
||||
|
||||
for component in REUSED_COMPONENTS:
|
||||
component_dir = component_source / component
|
||||
component_dir.mkdir(parents=True)
|
||||
(component_dir / "config.json").write_text("{}\n")
|
||||
|
||||
write_model_index(output, component_source, symlink_components=True)
|
||||
|
||||
model_index = json.loads((output / "model_index.json").read_text())
|
||||
assert model_index == MODEL_INDEX
|
||||
assert model_index["_class_name"] == "DreamXWorldPipeline"
|
||||
assert model_index["transformer"] == ["diffusers", "DreamXWorldTransformer3DModel"]
|
||||
for component in REUSED_COMPONENTS:
|
||||
assert (output / component).is_symlink()
|
||||
|
||||
|
||||
def test_dreamx_world_transformer_config_has_camera_adapter_enabled():
|
||||
assert TRANSFORMER_CONFIG["_class_name"] == "DreamXWorldTransformer3DModel"
|
||||
assert TRANSFORMER_CONFIG["add_control_adapter"] is True
|
||||
assert TRANSFORMER_CONFIG["cam_method"] == "prope"
|
||||
assert TRANSFORMER_CONFIG["num_layers"] == 30
|
||||
|
||||
|
||||
def test_dreamx_world_converter_transformer_config_matches_pipeline_dit_config():
|
||||
dit_config = make_dreamx_world_5b_cam_dit_config()
|
||||
|
||||
assert TRANSFORMER_CONFIG["num_attention_heads"] == dit_config.num_attention_heads
|
||||
assert TRANSFORMER_CONFIG["attention_head_dim"] == dit_config.attention_head_dim
|
||||
assert TRANSFORMER_CONFIG["in_channels"] == dit_config.in_channels
|
||||
assert TRANSFORMER_CONFIG["out_channels"] == dit_config.out_channels
|
||||
assert TRANSFORMER_CONFIG["ffn_dim"] == dit_config.ffn_dim
|
||||
assert TRANSFORMER_CONFIG["num_layers"] == dit_config.num_layers
|
||||
assert TRANSFORMER_CONFIG["cross_attn_norm"] == dit_config.cross_attn_norm
|
||||
assert TRANSFORMER_CONFIG["qk_norm"] == dit_config.qk_norm
|
||||
assert TRANSFORMER_CONFIG["add_control_adapter"] == dit_config.add_control_adapter
|
||||
assert TRANSFORMER_CONFIG["cam_method"] == dit_config.cam_method
|
||||
assert TRANSFORMER_CONFIG["attn_compress"] == dit_config.attn_compress
|
||||
assert TRANSFORMER_CONFIG["cam_self_attn_layers"] == dit_config.cam_self_attn_layers
|
||||
|
||||
|
||||
def test_dreamx_world_model_index_component_classes_are_registered():
|
||||
for component in ("scheduler", "text_encoder", "transformer", "vae"):
|
||||
class_name = MODEL_INDEX[component][1]
|
||||
assert class_name in _LEGACY_FAST_VIDEO_MODELS
|
||||
|
||||
|
||||
def test_dreamx_world_copy_or_link_component_keeps_broken_symlink(tmp_path):
|
||||
component_source = tmp_path / "wan22"
|
||||
src = component_source / "scheduler"
|
||||
src.mkdir(parents=True)
|
||||
output = tmp_path / "dreamx"
|
||||
output.mkdir()
|
||||
dst = output / "scheduler"
|
||||
dst.symlink_to(tmp_path / "missing_scheduler", target_is_directory=True)
|
||||
|
||||
_copy_or_link_component("scheduler", component_source, output, symlink=True)
|
||||
|
||||
assert dst.is_symlink()
|
||||
@@ -0,0 +1,306 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
"""DreamX-World pipeline config and conditioning smoke tests."""
|
||||
|
||||
from types import SimpleNamespace
|
||||
import json
|
||||
|
||||
import numpy as np
|
||||
import torch
|
||||
|
||||
from fastvideo.api.presets import get_preset
|
||||
from fastvideo.fastvideo_args import WorkloadType
|
||||
from fastvideo.models.schedulers.scheduling_flow_match_euler_discrete import FlowMatchEulerDiscreteScheduler
|
||||
from fastvideo.pipelines.pipeline_batch_info import ForwardBatch
|
||||
from fastvideo.pipelines.pipeline_registry import PipelineType, import_pipeline_classes
|
||||
from fastvideo.pipelines.basic.dreamx_world.camera_conditioning import (
|
||||
DreamXCamera,
|
||||
_interpolate_camera_poses,
|
||||
build_dreamx_camera_condition,
|
||||
)
|
||||
from fastvideo.configs.pipelines.dreamx_world import DreamXWorld5BARPipelineConfig, DreamXWorld5BCamPipelineConfig
|
||||
from fastvideo.pipelines.basic.dreamx_world.ar_denoising import DreamXWorldARCausalDenoisingStage
|
||||
from fastvideo.pipelines.basic.dreamx_world.dreamx_world_ar_pipeline import DreamXWorldARPipeline
|
||||
from fastvideo.pipelines.basic.dreamx_world.dreamx_world_pipeline import DreamXWorldPipeline
|
||||
from fastvideo.pipelines.basic.dreamx_world.stages import (
|
||||
DREAMX_Y_CAMERA_KEY,
|
||||
DreamXWorldCameraConditioningStage,
|
||||
)
|
||||
from fastvideo.pipelines.stages.denoising import DenoisingStage
|
||||
from fastvideo.registry import get_default_preset, get_model_info, get_pipeline_config_cls_from_name
|
||||
|
||||
|
||||
def test_dreamx_world_5b_cam_pipeline_config_wires_first_scope_components():
|
||||
config = DreamXWorld5BCamPipelineConfig()
|
||||
|
||||
assert config.flow_shift == 3.0
|
||||
assert config.ti2v_task is True
|
||||
assert config.expand_timesteps is True
|
||||
assert config.dit_config.expand_timesteps is True
|
||||
assert config.dit_config.num_layers == 30
|
||||
assert config.dit_config.add_control_adapter is True
|
||||
assert config.dit_config.cam_method == "prope"
|
||||
|
||||
assert config.vae_config.load_encoder is True
|
||||
assert config.vae_config.load_decoder is True
|
||||
assert config.vae_config.z_dim == 48
|
||||
assert config.vae_config.scale_factor_temporal == 4
|
||||
assert config.vae_config.scale_factor_spatial == 16
|
||||
|
||||
assert len(config.text_encoder_configs) == 1
|
||||
text_config = config.text_encoder_configs[0]
|
||||
assert text_config.prefix == "umt5"
|
||||
assert text_config.vocab_size == 256384
|
||||
assert text_config.d_model == 4096
|
||||
assert config.text_encoder_precisions == ("bf16",)
|
||||
|
||||
|
||||
def test_dreamx_world_pipeline_registry_discovers_entrypoint():
|
||||
pipelines = import_pipeline_classes(PipelineType.BASIC)
|
||||
|
||||
assert pipelines["basic"]["DreamXWorldPipeline"] is DreamXWorldPipeline
|
||||
|
||||
|
||||
|
||||
|
||||
def test_dreamx_world_local_model_index_resolves_model_info(tmp_path):
|
||||
model_dir = tmp_path / "DreamX-World-5B-Cam-converted"
|
||||
model_dir.mkdir()
|
||||
for component in ("scheduler", "text_encoder", "tokenizer", "transformer", "vae"):
|
||||
(model_dir / component).mkdir()
|
||||
(model_dir / "model_index.json").write_text(json.dumps({
|
||||
"_class_name": "DreamXWorldPipeline",
|
||||
"_diffusers_version": "0.31.0",
|
||||
"scheduler": ["diffusers", "FlowMatchEulerDiscreteScheduler"],
|
||||
"text_encoder": ["transformers", "UMT5EncoderModel"],
|
||||
"tokenizer": ["transformers", "AutoTokenizer"],
|
||||
"transformer": ["diffusers", "DreamXWorldTransformer3DModel"],
|
||||
"vae": ["diffusers", "AutoencoderKLWan"],
|
||||
}) + "\n")
|
||||
|
||||
info = get_model_info(str(model_dir), pipeline_type=PipelineType.BASIC, workload_type=WorkloadType.I2V)
|
||||
|
||||
assert info.pipeline_cls is DreamXWorldPipeline
|
||||
assert info.pipeline_config_cls is DreamXWorld5BCamPipelineConfig
|
||||
|
||||
def test_dreamx_world_model_path_resolves_pipeline_config():
|
||||
assert get_pipeline_config_cls_from_name("GD-ML/DreamX-World-5B-Cam") is DreamXWorld5BCamPipelineConfig
|
||||
|
||||
|
||||
|
||||
def test_dreamx_world_default_preset_is_registered():
|
||||
preset_name = get_default_preset("GD-ML/DreamX-World-5B-Cam")
|
||||
preset = get_preset(preset_name, "dreamx_world")
|
||||
|
||||
assert preset.name == "dreamx_world_5b_cam"
|
||||
assert preset.workload_type == "i2v"
|
||||
assert preset.defaults["height"] == 480
|
||||
assert preset.defaults["width"] == 832
|
||||
assert preset.defaults["num_frames"] == 161
|
||||
assert preset.defaults["num_inference_steps"] == 30
|
||||
assert preset.defaults["guidance_scale"] == 5.0
|
||||
|
||||
|
||||
|
||||
|
||||
def test_dreamx_world_pipeline_initializes_official_flow_scheduler():
|
||||
pipeline = DreamXWorldPipeline.__new__(DreamXWorldPipeline)
|
||||
pipeline.modules = {}
|
||||
fastvideo_args = SimpleNamespace(pipeline_config=DreamXWorld5BCamPipelineConfig())
|
||||
|
||||
pipeline.initialize_pipeline(fastvideo_args)
|
||||
|
||||
scheduler = pipeline.modules["scheduler"]
|
||||
assert isinstance(scheduler, FlowMatchEulerDiscreteScheduler)
|
||||
assert scheduler.config.shift == 3.0
|
||||
|
||||
def test_dreamx_world_camera_conditioning_stage_sets_y_camera_extra():
|
||||
batch = ForwardBatch(
|
||||
data_type="t2v",
|
||||
action_list=["wj", "d"],
|
||||
action_speed_list=[4, 6],
|
||||
num_frames=17,
|
||||
height=704,
|
||||
width=1280,
|
||||
latents=torch.zeros(1, 16, 5, 44, 80),
|
||||
)
|
||||
stage = DreamXWorldCameraConditioningStage()
|
||||
|
||||
out = stage.forward(batch, fastvideo_args=object())
|
||||
y_camera = out.extra[DREAMX_Y_CAMERA_KEY]
|
||||
expected = build_dreamx_camera_condition(
|
||||
["wj", "d"],
|
||||
[4, 6],
|
||||
num_frames=17,
|
||||
height=704,
|
||||
width=1280,
|
||||
dtype=torch.float32,
|
||||
device="cpu",
|
||||
)
|
||||
|
||||
assert set(y_camera) == {"viewmats", "K"}
|
||||
for key, expected_value in expected.items():
|
||||
assert y_camera[key].shape == (1, *expected_value.shape)
|
||||
torch.testing.assert_close(y_camera[key][0], expected_value)
|
||||
assert stage.verify_output(out, fastvideo_args=object()).is_valid()
|
||||
|
||||
|
||||
def test_dreamx_world_denoising_kwargs_filter_for_y_camera():
|
||||
y_camera = {"viewmats": torch.eye(4).reshape(1, 1, 4, 4), "K": torch.eye(3).reshape(1, 1, 3, 3)}
|
||||
stage = DenoisingStage.__new__(DenoisingStage)
|
||||
|
||||
def accepts_y_camera(hidden_states, encoder_hidden_states, timestep, y_camera=None):
|
||||
return y_camera
|
||||
|
||||
def no_y_camera(hidden_states, encoder_hidden_states, timestep):
|
||||
return hidden_states
|
||||
|
||||
assert stage.prepare_extra_func_kwargs(accepts_y_camera, {"y_camera": y_camera}) == {"y_camera": y_camera}
|
||||
assert stage.prepare_extra_func_kwargs(no_y_camera, {"y_camera": y_camera}) == {}
|
||||
|
||||
|
||||
|
||||
def test_dreamx_world_ar_pipeline_config_wires_components():
|
||||
config = DreamXWorld5BARPipelineConfig()
|
||||
assert config.is_causal is True
|
||||
assert config.flow_shift == 5.0
|
||||
assert config.dmd_denoising_steps == (1000, 750, 500, 250)
|
||||
assert config.warp_denoising_step is True
|
||||
assert config.context_noise == 0.1
|
||||
assert config.dit_config.arch_config.local_attn_size == 12
|
||||
assert config.dit_config.arch_config.sink_size == 3
|
||||
assert config.dit_config.arch_config.attn_compress == 4
|
||||
|
||||
|
||||
def test_dreamx_world_ar_pipeline_registry_and_preset():
|
||||
from fastvideo.api.presets import get_preset
|
||||
from fastvideo.registry import get_pipeline_config_cls_from_name, get_preset_selection
|
||||
|
||||
assert get_pipeline_config_cls_from_name("GD-ML/DreamX-World-5B") is DreamXWorld5BARPipelineConfig
|
||||
preset_name, family = get_preset_selection("GD-ML/DreamX-World-5B")
|
||||
assert (preset_name, family) == ("dreamx_world_5b_ar", "dreamx_world")
|
||||
preset = get_preset("dreamx_world_5b_ar", "dreamx_world")
|
||||
assert preset.defaults["num_inference_steps"] == 4
|
||||
assert DreamXWorldARPipeline.pipeline_config_cls is DreamXWorld5BARPipelineConfig
|
||||
|
||||
|
||||
def test_dreamx_world_camera_conditioning_stage_expands_scalar_speed():
|
||||
batch = ForwardBatch(
|
||||
data_type="t2v",
|
||||
action_list=["w", "d"],
|
||||
action_speed_list=2.0,
|
||||
num_frames=17,
|
||||
height=704,
|
||||
width=1280,
|
||||
latents=torch.zeros(1, 16, 5, 44, 80),
|
||||
)
|
||||
stage = DreamXWorldCameraConditioningStage()
|
||||
|
||||
out = stage.forward(batch, fastvideo_args=object())
|
||||
|
||||
assert set(out.extra[DREAMX_Y_CAMERA_KEY]) == {"viewmats", "K"}
|
||||
|
||||
|
||||
def test_dreamx_world_camera_interpolation_handles_single_camera():
|
||||
camera = DreamXCamera(
|
||||
fx=0.8,
|
||||
fy=0.8,
|
||||
cx=0.5,
|
||||
cy=0.5,
|
||||
w2c_mat=np.eye(4, dtype=np.float64),
|
||||
)
|
||||
|
||||
out = _interpolate_camera_poses(
|
||||
[camera],
|
||||
src_indices=np.array([0.0]),
|
||||
tgt_indices=np.array([0.0, 1.0, 2.0]),
|
||||
)
|
||||
|
||||
assert out == [camera, camera, camera]
|
||||
|
||||
|
||||
def test_dreamx_world_ar_cache_initializes_camera_self_attention_entries():
|
||||
cam_self_attn = SimpleNamespace(num_heads=3, head_dim=5)
|
||||
transformer = SimpleNamespace(
|
||||
blocks=[SimpleNamespace(cam_self_attn=cam_self_attn), SimpleNamespace(cam_self_attn=cam_self_attn)],
|
||||
num_attention_heads=2,
|
||||
attention_head_dim=4,
|
||||
)
|
||||
stage = DreamXWorldARCausalDenoisingStage.__new__(DreamXWorldARCausalDenoisingStage)
|
||||
stage.transformer = transformer
|
||||
stage.num_transformer_blocks = 2
|
||||
stage.local_attn_size = 6
|
||||
|
||||
caches = stage._initialize_kv_cache(
|
||||
batch_size=1,
|
||||
dtype=torch.float32,
|
||||
device=torch.device("cpu"),
|
||||
frame_seq_length=7,
|
||||
)
|
||||
|
||||
assert len(caches) == 2
|
||||
assert caches[0]["k"].shape == (1, 42, 2, 4)
|
||||
assert caches[0]["prope_k"].shape == (1, 42, 3, 5)
|
||||
assert caches[0]["prope_v"].shape == (1, 42, 3, 5)
|
||||
assert int(caches[0]["prope_global_end_index"].item()) == 0
|
||||
assert int(caches[0]["prope_local_end_index"].item()) == 0
|
||||
|
||||
|
||||
def test_dreamx_world_ar_context_noise_fraction_maps_to_scheduler_timestep():
|
||||
assert DreamXWorldARCausalDenoisingStage._context_noise_timestep(0.1) == 100
|
||||
assert DreamXWorldARCausalDenoisingStage._context_noise_timestep(100) == 100
|
||||
|
||||
|
||||
|
||||
def test_dreamx_world_ar_context_update_advances_camera_cache_indices():
|
||||
class DummyTransformer:
|
||||
def __call__(self, *, hidden_states, encoder_hidden_states, timestep, y_camera, kv_cache, crossattn_cache,
|
||||
current_start):
|
||||
del encoder_hidden_states, y_camera, crossattn_cache
|
||||
assert current_start == 0
|
||||
assert timestep.unique().tolist() == [100]
|
||||
new_tokens = timestep.shape[1]
|
||||
for cache in kv_cache:
|
||||
cache["local_end_index"] += new_tokens
|
||||
cache["global_end_index"] += new_tokens
|
||||
cache["prope_local_end_index"] += new_tokens
|
||||
cache["prope_global_end_index"] += new_tokens
|
||||
cache["k"][:, :new_tokens] = 1
|
||||
cache["prope_k"][:, :new_tokens] = 1
|
||||
return hidden_states
|
||||
|
||||
cam_self_attn = SimpleNamespace(num_heads=3, head_dim=5)
|
||||
cache_transformer = SimpleNamespace(
|
||||
blocks=[SimpleNamespace(cam_self_attn=cam_self_attn)],
|
||||
num_attention_heads=2,
|
||||
attention_head_dim=4,
|
||||
)
|
||||
stage = DreamXWorldARCausalDenoisingStage.__new__(DreamXWorldARCausalDenoisingStage)
|
||||
stage.transformer = cache_transformer
|
||||
stage.num_transformer_blocks = 1
|
||||
stage.local_attn_size = 6
|
||||
caches = stage._initialize_kv_cache(
|
||||
batch_size=1,
|
||||
dtype=torch.float32,
|
||||
device=torch.device("cpu"),
|
||||
frame_seq_length=2,
|
||||
)
|
||||
# Keep the cache allocation source separate from the callable transformer used by _update_context_cache.
|
||||
stage.transformer = DummyTransformer()
|
||||
|
||||
stage._update_context_cache(
|
||||
block_latents=torch.zeros(1, 4, 3, 2, 2),
|
||||
context=[torch.zeros(2, 4)],
|
||||
camera_block={"viewmats": torch.eye(4).reshape(1, 1, 4, 4), "K": torch.eye(3).reshape(1, 1, 3, 3)},
|
||||
kv_cache=caches,
|
||||
crossattn_cache=[{}],
|
||||
start=0,
|
||||
frame_seq_length=2,
|
||||
target_dtype=torch.float32,
|
||||
autocast_enabled=False,
|
||||
context_noise=0.1,
|
||||
)
|
||||
|
||||
assert int(caches[0]["local_end_index"].item()) == 6
|
||||
assert int(caches[0]["prope_local_end_index"].item()) == 6
|
||||
assert caches[0]["k"][:, :6].sum().item() == 48
|
||||
assert caches[0]["prope_k"][:, :6].sum().item() == 90
|
||||
@@ -0,0 +1,53 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
"""DreamX-World default Flow scheduler parity.
|
||||
|
||||
Coverage scope: implementation_subcomponent. DreamX-World-5B-Cam defaults to
|
||||
Diffusers FlowMatchEulerDiscreteScheduler for sampler_name=Flow. This test
|
||||
checks that FastVideo's native FlowMatchEulerDiscreteScheduler matches the
|
||||
timestep schedule and Euler step used by the official default path.
|
||||
"""
|
||||
from __future__ import annotations
|
||||
|
||||
from pathlib import Path
|
||||
import inspect
|
||||
|
||||
import torch
|
||||
from diffusers import FlowMatchEulerDiscreteScheduler as OfficialFlowMatchEulerDiscreteScheduler
|
||||
from omegaconf import OmegaConf
|
||||
from torch.testing import assert_close
|
||||
|
||||
from fastvideo.models.schedulers.scheduling_flow_match_euler_discrete import (
|
||||
FlowMatchEulerDiscreteScheduler as FastVideoFlowMatchEulerDiscreteScheduler,
|
||||
)
|
||||
|
||||
REPO_ROOT = Path(__file__).resolve().parents[3]
|
||||
PARITY_SCOPE = "implementation_subcomponent"
|
||||
|
||||
|
||||
def _scheduler_kwargs(cls):
|
||||
config_path = REPO_ROOT / "DreamX-World" / "configs" / "wan2.2" / "wan_ti2v_5b.yaml"
|
||||
config = OmegaConf.load(config_path)
|
||||
raw_kwargs = OmegaConf.to_container(config["scheduler_kwargs"])
|
||||
signature = inspect.signature(cls)
|
||||
return {key: value for key, value in raw_kwargs.items() if key in signature.parameters}
|
||||
|
||||
|
||||
def test_dreamx_world_flow_scheduler_timesteps_and_step_match():
|
||||
official = OfficialFlowMatchEulerDiscreteScheduler(**_scheduler_kwargs(OfficialFlowMatchEulerDiscreteScheduler))
|
||||
fastvideo = FastVideoFlowMatchEulerDiscreteScheduler(**_scheduler_kwargs(FastVideoFlowMatchEulerDiscreteScheduler))
|
||||
|
||||
official.set_timesteps(50, device="cpu", mu=1)
|
||||
fastvideo.set_timesteps(50, device="cpu", mu=1)
|
||||
assert_close(fastvideo.timesteps, official.timesteps, atol=0, rtol=0)
|
||||
assert_close(fastvideo.sigmas, official.sigmas, atol=0, rtol=0)
|
||||
|
||||
torch.manual_seed(7)
|
||||
sample = torch.randn(1, 4, 2, 8, 8)
|
||||
model_output = torch.randn_like(sample)
|
||||
timestep = official.timesteps[3]
|
||||
|
||||
official_prev = official.step(model_output, timestep, sample, return_dict=False)[0]
|
||||
fastvideo_prev = fastvideo.step(model_output, fastvideo.timesteps[3], sample, return_dict=False)[0]
|
||||
diff = (official_prev - fastvideo_prev).abs()
|
||||
print(f"scheduler diff_max={diff.max().item():.8f} diff_mean={diff.mean().item():.8f}")
|
||||
assert_close(fastvideo_prev, official_prev, atol=0, rtol=0)
|
||||
@@ -0,0 +1,159 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
"""DreamX-World Wan T5 encoder reuse parity scaffold.
|
||||
|
||||
Coverage scope: implementation_subcomponent. It records the official
|
||||
WanT5EncoderModel loading path and FastVideo T5 target for later activation
|
||||
with staged Wan2.2 base text encoder/tokenizer weights.
|
||||
"""
|
||||
from __future__ import annotations
|
||||
|
||||
import os
|
||||
from pathlib import Path
|
||||
import sys
|
||||
import types
|
||||
|
||||
import pytest
|
||||
import torch
|
||||
from omegaconf import OmegaConf
|
||||
from torch.testing import assert_close
|
||||
from transformers import AutoTokenizer
|
||||
|
||||
from fastvideo.fastvideo_args import FastVideoArgs
|
||||
from fastvideo.configs.pipelines.dreamx_world import (
|
||||
DreamXWorld5BCamPipelineConfig,
|
||||
make_dreamx_world_5b_cam_text_encoder_config,
|
||||
)
|
||||
from fastvideo.models.loader.component_loader import TextEncoderLoader
|
||||
|
||||
REPO_ROOT = Path(__file__).resolve().parents[3]
|
||||
OFFICIAL_REF_DIR = Path(os.getenv("DREAMX_WORLD_OFFICIAL_REF_DIR", REPO_ROOT / "DreamX-World"))
|
||||
WAN_BASE_DIR = Path(os.getenv("DREAMX_WORLD_WAN_BASE_DIR", REPO_ROOT / "official_weights" / "Wan2.2-TI2V-5B"))
|
||||
WAN_DIFFUSERS_DIR = Path(os.getenv("DREAMX_WORLD_WAN_DIFFUSERS_DIR", REPO_ROOT / "official_weights" / "Wan2.2-TI2V-5B-Diffusers"))
|
||||
PARITY_SCOPE = "implementation_subcomponent"
|
||||
|
||||
|
||||
def _add_official_to_path():
|
||||
if not OFFICIAL_REF_DIR.exists():
|
||||
pytest.skip(f"Official reference missing: {OFFICIAL_REF_DIR}")
|
||||
if str(OFFICIAL_REF_DIR) not in sys.path:
|
||||
sys.path.insert(0, str(OFFICIAL_REF_DIR))
|
||||
|
||||
|
||||
def _install_xfuser_stub() -> None:
|
||||
if "xfuser" in sys.modules:
|
||||
return
|
||||
xfuser = types.ModuleType("xfuser")
|
||||
core = types.ModuleType("xfuser.core")
|
||||
distributed = types.ModuleType("xfuser.core.distributed")
|
||||
long_ctx = types.ModuleType("xfuser.core.long_ctx_attention")
|
||||
distributed.get_sequence_parallel_rank = lambda: 0
|
||||
distributed.get_sequence_parallel_world_size = lambda: 1
|
||||
distributed.get_sp_group = lambda: None
|
||||
distributed.get_world_group = lambda: types.SimpleNamespace(local_rank=0, rank=0)
|
||||
distributed.init_distributed_environment = lambda *args, **kwargs: None
|
||||
distributed.initialize_model_parallel = lambda *args, **kwargs: None
|
||||
distributed.model_parallel_is_initialized = lambda: False
|
||||
|
||||
class XFuserLongContextAttention:
|
||||
def __call__(self, *args, **kwargs):
|
||||
raise RuntimeError("xfuser stub cannot execute attention")
|
||||
|
||||
long_ctx.xFuserLongContextAttention = XFuserLongContextAttention
|
||||
sys.modules.update({
|
||||
"xfuser": xfuser,
|
||||
"xfuser.core": core,
|
||||
"xfuser.core.distributed": distributed,
|
||||
"xfuser.core.long_ctx_attention": long_ctx,
|
||||
})
|
||||
|
||||
|
||||
def _text_kwargs():
|
||||
config = OmegaConf.load(OFFICIAL_REF_DIR / "configs" / "wan2.2" / "wan_ti2v_5b.yaml")
|
||||
return OmegaConf.to_container(config["text_encoder_kwargs"])
|
||||
|
||||
|
||||
def _patch_single_process_text_parallel(monkeypatch):
|
||||
import fastvideo.layers.linear as fastvideo_linear
|
||||
import fastvideo.layers.vocab_parallel_embedding as fastvideo_embedding
|
||||
import fastvideo.models.encoders.t5 as fastvideo_t5
|
||||
|
||||
for module in (fastvideo_t5, fastvideo_embedding, fastvideo_linear):
|
||||
if hasattr(module, "get_tp_rank"):
|
||||
monkeypatch.setattr(module, "get_tp_rank", lambda: 0)
|
||||
if hasattr(module, "get_tp_world_size"):
|
||||
monkeypatch.setattr(module, "get_tp_world_size", lambda: 1)
|
||||
monkeypatch.setattr(fastvideo_embedding, "tensor_model_parallel_all_reduce", lambda x: x)
|
||||
|
||||
|
||||
def _load_official_text_encoder(device, dtype):
|
||||
_add_official_to_path()
|
||||
_install_xfuser_stub()
|
||||
text_path = WAN_BASE_DIR / "models_t5_umt5-xxl-enc-bf16.pth"
|
||||
if not text_path.exists():
|
||||
pytest.skip(f"Wan2.2 text encoder weights missing: {text_path}")
|
||||
try:
|
||||
from models import WanT5EncoderModel
|
||||
except Exception as exc: # noqa: BLE001
|
||||
pytest.skip(f"Cannot import official DreamX text encoder: {exc}")
|
||||
model = WanT5EncoderModel.from_pretrained(
|
||||
str(text_path), additional_kwargs=_text_kwargs(), low_cpu_mem_usage=True, torch_dtype=dtype
|
||||
)
|
||||
return model.to(device=device, dtype=dtype).eval()
|
||||
|
||||
|
||||
def _load_fastvideo_text_encoder(device, dtype, monkeypatch):
|
||||
text_encoder_path = WAN_DIFFUSERS_DIR / "text_encoder"
|
||||
if not text_encoder_path.exists():
|
||||
pytest.skip(f"Wan2.2 Diffusers text encoder missing: {text_encoder_path}")
|
||||
_patch_single_process_text_parallel(monkeypatch)
|
||||
pipeline_config = DreamXWorld5BCamPipelineConfig()
|
||||
pipeline_config.text_encoder_configs[0]._fsdp_shard_conditions = []
|
||||
args = FastVideoArgs(
|
||||
model_path=str(text_encoder_path),
|
||||
pipeline_config=pipeline_config,
|
||||
text_encoder_cpu_offload=(device.type == "cpu"),
|
||||
)
|
||||
args.model_paths = {}
|
||||
return TextEncoderLoader().load(str(text_encoder_path), args).to(device=device, dtype=dtype).eval()
|
||||
|
||||
|
||||
def test_dreamx_world_text_encoder_config_matches_umt5_xxl_shape():
|
||||
config = make_dreamx_world_5b_cam_text_encoder_config()
|
||||
assert config.vocab_size == 256384
|
||||
assert config.d_model == 4096
|
||||
assert config.d_kv == 64
|
||||
assert config.d_ff == 10240
|
||||
assert config.num_heads == 64
|
||||
assert config.num_layers == 24
|
||||
assert config.relative_attention_num_buckets == 32
|
||||
assert config.dropout_rate == 0.0
|
||||
assert config.text_len == 512
|
||||
assert config.prefix == "umt5"
|
||||
|
||||
|
||||
def test_dreamx_world_fastvideo_text_encoder_loads_staged_weights(monkeypatch):
|
||||
device = torch.device("cuda:0" if torch.cuda.is_available() else "cpu")
|
||||
model = _load_fastvideo_text_encoder(device, torch.bfloat16, monkeypatch)
|
||||
assert model.__class__.__name__ == "UMT5EncoderModel"
|
||||
assert next(model.parameters()).device.type == device.type
|
||||
assert next(model.parameters()).dtype == torch.bfloat16
|
||||
|
||||
|
||||
@pytest.mark.skipif(not torch.cuda.is_available(), reason="CUDA required for text encoder parity.")
|
||||
def test_dreamx_world_text_encoder_parity_scaffold(monkeypatch):
|
||||
device = torch.device("cuda:0")
|
||||
dtype = torch.bfloat16
|
||||
official = _load_official_text_encoder(device, dtype)
|
||||
fastvideo = _load_fastvideo_text_encoder(device, dtype, monkeypatch)
|
||||
tokenizer_path = WAN_BASE_DIR / "google" / "umt5-xxl"
|
||||
if not tokenizer_path.exists():
|
||||
pytest.skip(f"Wan2.2 tokenizer missing: {tokenizer_path}")
|
||||
tokenizer = AutoTokenizer.from_pretrained(str(tokenizer_path))
|
||||
batch = tokenizer(["A quiet forest trail at sunrise."], padding="max_length", max_length=512, return_tensors="pt")
|
||||
input_ids = batch.input_ids.to(device)
|
||||
attention_mask = batch.attention_mask.to(device)
|
||||
with torch.inference_mode():
|
||||
official_hidden = official(input_ids, attention_mask=attention_mask)[0].float().cpu()
|
||||
fastvideo_hidden = fastvideo(input_ids, attention_mask=attention_mask).last_hidden_state.float().cpu()
|
||||
assert official_hidden.shape == fastvideo_hidden.shape
|
||||
assert_close(fastvideo_hidden, official_hidden, atol=1e-3, rtol=1e-3)
|
||||
@@ -0,0 +1,335 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
"""DreamX-World transformer parity scaffold.
|
||||
|
||||
Coverage scope: both. The official side loads DreamX-World-5B-Cam through
|
||||
Wan2_2Transformer3DModel.from_pretrained with PRoPE camera control enabled.
|
||||
The FastVideo side strict-loads the converted DreamX transformer weights into
|
||||
the native DreamX-World DiT implementation.
|
||||
"""
|
||||
from __future__ import annotations
|
||||
|
||||
import math
|
||||
import os
|
||||
from pathlib import Path
|
||||
import sys
|
||||
import types
|
||||
|
||||
import pytest
|
||||
import torch
|
||||
from omegaconf import OmegaConf
|
||||
from torch.testing import assert_close
|
||||
|
||||
from fastvideo.forward_context import set_forward_context
|
||||
from fastvideo.configs.models.dits.dreamx_world import (
|
||||
DreamXWorldArchConfig, DreamXWorldConfig)
|
||||
from fastvideo.pipelines.basic.dreamx_world.camera_conditioning import build_dreamx_camera_condition
|
||||
from fastvideo.configs.pipelines.dreamx_world import make_dreamx_world_5b_cam_dit_config
|
||||
from fastvideo.platforms import AttentionBackendEnum
|
||||
from fastvideo.models.dits.dreamx_world import (
|
||||
DreamXPropeSelfAttention, DreamXWorldTransformer3DModel,
|
||||
DreamXWorldTransformerBlock)
|
||||
from fastvideo.models.loader.fsdp_load import load_model_from_full_model_state_dict
|
||||
from fastvideo.models.loader.utils import get_param_names_mapping
|
||||
from fastvideo.models.loader.weight_utils import resolve_safetensors_files, safetensors_weights_iterator
|
||||
from scripts.checkpoint_conversion.dreamx_world_to_diffusers import map_transformer_key
|
||||
|
||||
REPO_ROOT = Path(__file__).resolve().parents[3]
|
||||
OFFICIAL_REF_DIR = Path(os.getenv("DREAMX_WORLD_OFFICIAL_REF_DIR", REPO_ROOT / "DreamX-World"))
|
||||
LOCAL_WEIGHTS_DIR = Path(os.getenv("DREAMX_WORLD_LOCAL_WEIGHTS_DIR", REPO_ROOT / "official_weights" / "dreamx_world"))
|
||||
WAN_BASE_DIR = Path(os.getenv("DREAMX_WORLD_WAN_BASE_DIR", REPO_ROOT / "official_weights" / "Wan2.2-TI2V-5B"))
|
||||
CONVERTED_WEIGHTS_DIR = Path(os.getenv("DREAMX_WORLD_CONVERTED_WEIGHTS_DIR", REPO_ROOT / "converted_weights" / "dreamx_world"))
|
||||
CONVERTED_HF_REPO = "FastVideo/DreamX-World-5B-Cam-Diffusers"
|
||||
PARITY_SCOPE = "both"
|
||||
|
||||
|
||||
def _install_xfuser_stub() -> None:
|
||||
if "xfuser" in sys.modules:
|
||||
return
|
||||
xfuser = types.ModuleType("xfuser")
|
||||
core = types.ModuleType("xfuser.core")
|
||||
distributed = types.ModuleType("xfuser.core.distributed")
|
||||
long_ctx = types.ModuleType("xfuser.core.long_ctx_attention")
|
||||
distributed.get_sequence_parallel_rank = lambda: 0
|
||||
distributed.get_sequence_parallel_world_size = lambda: 1
|
||||
distributed.get_sp_group = lambda: None
|
||||
distributed.get_world_group = lambda: types.SimpleNamespace(local_rank=0, rank=0)
|
||||
distributed.init_distributed_environment = lambda *args, **kwargs: None
|
||||
distributed.initialize_model_parallel = lambda *args, **kwargs: None
|
||||
distributed.model_parallel_is_initialized = lambda: False
|
||||
|
||||
class XFuserLongContextAttention:
|
||||
def __call__(self, *args, **kwargs):
|
||||
raise RuntimeError("xfuser stub cannot execute attention")
|
||||
|
||||
long_ctx.xFuserLongContextAttention = XFuserLongContextAttention
|
||||
sys.modules.update({
|
||||
"xfuser": xfuser,
|
||||
"xfuser.core": core,
|
||||
"xfuser.core.distributed": distributed,
|
||||
"xfuser.core.long_ctx_attention": long_ctx,
|
||||
})
|
||||
|
||||
|
||||
def _make_tiny_dreamx_config() -> DreamXWorldConfig:
|
||||
return DreamXWorldConfig(
|
||||
arch_config=DreamXWorldArchConfig(
|
||||
num_attention_heads=1,
|
||||
attention_head_dim=8,
|
||||
in_channels=16,
|
||||
out_channels=16,
|
||||
ffn_dim=32,
|
||||
num_layers=1,
|
||||
cross_attn_norm=True,
|
||||
qk_norm="rms_norm_across_heads",
|
||||
add_control_adapter=True,
|
||||
cam_method="prope",
|
||||
attn_compress=1,
|
||||
cam_self_attn_layers=None,
|
||||
))
|
||||
|
||||
|
||||
def _add_official_to_path() -> None:
|
||||
if not OFFICIAL_REF_DIR.exists():
|
||||
pytest.skip(f"Official reference missing: {OFFICIAL_REF_DIR}")
|
||||
if str(OFFICIAL_REF_DIR) not in sys.path:
|
||||
sys.path.insert(0, str(OFFICIAL_REF_DIR))
|
||||
|
||||
|
||||
def _official_transformer_kwargs() -> dict:
|
||||
config_path = OFFICIAL_REF_DIR / "configs" / "wan2.2" / "wan_ti2v_5b.yaml"
|
||||
if not config_path.exists():
|
||||
pytest.skip(f"DreamX Wan config missing: {config_path}")
|
||||
config = OmegaConf.load(config_path)
|
||||
kwargs = OmegaConf.to_container(config["transformer_additional_kwargs"])
|
||||
kwargs["cam_method"] = "prope"
|
||||
kwargs["add_control_adapter"] = True
|
||||
return kwargs
|
||||
|
||||
|
||||
def _load_official_transformer(device: torch.device, dtype: torch.dtype) -> torch.nn.Module:
|
||||
_add_official_to_path()
|
||||
_install_xfuser_stub()
|
||||
if not LOCAL_WEIGHTS_DIR.exists():
|
||||
pytest.skip(f"DreamX transformer weights missing: {LOCAL_WEIGHTS_DIR}")
|
||||
try:
|
||||
from models import Wan2_2Transformer3DModel
|
||||
except Exception as exc: # noqa: BLE001 - local parity should skip missing refs.
|
||||
pytest.skip(f"Cannot import official DreamX transformer: {exc}")
|
||||
model = Wan2_2Transformer3DModel.from_pretrained(
|
||||
str(LOCAL_WEIGHTS_DIR),
|
||||
transformer_additional_kwargs=_official_transformer_kwargs(),
|
||||
torch_dtype=dtype,
|
||||
)
|
||||
return model.to(device=device, dtype=dtype).eval()
|
||||
|
||||
|
||||
def _load_fastvideo_transformer(device: torch.device, dtype: torch.dtype) -> torch.nn.Module:
|
||||
model = _load_fastvideo_transformer_strict(torch.device("cpu"), dtype)
|
||||
return model.to(device=device, dtype=dtype).eval()
|
||||
|
||||
|
||||
def _load_fastvideo_transformer_strict(device: torch.device, dtype: torch.dtype) -> torch.nn.Module:
|
||||
transformer_dir = CONVERTED_WEIGHTS_DIR / "transformer"
|
||||
if not transformer_dir.exists():
|
||||
# No local conversion: pull the published Diffusers transformer from the hub.
|
||||
from huggingface_hub import snapshot_download
|
||||
|
||||
transformer_dir = Path(snapshot_download(CONVERTED_HF_REPO, allow_patterns=["transformer/*"])) / "transformer"
|
||||
safetensors_files = resolve_safetensors_files(str(transformer_dir))
|
||||
config = make_dreamx_world_5b_cam_dit_config()
|
||||
import fastvideo.models.dits.dreamx_world as fastvideo_dreamx
|
||||
original_get_sp_world_size = fastvideo_dreamx.get_sp_world_size
|
||||
fastvideo_dreamx.get_sp_world_size = lambda: 1
|
||||
try:
|
||||
with torch.device("meta"):
|
||||
model = DreamXWorldTransformer3DModel(config=config, hf_config={})
|
||||
finally:
|
||||
fastvideo_dreamx.get_sp_world_size = original_get_sp_world_size
|
||||
incompatible = load_model_from_full_model_state_dict(
|
||||
model,
|
||||
safetensors_weights_iterator(safetensors_files, to_cpu=True),
|
||||
device=device,
|
||||
param_dtype=dtype,
|
||||
strict=True,
|
||||
param_names_mapping=get_param_names_mapping(model.param_names_mapping),
|
||||
training_mode=False,
|
||||
)
|
||||
assert incompatible.missing_keys == []
|
||||
assert incompatible.unexpected_keys == []
|
||||
assert not any(param.is_meta for param in model.parameters())
|
||||
return model.to(device=device, dtype=dtype).eval()
|
||||
|
||||
|
||||
def _make_inputs(device: torch.device, dtype: torch.dtype):
|
||||
torch.manual_seed(1234)
|
||||
num_frames = 5
|
||||
height = 64
|
||||
width = 64
|
||||
latent_frames = (num_frames - 1) // 4 + 1
|
||||
latent_h = height // 16
|
||||
latent_w = width // 16
|
||||
x = torch.randn(1, 48, latent_frames, latent_h, latent_w, device=device, dtype=dtype)
|
||||
context = [torch.randn(16, 4096, device=device, dtype=dtype)]
|
||||
seq_len = math.ceil((latent_h * latent_w) / 4 * latent_frames)
|
||||
timestep = torch.full((1, seq_len), 250, device=device, dtype=torch.long)
|
||||
camera = build_dreamx_camera_condition(
|
||||
["w"], [4], num_frames=num_frames, height=height, width=width, dtype=dtype, device=device
|
||||
)
|
||||
camera = {key: value.unsqueeze(0) for key, value in camera.items()}
|
||||
return {"x": [x[0]], "context": context, "t": timestep, "seq_len": seq_len, "y_camera": camera}
|
||||
|
||||
|
||||
def _run_official(model: torch.nn.Module, inputs: dict) -> torch.Tensor:
|
||||
with torch.inference_mode():
|
||||
output = model(**inputs)
|
||||
if isinstance(output, list):
|
||||
output = torch.stack(output, dim=0)
|
||||
assert torch.is_tensor(output), f"official output is not a tensor: {type(output)}"
|
||||
return output.detach().float().cpu()
|
||||
|
||||
|
||||
def _run_fastvideo(model: torch.nn.Module, inputs: dict) -> torch.Tensor:
|
||||
hidden_states = torch.stack(inputs["x"], dim=0)
|
||||
encoder_hidden_states = torch.stack([
|
||||
torch.cat([inputs["context"][0], inputs["context"][0].new_zeros(512 - inputs["context"][0].shape[0], 4096)])
|
||||
])
|
||||
with torch.inference_mode(), set_forward_context(current_timestep=0, attn_metadata=None):
|
||||
output = model(
|
||||
hidden_states=hidden_states,
|
||||
encoder_hidden_states=encoder_hidden_states,
|
||||
timestep=inputs["t"],
|
||||
y_camera=inputs["y_camera"],
|
||||
)
|
||||
assert torch.is_tensor(output), f"FastVideo output is not a tensor: {type(output)}"
|
||||
return output.detach().float().cpu()
|
||||
|
||||
|
||||
def test_dreamx_world_conversion_mapping_strict_load_smoke(monkeypatch):
|
||||
_add_official_to_path()
|
||||
_install_xfuser_stub()
|
||||
try:
|
||||
from models import Wan2_2Transformer3DModel
|
||||
except Exception as exc: # noqa: BLE001
|
||||
pytest.skip(f"Cannot import official DreamX transformer: {exc}")
|
||||
|
||||
official = Wan2_2Transformer3DModel(
|
||||
dim=8,
|
||||
ffn_dim=32,
|
||||
num_heads=1,
|
||||
num_layers=1,
|
||||
add_control_adapter=True,
|
||||
cam_method="prope",
|
||||
)
|
||||
official_state = official.state_dict()
|
||||
diffusers_like_state = {
|
||||
map_transformer_key(key): value.detach().clone()
|
||||
for key, value in official_state.items()
|
||||
}
|
||||
|
||||
import fastvideo.models.dits.dreamx_world as fastvideo_dreamx
|
||||
monkeypatch.setattr(fastvideo_dreamx, "get_sp_world_size", lambda: 1)
|
||||
with torch.device("meta"):
|
||||
fastvideo = DreamXWorldTransformer3DModel(config=_make_tiny_dreamx_config(), hf_config={})
|
||||
|
||||
incompatible = load_model_from_full_model_state_dict(
|
||||
fastvideo,
|
||||
iter(diffusers_like_state.items()),
|
||||
device=torch.device("cpu"),
|
||||
param_dtype=torch.float32,
|
||||
strict=True,
|
||||
param_names_mapping=get_param_names_mapping(fastvideo.param_names_mapping),
|
||||
training_mode=False,
|
||||
)
|
||||
assert incompatible.missing_keys == []
|
||||
assert incompatible.unexpected_keys == []
|
||||
assert not any(param.is_meta for param in fastvideo.parameters())
|
||||
|
||||
|
||||
def test_dreamx_world_5b_cam_dit_config_matches_official_shape():
|
||||
config = make_dreamx_world_5b_cam_dit_config()
|
||||
assert config.num_layers == 30
|
||||
assert config.num_attention_heads == 24
|
||||
assert config.attention_head_dim == 128
|
||||
assert config.hidden_size == 3072
|
||||
assert config.ffn_dim == 14336
|
||||
assert config.add_control_adapter is True
|
||||
assert config.cam_method == "prope"
|
||||
assert config.attn_compress == 1
|
||||
|
||||
|
||||
def test_dreamx_world_converted_5b_transformer_strict_loads():
|
||||
device = torch.device("cuda:0" if torch.cuda.is_available() else "cpu")
|
||||
_load_fastvideo_transformer_strict(device, torch.bfloat16)
|
||||
|
||||
|
||||
def test_dreamx_world_fastvideo_prope_branch_smoke():
|
||||
block = DreamXWorldTransformerBlock(
|
||||
8,
|
||||
32,
|
||||
1,
|
||||
cross_attn_norm=True,
|
||||
add_control_adapter=True,
|
||||
cam_method="prope",
|
||||
attn_compress=1,
|
||||
layer_idx=0,
|
||||
supported_attention_backends=(AttentionBackendEnum.TORCH_SDPA,),
|
||||
)
|
||||
assert block.cam_self_attn is not None
|
||||
assert [
|
||||
name for name, _ in block.named_parameters()
|
||||
if name.startswith("cam_self_attn.")
|
||||
][:8] == [
|
||||
"cam_self_attn.q_proj.weight",
|
||||
"cam_self_attn.q_proj.bias",
|
||||
"cam_self_attn.k_proj.weight",
|
||||
"cam_self_attn.k_proj.bias",
|
||||
"cam_self_attn.v_proj.weight",
|
||||
"cam_self_attn.v_proj.bias",
|
||||
"cam_self_attn.out_proj.weight",
|
||||
"cam_self_attn.out_proj.bias",
|
||||
]
|
||||
|
||||
module = DreamXPropeSelfAttention(
|
||||
dim=8,
|
||||
attn_dim=8,
|
||||
num_heads=1,
|
||||
qk_norm="rms_norm_across_heads",
|
||||
).eval()
|
||||
assert module.num_heads == 1
|
||||
assert module.head_dim == 8
|
||||
assert tuple(module.out_proj.weight.shape) == (8, 8)
|
||||
assert torch.count_nonzero(module.out_proj.weight) == 0
|
||||
|
||||
|
||||
@pytest.mark.skipif(not torch.cuda.is_available(), reason="CUDA required for transformer parity.")
|
||||
def test_dreamx_world_transformer_parity_scaffold(monkeypatch):
|
||||
device = torch.device("cuda:0")
|
||||
dtype = torch.float32
|
||||
inputs = _make_inputs(device, dtype)
|
||||
|
||||
official = _load_official_transformer(device, dtype)
|
||||
official_out = _run_official(official, inputs)
|
||||
del official
|
||||
torch.cuda.empty_cache()
|
||||
|
||||
import fastvideo.attention.layer as attention_layer
|
||||
import fastvideo.distributed.communication_op as communication_op
|
||||
import fastvideo.models.dits.dreamx_world as fastvideo_dreamx
|
||||
monkeypatch.setattr(attention_layer, "get_sp_parallel_rank", lambda: 0)
|
||||
monkeypatch.setattr(attention_layer, "get_sp_world_size", lambda: 1)
|
||||
monkeypatch.setattr(attention_layer, "sequence_model_parallel_all_to_all_4D", lambda tensor, scatter_dim=2, gather_dim=1: tensor)
|
||||
monkeypatch.setattr(attention_layer, "sequence_model_parallel_all_gather", lambda tensor, dim=-1: tensor)
|
||||
monkeypatch.setattr(fastvideo_dreamx, "get_sp_world_size", lambda: 1)
|
||||
monkeypatch.setattr(communication_op, "get_sp_world_size", lambda: 1)
|
||||
monkeypatch.setattr(fastvideo_dreamx, "sequence_model_parallel_shard", lambda tensor, dim=1: (tensor, tensor.shape[dim]))
|
||||
monkeypatch.setattr(
|
||||
fastvideo_dreamx,
|
||||
"sequence_model_parallel_all_gather_with_unpad",
|
||||
lambda tensor, original_seq_len, dim=1: tensor.narrow(dim, 0, original_seq_len),
|
||||
)
|
||||
fastvideo = _load_fastvideo_transformer(device, dtype)
|
||||
fastvideo_out = _run_fastvideo(fastvideo, inputs)
|
||||
assert official_out.shape == fastvideo_out.shape
|
||||
diff = (official_out - fastvideo_out).abs()
|
||||
print(f"diff_max={diff.max().item():.6f} diff_mean={diff.mean().item():.6f}")
|
||||
assert_close(fastvideo_out, official_out, atol=1e-1, rtol=1e-1)
|
||||
@@ -0,0 +1,220 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
"""DreamX-World Wan2.2 VAE reuse parity scaffold.
|
||||
|
||||
Coverage scope: implementation_subcomponent. The official side uses
|
||||
AutoencoderKLWan3_8 from DreamX, while the FastVideo side targets the native
|
||||
Wan VAE. This remains a scaffold until Wan2.2 base VAE weights are staged.
|
||||
"""
|
||||
from __future__ import annotations
|
||||
|
||||
import os
|
||||
from pathlib import Path
|
||||
import sys
|
||||
import re
|
||||
import types
|
||||
|
||||
import pytest
|
||||
import torch
|
||||
from omegaconf import OmegaConf
|
||||
from torch.testing import assert_close
|
||||
|
||||
from fastvideo.configs.pipelines.dreamx_world import make_dreamx_world_5b_cam_vae_config
|
||||
from fastvideo.models.vaes.wanvae import AutoencoderKLWan
|
||||
|
||||
REPO_ROOT = Path(__file__).resolve().parents[3]
|
||||
OFFICIAL_REF_DIR = Path(os.getenv("DREAMX_WORLD_OFFICIAL_REF_DIR", REPO_ROOT / "DreamX-World"))
|
||||
WAN_BASE_DIR = Path(os.getenv("DREAMX_WORLD_WAN_BASE_DIR", REPO_ROOT / "official_weights" / "Wan2.2-TI2V-5B"))
|
||||
PARITY_SCOPE = "implementation_subcomponent"
|
||||
|
||||
|
||||
def _add_official_to_path():
|
||||
if not OFFICIAL_REF_DIR.exists():
|
||||
pytest.skip(f"Official reference missing: {OFFICIAL_REF_DIR}")
|
||||
if str(OFFICIAL_REF_DIR) not in sys.path:
|
||||
sys.path.insert(0, str(OFFICIAL_REF_DIR))
|
||||
|
||||
|
||||
def _install_xfuser_stub() -> None:
|
||||
if "xfuser" in sys.modules:
|
||||
return
|
||||
xfuser = types.ModuleType("xfuser")
|
||||
core = types.ModuleType("xfuser.core")
|
||||
distributed = types.ModuleType("xfuser.core.distributed")
|
||||
long_ctx = types.ModuleType("xfuser.core.long_ctx_attention")
|
||||
distributed.get_sequence_parallel_rank = lambda: 0
|
||||
distributed.get_sequence_parallel_world_size = lambda: 1
|
||||
distributed.get_sp_group = lambda: None
|
||||
distributed.get_world_group = lambda: types.SimpleNamespace(local_rank=0, rank=0)
|
||||
distributed.init_distributed_environment = lambda *args, **kwargs: None
|
||||
distributed.initialize_model_parallel = lambda *args, **kwargs: None
|
||||
distributed.model_parallel_is_initialized = lambda: False
|
||||
|
||||
class XFuserLongContextAttention:
|
||||
def __call__(self, *args, **kwargs):
|
||||
raise RuntimeError("xfuser stub cannot execute attention")
|
||||
|
||||
long_ctx.xFuserLongContextAttention = XFuserLongContextAttention
|
||||
sys.modules.update({
|
||||
"xfuser": xfuser,
|
||||
"xfuser.core": core,
|
||||
"xfuser.core.distributed": distributed,
|
||||
"xfuser.core.long_ctx_attention": long_ctx,
|
||||
})
|
||||
|
||||
|
||||
def _map_residual_subkey(prefix: str, sub: str) -> str | None:
|
||||
if sub == "residual.0.gamma":
|
||||
return f"{prefix}.norm1.gamma"
|
||||
match = re.match(r"^residual\.2\.(weight|bias)$", sub)
|
||||
if match:
|
||||
return f"{prefix}.conv1.{match.group(1)}"
|
||||
if sub == "residual.3.gamma":
|
||||
return f"{prefix}.norm2.gamma"
|
||||
match = re.match(r"^residual\.6\.(weight|bias)$", sub)
|
||||
if match:
|
||||
return f"{prefix}.conv2.{match.group(1)}"
|
||||
match = re.match(r"^shortcut\.(weight|bias)$", sub)
|
||||
if match:
|
||||
return f"{prefix}.conv_shortcut.{match.group(1)}"
|
||||
return None
|
||||
|
||||
|
||||
def _map_attention_subkey(prefix: str, sub: str) -> str | None:
|
||||
if sub == "norm.gamma":
|
||||
return f"{prefix}.norm.gamma"
|
||||
match = re.match(r"^(to_qkv|proj)\.(weight|bias)$", sub)
|
||||
if match:
|
||||
return f"{prefix}.{match.group(1)}.{match.group(2)}"
|
||||
return None
|
||||
|
||||
|
||||
def _map_resample_subkey(prefix: str, sub: str) -> str | None:
|
||||
match = re.match(r"^resample\.1\.(weight|bias)$", sub)
|
||||
if match:
|
||||
return f"{prefix}.resample.1.{match.group(1)}"
|
||||
match = re.match(r"^time_conv\.(weight|bias)$", sub)
|
||||
if match:
|
||||
return f"{prefix}.time_conv.{match.group(1)}"
|
||||
return None
|
||||
|
||||
|
||||
def _map_dreamx_raw_vae_key(key: str) -> str | None:
|
||||
match = re.match(r"^(conv1|conv2)\.(weight|bias)$", key)
|
||||
if match:
|
||||
prefix = "quant_conv" if match.group(1) == "conv1" else "post_quant_conv"
|
||||
return f"{prefix}.{match.group(2)}"
|
||||
match = re.match(r"^(encoder|decoder)\.conv1\.(weight|bias)$", key)
|
||||
if match:
|
||||
return f"{match.group(1)}.conv_in.{match.group(2)}"
|
||||
match = re.match(r"^(encoder|decoder)\.head\.0\.gamma$", key)
|
||||
if match:
|
||||
return f"{match.group(1)}.norm_out.gamma"
|
||||
match = re.match(r"^(encoder|decoder)\.head\.2\.(weight|bias)$", key)
|
||||
if match:
|
||||
return f"{match.group(1)}.conv_out.{match.group(2)}"
|
||||
match = re.match(r"^(encoder|decoder)\.middle\.0\.(.*)$", key)
|
||||
if match:
|
||||
return _map_residual_subkey(f"{match.group(1)}.mid_block.resnets.0", match.group(2))
|
||||
match = re.match(r"^(encoder|decoder)\.middle\.1\.(.*)$", key)
|
||||
if match:
|
||||
return _map_attention_subkey(f"{match.group(1)}.mid_block.attentions.0", match.group(2))
|
||||
match = re.match(r"^(encoder|decoder)\.middle\.2\.(.*)$", key)
|
||||
if match:
|
||||
return _map_residual_subkey(f"{match.group(1)}.mid_block.resnets.1", match.group(2))
|
||||
match = re.match(r"^encoder\.downsamples\.(\d+)\.downsamples\.(\d+)\.(.*)$", key)
|
||||
if match:
|
||||
stage = int(match.group(1))
|
||||
block = int(match.group(2))
|
||||
sub = match.group(3)
|
||||
if block in (0, 1):
|
||||
return _map_residual_subkey(f"encoder.down_blocks.{stage}.resnets.{block}", sub)
|
||||
if block == 2:
|
||||
return _map_resample_subkey(f"encoder.down_blocks.{stage}.downsampler", sub)
|
||||
return None
|
||||
match = re.match(r"^decoder\.upsamples\.(\d+)\.upsamples\.(\d+)\.(.*)$", key)
|
||||
if match:
|
||||
stage = int(match.group(1))
|
||||
block = int(match.group(2))
|
||||
sub = match.group(3)
|
||||
if block in (0, 1, 2):
|
||||
return _map_residual_subkey(f"decoder.up_blocks.{stage}.resnets.{block}", sub)
|
||||
if block == 3:
|
||||
return _map_resample_subkey(f"decoder.up_blocks.{stage}.upsampler", sub)
|
||||
return None
|
||||
return None
|
||||
|
||||
|
||||
def _vae_kwargs():
|
||||
config = OmegaConf.load(OFFICIAL_REF_DIR / "configs" / "wan2.2" / "wan_ti2v_5b.yaml")
|
||||
return OmegaConf.to_container(config["vae_kwargs"])
|
||||
|
||||
|
||||
def _load_official_vae(device, dtype):
|
||||
_add_official_to_path()
|
||||
_install_xfuser_stub()
|
||||
vae_path = WAN_BASE_DIR / "Wan2.2_VAE.pth"
|
||||
if not vae_path.exists():
|
||||
pytest.skip(f"Wan2.2 base VAE weights missing: {vae_path}")
|
||||
try:
|
||||
from models import AutoencoderKLWan3_8
|
||||
except Exception as exc: # noqa: BLE001
|
||||
pytest.skip(f"Cannot import official DreamX VAE: {exc}")
|
||||
model = AutoencoderKLWan3_8.from_pretrained(str(vae_path), additional_kwargs=_vae_kwargs())
|
||||
return model.to(device=device, dtype=dtype).eval()
|
||||
|
||||
|
||||
def _load_fastvideo_vae(device, dtype):
|
||||
vae_path = WAN_BASE_DIR / "Wan2.2_VAE.pth"
|
||||
if not vae_path.exists():
|
||||
pytest.skip(f"Wan2.2 raw VAE weights missing: {vae_path}")
|
||||
config = make_dreamx_world_5b_cam_vae_config()
|
||||
config.load_encoder = True
|
||||
config.load_decoder = True
|
||||
model = AutoencoderKLWan(config).to(device=device, dtype=dtype)
|
||||
raw_state = torch.load(str(vae_path), map_location="cpu", weights_only=True)
|
||||
mapped_state = {}
|
||||
for key, value in raw_state.items():
|
||||
mapped_key = _map_dreamx_raw_vae_key(key)
|
||||
if mapped_key is None:
|
||||
raise AssertionError(f"Unmapped DreamX raw VAE key: {key}")
|
||||
mapped_state[mapped_key] = value
|
||||
model.load_state_dict(mapped_state, strict=True)
|
||||
return model.eval()
|
||||
|
||||
|
||||
def _normalize_fastvideo_vae_latent(latent: torch.Tensor) -> torch.Tensor:
|
||||
config = make_dreamx_world_5b_cam_vae_config()
|
||||
mean = torch.tensor(config.latents_mean, device=latent.device, dtype=latent.dtype).view(1, -1, 1, 1, 1)
|
||||
std = torch.tensor(config.latents_std, device=latent.device, dtype=latent.dtype).view(1, -1, 1, 1, 1)
|
||||
return (latent - mean) / std
|
||||
|
||||
|
||||
def test_dreamx_world_vae_config_matches_wan22_shape():
|
||||
config = make_dreamx_world_5b_cam_vae_config()
|
||||
assert config.z_dim == 48
|
||||
assert config.in_channels == 12
|
||||
assert config.out_channels == 12
|
||||
assert config.base_dim == 160
|
||||
assert config.decoder_base_dim == 256
|
||||
assert config.scale_factor_temporal == 4
|
||||
assert config.scale_factor_spatial == 16
|
||||
assert config.patch_size == 2
|
||||
assert config.is_residual is True
|
||||
assert config.clip_output is False
|
||||
assert len(config.latents_mean) == 48
|
||||
assert len(config.latents_std) == 48
|
||||
|
||||
|
||||
@pytest.mark.skipif(not torch.cuda.is_available(), reason="CUDA required for VAE parity.")
|
||||
def test_dreamx_world_vae_encode_parity_scaffold():
|
||||
device = torch.device("cuda:0")
|
||||
dtype = torch.bfloat16
|
||||
official = _load_official_vae(device, dtype)
|
||||
fastvideo = _load_fastvideo_vae(device, dtype)
|
||||
torch.manual_seed(123)
|
||||
video = torch.randn(1, 3, 5, 64, 64, device=device, dtype=dtype).clamp(-1, 1)
|
||||
with torch.inference_mode():
|
||||
official_latent = official.encode(video).latent_dist.mean.float().cpu()
|
||||
fastvideo_latent = _normalize_fastvideo_vae_latent(fastvideo.encode(video).mean).float().cpu()
|
||||
assert official_latent.shape == fastvideo_latent.shape
|
||||
assert_close(fastvideo_latent, official_latent, atol=5e-2, rtol=5e-2)
|
||||
@@ -0,0 +1,110 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
"""DreamX-World pipeline parity checks.
|
||||
|
||||
This test compares the FastVideo pipeline's DreamX-specific conditioning and
|
||||
single-step scheduler path against an explicit hand-rolled pass using the same
|
||||
loaded modules. It is intentionally local and deterministic: component parity
|
||||
against the official DreamX repository lives in ``tests/local_tests/dreamx_world``.
|
||||
"""
|
||||
from __future__ import annotations
|
||||
|
||||
import gc
|
||||
import os
|
||||
from pathlib import Path
|
||||
from typing import Any, cast
|
||||
|
||||
import pytest
|
||||
import torch
|
||||
from torch.testing import assert_close
|
||||
|
||||
|
||||
os.environ.setdefault("FASTVIDEO_ATTENTION_BACKEND", "TORCH_SDPA")
|
||||
|
||||
# Local converted dir or HF repo id; the loader downloads hub ids itself.
|
||||
MODEL_DIR = os.getenv("DREAMX_WORLD_MODEL_DIR", "FastVideo/DreamX-World-5B-Cam-Diffusers")
|
||||
|
||||
|
||||
pytestmark = pytest.mark.skipif(
|
||||
not torch.cuda.is_available(),
|
||||
reason="DreamX-World pipeline parity requires CUDA",
|
||||
)
|
||||
|
||||
|
||||
|
||||
|
||||
def _run_worker_forward_batch(worker_wrapper: Any, request_kwargs: dict[str, Any]) -> torch.Tensor:
|
||||
from fastvideo.api.sampling_param import SamplingParam
|
||||
from fastvideo.pipelines.pipeline_batch_info import ForwardBatch
|
||||
from fastvideo.utils import shallow_asdict
|
||||
|
||||
fastvideo_args = worker_wrapper.worker.fastvideo_args
|
||||
sampling_param = SamplingParam.from_pretrained(fastvideo_args.model_path)
|
||||
sampling_param.update({
|
||||
key: value
|
||||
for key, value in request_kwargs.items()
|
||||
if key not in {"prompt", "output_path"}
|
||||
})
|
||||
sampling_param.prompt = request_kwargs["prompt"]
|
||||
|
||||
latents_size = [
|
||||
(sampling_param.num_frames - 1) // 4 + 1,
|
||||
sampling_param.height // 8,
|
||||
sampling_param.width // 8,
|
||||
]
|
||||
n_tokens = latents_size[0] * latents_size[1] * latents_size[2]
|
||||
batch = ForwardBatch(
|
||||
**shallow_asdict(sampling_param),
|
||||
eta=0.0,
|
||||
n_tokens=n_tokens,
|
||||
VSA_sparsity=fastvideo_args.VSA_sparsity,
|
||||
)
|
||||
output_batch = worker_wrapper.worker.pipeline.forward(batch, fastvideo_args)
|
||||
assert output_batch.output is not None
|
||||
return output_batch.output.detach().cpu()
|
||||
|
||||
def _close_generator(generator: Any) -> None:
|
||||
generator.shutdown()
|
||||
gc.collect()
|
||||
torch.cuda.empty_cache()
|
||||
|
||||
|
||||
def test_dreamx_world_one_step_pipeline_latent_matches_manual_pass() -> None:
|
||||
from fastvideo import VideoGenerator
|
||||
common_kwargs = dict(
|
||||
prompt="a quiet road through a futuristic city at sunrise",
|
||||
output_path="outputs_video/dreamx_world_parity",
|
||||
save_video=False,
|
||||
return_frames=True,
|
||||
height=64,
|
||||
width=64,
|
||||
num_frames=9,
|
||||
num_inference_steps=1,
|
||||
guidance_scale=1.0,
|
||||
action_list=["w"],
|
||||
action_speed_list=[2.0],
|
||||
seed=123,
|
||||
)
|
||||
|
||||
generator = VideoGenerator.from_pretrained(
|
||||
MODEL_DIR,
|
||||
num_gpus=1,
|
||||
use_fsdp_inference=False,
|
||||
dit_cpu_offload=False,
|
||||
vae_cpu_offload=True,
|
||||
text_encoder_cpu_offload=True,
|
||||
pin_cpu_memory=False,
|
||||
output_type="latent",
|
||||
override_pipeline_cls_name="DreamXWorldPipeline",
|
||||
)
|
||||
try:
|
||||
result = cast(dict[str, Any], generator.generate_video(**common_kwargs))
|
||||
pipeline_latents = cast(torch.Tensor, result["samples"]).detach().cpu()
|
||||
|
||||
manual_latents = generator.executor.collective_rpc(
|
||||
_run_worker_forward_batch,
|
||||
kwargs={"request_kwargs": common_kwargs},
|
||||
)[0]
|
||||
finally:
|
||||
_close_generator(generator)
|
||||
|
||||
assert_close(pipeline_latents, manual_latents, atol=0.0, rtol=0.0)
|
||||
@@ -0,0 +1,157 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
"""Smoke tests for the DreamX-World-5B-Cam pipeline."""
|
||||
from __future__ import annotations
|
||||
|
||||
import os
|
||||
from pathlib import Path
|
||||
from typing import Any, cast
|
||||
|
||||
import pytest
|
||||
import torch
|
||||
from PIL import Image, ImageDraw
|
||||
|
||||
|
||||
os.environ.setdefault("FASTVIDEO_ATTENTION_BACKEND", "TORCH_SDPA")
|
||||
|
||||
# Local converted dir or HF repo id; the loader downloads hub ids itself.
|
||||
MODEL_DIR = os.getenv("DREAMX_WORLD_MODEL_DIR", "FastVideo/DreamX-World-5B-Cam-Diffusers")
|
||||
|
||||
|
||||
|
||||
|
||||
def _write_smoke_image(path: Path) -> None:
|
||||
image = Image.new("RGB", (96, 96), color=(42, 76, 112))
|
||||
draw = ImageDraw.Draw(image)
|
||||
draw.rectangle((0, 56, 96, 96), fill=(36, 44, 52))
|
||||
draw.polygon([(0, 56), (48, 30), (96, 56)], fill=(120, 142, 158))
|
||||
draw.rectangle((34, 42, 62, 72), fill=(178, 198, 212))
|
||||
image.save(path)
|
||||
|
||||
pytestmark = pytest.mark.skipif(
|
||||
not torch.cuda.is_available(),
|
||||
reason="DreamX-World pipeline smoke requires CUDA",
|
||||
)
|
||||
|
||||
|
||||
def test_dreamx_world_typed_surface_preflight() -> None:
|
||||
import fastvideo.registry as registry
|
||||
from fastvideo.api.presets import get_preset, get_presets_for_family
|
||||
from fastvideo.configs.pipelines.dreamx_world import DreamXWorld5BCamPipelineConfig
|
||||
from fastvideo.fastvideo_args import WorkloadType
|
||||
from fastvideo.pipelines.basic.dreamx_world.dreamx_world_pipeline import (
|
||||
DreamXWorldPipeline,
|
||||
EntryClass,
|
||||
)
|
||||
|
||||
assert DreamXWorldPipeline.__name__ == "DreamXWorldPipeline"
|
||||
assert EntryClass is DreamXWorldPipeline
|
||||
assert DreamXWorldPipeline._required_config_modules == [
|
||||
"text_encoder",
|
||||
"tokenizer",
|
||||
"vae",
|
||||
"transformer",
|
||||
"scheduler",
|
||||
]
|
||||
|
||||
default_preset, model_family = registry.get_preset_selection(
|
||||
"GD-ML/DreamX-World-5B-Cam"
|
||||
)
|
||||
assert model_family == "dreamx_world"
|
||||
assert default_preset == "dreamx_world_5b_cam"
|
||||
|
||||
info = registry.get_model_info(
|
||||
"GD-ML/DreamX-World-5B-Cam",
|
||||
workload_type=WorkloadType.I2V,
|
||||
override_pipeline_cls_name="DreamXWorldPipeline",
|
||||
)
|
||||
assert info.pipeline_cls is DreamXWorldPipeline
|
||||
assert info.pipeline_config_cls is DreamXWorld5BCamPipelineConfig
|
||||
|
||||
names = {p.name for p in get_presets_for_family("dreamx_world")}
|
||||
assert "dreamx_world_5b_cam" in names
|
||||
preset = get_preset("dreamx_world_5b_cam", "dreamx_world")
|
||||
assert preset.defaults["num_inference_steps"] == 30
|
||||
assert preset.defaults["height"] == 480
|
||||
assert preset.defaults["width"] == 832
|
||||
assert preset.defaults["num_frames"] == 161
|
||||
assert preset.defaults["guidance_scale"] == 5.0
|
||||
|
||||
cfg = DreamXWorld5BCamPipelineConfig()
|
||||
assert cfg.flow_shift == 3.0
|
||||
assert cfg.ti2v_task is True
|
||||
assert cfg.expand_timesteps is True
|
||||
assert cfg.dit_config.arch_config.add_control_adapter is True
|
||||
assert cfg.dit_config.arch_config.cam_method == "prope"
|
||||
|
||||
|
||||
def test_dreamx_world_camera_stage_writes_y_camera() -> None:
|
||||
from types import SimpleNamespace
|
||||
|
||||
from fastvideo.pipelines.basic.dreamx_world.stages import (
|
||||
DREAMX_Y_CAMERA_KEY,
|
||||
DreamXWorldCameraConditioningStage,
|
||||
)
|
||||
from fastvideo.pipelines.pipeline_batch_info import ForwardBatch
|
||||
|
||||
batch = ForwardBatch(
|
||||
data_type="video",
|
||||
prompt="camera smoke",
|
||||
latents=torch.zeros(1, 48, 3, 8, 8, dtype=torch.bfloat16, device="cuda"),
|
||||
num_frames=9,
|
||||
height=64,
|
||||
width=64,
|
||||
action_list=["w", "d"],
|
||||
action_speed_list=[2.0, 1.0],
|
||||
)
|
||||
out = DreamXWorldCameraConditioningStage().forward(batch, cast(Any, SimpleNamespace()))
|
||||
|
||||
y_camera = out.extra[DREAMX_Y_CAMERA_KEY]
|
||||
assert set(y_camera) == {"viewmats", "K"}
|
||||
assert y_camera["viewmats"].shape == (1, 3, 4, 4)
|
||||
assert y_camera["K"].shape == (1, 3, 3, 3)
|
||||
assert y_camera["viewmats"].device.type == "cuda"
|
||||
assert y_camera["viewmats"].dtype == torch.bfloat16
|
||||
|
||||
|
||||
def test_dreamx_world_pipeline_load_generate_latent_smoke(tmp_path: Path) -> None:
|
||||
from fastvideo import VideoGenerator
|
||||
|
||||
image_path = tmp_path / "dreamx_world_smoke_input.png"
|
||||
_write_smoke_image(image_path)
|
||||
|
||||
generator = VideoGenerator.from_pretrained(
|
||||
MODEL_DIR,
|
||||
num_gpus=1,
|
||||
use_fsdp_inference=False,
|
||||
dit_cpu_offload=False,
|
||||
vae_cpu_offload=True,
|
||||
text_encoder_cpu_offload=True,
|
||||
pin_cpu_memory=False,
|
||||
output_type="latent",
|
||||
override_pipeline_cls_name="DreamXWorldPipeline",
|
||||
)
|
||||
try:
|
||||
result = generator.generate_video(
|
||||
prompt="a quiet road through a futuristic city at sunrise",
|
||||
output_path="outputs_video/dreamx_world_smoke",
|
||||
save_video=False,
|
||||
return_frames=True,
|
||||
height=64,
|
||||
width=64,
|
||||
num_frames=9,
|
||||
num_inference_steps=1,
|
||||
guidance_scale=1.0,
|
||||
image_path=str(image_path),
|
||||
action_list=["w"],
|
||||
action_speed_list=[2.0],
|
||||
seed=0,
|
||||
)
|
||||
finally:
|
||||
generator.shutdown()
|
||||
|
||||
assert isinstance(result, dict)
|
||||
samples = cast(dict[str, Any], result)["samples"]
|
||||
assert torch.is_tensor(samples)
|
||||
assert samples.ndim == 5
|
||||
assert samples.shape[1] == 48
|
||||
assert torch.isfinite(samples).all()
|
||||
Reference in New Issue
Block a user