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
62 changes: 21 additions & 41 deletions src/diffusers/callbacks.py
Original file line number Diff line number Diff line change
Expand Up @@ -86,12 +86,14 @@ def callback_fn(self, pipeline, step_index, timestep, callback_kwargs) -> dict[s
)

if step_index == cutoff_step:
prompt_embeds = callback_kwargs[self.tensor_inputs[0]]
prompt_embeds = prompt_embeds[-1:] # "-1" denotes the embeddings for conditional text tokens.
if pipeline.do_classifier_free_guidance:
# the CFG batches are `[negative, conditional]`; keep the conditional half of each tensor
batch_size = callback_kwargs[self.tensor_inputs[0]].shape[0] // 2
for name in self.tensor_inputs:
callback_kwargs[name] = callback_kwargs[name][-batch_size:]

pipeline._guidance_scale = 0.0

callback_kwargs[self.tensor_inputs[0]] = prompt_embeds
return callback_kwargs


Expand Down Expand Up @@ -119,21 +121,14 @@ def callback_fn(self, pipeline, step_index, timestep, callback_kwargs) -> dict[s
)

if step_index == cutoff_step:
prompt_embeds = callback_kwargs[self.tensor_inputs[0]]
prompt_embeds = prompt_embeds[-1:] # "-1" denotes the embeddings for conditional text tokens.

add_text_embeds = callback_kwargs[self.tensor_inputs[1]]
add_text_embeds = add_text_embeds[-1:] # "-1" denotes the embeddings for conditional pooled text tokens

add_time_ids = callback_kwargs[self.tensor_inputs[2]]
add_time_ids = add_time_ids[-1:] # "-1" denotes the embeddings for conditional added time vector
if pipeline.do_classifier_free_guidance:
# the CFG batches are `[negative, conditional]`; keep the conditional half of each tensor
batch_size = callback_kwargs[self.tensor_inputs[0]].shape[0] // 2
for name in self.tensor_inputs:
callback_kwargs[name] = callback_kwargs[name][-batch_size:]

pipeline._guidance_scale = 0.0

callback_kwargs[self.tensor_inputs[0]] = prompt_embeds
callback_kwargs[self.tensor_inputs[1]] = add_text_embeds
callback_kwargs[self.tensor_inputs[2]] = add_time_ids

return callback_kwargs


Expand Down Expand Up @@ -162,26 +157,15 @@ def callback_fn(self, pipeline, step_index, timestep, callback_kwargs) -> dict[s
)

if step_index == cutoff_step:
prompt_embeds = callback_kwargs[self.tensor_inputs[0]]
prompt_embeds = prompt_embeds[-1:] # "-1" denotes the embeddings for conditional text tokens.

add_text_embeds = callback_kwargs[self.tensor_inputs[1]]
add_text_embeds = add_text_embeds[-1:] # "-1" denotes the embeddings for conditional pooled text tokens

add_time_ids = callback_kwargs[self.tensor_inputs[2]]
add_time_ids = add_time_ids[-1:] # "-1" denotes the embeddings for conditional added time vector

# For Controlnet
image = callback_kwargs[self.tensor_inputs[3]]
image = image[-1:]
if pipeline.do_classifier_free_guidance:
# the CFG batches are `[negative, conditional]`; keep the conditional half of each tensor. In guess
# mode the controlnet `image` is not duplicated, so the slice keeps it whole.
batch_size = callback_kwargs[self.tensor_inputs[0]].shape[0] // 2
for name in self.tensor_inputs:
callback_kwargs[name] = callback_kwargs[name][-batch_size:]

pipeline._guidance_scale = 0.0

callback_kwargs[self.tensor_inputs[0]] = prompt_embeds
callback_kwargs[self.tensor_inputs[1]] = add_text_embeds
callback_kwargs[self.tensor_inputs[2]] = add_time_ids
callback_kwargs[self.tensor_inputs[3]] = image

return callback_kwargs


Expand Down Expand Up @@ -229,16 +213,12 @@ def callback_fn(self, pipeline, step_index, timestep, callback_kwargs) -> dict[s
)

if step_index == cutoff_step:
prompt_embeds = callback_kwargs[self.tensor_inputs[0]]
prompt_embeds = prompt_embeds[-1:] # "-1" denotes the embeddings for conditional text tokens.

pooled_prompt_embeds = callback_kwargs[self.tensor_inputs[1]]
pooled_prompt_embeds = pooled_prompt_embeds[
-1:
] # "-1" denotes the embeddings for conditional pooled text tokens.
if pipeline.do_classifier_free_guidance:
# the CFG batches are `[negative, conditional]`; keep the conditional half of each tensor
batch_size = callback_kwargs[self.tensor_inputs[0]].shape[0] // 2
for name in self.tensor_inputs:
callback_kwargs[name] = callback_kwargs[name][-batch_size:]

