Skip to content

dreamx_world

Classes

fastvideo.pipelines.basic.dreamx_world.DreamXWorld5BARPipelineConfig dataclass

DreamXWorld5BARPipelineConfig(model_path: str = '', pipeline_config_path: str | None = None, embedded_cfg_scale: float = 6.0, flow_shift: float | None = 5.0, flow_shift_sr: float | None = None, disable_autocast: bool = False, scheduler_step_in_fp32: bool = False, is_causal: bool = True, dit_config: DiTConfig = make_dreamx_world_5b_ar_dit_config(), dit_precision: str = 'bf16', upsampler_config: UpsamplerConfig = UpsamplerConfig(), upsampler_precision: str = 'fp32', vae_config: VAEConfig = make_dreamx_world_5b_cam_vae_config(), vae_precision: str = 'fp32', vae_decode_precision: str | None = 'bf16', vae_tiling: bool = False, vae_sp: bool = False, image_encoder_config: EncoderConfig = EncoderConfig(), image_encoder_precision: str = 'fp32', text_encoder_configs: tuple[EncoderConfig, ...] = (lambda: (make_dreamx_world_5b_cam_text_encoder_config(),))(), text_encoder_precisions: tuple[str, ...] = (lambda: ('bf16',))(), preprocess_text_funcs: tuple[Callable[[str], str], ...] = (lambda: (preprocess_text,))(), postprocess_text_funcs: tuple[Callable[[BaseEncoderOutput], Tensor], ...] = (lambda: (t5_postprocess_text,))(), dmd_denoising_steps: tuple[int, ...] = (1000, 750, 500, 250), ti2v_task: bool = True, lucy_edit_task: bool = False, boundary_ratio: float | None = None, expand_timesteps: bool = True, warp_denoising_step: bool = True, context_noise: float = 0.1, num_frames_per_block: int = 3, color_correction_strength: float = 1.0)

Bases: DreamXWorld5BCamPipelineConfig

Pipeline config for DreamX-World-5B autoregressive forcing.

fastvideo.pipelines.basic.dreamx_world.DreamXWorld5BCamPipelineConfig dataclass

DreamXWorld5BCamPipelineConfig(model_path: str = '', pipeline_config_path: str | None = None, embedded_cfg_scale: float = 6.0, flow_shift: float | None = 3.0, flow_shift_sr: float | None = None, disable_autocast: bool = False, scheduler_step_in_fp32: bool = False, is_causal: bool = False, dit_config: DiTConfig = make_dreamx_world_5b_cam_dit_config(), dit_precision: str = 'bf16', upsampler_config: UpsamplerConfig = UpsamplerConfig(), upsampler_precision: str = 'fp32', vae_config: VAEConfig = make_dreamx_world_5b_cam_vae_config(), vae_precision: str = 'fp32', vae_decode_precision: str | None = 'bf16', vae_tiling: bool = False, vae_sp: bool = False, image_encoder_config: EncoderConfig = EncoderConfig(), image_encoder_precision: str = 'fp32', text_encoder_configs: tuple[EncoderConfig, ...] = (lambda: (make_dreamx_world_5b_cam_text_encoder_config(),))(), text_encoder_precisions: tuple[str, ...] = (lambda: ('bf16',))(), preprocess_text_funcs: tuple[Callable[[str], str], ...] = (lambda: (preprocess_text,))(), postprocess_text_funcs: tuple[Callable[[BaseEncoderOutput], Tensor], ...] = (lambda: (t5_postprocess_text,))(), dmd_denoising_steps: list[int] | None = None, ti2v_task: bool = True, lucy_edit_task: bool = False, boundary_ratio: float | None = None, expand_timesteps: bool = True)

Bases: PipelineConfig

Pipeline config for the first-scope DreamX-World-5B-Cam mode.

fastvideo.pipelines.basic.dreamx_world.DreamXWorldARPipeline

DreamXWorldARPipeline(*args, **kwargs)

