Skip to content

scheduling_helios_dmd

Native scheduler for the distilled three-stage Helios pyramid.

Classes

fastvideo.models.schedulers.scheduling_helios_dmd.HeliosDMDScheduler

HeliosDMDScheduler(num_train_timesteps: int = 1000, shift: float = 1.0, stages: int = 3, stage_range: list[float] | None = None, gamma: float = 1 / 3, prediction_type: str = 'flow_prediction', use_flow_sigmas: bool = True, use_dynamic_shifting: bool = False, time_shift_type: Literal['exponential', 'linear'] = 'linear', base_image_seq_len: int = 256, max_image_seq_len: int = 4096, base_shift: float = 0.5, max_shift: float = 1.15, scheduler_type: str = 'dmd', _diffusers_version: str | None = None, **kwargs)

Bases: SchedulerMixin, ConfigMixin, BaseScheduler

DMD flow scheduler used by BestWishYsh/Helios-Distilled.

Source code in fastvideo/models/schedulers/scheduling_helios_dmd.py
@register_to_config
def __init__(
    self,
    num_train_timesteps: int = 1000,
    shift: float = 1.0,
    stages: int = 3,
    stage_range: list[float] | None = None,
    gamma: float = 1 / 3,
    prediction_type: str = "flow_prediction",
    use_flow_sigmas: bool = True,
    use_dynamic_shifting: bool = False,
    time_shift_type: Literal["exponential", "linear"] = "linear",
    # Read off ``self.config`` by ``HeliosPyramidDenoisingStage`` to derive
    # ``mu``; they must be declared here or checkpoint values are swallowed
    # by ``**kwargs`` and replaced by these defaults.
    base_image_seq_len: int = 256,
    max_image_seq_len: int = 4096,
    base_shift: float = 0.5,
    max_shift: float = 1.15,
    scheduler_type: str = "dmd",
    _diffusers_version: str | None = None,
    **kwargs,
) -> None:
    del scheduler_type, _diffusers_version, kwargs
    if stage_range is None:
        stage_range = [0, 1 / 3, 2 / 3, 1]
        self.register_to_config(stage_range=stage_range)
    self.num_train_timesteps = num_train_timesteps
    self.timestep_ratios: dict[int, tuple[float, float]] = {}
    self.timesteps_per_stage: dict[int, torch.Tensor] = {}
    self.sigmas_per_stage: dict[int, torch.Tensor] = {}
    self.start_sigmas: dict[int, float] = {}
    self.end_sigmas: dict[int, float] = {}
    self.ori_start_sigmas: dict[int, float] = {}

    self.init_sigmas_for_each_stage()
    self.sigma_min = self.sigmas[-1].item()
    self.sigma_max = self.sigmas[0].item()
    self.gamma = gamma
    self.last_sample = None
    self._step_index = None
    self._begin_index = None
    BaseScheduler.__init__(self)