pipeline._guidance_scale = 0.0

callback_kwargs[self.tensor_inputs[0]] = prompt_embeds
callback_kwargs[self.tensor_inputs[1]] = pooled_prompt_embeds
return callback_kwargs
37 changes: 37 additions & 0 deletions tests/pipelines/controlnet/test_controlnet_sdxl.py
Original file line number Diff line number Diff line change
Expand Up @@ -31,6 +31,7 @@
StableDiffusionXLImg2ImgPipeline,
UNet2DConditionModel,
)
from diffusers.callbacks import MultiPipelineCallbacks, PipelineCallback, SDXLControlnetCFGCutoffCallback
from diffusers.models.unets.unet_2d_blocks import UNetMidBlock2D
from diffusers.pipelines.controlnet.pipeline_controlnet import MultiControlNetModel
from diffusers.utils.torch_utils import randn_tensor
Expand Down Expand Up @@ -252,6 +253,42 @@ def test_controlnet_sdxl_lcm(self):
image_slice = image[0, -1, -3:, -3:]
assert_tensors_close(image_slice.flatten().cpu(), self.expected_lcm_slice, atol=1e-2)

def test_cfg_cutoff_callback(self):
# after the cutoff the callback must keep one conditional embedding per sample, not only the last row
cutoff_callback = SDXLControlnetCFGCutoffCallback(cutoff_step_ratio=None, cutoff_step_index=1)

class CheckBatchCallback(PipelineCallback):
tensor_inputs = ["latents", "prompt_embeds"]

def callback_fn(self, pipeline, step_index, timestep, callback_kwargs):
if step_index >= 1:
assert callback_kwargs["prompt_embeds"].shape[0] == callback_kwargs["latents"].shape[0]
return callback_kwargs

# Run on CPU: guess mode calls `torch.logspace`, which has no MPS kernel.
pipe = self.get_pipeline()
inputs = self.get_dummy_inputs()
inputs["prompt"] = [inputs["prompt"], "a different prompt"]
inputs["num_inference_steps"] = 3
inputs["callback_on_step_end"] = MultiPipelineCallbacks([cutoff_callback, CheckBatchCallback()])
image = pipe(**inputs).images

assert image.shape == (2, *self.output_shape)
assert pipe.guidance_scale == 0.0

# without CFG there is no negative batch to drop
inputs["guidance_scale"] = 1.0
image = pipe(**inputs).images

assert image.shape == (2, *self.output_shape)

# in guess mode the conditioning image is not duplicated for CFG and has to be kept whole
inputs["guidance_scale"] = 6.0
inputs["guess_mode"] = True
image = pipe(**inputs).images

assert image.shape == (2, *self.output_shape)


class TestStableDiffusionXLControlNetPipeline(
StableDiffusionXLControlNetPipelineTesterConfig, StableDiffusionXLControlNetPipelineTests, PipelineTesterMixin
Expand Down
29 changes: 29 additions & 0 deletions tests/pipelines/stable_diffusion/test_stable_diffusion.py
Original file line number Diff line number Diff line change
Expand Up @@ -45,6 +45,7 @@
UNet2DConditionModel,
logging,
)
from diffusers.callbacks import MultiPipelineCallbacks, PipelineCallback, SDCFGCutoffCallback
from diffusers.utils.import_utils import is_accelerate_available

from ...models.testing_utils.lora import check_if_lora_correctly_set
Expand Down Expand Up @@ -603,6 +604,34 @@ def callback_on_step_end(pipe, i, t, callback_kwargs):
# they should be the same
assert torch.allclose(intermediate_latent, output_interrupted, atol=1e-4)

def test_cfg_cutoff_callback(self):
# after the cutoff the callback must keep one conditional embedding per sample, not only the last row
cutoff_callback = SDCFGCutoffCallback(cutoff_step_ratio=None, cutoff_step_index=1)

class CheckBatchCallback(PipelineCallback):
tensor_inputs = ["latents", "prompt_embeds"]

def callback_fn(self, pipeline, step_index, timestep, callback_kwargs):
if step_index >= 1:
assert callback_kwargs["prompt_embeds"].shape[0] == callback_kwargs["latents"].shape[0]
return callback_kwargs

pipe = self.get_pipeline().to(torch_device)
inputs = self.get_dummy_inputs()
inputs["prompt"] = [inputs["prompt"], "a different prompt"]
inputs["num_inference_steps"] = 3
inputs["callback_on_step_end"] = MultiPipelineCallbacks([cutoff_callback, CheckBatchCallback()])
image = pipe(**inputs).images

assert image.shape == (2, *self.output_shape)
assert pipe.guidance_scale == 0.0