Bases: LoRAPipeline, ComposedPipelineBase

DreamX-World-5B autoregressive causal camera pipeline.

Source code in fastvideo/pipelines/lora_pipeline.py
def __init__(self, *args, **kwargs) -> None:
    super().__init__(*args, **kwargs)
    self.device = get_local_torch_device()
    self.lora_adapter_paths = {}
    # build list of trainable transformers
    for transformer_name in self.trainable_transformer_names:
        if (transformer_name in self.modules and self.modules[transformer_name] is not None):
            self.trainable_transformer_modules[transformer_name] = (self.modules[transformer_name])
        # check for transformer_2 in case of Wan2.2 MoE or fake_score_transformer_2
        if transformer_name.endswith("_2"):
            raise ValueError(
                f"trainable_transformer_name override in pipelines should not include _2 suffix: {transformer_name}"
            )

        secondary_transformer_name = transformer_name + "_2"
        if (secondary_transformer_name in self.modules and self.modules[secondary_transformer_name] is not None):
            self.trainable_transformer_modules[secondary_transformer_name] = self.modules[
                secondary_transformer_name]

    logger.info(
        "trainable_transformer_modules: %s",
        self.trainable_transformer_modules.keys(),
    )

    for (
            transformer_name,
            transformer_module,
    ) in self.trainable_transformer_modules.items():
        self.exclude_lora_layers[transformer_name] = (transformer_module.config.arch_config.exclude_lora_layers)
    self.lora_target_modules = self.fastvideo_args.lora_target_modules
    self.lora_path = self.fastvideo_args.lora_path
    self.lora_nickname = self.fastvideo_args.lora_nickname
    self.training_mode = self.fastvideo_args.training_mode
    if self.training_mode and getattr(self.fastvideo_args, "lora_training", False):
        assert isinstance(self.fastvideo_args, TrainingArgs)
        if self.fastvideo_args.lora_alpha is None:
            self.fastvideo_args.lora_alpha = self.fastvideo_args.lora_rank
        self.lora_rank = self.fastvideo_args.lora_rank  # type: ignore
        self.lora_alpha = self.fastvideo_args.lora_alpha  # type: ignore
        logger.info(
            "Using LoRA training with rank %d and alpha %d",
            self.lora_rank,
            self.lora_alpha,
        )
        if self.lora_target_modules is None:
            self.lora_target_modules = [
                "q_proj",
                "k_proj",
                "v_proj",
                "o_proj",
                "to_q",
                "to_k",
                "to_v",
                "to_out",
                "to_qkv",
                "to_gate_compress",
            ]
            logger.info(
                "Using default lora_target_modules for all transformers: %s",
                self.lora_target_modules,
            )
        else:
            logger.warning(
                "Using custom lora_target_modules for all transformers, which may not be intended: %s",
                self.lora_target_modules,
            )

        self.convert_to_lora_layers()
    # Inference
    elif not self.training_mode and self.lora_path is not None:
        self.convert_to_lora_layers()
        self.set_lora_adapter(
            self.lora_nickname,  # type: ignore
            self.lora_path,
        )  # type: ignore

fastvideo.pipelines.basic.dreamx_world.DreamXWorldCameraConditioningStage

Bases: PipelineStage

Build PRoPE camera conditioning for DreamX-World-5B-Cam.

fastvideo.pipelines.basic.dreamx_world.DreamXWorldPipeline

DreamXWorldPipeline(*args, **kwargs)

Bases: LoRAPipeline, ComposedPipelineBase

DreamX-World-5B-Cam pipeline with native FastVideo camera conditioning.

