Skip to content

Qwen Image 2.1: baseline sampling crashes when text embeddings are cached ("prompt embeddings hold 741 reference image slots but 3 reference image(s) were passed") #1059

Description

@KJeon10

This is for bugs only

Did you already ask in the discord?

No. I searched #qwen_image_2 first. There are two reports of the same slot mismatch error from Sep 21 and Sep 22 (both with cached text embeddings, one of them clearly in the training step, and both on code from before 07abdbe). The workaround there was turning off text embedding caching. I couldn't find anything about the sampling path on current main, which is what this issue is about.

You verified that this is a bug and not a feature request or question by asking in the discord?

No, verified by reading the code and testing the fix locally (details below).

Describe the bug

What happens

Training a Qwen Image 2.1 LoRA with 3 control images and cache_text_embeddings: true. Text embedding caching finishes fine (2364 items, ~46 min), then the job dies on the very first baseline sample:

Generating baseline samples before training
Generating Samples:   0%|          | 0/7
...
  File ".../qwen_image_2/qwen_image_2.py", line 579, in generate_single_image
    return pipeline(
  File ".../qwen_image_2/src/pipeline.py", line 389, in __call__
    noise_pred = run_transformer(
  File ".../qwen_image_2/src/pipeline.py", line 268, in run_transformer
    raise ValueError(
ValueError: the prompt embeddings hold 741 reference image slots but 3 reference image(s) were passed, worth 4653 slots. They must be encoded together -- an embedding cached against other references (caption dropout, a blank unconditional embedding) cannot be reused here.

Relevant config:

  • arch: qwen_image_2, model_kwargs.match_target_res: true
  • cache_text_embeddings: true
  • 3 control paths, sample width: 1056, height: 1504, 7 samples each with ctrl_img_1/2/3
  • ai-toolkit main @ 0bd3411

Why

When text embeddings are cached, the TE gets unloaded, so SDTrainer.cache_sample_prompts() pre-encodes the sample prompts first. It builds a GenerateImageConfig there without passing width/height, so the config falls back to its default 512x512. That gets passed on as target_size:

target_size = (gen_img_config.width, gen_img_config.height)   # (512, 512)
positive = self.sd.encode_prompt(..., control_images=ctrl_img, target_size=target_size)

With match_target_res, each reference is scaled to the target's pixel area, so the sample prompt embeddings are built with references scaled to 512x512 area. That's 247 slots per image, 741 for 3.

At actual generation time generate_single_image() scales the references to the real sample size (1056x1504), which is 1551 slots per image, 4653 for 3. So the slot check in run_transformer fails.

Without cached text embeddings this doesn't happen, because base_model.generate_images() encodes with target_size=(gen_config.width, gen_config.height) directly. And with match_target_res: false it doesn't happen either, since both sides use the fixed pixel cap. It also wouldn't show up if the samples happen to be 512x512.

The training path is fine on current main. The dataloader already keys the cached embeddings on the bucket size, so only the sample prompt cache is affected.

Fix

Pass the sample size through, same as BaseSDTrainProcess.sample() already does when it builds its GenerateImageConfig:

--- a/extensions_built_in/sd_trainer/SDTrainer.py
+++ b/extensions_built_in/sd_trainer/SDTrainer.py
@@ -171,6 +171,8 @@ class SDTrainer(BaseSDTrainProcess):
                 gen_img_config = GenerateImageConfig(
                     prompt=prompt, # it will autoparse the prompt
                     negative_prompt=sample_item.neg,
+                    width=sample_item.width,
+                    height=sample_item.height,
                     output_path=output_path,
                     ctrl_img=sample_item.ctrl_img,
                     ctrl_img_1=sample_item.ctrl_img_1,

After this the sample prompt embeddings are encoded against references at the same size generation uses. For models that ignore target_size nothing changes, since width/height in this config aren't used for anything else here. --w/--h flags in the prompt still override it, same as in the regular sampling path.

Tested

Applied locally and restarted the job. The existing text embedding cache was reused (no re-caching needed), the baseline samples generated fine, and so did the step 250 samples.

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