Skip to content

scheduling_piflow

Diffusers scheduler for distilled Kandinsky 6 PiFlow checkpoints.

Classes

fastvideo.models.schedulers.scheduling_piflow.DXPolicy

DXPolicy(denoising_output: Tensor, x_t_src: Tensor, sigma_t_src: Tensor, segment_size: float | Tensor = 1.0, shift: float = 1.0, mode: str = 'grid', eps: float = 0.0001)

Network-free DX policy over one flow-matching segment.

Source code in fastvideo/models/schedulers/scheduling_piflow.py
def __init__(  # noqa: PLR0913
    self,
    denoising_output: torch.Tensor,
    x_t_src: torch.Tensor,
    sigma_t_src: torch.Tensor,
    segment_size: float | torch.Tensor = 1.0,
    shift: float = 1.0,
    mode: str = "grid",
    eps: float = 1e-4,
) -> None:
    self.ndim = x_t_src.dim()
    self.shift = shift
    self.eps = eps
    if mode not in ("grid", "polynomial"):
        raise ValueError(f"Unknown mode: {mode}")
    self.mode = mode

    sigma_t_src = sigma_t_src.reshape(*sigma_t_src.size(), *((self.ndim - sigma_t_src.dim()) * [1]))
    self.raw_t_src = self._unwarp_t(sigma_t_src)
    segment = segment_size
    if isinstance(segment, torch.Tensor) and segment.dim() < self.raw_t_src.dim():
        segment = segment.reshape(*segment.size(), *((self.raw_t_src.dim() - segment.dim()) * [1]))
    self.raw_t_dst = (self.raw_t_src - segment).clamp(min=0)
    self.segment_size = (self.raw_t_src - self.raw_t_dst).clamp(min=eps)
    self.denoising_output_x_0 = x_t_src.unsqueeze(1) - sigma_t_src.unsqueeze(1) * denoising_output

fastvideo.models.schedulers.scheduling_piflow.PiflowScheduler

PiflowScheduler(num_train_timesteps: int = 1000, shift: float = 5.0, n_grid: int = 10, nfe: int | None = None, eps: float = 1e-06, final_step_size_scale: float = 0.5, num_policy_substeps: int = 128)

Bases: FlowMatchEulerDiscreteScheduler

Few-step PiFlow scheduler for widened-output diffusion transformers.

PiFlow evaluates the denoising model at a small number of grid points and integrates a network-free policy between those evaluations. The scheduler is intended for distilled Kandinsky 6 checkpoints, including the main video/audio model and the video super-resolution model. Their model output contains n_grid predictions per sample channel.

Parameters:

Name Type Description Default
num_train_timesteps `int`, *optional*, defaults to 1000

Number of training diffusion steps.

1000
shift `float`, *optional*, defaults to 5.0

Flow-matching timestep shift.

5.0
n_grid `int`, *optional*, defaults to 10

Number of predictions in the widened model output.

10
nfe `int`, *optional*

Number of model evaluations used at inference.

None
eps `float`, *optional*, defaults to 1e-6

Minimum timestep and policy denominator.

1e-06
final_step_size_scale `float`, *optional*, defaults to 0.5

Relative size of the final raw-timestep segment.

0.5
num_policy_substeps `int`, *optional*, defaults to 128

Maximum policy integration substeps per raw-timestep unit.

128
Source code in fastvideo/models/schedulers/scheduling_piflow.py
@register_to_config
def __init__(
    self,
    num_train_timesteps: int = 1000,
    shift: float = 5.0,
    n_grid: int = 10,
    nfe: int | None = None,
    eps: float = 1e-6,
    final_step_size_scale: float = 0.5,
    num_policy_substeps: int = 128,
) -> None:
    if n_grid < 2:
        raise ValueError(f"PiflowScheduler requires n_grid >= 2, got {n_grid}")
    if eps <= 0:
        raise ValueError(f"PiflowScheduler requires eps > 0, got {eps}")
    if not 0 < final_step_size_scale <= 1:
        raise ValueError("PiflowScheduler requires 0 < final_step_size_scale <= 1")
    if num_policy_substeps < 1:
        raise ValueError("PiflowScheduler requires num_policy_substeps >= 1")
    super().__init__(
        num_train_timesteps=num_train_timesteps,
        shift=shift,
    )
    self.n_grid = int(n_grid)
    self.eps = float(eps)
    self.final_step_size_scale = float(final_step_size_scale)
    self.num_policy_substeps = int(num_policy_substeps)
    self._piflow_raw_timesteps = torch.empty(0)

