Skip to content

causal_denoising

Compatibility imports for Wan's family-owned causal sampling stages.

Classes

fastvideo.pipelines.stages.causal_denoising.CausalDMDDenosingStage

CausalDMDDenosingStage(transformer, scheduler, transformer_2=None, vae=None)

Bases: WanCausalDenoisingBase

Denoising stage for causal diffusion.

Source code in fastvideo/pipelines/basic/wan/stages/causal_denoising.py
def __init__(self, transformer, scheduler, transformer_2=None, vae=None) -> None:
    super().__init__(transformer, scheduler, transformer_2=transformer_2)
    # KV and cross-attention cache state (initialized on first forward)
    self.transformer = transformer
    self.transformer_2 = transformer_2
    # Model-dependent constants (aligned with causal_inference.py assumptions)
    self.num_transformer_blocks = len(self.transformer.blocks)
    self.num_frames_per_block = self.transformer.config.arch_config.num_frames_per_block
    self.sliding_window_num_frames = self.transformer.config.arch_config.sliding_window_num_frames

    self.local_attn_size = _get_transformer_attr(self.transformer, "local_attn_size", -1)
    self.sink_size = _get_transformer_attr(self.transformer, "sink_size", 0)

fastvideo.pipelines.stages.causal_denoising.CausalDenoisingStage

CausalDenoisingStage(transformer, scheduler, transformer_2=None, vae=None)

Bases: WanCausalDenoisingBase

Causal block-by-block denoising with standard multi-step flow matching (scheduler.step), not DMD few-step.

Each block is fully denoised through all scheduler timesteps before moving to the next block. After each block is denoised, the KV cache is updated with clean context so subsequent blocks can attend to prior clean frames.

Source code in fastvideo/pipelines/basic/wan/stages/causal_denoising.py
def __init__(self, transformer, scheduler, transformer_2=None, vae=None) -> None:
    super().__init__(transformer, scheduler, transformer_2=transformer_2)
    # KV and cross-attention cache state (initialized on first forward)
    self.transformer = transformer
    self.transformer_2 = transformer_2
    # Model-dependent constants (aligned with causal_inference.py assumptions)
    self.num_transformer_blocks = len(self.transformer.blocks)
    self.num_frames_per_block = self.transformer.config.arch_config.num_frames_per_block
    self.sliding_window_num_frames = self.transformer.config.arch_config.sliding_window_num_frames

    self.local_attn_size = _get_transformer_attr(self.transformer, "local_attn_size", -1)
    self.sink_size = _get_transformer_attr(self.transformer, "sink_size", 0)