[diffusion] fix: fix single-step flow-match timesteps (#24708)
This commit is contained in:
@@ -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:
|
||||||
|
|||||||
+2
@@ -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()
|
||||||
)
|
)
|
||||||
|
|||||||
Reference in New Issue
Block a user