Methods:

fastvideo.models.schedulers.scheduling_piflow.PiflowScheduler.set_timesteps
set_timesteps(num_inference_steps: int | None = None, device: str | device | None = None, sigmas: list[float] | None = None, mu: float | None = None, timesteps: list[float] | None = None) -> None

Set the distilled PiFlow timestep schedule.

Parameters:

Name Type Description Default
num_inference_steps `int`

Number of model evaluations.

None
device `str` or `torch.device`, *optional*

Device for the schedule.

None
sigmas `list[float]`, *optional*

Unsupported custom sigma schedule.

None
mu `float`, *optional*

Unsupported dynamic-shift parameter.

None
timesteps `list[float]`, *optional*

Unsupported custom timestep schedule.

None
Source code in fastvideo/models/schedulers/scheduling_piflow.py
def set_timesteps(
    self,
    num_inference_steps: int | None = None,
    device: str | torch.device | None = None,
    sigmas: list[float] | None = None,
    mu: float | None = None,
    timesteps: list[float] | None = None,
) -> None:
    """Set the distilled PiFlow timestep schedule.

    Args:
        num_inference_steps (`int`): Number of model evaluations.
        device (`str` or `torch.device`, *optional*): Device for the schedule.
        sigmas (`list[float]`, *optional*): Unsupported custom sigma schedule.
        mu (`float`, *optional*): Unsupported dynamic-shift parameter.
        timesteps (`list[float]`, *optional*): Unsupported custom timestep schedule.
    """
    if sigmas is not None or mu is not None or timesteps is not None:
        raise ValueError("PiflowScheduler only supports its configured distilled timestep schedule")
    if num_inference_steps is None or num_inference_steps < 1:
        raise ValueError(f"num_inference_steps must be positive, got {num_inference_steps}")
    one_minus_final = 1.0 - self.final_step_size_scale
    segment = 1.0 / (num_inference_steps - one_minus_final)
    raw = 1.0 - torch.arange(num_inference_steps, dtype=torch.float32, device=device) * segment
    sigmas = shift_timesteps(raw, float(self.config.shift))
    self.num_inference_steps = int(num_inference_steps)
    self._piflow_raw_timesteps = raw
    self.timesteps = sigmas * self.config.num_train_timesteps
    self.sigmas = torch.cat([sigmas, sigmas.new_zeros(1)])
    self._step_index = None
    self._begin_index = None
fastvideo.models.schedulers.scheduling_piflow.PiflowScheduler.step
step(model_output: FloatTensor, timestep: float | FloatTensor, sample: FloatTensor, return_dict: bool = True) -> FlowMatchEulerDiscreteSchedulerOutput | tuple

Advance one step by integrating the PiFlow policy.

Parameters:

Name Type Description Default
model_output `torch.FloatTensor`

Widened model output containing n_grid predictions per sample channel.

required
timestep `float` or `torch.FloatTensor`

Current scheduler timestep.

required
sample `torch.FloatTensor`

Current noisy sample.

required
return_dict `bool`, *optional*, defaults to True

Whether to return a [FlowMatchEulerDiscreteSchedulerOutput].

True

Returns:

Type Description
FlowMatchEulerDiscreteSchedulerOutput | tuple

[FlowMatchEulerDiscreteSchedulerOutput] or tuple: Updated sample.

