[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,6 +82,7 @@ class FlowMatchScheduler(BaseScheduler):
if self.shift_terminal is not None: if self.shift_terminal is not None:
one_minus_z = 1 - self.sigmas one_minus_z = 1 - self.sigmas
scale_factor = one_minus_z[-1] / (1 - self.shift_terminal) scale_factor = one_minus_z[-1] / (1 - self.shift_terminal)
if scale_factor != 0:
self.sigmas = 1 - (one_minus_z / scale_factor) self.sigmas = 1 - (one_minus_z / scale_factor)
if self.reverse_sigmas: if self.reverse_sigmas:
self.sigmas = 1 - self.sigmas self.sigmas = 1 - self.sigmas
@@ -421,6 +422,7 @@ class FlowMatchPairScheduler(FlowMatchScheduler):
if self.shift_terminal is not None: if self.shift_terminal is not None:
one_minus_z = 1 - base one_minus_z = 1 - base
scale_factor = one_minus_z[-1] / (1 - self.shift_terminal) scale_factor = one_minus_z[-1] / (1 - self.shift_terminal)
if scale_factor != 0:
base = 1 - (one_minus_z / scale_factor) base = 1 - (one_minus_z / scale_factor)
if self.reverse_sigmas: if self.reverse_sigmas:
@@ -266,6 +266,8 @@ class FlowMatchEulerDiscreteScheduler(
""" """
one_minus_z = 1 - t one_minus_z = 1 - t
scale_factor = one_minus_z[-1] / (1 - self.config.shift_terminal) 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) stretched_t = 1 - (one_minus_z / scale_factor)
return stretched_t return stretched_t
@@ -191,8 +191,7 @@ class TimestepPreparationStage(PipelineStage):
and isinstance(batch.timesteps, torch.Tensor) and isinstance(batch.timesteps, torch.Tensor)
and torch.isnan(batch.timesteps).any() and torch.isnan(batch.timesteps).any()
): ):
# when num-inference-steps == 1, the last sigma being 1, the 1 / last_sigma could be nan # diffusers flow-match scheduler can emit NaN for one-step warmup
# this a workaround for warmup req only
batch.timesteps = torch.ones( batch.timesteps = torch.ones(
(1,), dtype=torch.float32, device=get_local_torch_device() (1,), dtype=torch.float32, device=get_local_torch_device()
) )