[diffusion] fix: fix single-step flow-match timesteps (#24708)

This commit is contained in:
Mick
2026-05-11 18:09:24 +08:00
committed by GitHub
parent c027ae677c
commit 6d30b571b2
3 changed files with 7 additions and 4 deletions
@@ -82,7 +82,8 @@ class FlowMatchScheduler(BaseScheduler):
if self.shift_terminal is not None:
one_minus_z = 1 - self.sigmas
scale_factor = one_minus_z[-1] / (1 - self.shift_terminal)
self.sigmas = 1 - (one_minus_z / scale_factor)
if scale_factor != 0:
self.sigmas = 1 - (one_minus_z / scale_factor)
if self.reverse_sigmas:
self.sigmas = 1 - self.sigmas
self.timesteps = self.sigmas * self.num_train_timesteps
@@ -421,7 +422,8 @@ class FlowMatchPairScheduler(FlowMatchScheduler):
if self.shift_terminal is not None:
one_minus_z = 1 - base
scale_factor = one_minus_z[-1] / (1 - self.shift_terminal)
base = 1 - (one_minus_z / scale_factor)
if scale_factor != 0:
base = 1 - (one_minus_z / scale_factor)
if self.reverse_sigmas:
base = 1 - base
@@ -266,6 +266,8 @@ class FlowMatchEulerDiscreteScheduler(
"""
one_minus_z = 1 - t
scale_factor = one_minus_z[-1] / (1 - self.config.shift_terminal)
if scale_factor == 0:
return t
stretched_t = 1 - (one_minus_z / scale_factor)
return stretched_t
@@ -191,8 +191,7 @@ class TimestepPreparationStage(PipelineStage):
and isinstance(batch.timesteps, torch.Tensor)
and torch.isnan(batch.timesteps).any()
):
# when num-inference-steps == 1, the last sigma being 1, the 1 / last_sigma could be nan
# this a workaround for warmup req only
# diffusers flow-match scheduler can emit NaN for one-step warmup
batch.timesteps = torch.ones(
(1,), dtype=torch.float32, device=get_local_torch_device()
)