Skip to content
Open
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
7 changes: 6 additions & 1 deletion src/diffusers/pipelines/qwenimage21/pipeline_qwenimage21.py
Original file line number Diff line number Diff line change
Expand Up @@ -720,7 +720,12 @@ def __call__(
] * batch_size

# 4. Prepare timesteps
sigmas = np.linspace(1.0, 1 / num_inference_steps, num_inference_steps) if sigmas is None else sigmas
if sigmas is None:
sample_sigmas = self.scheduler.config.get("sample_sigmas")
if sample_sigmas is not None:
num_inference_steps = len(sample_sigmas)
else:
sigmas = np.linspace(1.0, 1 / num_inference_steps, num_inference_steps)
mu = calculate_shift(
latents.shape[1],
self.scheduler.config.get("base_image_seq_len", 256),
Expand Down
34 changes: 34 additions & 0 deletions src/diffusers/schedulers/scheduling_flow_match_euler_discrete.py
Original file line number Diff line number Diff line change
Expand Up @@ -83,6 +83,12 @@ class FlowMatchEulerDiscreteScheduler(SchedulerMixin, ConfigMixin):
The type of dynamic resolution-dependent timestep shifting to apply. Either "exponential" or "linear".
stochastic_sampling (`bool`, defaults to False):
Whether to use stochastic sampling.
sample_sigmas (`list[float]`, *optional*):
A fixed list of pre-computed sigma values (descending, excluding the terminal 0) to use as the sampling
schedule. When set, `set_timesteps` uses these values directly, **bypassing** dynamic shifting,
shift-terminal stretching, and karras/exponential/beta sigma conversions. This is useful for distilled
models whose training grid was derived from a specific set of sigma points. The terminal sigma (0) is
appended automatically. If `None` (default), sigmas are computed on the fly as usual.
"""

_compatibles = []
Expand All @@ -105,6 +111,7 @@ def __init__(
use_beta_sigmas: bool = False,
time_shift_type: Literal["exponential", "linear"] = "exponential",
stochastic_sampling: bool = False,
sample_sigmas: list[float] | None = None,
):
if self.config.use_beta_sigmas and not is_scipy_available():
raise ImportError("Make sure to install scipy if you want to use beta sigmas.")
Expand Down Expand Up @@ -306,6 +313,33 @@ def set_timesteps(
Custom values for timesteps to be used for each diffusion step. If `None`, the timesteps are computed
automatically.
"""
# Fast path: if `sample_sigmas` is set in config and the caller did not pass explicit

Copy link
Copy Markdown

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

i think it's easier to make the sample_sigmas a pipeline config, instead of a scheduler one, because we try to keep our our scheduler to a standard set of config/arguments so that people can swap them.

also, I think we would not need to make any changes to scheduler when we pass sigmas directly if we just also disable dynamic shifting/shift-terminal stretching etc in the scheduler config for distilled checkpoint here (but for the distilled checkpoint instead) https://huggingface.co/Qwen/Qwen-Image-2.1/blob/main/scheduler/scheduler_config.json, e.g. we can just set use_dynamic_shifting=False shift=1.0 etc

# sigmas/timesteps, use the pre-computed values directly. This bypasses dynamic shifting,
# shift-terminal stretching, and karras/exponential/beta conversions — the stored sigmas are
# assumed to be final (e.g. a distilled model's fixed sampling grid).
if self.config.sample_sigmas is not None and sigmas is None and timesteps is None:
sample_sigmas = np.array(self.config.sample_sigmas, dtype=np.float32)
if num_inference_steps is not None and num_inference_steps != len(sample_sigmas):
raise ValueError(
f"`num_inference_steps` ({num_inference_steps}) does not match the length of "
f"`sample_sigmas` ({len(sample_sigmas)}) in the scheduler config. Either omit "
f"`num_inference_steps` or pass a value that matches."
)
self.num_inference_steps = len(sample_sigmas)
sigmas_tensor = torch.from_numpy(sample_sigmas).to(dtype=torch.float32, device=device)
timesteps = sigmas_tensor * self.config.num_train_timesteps
if self.config.invert_sigmas:
sigmas_tensor = 1.0 - sigmas_tensor
timesteps = sigmas_tensor * self.config.num_train_timesteps
sigmas_tensor = torch.cat([sigmas_tensor, torch.ones(1, device=sigmas_tensor.device)])
else:
sigmas_tensor = torch.cat([sigmas_tensor, torch.zeros(1, device=sigmas_tensor.device)])
self.timesteps = timesteps
self.sigmas = sigmas_tensor
self._step_index = None
self._begin_index = None
return

if self.config.use_dynamic_shifting and mu is None:
raise ValueError("`mu` must be passed when `use_dynamic_shifting` is set to be `True`")

Expand Down
45 changes: 45 additions & 0 deletions tests/pipelines/qwenimage21/test_qwenimage21.py
Original file line number Diff line number Diff line change
Expand Up @@ -182,6 +182,51 @@ def get_dummy_inputs(self):


class TestQwenImage21Pipeline(QwenImage21PipelineTesterConfig, PipelineTesterMixin):
@pytest.mark.parametrize("num_inference_steps", [None, 2])
def test_sample_sigmas_from_scheduler_config(self, num_inference_steps):
pipe = self.get_pipeline()
sample_sigmas = [1.0, 0.978453, 0.954180, 0.926626, 0.895080, 0.845148, 0.704534, 0.414568]
pipe.scheduler = FlowMatchEulerDiscreteScheduler.from_config(
pipe.scheduler.config,
sample_sigmas=sample_sigmas,
use_dynamic_shifting=True,
shift_terminal=0.02,
)
inputs = self.get_dummy_inputs()
inputs["output_type"] = "latent"
if num_inference_steps is None:
inputs.pop("num_inference_steps")
else:
inputs["num_inference_steps"] = num_inference_steps
seen_timesteps = []

def callback(pipe, step, timestep, callback_kwargs):
seen_timesteps.append(timestep.clone())
return callback_kwargs

pipe(**inputs, callback_on_step_end=callback)
expected_sigmas = torch.tensor(sample_sigmas + [0.0], device=pipe.scheduler.sigmas.device)
expected_timesteps = expected_sigmas[:-1] * pipe.scheduler.config.num_train_timesteps
assert torch.equal(pipe.scheduler.sigmas, expected_sigmas)
assert torch.equal(pipe.scheduler.timesteps, expected_timesteps)
assert torch.equal(torch.stack(seen_timesteps), expected_timesteps)
assert pipe.scheduler.num_inference_steps == len(sample_sigmas)

def test_explicit_sigmas_override_scheduler_config(self):
pipe = self.get_pipeline()
pipe.scheduler = FlowMatchEulerDiscreteScheduler.from_config(
pipe.scheduler.config, use_dynamic_shifting=True, shift_terminal=0.02
)
inputs = self.get_dummy_inputs()
inputs.update(sigmas=[1.0, 0.6, 0.2], output_type="latent")
expected = pipe(**inputs).images
expected_sigmas = pipe.scheduler.sigmas.clone()
pipe.scheduler.register_to_config(sample_sigmas=[1.0, 0.8])
inputs["generator"] = self.get_generator(0)
actual = pipe(**inputs).images
assert torch.equal(pipe.scheduler.sigmas, expected_sigmas)
assert torch.equal(actual, expected)

@PROCESSOR_REQUIRED_AT_INIT
def test_encode_prompt_works_in_isolation(self):
super().test_encode_prompt_works_in_isolation()
Expand Down
Loading