Source code in fastvideo/models/schedulers/scheduling_piflow.py
def step(
    self,
    model_output: torch.FloatTensor,
    timestep: float | torch.FloatTensor,
    sample: torch.FloatTensor,
    return_dict: bool = True,
) -> FlowMatchEulerDiscreteSchedulerOutput | tuple:
    """Advance one step by integrating the PiFlow policy.

    Args:
        model_output (`torch.FloatTensor`): Widened model output containing
            ``n_grid`` predictions per sample channel.
        timestep (`float` or `torch.FloatTensor`): Current scheduler timestep.
        sample (`torch.FloatTensor`): Current noisy sample.
        return_dict (`bool`, *optional*, defaults to True): Whether to return
            a [`FlowMatchEulerDiscreteSchedulerOutput`].

    Returns:
        [`FlowMatchEulerDiscreteSchedulerOutput`] or `tuple`: Updated sample.
    """
    if isinstance(timestep, int) or isinstance(timestep, (torch.IntTensor, torch.LongTensor)):
        raise ValueError(
            "Passing integer indices as timesteps to PiflowScheduler.step() is not supported; "
            "pass a value from scheduler.timesteps instead"
        )
    step_index = self._step_index_for(timestep)
    # PiFlow's policy rollout performs its update in float32. Keep that
    # precision across outer steps, matching the native sampler.
    updated = self._policy_step(model_output, sample.to(torch.float32), step_index)
    self._step_index += 1
    if return_dict:
        return FlowMatchEulerDiscreteSchedulerOutput(prev_sample=updated)
    return (updated,)

Functions:

fastvideo.models.schedulers.scheduling_piflow.policy_rollout_fm

policy_rollout_fm(x_t_start: Tensor, sigma_t_start: Tensor, raw_t_start: Tensor, raw_t_end: Tensor, total_substeps: int, policy: DXPolicy) -> tuple[Tensor, Tensor, Tensor]

Integrate policy.pi from raw_t_start to raw_t_end.

Source code in fastvideo/models/schedulers/scheduling_piflow.py
def policy_rollout_fm(  # noqa: PLR0913
    x_t_start: torch.Tensor,
    sigma_t_start: torch.Tensor,
    raw_t_start: torch.Tensor,
    raw_t_end: torch.Tensor,
    total_substeps: int,
    policy: DXPolicy,
) -> tuple[torch.Tensor, torch.Tensor, torch.Tensor]:
    """Integrate ``policy.pi`` from ``raw_t_start`` to ``raw_t_end``."""
    num_batches = x_t_start.size(0)
    ndim = x_t_start.dim()
    shape = (num_batches, *((ndim - 1) * [1]))
    raw_t_start = raw_t_start.reshape(shape)
    raw_t_end = raw_t_end.reshape(shape)
    sigma_t = sigma_t_start.reshape(shape)

    delta_raw_t = raw_t_start - raw_t_end
    num_substeps = (delta_raw_t * total_substeps).round().to(torch.long).clamp(min=1)
    substep_size = delta_raw_t / num_substeps
    max_num_substeps = num_substeps.max()

    raw_t = raw_t_start
    x_t = x_t_start
    for substep_id in range(max_num_substeps.item()):
        velocity = policy.pi(x_t, sigma_t)
        raw_t_minus = (raw_t - substep_size).clamp(min=0)
        sigma_t_minus = shift_timesteps(raw_t_minus, policy.shift)
        x_t_minus = x_t + velocity * (sigma_t_minus - sigma_t)

        active_mask = num_substeps > substep_id
        x_t = torch.where(active_mask, x_t_minus, x_t)
        sigma_t = torch.where(active_mask, sigma_t_minus, sigma_t)
        raw_t = torch.where(active_mask, raw_t_minus, raw_t)

    return x_t, sigma_t, sigma_t.flatten() * 1_000

fastvideo.models.schedulers.scheduling_piflow.shift_timesteps

shift_timesteps(t: Tensor, shift: float) -> Tensor

Map raw flow-matching time to the shifted DiT time.

Source code in fastvideo/models/schedulers/scheduling_piflow.py
def shift_timesteps(t: torch.Tensor, shift: float) -> torch.Tensor:
    """Map raw flow-matching time to the shifted DiT time."""
    return shift * t / (1 + (shift - 1) * t)