Skip to content

s2v_stages

Wan-S2V-specific pipeline stages and the clip plan the pipeline loops over.

These replicate the steps of the official runner (wan/speech2video.py) that the shared stages do differently:

  • The reference image is VAE-encoded alone as one frame. The shared ImageVAEEncodingStage instead builds an I2V-style zero-padded video and encodes all of it, which is a different conditioning format entirely.
  • Decoding prepends temporal context so the causal VAE does not start cold: the reference latent on the first clip, the previous clip's motion latents afterwards. Only the generated span is kept, and the first clip also drops WARMUP_FRAMES frames that decode with too little context.
  • Long audio is covered by generating several clips back to back. Each clip re-encodes the last motion_frames pixels of the video so far as its motion history (plan_clips decides how many clips a request needs).

Classes

fastvideo.pipelines.basic.wan.s2v_stages.S2VClipPlan dataclass

S2VClipPlan(num_frames: int, infer_frames: int, num_clips: int)

How one request splits into clips.

num_frames is what the caller asked for and what they get back: the usual FastVideo 4k+1 count. infer_frames is what the transformer generates per clip (4n, the official runner's infer_frames). The first clip shows infer_frames - WARMUP_FRAMES of those, later clips all of them, and the concatenation is cut down to num_frames.

fastvideo.pipelines.basic.wan.s2v_stages.S2VDecodingStage

S2VDecodingStage(vae, pipeline=None)

Bases: DecodingStage

Decode with temporal context prepended, official-runner style.

The Wan VAE is causal in time: the first frames decode with less context and come out degraded. The official runner therefore decodes [context | generated latents] and keeps the trailing infer_frames pixels. The context is the previous clip's motion latents when there are any, else the reference latent -- and in that first-clip case 3 more warm-up frames are dropped. With output_type='latent' nothing is decoded, so the prepended context is sliced back off instead.

Source code in fastvideo/pipelines/stages/decoding.py
def __init__(self, vae, pipeline=None) -> None:
    self.vae: ParallelTiledVAE = vae
    self.pipeline = weakref.ref(pipeline) if pipeline else None

Methods:

fastvideo.pipelines.basic.wan.s2v_stages.S2VDecodingStage.trim staticmethod
trim(frames: Tensor, infer_frames: int, first_clip: bool) -> Tensor

Keep the generated span of a decoded [B, C, T, H, W] tensor.

Source code in fastvideo/pipelines/basic/wan/s2v_stages.py
@staticmethod
def trim(frames: torch.Tensor, infer_frames: int, first_clip: bool) -> torch.Tensor:
    """Keep the generated span of a decoded ``[B, C, T, H, W]`` tensor."""
    frames = frames[:, :, -infer_frames:]
    return frames[:, :, WARMUP_FRAMES:] if first_clip else frames

fastvideo.pipelines.basic.wan.s2v_stages.S2VRefImageEncodingStage

S2VRefImageEncodingStage(vae: ParallelTiledVAE)

Bases: ImageVAEEncodingStage

VAE-encode the reference image as a single latent frame.

Writes batch.image_latent with shape [B, C, 1, h, w]. Deterministic (distribution mode, not a sample): the reference is ground truth to preserve, and the official runner's native VAE encode is deterministic too. encode_pixels is shared with the pipeline's motion-history encode so both conditioning latents go through exactly the same normalisation.

Source code in fastvideo/pipelines/stages/image_encoding.py
def __init__(self, vae: ParallelTiledVAE) -> None:
    self.vae: ParallelTiledVAE = vae

Methods:

fastvideo.pipelines.basic.wan.s2v_stages.S2VRefImageEncodingStage.encode_pixels
encode_pixels(pixels: Tensor, fastvideo_args: FastVideoArgs) -> Tensor

[B, 3, T, H, W] pixels in [-1, 1] -> normalised latents [B, C, t, h, w].

Source code in fastvideo/pipelines/basic/wan/s2v_stages.py
def encode_pixels(self, pixels: torch.Tensor, fastvideo_args: FastVideoArgs) -> torch.Tensor:
    """[B, 3, T, H, W] pixels in [-1, 1] -> normalised latents [B, C, t, h, w]."""
    self.vae.to(get_local_torch_device())  # Module.to is in-place; no rebind, keeps mypy able to type self.vae
    pixels = pixels.to(get_local_torch_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
    with torch.autocast(device_type="cuda", dtype=vae_dtype, enabled=vae_autocast_enabled):
        if not vae_autocast_enabled:
            pixels = pixels.to(vae_dtype)
        latent = self.retrieve_latents(self.vae.encode(pixels), generator=None, sample_mode="argmax")

    # Same normalisation the other latents go through.
    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)

    if fastvideo_args.vae_cpu_offload:
        self.vae.to("cpu")
    return latent

Functions:

fastvideo.pipelines.basic.wan.s2v_stages.plan_clips

plan_clips(num_frames: int, clip_frames: int) -> S2VClipPlan

Split num_frames output frames into clips of at most clip_frames.

A single clip covers clip_frames - WARMUP_FRAMES visible frames, so the default 84-frame clip yields exactly the 81-frame default request. Shorter requests shrink the clip instead of generating frames that get thrown away.

Source code in fastvideo/pipelines/basic/wan/s2v_stages.py
def plan_clips(num_frames: int, clip_frames: int) -> S2VClipPlan:
    """Split ``num_frames`` output frames into clips of at most ``clip_frames``.

    A single clip covers ``clip_frames - WARMUP_FRAMES`` visible frames, so the
    default 84-frame clip yields exactly the 81-frame default request. Shorter
    requests shrink the clip instead of generating frames that get thrown away.
    """
    if num_frames < 1 or (num_frames - 1) % 4 != 0:
        raise ValueError(f"Wan-S2V needs num_frames = 4k+1 (e.g. 81), got {num_frames}. The VAE turns "
                         "4 pixel frames into 1 latent frame, plus one for the first frame.")
    if clip_frames < 4 or clip_frames % 4 != 0:
        raise ValueError(f"clip_frames must be a positive multiple of 4, got {clip_frames}")
    infer_frames = min(clip_frames, num_frames + WARMUP_FRAMES)
    remaining = max(0, num_frames - (infer_frames - WARMUP_FRAMES))
    num_clips = 1 + -(-remaining // infer_frames)  # ceil division
    return S2VClipPlan(num_frames=num_frames, infer_frames=infer_frames, num_clips=num_clips)