Source code in fastvideo/pipelines/lora_pipeline.py
def __init__(self, *args, **kwargs) -> None:
    super().__init__(*args, **kwargs)
    self.device = get_local_torch_device()
    self.lora_adapter_paths = {}
    # build list of trainable transformers
    for transformer_name in self.trainable_transformer_names:
        if (transformer_name in self.modules and self.modules[transformer_name] is not None):
            self.trainable_transformer_modules[transformer_name] = (self.modules[transformer_name])
        # check for transformer_2 in case of Wan2.2 MoE or fake_score_transformer_2
        if transformer_name.endswith("_2"):
            raise ValueError(
                f"trainable_transformer_name override in pipelines should not include _2 suffix: {transformer_name}"
            )

        secondary_transformer_name = transformer_name + "_2"
        if (secondary_transformer_name in self.modules and self.modules[secondary_transformer_name] is not None):
            self.trainable_transformer_modules[secondary_transformer_name] = self.modules[
                secondary_transformer_name]

    logger.info(
        "trainable_transformer_modules: %s",
        self.trainable_transformer_modules.keys(),
    )

    for (
            transformer_name,
            transformer_module,
    ) in self.trainable_transformer_modules.items():
        self.exclude_lora_layers[transformer_name] = (transformer_module.config.arch_config.exclude_lora_layers)
    self.lora_target_modules = self.fastvideo_args.lora_target_modules
    self.lora_path = self.fastvideo_args.lora_path
    self.lora_nickname = self.fastvideo_args.lora_nickname
    self.training_mode = self.fastvideo_args.training_mode
    if self.training_mode and getattr(self.fastvideo_args, "lora_training", False):
        assert isinstance(self.fastvideo_args, TrainingArgs)
        if self.fastvideo_args.lora_alpha is None:
            self.fastvideo_args.lora_alpha = self.fastvideo_args.lora_rank
        self.lora_rank = self.fastvideo_args.lora_rank  # type: ignore
        self.lora_alpha = self.fastvideo_args.lora_alpha  # type: ignore
        logger.info(
            "Using LoRA training with rank %d and alpha %d",
            self.lora_rank,
            self.lora_alpha,
        )
        if self.lora_target_modules is None:
            self.lora_target_modules = [
                "q_proj",
                "k_proj",
                "v_proj",
                "o_proj",
                "to_q",
                "to_k",
                "to_v",
                "to_out",
                "to_qkv",
                "to_gate_compress",
            ]
            logger.info(
                "Using default lora_target_modules for all transformers: %s",
                self.lora_target_modules,
            )
        else:
            logger.warning(
                "Using custom lora_target_modules for all transformers, which may not be intended: %s",
                self.lora_target_modules,
            )

        self.convert_to_lora_layers()
    # Inference
    elif not self.training_mode and self.lora_path is not None:
        self.convert_to_lora_layers()
        self.set_lora_adapter(
            self.lora_nickname,  # type: ignore
            self.lora_path,
        )  # type: ignore

Functions:

fastvideo.pipelines.basic.dreamx_world.make_dreamx_world_5b_ar_dit_config

make_dreamx_world_5b_ar_dit_config() -> DreamXWorldARConfig

Return the DreamX-World-5B autoregressive causal DiT config.

Source code in fastvideo/configs/pipelines/dreamx_world.py
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,
    ))

fastvideo.pipelines.basic.dreamx_world.make_dreamx_world_5b_cam_dit_config

make_dreamx_world_5b_cam_dit_config() -> DreamXWorldConfig

Return the DreamX-World DiT config matching DreamX-World-5B-Cam.

Source code in fastvideo/configs/pipelines/dreamx_world.py
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,
    ))

fastvideo.pipelines.basic.dreamx_world.make_dreamx_world_5b_cam_text_encoder_config

make_dreamx_world_5b_cam_text_encoder_config() -> T5Config

Return the UMT5-XXL text encoder config used by DreamX-World-5B-Cam.

Source code in fastvideo/configs/pipelines/dreamx_world.py
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",
    )

fastvideo.pipelines.basic.dreamx_world.make_dreamx_world_5b_cam_vae_config

make_dreamx_world_5b_cam_vae_config() -> WanVAEConfig

Return the Wan2.2 48-channel VAE config used by DreamX-World-5B-Cam.

Source code in fastvideo/configs/pipelines/dreamx_world.py
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