# without CFG there is no negative batch to drop
inputs["guidance_scale"] = 1.0
image = pipe(**inputs).images

assert image.shape == (2, *self.output_shape)

def test_pipeline_accept_tuple_type_unet_sample_size(self):
# the purpose of this test is to see whether the pipeline would accept a unet with the tuple-typed sample size
sd_repo_id = "stable-diffusion-v1-5/stable-diffusion-v1-5"
Expand Down
Original file line number Diff line number Diff line change
Expand Up @@ -13,6 +13,7 @@
)

from diffusers import AutoencoderKL, FlowMatchEulerDiscreteScheduler, SD3Transformer2DModel, StableDiffusion3Pipeline
from diffusers.callbacks import MultiPipelineCallbacks, PipelineCallback, SD3CFGCutoffCallback

from ...testing_utils import (
assert_tensors_close,
Expand Down Expand Up @@ -205,6 +206,34 @@ def test_skip_guidance_layers(self):
assert not torch.allclose(output_full, output_skip, atol=1e-5), "Outputs should differ when layers are skipped"
assert output_full.shape == output_skip.shape, "Outputs should have the same shape"

def test_cfg_cutoff_callback(self):
# after the cutoff the callback must keep one conditional embedding per sample, not only the last row
cutoff_callback = SD3CFGCutoffCallback(cutoff_step_ratio=None, cutoff_step_index=1)

class CheckBatchCallback(PipelineCallback):
tensor_inputs = ["latents", "prompt_embeds"]

def callback_fn(self, pipeline, step_index, timestep, callback_kwargs):
if step_index >= 1:
assert callback_kwargs["prompt_embeds"].shape[0] == callback_kwargs["latents"].shape[0]
return callback_kwargs

pipe = self.get_pipeline().to(torch_device)
inputs = self.get_dummy_inputs()
inputs["prompt"] = [inputs["prompt"], "a different prompt"]
inputs["num_inference_steps"] = 3
inputs["callback_on_step_end"] = MultiPipelineCallbacks([cutoff_callback, CheckBatchCallback()])
image = pipe(**inputs).images

assert image.shape == (2, *self.output_shape)
assert pipe.guidance_scale == 0.0

# without CFG there is no negative batch to drop
inputs["guidance_scale"] = 1.0
image = pipe(**inputs).images

assert image.shape == (2, *self.output_shape)


class TestStableDiffusion3PipelineMemory(StableDiffusion3PipelineTesterConfig, MemoryTesterMixin):
"""Memory optimization tests (CPU offload, group offload, layerwise casting) for the SD3 pipeline."""
Expand Down
29 changes: 29 additions & 0 deletions tests/pipelines/stable_diffusion_xl/test_stable_diffusion_xl.py
Original file line number Diff line number Diff line change
Expand Up @@ -40,6 +40,7 @@
UNet2DConditionModel,
UniPCMultistepScheduler,
)
from diffusers.callbacks import MultiPipelineCallbacks, PipelineCallback, SDXLCFGCutoffCallback
from diffusers.utils import logging

from ...models.testing_utils.lora import check_if_lora_correctly_set
Expand Down Expand Up @@ -993,6 +994,34 @@ def callback_on_step_end(pipe, i, t, callback_kwargs):
# they should be the same
assert_tensors_close(intermediate_latent, output_interrupted, atol=1e-4)

def test_cfg_cutoff_callback(self):
# after the cutoff the callback must keep one conditional embedding per sample, not only the last row
cutoff_callback = SDXLCFGCutoffCallback(cutoff_step_ratio=None, cutoff_step_index=1)

class CheckBatchCallback(PipelineCallback):
tensor_inputs = ["latents", "prompt_embeds"]

def callback_fn(self, pipeline, step_index, timestep, callback_kwargs):
if step_index >= 1:
assert callback_kwargs["prompt_embeds"].shape[0] == callback_kwargs["latents"].shape[0]
return callback_kwargs

pipe = self.get_pipeline().to(torch_device)
inputs = self.get_dummy_inputs()
inputs["prompt"] = [inputs["prompt"], "a different prompt"]
inputs["num_inference_steps"] = 3
inputs["callback_on_step_end"] = MultiPipelineCallbacks([cutoff_callback, CheckBatchCallback()])
image = pipe(**inputs).images

assert image.shape == (2, *self.output_shape)
assert pipe.guidance_scale == 0.0

# without CFG there is no negative batch to drop
inputs["guidance_scale"] = 1.0
image = pipe(**inputs).images

assert image.shape == (2, *self.output_shape)


class TestStableDiffusionXLPipelineMemory(StableDiffusionXLPipelineTesterConfig, MemoryTesterMixin):
"""Memory optimization tests (CPU offload, group offload, layerwise casting) for the SDXL pipeline."""
Expand Down
Loading