Skip to content

dreamx_world_ar_pipeline

DreamX-World-5B autoregressive pipeline entrypoint.

Classes

fastvideo.pipelines.basic.dreamx_world.dreamx_world_ar_pipeline.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()
    # Adapter tensors and wrapped model layers belong to this pipeline's module
    # instances. Sharing either cache across two generators can apply one model's
    # adapter to another model's layers.
    self.trainable_transformer_modules = {}
    self.lora_adapters = defaultdict(dict)
    self.lora_adapter_paths = {}
    self.lora_layers = {}
    self.exclude_lora_layers = {}
    self.cur_adapter_name = ""
    self.cur_adapter_path = ""
    self.cur_adapter_strength = 1.0
    # 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(),
    )

    # Only override the pipeline class's own default when the caller actually set
    # one. Assigning unconditionally erases per-model defaults, and a model that
    # declares one usually does so because wrapping every linear breaks its forward.
    if self.fastvideo_args.lora_target_modules is not None:
        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.lora_strength = self.fastvideo_args.lora_strength
    constructor_patch = DenseLoRAPatch.from_adapter(self.lora_path) if self.lora_path else None
    self._constructor_dense_lora_path = self.lora_path if constructor_patch is not None else None
    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()
        if not any(is_lazy_module(module) for module in self.trainable_transformer_modules.values()):
            self._setting_constructor_adapter = True
            try:
                self.set_lora_adapter(
                    self.lora_nickname,  # type: ignore
                    self.lora_path,
                    strength=self.lora_strength,
                )  # type: ignore
            finally:
                self._setting_constructor_adapter = False

Functions: