Skip to content

CFG cutoff callbacks keep a single embedding row, so they break for num_images_per_prompt > 1 or several prompts #14907

Description

@JoeyTan21

Describe the bug

The four official CFG cutoff callbacks in src/diffusers/callbacks.py (SDCFGCutoffCallback, SDXLCFGCutoffCallback, SDXLControlnetCFGCutoffCallback, SD3CFGCutoffCallback) disable CFG at the cutoff step by setting pipeline._guidance_scale = 0.0 and slicing the conditioning tensors with [-1:] (# "-1" denotes the embeddings for conditional text tokens).

With CFG the pipelines build prompt_embeds = torch.cat([negative_prompt_embeds, prompt_embeds]) (same for add_text_embeds, add_time_ids, pooled_prompt_embeds and the ControlNet image), so the batch is [negative, conditional] with 2 * batch_size rows. [-1:] keeps a single row instead of the conditional batch_size rows, which only works for a batch of one:

  • StableDiffusionPipeline, StableDiffusionXLPipeline, StableDiffusionXLControlNetPipeline: the step after the cutoff runs the UNet with batch_size latents and one encoder hidden state, and fails with RuntimeError: The size of tensor a (512) must match the size of tensor b (256) at non-singleton dimension 1 (for num_images_per_prompt=2 as well as for two prompts).
  • StableDiffusion3Pipeline: no crash, because the joint attention broadcasts the single remaining row, but after the cutoff every sample is conditioned on the last prompt of the batch. A second callback in the same run shows prompt_embeds.shape[0] == 1 while latents.shape[0] == 2.

docs/source/en/using-diffusers/callback.md presents these as the official way to drop CFG after a number of steps and nothing there limits them to a batch size of one. There are currently no tests for these callbacks in tests/.

I have a fix ready and can open the PR if the approach is fine: keep the conditional half of each tensor (x[-batch_size:] with batch_size = prompt_embeds.shape[0] // 2, which also keeps a non-duplicated guess_mode ControlNet image whole), skip the slicing when pipeline.do_classifier_free_guidance is false (there is no negative batch to drop then), and add one test_cfg_cutoff_callback per pipeline test file.

Reproduction

Tiny random SD components on CPU; the only download is the hf-internal-testing/tiny-random-clip tokenizer.

import torch
from transformers import CLIPTextConfig, CLIPTextModel, CLIPTokenizer

from diffusers import AutoencoderKL, DDIMScheduler, StableDiffusionPipeline, UNet2DConditionModel
from diffusers.callbacks import SDCFGCutoffCallback


torch.manual_seed(0)
unet = UNet2DConditionModel(
    block_out_channels=(4, 8),
    layers_per_block=1,
    sample_size=32,
    in_channels=4,
    out_channels=4,
    down_block_types=("DownBlock2D", "CrossAttnDownBlock2D"),
    up_block_types=("CrossAttnUpBlock2D", "UpBlock2D"),
    cross_attention_dim=8,
    norm_num_groups=2,
)
vae = AutoencoderKL(
    block_out_channels=[4, 8],
    in_channels=3,
    out_channels=3,
    down_block_types=["DownEncoderBlock2D", "DownEncoderBlock2D"],
    up_block_types=["UpDecoderBlock2D", "UpDecoderBlock2D"],
    latent_channels=4,
    norm_num_groups=2,
)
text_encoder = CLIPTextModel(
    CLIPTextConfig(
        bos_token_id=0,
        eos_token_id=2,
        hidden_size=8,
        intermediate_size=16,
        layer_norm_eps=1e-05,
        num_attention_heads=2,
        num_hidden_layers=2,
        pad_token_id=1,
        vocab_size=1000,
    )
)
tokenizer = CLIPTokenizer.from_pretrained("hf-internal-testing/tiny-random-clip")
scheduler = DDIMScheduler(beta_start=0.00085, beta_end=0.012, beta_schedule="scaled_linear", clip_sample=False, steps_offset=1)
pipe = StableDiffusionPipeline(
    unet=unet,
    scheduler=scheduler,
    vae=vae,
    text_encoder=text_encoder,
    tokenizer=tokenizer,
    safety_checker=None,
    feature_extractor=None,
    image_encoder=None,
    requires_safety_checker=False,
)
pipe.set_progress_bar_config(disable=True)


def run(**kwargs):
    callback = SDCFGCutoffCallback(cutoff_step_ratio=0.5)
    return pipe(
        num_inference_steps=4,
        guidance_scale=6.0,
        output_type="np",
        callback_on_step_end=callback,
        callback_on_step_end_tensor_inputs=callback.tensor_inputs,
        generator=torch.Generator().manual_seed(0),
        **kwargs,
    ).images


print("batch=1                 ->", run(prompt="a").shape)
for label, kwargs in [
    ("two prompts            ", {"prompt": ["a", "b"]}),
    ("num_images_per_prompt=2", {"prompt": "a", "num_images_per_prompt": 2}),
]:
    try:
        print(f"{label} ->", run(**kwargs).shape)
    except Exception as e:
        print(f"{label} -> {type(e).__name__}: {e}")

Logs

batch=1                 -> (1, 64, 64, 3)
two prompts             -> RuntimeError: The size of tensor a (512) must match the size of tensor b (256) at non-singleton dimension 1
num_images_per_prompt=2 -> RuntimeError: The size of tensor a (512) must match the size of tensor b (256) at non-singleton dimension 1

System Info

  • 🤗 Diffusers version: 0.41.0.dev0 (main, fef717f)
  • Platform: macOS-15.3.1-arm64-arm-64bit
  • Running on Google Colab?: No
  • Python version: 3.12.13
  • PyTorch version (GPU?): 2.14.0 (False)
  • Huggingface_hub version: 1.33.0
  • Transformers version: 5.17.0
  • Accelerate version: 1.15.0
  • PEFT version: 0.21.1
  • Safetensors version: 0.8.0
  • xFormers version: not installed
  • Accelerator: Apple M3
  • Using GPU in script?: No
  • Using distributed or parallel set-up in script?: No

Who can help?

@asomoza @yiyixuxu

Disclosure: found, reproduced and the draft fix tested locally with the help of an AI coding agent (Claude Code); reviewed before filing.

Activity

Sign up for free to join this conversation on GitHub. Already have an account? Sign in to comment

Metadata

Metadata

Assignees

No one assigned

    Labels

    No labels
    No labels

    Type

    No type

    Projects

    No projects

      Milestone

      No milestone

      Relationships

      None yet

      Development

      No branches or pull requests

      Issue actions