Skip to content

RuntimeError: The size of tensor a (4608) must match the size of tensor b (5120) at non-singleton dimension 2 during DreamBooth Training with Prior Preservation #10722

Description

@yinguoweiOvO

Describe the bug

I am trying to run "train_dreambooth_lora_flux.py" on my dataset, but the error will happen if --with_prior_preservation is used.

Who can help me? Thanks!

Reproduction

python ./examples/dreambooth/train_dreambooth_lora_flux.py
--pretrained_model_name_or_path=$MODEL_NAME
--instance_data_dir=$INSTANCE_DIR
--output_dir=$OUTPUT_DIR
--with_prior_preservation
--class_data_dir="my_file"
--class_prompt="A photo"
--instance_prompt="A sks photo"
--resolution=1024
--rank=32
--max_train_steps=5000
--checkpointing_steps=100
--seed="0"
--mixed_precision="bf16"
--train_batch_size=1
--guidance_scale=1
--gradient_accumulation_steps=4
--optimizer="prodigy"
--learning_rate=1.
--report_to="tensorboard"
--lr_scheduler="constant"
--lr_warmup_steps=0

Logs

Traceback (most recent call last):
  File "/data4/work/yinguowei/code/diffusers/./examples/dreambooth/train_dreambooth_lora_flux.py", line 1926, in <module>
    main(args)
  File "/data4/work/yinguowei/code/diffusers/./examples/dreambooth/train_dreambooth_lora_flux.py", line 1720, in main
    model_pred = transformer(
  File "/data/miniconda3/envs/diffusers_ygw/lib/python3.10/site-packages/torch/nn/modules/module.py", line 1736, in _wrapped_call_impl
    return self._call_impl(*args, **kwargs)
  File "/data/miniconda3/envs/diffusers_ygw/lib/python3.10/site-packages/torch/nn/modules/module.py", line 1747, in _call_impl
    return forward_call(*args, **kwargs)
  File "/data/miniconda3/envs/diffusers_ygw/lib/python3.10/site-packages/accelerate/utils/operations.py", line 819, in forward
    return model_forward(*args, **kwargs)
  File "/data/miniconda3/envs/diffusers_ygw/lib/python3.10/site-packages/accelerate/utils/operations.py", line 807, in __call__
    return convert_to_fp32(self.model_forward(*args, **kwargs))
  File "/data/miniconda3/envs/diffusers_ygw/lib/python3.10/site-packages/torch/amp/autocast_mode.py", line 44, in decorate_autocast
    return func(*args, **kwargs)
  File "/data4/work/yinguowei/code/diffusers/src/diffusers/models/transformers/transformer_flux.py", line 529, in forward
    encoder_hidden_states, hidden_states = block(
  File "/data/miniconda3/envs/diffusers_ygw/lib/python3.10/site-packages/torch/nn/modules/module.py", line 1736, in _wrapped_call_impl
    return self._call_impl(*args, **kwargs)
  File "/data/miniconda3/envs/diffusers_ygw/lib/python3.10/site-packages/torch/nn/modules/module.py", line 1747, in _call_impl
    return forward_call(*args, **kwargs)
  File "/data4/work/yinguowei/code/diffusers/src/diffusers/models/transformers/transformer_flux.py", line 188, in forward
    attention_outputs = self.attn(
  File "/data/miniconda3/envs/diffusers_ygw/lib/python3.10/site-packages/torch/nn/modules/module.py", line 1736, in _wrapped_call_impl
    return self._call_impl(*args, **kwargs)
  File "/data/miniconda3/envs/diffusers_ygw/lib/python3.10/site-packages/torch/nn/modules/module.py", line 1747, in _call_impl
    return forward_call(*args, **kwargs)
  File "/data4/work/yinguowei/code/diffusers/src/diffusers/models/attention_processor.py", line 595, in forward
    return self.processor(
  File "/data4/work/yinguowei/code/diffusers/src/diffusers/models/attention_processor.py", line 2325, in __call__
    query = apply_rotary_emb(query, image_rotary_emb)
  File "/data4/work/yinguowei/code/diffusers/src/diffusers/models/embeddings.py", line 1204, in apply_rotary_emb
    out = (x.float() * cos + x_rotated.float() * sin).to(x.dtype)
RuntimeError: The size of tensor a (4608) must match the size of tensor b (5120) at non-singleton dimension 2

System Info

  • 🤗 Diffusers version: 0.33.0.dev0
  • Platform: Linux-5.4.119-19.0009.44-x86_64-with-glibc2.28
  • Running on Google Colab?: No
  • Python version: 3.10.16
  • PyTorch version (GPU?): 2.5.1+cu124 (True)
  • Flax version (CPU?/GPU?/TPU?): not installed (NA)
  • Jax version: not installed
  • JaxLib version: not installed
  • Huggingface_hub version: 0.27.1
  • Transformers version: 4.48.1
  • Accelerate version: 1.3.0
  • PEFT version: 0.14.0
  • Bitsandbytes version: not installed
  • Safetensors version: 0.5.2
  • xFormers version: not installed
  • Accelerator: NVIDIA A800-SXM4-80GB, 81920 MiB
    NVIDIA A800-SXM4-80GB, 81920 MiB
    NVIDIA A800-SXM4-80GB, 81920 MiB
    NVIDIA A800-SXM4-80GB, 81920 MiB
    NVIDIA A800-SXM4-80GB, 81920 MiB
    NVIDIA A800-SXM4-80GB, 81920 MiB
    NVIDIA A800-SXM4-80GB, 81920 MiB
    NVIDIA A800-SXM4-80GB, 81920 MiB
  • Using GPU in script?:
  • Using distributed or parallel set-up in script?:

Who can help?

No response

Activity

  1. yinguoweiOvO commented on Feb 5, 2025

    @yinguoweiOvO
    Author

    And I guess this bug maybe is caused by following code, these code make the shape of text_ids from (512, 3) to (1024, 3). How can I fix it? Please help me
    if not train_dataset.custom_instance_prompts: if not args.train_text_encoder: prompt_embeds = instance_prompt_hidden_states pooled_prompt_embeds = instance_pooled_prompt_embeds text_ids = instance_text_ids if args.with_prior_preservation: prompt_embeds = torch.cat([prompt_embeds, class_prompt_hidden_states], dim=0) pooled_prompt_embeds = torch.cat([pooled_prompt_embeds, class_pooled_prompt_embeds], dim=0) text_ids = torch.cat([text_ids, class_text_ids], dim=0)

  2. Fir-lat commented on Feb 15, 2025

    @Fir-lat

    exactly same issue :(

  3. Fir-lat commented on Feb 15, 2025

    @Fir-lat

    And I guess this bug maybe is caused by following code, these code make the shape of text_ids from (512, 3) to (1024, 3). How can I fix it? Please help me if not train_dataset.custom_instance_prompts: if not args.train_text_encoder: prompt_embeds = instance_prompt_hidden_states pooled_prompt_embeds = instance_pooled_prompt_embeds text_ids = instance_text_ids if args.with_prior_preservation: prompt_embeds = torch.cat([prompt_embeds, class_prompt_hidden_states], dim=0) pooled_prompt_embeds = torch.cat([pooled_prompt_embeds, class_pooled_prompt_embeds], dim=0) text_ids = torch.cat([text_ids, class_text_ids], dim=0)

    It seems that this code doesn't make the shape of '''prompt_embeds''', '''pooled_prompt_embeds''', and '''text_ids''' match the model_inputs. As the '''prompt''' is contained in every batch, I managed to solve this issue by simply modifying the line 1538-1546 to

    else:
        elems_to_repeat = len(prompts)
        if args.train_text_encoder:
            prompt_embeds, pooled_prompt_embeds, text_ids = encode_prompt(
                text_encoders=[text_encoder_one, text_encoder_two],
                tokenizers=[None, None],
                text_input_ids_list=[
                    tokens_one.repeat(elems_to_repeat, 1),
                    tokens_two.repeat(elems_to_repeat, 1),
                ],
                max_sequence_length=args.max_sequence_length,
                device=accelerator.device,
                prompt=args.instance_prompt,
            )
        else:
            prompt_embeds, pooled_prompt_embeds, text_ids = compute_text_embeddings(
                prompts, text_encoders, tokenizers
            )
    

    This intended to encode the prompts in every training step to make sure that the shape of text embeddings match that of model_inputs

  4. yinguoweiOvO commented on Feb 17, 2025

    @yinguoweiOvO
    Author

    Yeah, I also find that because of the problem described in this issue, the text ids created no longer have the batch size dimension, so concatenating the class with the instance text ids in the following code results in a subsequent dimension error.

    if not args.train_text_encoder:
        prompt_embeds = instance_prompt_hidden_states
        pooled_prompt_embeds = instance_pooled_prompt_embeds
        text_ids = instance_text_ids
        if args.with_prior_preservation:
            prompt_embeds = torch.cat([prompt_embeds, class_prompt_hidden_states], dim=0)
            pooled_prompt_embeds = torch.cat([pooled_prompt_embeds, class_pooled_prompt_embeds], dim=0)
            text_ids = torch.cat([text_ids, class_text_ids], dim=0)
    

    Since text ids are all 0 vectors (as set in the method), this bug can be fixed by not concat the text ids of class and instance, and the rest of the code is not a problem because python has a broadcast mechanism

  5. github-actions commented on Mar 13, 2025

    @github-actions
    Contributor

    This issue has been automatically marked as stale because it has not had recent activity. If you think this still needs to be addressed please comment on this thread.

    Please note that issues that do not follow the contributing guidelines are likely to be ignored.

  6. qqzy commented on May 24, 2025

    @qqzy

    text_ids = torch.cat([text_ids, class_text_ids], dim=0) causes incorrect concatenation. Removing this line will make it work properly.

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

    bugSomething isn't workingstaleIssues that haven't received updates

    Type

    No type

    Projects

    No projects

      Milestone

      No milestone

      Relationships

      None yet

      Development

      No branches or pull requests

      Issue actions