diff --git a/examples/dreambooth/README_qwenimage21.md b/examples/dreambooth/README_qwenimage21.md new file mode 100644 index 000000000000..accc02a923d9 --- /dev/null +++ b/examples/dreambooth/README_qwenimage21.md @@ -0,0 +1,216 @@ +# DreamBooth training example for Qwen-Image 2.1 + +[DreamBooth](https://huggingface.co/papers/2208.12242) is a method to personalize text-to-image models given just a few (3~5) images of a subject. + +The `train_dreambooth_lora_qwenimage21.py` script shows how to implement the training procedure with [LoRA](https://huggingface.co/docs/peft/conceptual_guides/adapter#low-rank-adaptation-lora) and adapt it for [Qwen-Image 2.1](https://huggingface.co/Qwen/Qwen-Image-2.1). + +This will also allow us to push the trained model parameters to the Hugging Face Hub platform. + +Qwen-Image 2.1 also takes condition images. That task has its own script, `train_dreambooth_lora_qwenimage21_img2img.py`, +described in [Image-to-image (editing)](#image-to-image-editing) below. + +## Running locally with PyTorch + +### Installing the dependencies + +Before running the scripts, make sure to install the library's training dependencies: + +**Important** + +To make sure you can successfully run the latest versions of the example scripts, we highly recommend **installing from source** and keeping the install up to date as we update the example scripts frequently and install some example-specific requirements. To do this, execute the following steps in a new virtual environment: + +```bash +git clone https://github.com/huggingface/diffusers +cd diffusers +pip install -e . +``` + +Then cd in the `examples/dreambooth` folder and run + +```bash +pip install -r requirements_flux.txt +``` + +And initialize an [🤗Accelerate](https://github.com/huggingface/accelerate/) environment with: + +```bash +accelerate config +``` + +Or for a default accelerate configuration without answering questions about your environment + +```bash +accelerate config default +``` + +Or if your environment doesn't support an interactive shell (e.g., a notebook) + +```python +from accelerate.utils import write_basic_config +write_basic_config() +``` + +When running `accelerate config`, if we specify torch compile mode to True there can be dramatic speedups. +Note also that we use PEFT library as backend for LoRA training, make sure to have `peft>=0.14.0` installed in your environment. + +### Dog toy example + +Now let's get our dataset. For this example we will use some dog images: https://huggingface.co/datasets/diffusers/dog-example. + +Let's first download it locally: + +```python +from huggingface_hub import snapshot_download + +local_dir = "./dog" +snapshot_download( + "diffusers/dog-example", + local_dir=local_dir, repo_type="dataset", + ignore_patterns=".gitattributes", +) +``` + +This will also allow us to push the trained LoRA parameters to the Hugging Face Hub platform. + +Now, we can launch training using: + +```bash +export MODEL_NAME="Qwen/Qwen-Image-2.1" +export INSTANCE_DIR="dog" +export OUTPUT_DIR="trained-qwenimage21-lora" + +accelerate launch train_dreambooth_lora_qwenimage21.py \ + --pretrained_model_name_or_path=$MODEL_NAME \ + --instance_data_dir=$INSTANCE_DIR \ + --output_dir=$OUTPUT_DIR \ + --mixed_precision="bf16" \ + --instance_prompt="a photo of sks dog" \ + --resolution=1024 \ + --train_batch_size=1 \ + --gradient_accumulation_steps=4 \ + --use_8bit_adam \ + --rank=4 \ + --lora_alpha=4 \ + --learning_rate=2e-4 \ + --report_to="wandb" \ + --lr_scheduler="constant" \ + --lr_warmup_steps=0 \ + --max_train_steps=500 \ + --validation_prompt="A photo of sks dog in a bucket" \ + --validation_epochs=25 \ + --seed="0" \ + --push_to_hub +``` + +For using `push_to_hub`, make you're logged into your Hugging Face account: + +```bash +hf auth login +``` + +To better track our training experiments, we're using the following flags in the command above: + +* `report_to="wandb` will ensure the training runs are tracked on [Weights and Biases](https://wandb.ai/site). To use it, be sure to install `wandb` with `pip install wandb`. Don't forget to call `wandb login ` before training if you haven't done it before. +* `validation_prompt` and `validation_epochs` to allow the script to do a few validation inference runs. This allows us to qualitatively check if the training is progressing as expected. + +### LoRA rank and alpha + +`--rank` sets the dimension of the trainable LoRA matrices, and `--lora_alpha` scales what they contribute: +PEFT multiplies the LoRA update by `lora_alpha / rank`. Both default to 16 here, so the update is applied at +full strength out of the box. + +Change one and the ratio moves with it: + +* `lora_alpha == rank` - scale 1, the LoRA is applied at the strength it learned. +* `lora_alpha < rank` - scale below 1, a weaker LoRA. `--rank 16` on its own with `--lora_alpha 4` is scale + 0.25, which mostly shows up as a run that looks undertrained at a step count that should have been enough. +* `lora_alpha > rank` - scale above 1, a stronger effect without adding parameters. + +> [!TIP] +> Raise `--rank` for capacity, and raise `--lora_alpha` with it unless you mean to change the strength. +> If the style takes but subjects start losing their shape, the run is overcooked: cut the steps or the +> learning rate before reaching for a smaller alpha. + +## Model specifics + +A few things differ from the other DreamBooth LoRA trainers, all of them following the model rather than a choice made here: + +* **Resolutions are multiples of 32.** One latent token covers a 16x16 pixel tile and the transformer groups latents into 2x2 slots, so `--resolution` and every `--aspect_ratio_buckets` entry must divide by 32. The script raises on anything else rather than resizing silently. +* **Images are read as RGBA.** This VAE takes and returns four channels, so a three-channel tensor fails at its first convolution. Images without an alpha channel get an opaque one. +* **The prompt goes through Qwen3-VL**, as a processor rather than a tokenizer, and comes back as variable-length embeddings with a mask. There is no `--max_sequence_length`: the checkpoint's processor does not truncate. +* **The scheduler is used as shipped.** It sets `use_dynamic_shifting`, so the shift is derived from the sequence length at sampling time and the training sigmas stay unshifted. +* **`flex_attention` is worth having.** With it the block-causal mask runs as one compiled block-sparse pass; without it the model falls back to an exact multi-pass SDPA prefill, which gives the same results but costs more. + +Validation images are generated at `--resolution`, with `--validation_num_inference_steps` (default 40) and classifier-free guidance off, which is the recipe the model's own docs use. + +## Image-to-image (editing) + +`train_dreambooth_lora_qwenimage21_img2img.py` trains the same transformer on pairs: a condition image, the image it +should become, and the instruction that describes the change. The condition image enters twice, as vision tokens in the +prompt and as clean latents ahead of the noisy target in the sequence, and the loss is taken on the target alone. + +It needs a dataset that holds both images, so `--dataset_name` and `--cond_image_column` are required: + +```bash +accelerate launch train_dreambooth_lora_qwenimage21_img2img.py \ + --pretrained_model_name_or_path="Qwen/Qwen-Image-2.1" \ + --dataset_name="my-username/my-edit-pairs" \ + --cond_image_column="cond_image" \ + --image_column="image" \ + --caption_column="caption" \ + --instance_prompt="make it snow" \ + --output_dir="trained-qwenimage21-edit-lora" \ + --mixed_precision="bf16" \ + --resolution=1024 \ + --train_batch_size=1 \ + --learning_rate=1e-4 \ + --lr_scheduler="constant" \ + --lr_warmup_steps=0 \ + --max_train_steps=1000 \ + --validation_prompt="make it snow" \ + --validation_image="path/to/a/photo.png" \ + --validation_epochs=25 \ + --seed="0" +``` + +What differs from the text-to-image script: + +* **A batch shares one image-pad layout.** The transformer reads the layout from the first row of `img_mask`, so + every sample in a batch has to place the condition image's tokens identically. That holds when the samples share a + prompt and a bucket; otherwise train with `--train_batch_size 1`. The script raises rather than training on a + misaligned batch. +* **Condition images cannot be small.** The vision-language processor upsamples images below its minimum pixel count, + and then produces more vision tokens than the transformer has slots for. The script checks the two counts and says + so. 256px is the smallest size that lines up; train at 1024 in practice. +* **Prompt embeddings are per sample**, since each one is encoded together with its own condition image, and they are + cached that way. +* **A pair keeps its geometry.** The condition image is resized to the target's grid and takes the same crop and flip, + so pairs that were aligned stay aligned. +* `--with_prior_preservation` and `--caption_dropout` are rejected: both introduce a second prompt layout in a batch. + +## Notes + +Additionally, we welcome you to explore the following CLI arguments: + +* `--lora_layers`: The transformer modules to apply LoRA training on. Please specify the layers in a comma separated. E.g. - "to_k,to_q,to_v" will result in lora training of attention layers only. The default is `to_k,to_q,to_v,to_out.0`; the feed-forward layers are named `img_mlp.proj`, `img_mlp.gate_layer` and `img_mlp.out`. +* `--use_aspect_ratio_buckets` / `--aspect_ratio_buckets`: train on a set of aspect ratios instead of one square crop. Each batch is drawn from a single bucket. +* `--caption_dropout`: drop an instance caption in favour of the empty prompt with this probability. + +We provide several options for optimizing memory optimization: + +* `--offload`: When enabled, we will offload the text encoder and VAE to CPU, when they are not used. +* `cache_latents`: When enabled, we will pre-compute the latents from the input images with the VAE and remove the VAE from memory once done. +* `--use_8bit_adam`: When enabled, we will use the 8bit version of AdamW provided by the `bitsandbytes` library. + +Refer to the [official documentation](https://huggingface.co/docs/diffusers/main/en/api/pipelines/qwenimage21) of the `QwenImage21Pipeline` to know more about the model and its preferred dtypes during inference. + +## Using quantization + +You can quantize the base model with [`bitsandbytes`](https://huggingface.co/docs/bitsandbytes/index) to reduce memory usage. To do so, pass a JSON file path to `--bnb_quantization_config_path`. This file should hold the configuration to initialize `BitsAndBytesConfig`. Below is an example JSON file: + +```json +{ + "load_in_4bit": true, + "bnb_4bit_quant_type": "nf4" +} +``` diff --git a/examples/dreambooth/test_dreambooth_lora_qwenimage21.py b/examples/dreambooth/test_dreambooth_lora_qwenimage21.py new file mode 100644 index 000000000000..743bc8260496 --- /dev/null +++ b/examples/dreambooth/test_dreambooth_lora_qwenimage21.py @@ -0,0 +1,363 @@ +# coding=utf-8 +# Copyright 2026 HuggingFace Inc. +# +# Licensed under the Apache License, Version 2.0 (the "License"); +# you may not use this file except in compliance with the License. +# You may obtain a copy of the License at +# +# http://www.apache.org/licenses/LICENSE-2.0 +# +# Unless required by applicable law or agreed to in writing, software +# distributed under the License is distributed on an "AS IS" BASIS, +# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. +# See the License for the specific language governing permissions and +# limitations under the License. + +import json +import logging +import os +import sys +import tempfile + +import numpy as np +import pytest +import safetensors +from PIL import Image + +from diffusers.loaders.lora_base import LORA_ADAPTER_METADATA_KEY + + +sys.path.append("..") +from test_examples_utils import ExamplesTestsAccelerate, run_command # noqa: E402 + + +logging.basicConfig(level=logging.DEBUG) + +logger = logging.getLogger() +stream_handler = logging.StreamHandler(sys.stdout) +logger.addHandler(stream_handler) + + +class TestDreamBoothLoRAQwenImage21(ExamplesTestsAccelerate): + instance_data_dir = "docs/source/en/imgs" + instance_prompt = "photo" + pretrained_model_name_or_path = "hf-internal-testing/tiny-qwenimage21-pipe" + script_path = "examples/dreambooth/train_dreambooth_lora_qwenimage21.py" + transformer_layer_type = "transformer_blocks.0.attn.to_k" + + def test_dreambooth_lora_qwenimage21(self): + with tempfile.TemporaryDirectory() as tmpdir: + test_args = f""" + {self.script_path} + --pretrained_model_name_or_path {self.pretrained_model_name_or_path} + --instance_data_dir {self.instance_data_dir} + --instance_prompt {self.instance_prompt} + --resolution 64 + --train_batch_size 1 + --gradient_accumulation_steps 1 + --max_train_steps 2 + --learning_rate 5.0e-04 + --scale_lr + --lr_scheduler constant + --lr_warmup_steps 0 + --output_dir {tmpdir} + """.split() + + run_command(self._launch_args + test_args) + # save_pretrained smoke test + assert os.path.isfile(os.path.join(tmpdir, "pytorch_lora_weights.safetensors")) + + # make sure the state_dict has the correct naming in the parameters. + lora_state_dict = safetensors.torch.load_file(os.path.join(tmpdir, "pytorch_lora_weights.safetensors")) + is_lora = all("lora" in k for k in lora_state_dict.keys()) + assert is_lora + + # when not training the text encoder, all the parameters in the state dict should start + # with `"transformer"` in their names. + starts_with_transformer = all(key.startswith("transformer") for key in lora_state_dict.keys()) + assert starts_with_transformer + + def test_dreambooth_lora_offload_without_latent_caching(self): + # `--offload` without `--cache_latents` is the one path that encodes pixels inside the training + # loop while the VAE is being moved on and off the accelerator. It regressed once, by encoding + # after the offload context had already put the VAE back on the CPU. + with tempfile.TemporaryDirectory() as tmpdir: + test_args = f""" + {self.script_path} + --pretrained_model_name_or_path {self.pretrained_model_name_or_path} + --instance_data_dir {self.instance_data_dir} + --instance_prompt {self.instance_prompt} + --resolution 64 + --offload + --train_batch_size 1 + --gradient_accumulation_steps 1 + --max_train_steps 2 + --learning_rate 5.0e-04 + --lr_scheduler constant + --lr_warmup_steps 0 + --output_dir {tmpdir} + """.split() + + run_command(self._launch_args + test_args) + assert os.path.isfile(os.path.join(tmpdir, "pytorch_lora_weights.safetensors")) + + def test_dreambooth_lora_custom_captions(self): + # `--caption_column` caches one prompt embedding per sample. `encode_prompt` returns no mask when + # nothing in the batch is padded — the usual case, since captions often tokenize to equal length — + # so the cache has to store a dense mask rather than slice a `None`. + with tempfile.TemporaryDirectory() as tmpdir: + from datasets import Dataset, Features, Value + from datasets import Image as ImageFeature + + rng = np.random.default_rng(0) + rows = { + "image": [Image.fromarray(rng.integers(0, 255, (64, 64, 3), dtype=np.uint8)) for _ in range(2)], + "caption": ["a photo"] * 2, + } + dataset = Dataset.from_dict(rows, features=Features({"image": ImageFeature(), "caption": Value("string")})) + dataset_dir = os.path.join(tmpdir, "dataset") + os.makedirs(dataset_dir, exist_ok=True) + dataset.to_parquet(os.path.join(dataset_dir, "data.parquet")) + + test_args = f""" + {self.script_path} + --pretrained_model_name_or_path {self.pretrained_model_name_or_path} + --dataset_name {dataset_dir} + --caption_column caption + --instance_prompt {self.instance_prompt} + --resolution 64 + --train_batch_size 1 + --gradient_accumulation_steps 1 + --max_train_steps 2 + --learning_rate 5.0e-04 + --lr_scheduler constant + --lr_warmup_steps 0 + --output_dir {tmpdir} + """.split() + + run_command(self._launch_args + test_args) + assert os.path.isfile(os.path.join(tmpdir, "pytorch_lora_weights.safetensors")) + + def test_dreambooth_lora_latent_caching(self): + with tempfile.TemporaryDirectory() as tmpdir: + test_args = f""" + {self.script_path} + --pretrained_model_name_or_path {self.pretrained_model_name_or_path} + --instance_data_dir {self.instance_data_dir} + --instance_prompt {self.instance_prompt} + --resolution 64 + --train_batch_size 1 + --gradient_accumulation_steps 1 + --max_train_steps 2 + --cache_latents + --learning_rate 5.0e-04 + --scale_lr + --lr_scheduler constant + --lr_warmup_steps 0 + --output_dir {tmpdir} + """.split() + + run_command(self._launch_args + test_args) + # save_pretrained smoke test + assert os.path.isfile(os.path.join(tmpdir, "pytorch_lora_weights.safetensors")) + + # make sure the state_dict has the correct naming in the parameters. + lora_state_dict = safetensors.torch.load_file(os.path.join(tmpdir, "pytorch_lora_weights.safetensors")) + is_lora = all("lora" in k for k in lora_state_dict.keys()) + assert is_lora + + # when not training the text encoder, all the parameters in the state dict should start + # with `"transformer"` in their names. + starts_with_transformer = all(key.startswith("transformer") for key in lora_state_dict.keys()) + assert starts_with_transformer + + def test_dreambooth_lora_layers(self): + with tempfile.TemporaryDirectory() as tmpdir: + test_args = f""" + {self.script_path} + --pretrained_model_name_or_path {self.pretrained_model_name_or_path} + --instance_data_dir {self.instance_data_dir} + --instance_prompt {self.instance_prompt} + --resolution 64 + --train_batch_size 1 + --gradient_accumulation_steps 1 + --max_train_steps 2 + --cache_latents + --learning_rate 5.0e-04 + --scale_lr + --lora_layers {self.transformer_layer_type} + --lr_scheduler constant + --lr_warmup_steps 0 + --output_dir {tmpdir} + """.split() + + run_command(self._launch_args + test_args) + # save_pretrained smoke test + assert os.path.isfile(os.path.join(tmpdir, "pytorch_lora_weights.safetensors")) + + # make sure the state_dict has the correct naming in the parameters. + lora_state_dict = safetensors.torch.load_file(os.path.join(tmpdir, "pytorch_lora_weights.safetensors")) + is_lora = all("lora" in k for k in lora_state_dict.keys()) + assert is_lora + + # when not training the text encoder, all the parameters in the state dict should start + # with `"transformer"` in their names. In this test, we only params of + # transformer.transformer_blocks.0.attn.to_k should be in the state dict + starts_with_transformer = all( + key.startswith(f"transformer.{self.transformer_layer_type}") for key in lora_state_dict.keys() + ) + assert starts_with_transformer + + def test_dreambooth_lora_qwenimage21_checkpointing_checkpoints_total_limit(self): + with tempfile.TemporaryDirectory() as tmpdir: + test_args = f""" + {self.script_path} + --pretrained_model_name_or_path={self.pretrained_model_name_or_path} + --instance_data_dir={self.instance_data_dir} + --output_dir={tmpdir} + --instance_prompt={self.instance_prompt} + --resolution=64 + --train_batch_size=1 + --gradient_accumulation_steps=1 + --max_train_steps=6 + --checkpoints_total_limit=2 + --checkpointing_steps=2 + """.split() + + run_command(self._launch_args + test_args) + + assert {x for x in os.listdir(tmpdir) if "checkpoint" in x} == {"checkpoint-4", "checkpoint-6"} + + def test_dreambooth_lora_qwenimage21_checkpointing_checkpoints_total_limit_removes_multiple_checkpoints(self): + with tempfile.TemporaryDirectory() as tmpdir: + test_args = f""" + {self.script_path} + --pretrained_model_name_or_path={self.pretrained_model_name_or_path} + --instance_data_dir={self.instance_data_dir} + --output_dir={tmpdir} + --instance_prompt={self.instance_prompt} + --resolution=64 + --train_batch_size=1 + --gradient_accumulation_steps=1 + --max_train_steps=4 + --checkpointing_steps=2 + """.split() + + run_command(self._launch_args + test_args) + + assert {x for x in os.listdir(tmpdir) if "checkpoint" in x} == {"checkpoint-2", "checkpoint-4"} + + resume_run_args = f""" + {self.script_path} + --pretrained_model_name_or_path={self.pretrained_model_name_or_path} + --instance_data_dir={self.instance_data_dir} + --output_dir={tmpdir} + --instance_prompt={self.instance_prompt} + --resolution=64 + --train_batch_size=1 + --gradient_accumulation_steps=1 + --max_train_steps=8 + --checkpointing_steps=2 + --resume_from_checkpoint=checkpoint-4 + --checkpoints_total_limit=2 + """.split() + + run_command(self._launch_args + resume_run_args) + + assert {x for x in os.listdir(tmpdir) if "checkpoint" in x} == {"checkpoint-6", "checkpoint-8"} + + def test_dreambooth_lora_with_metadata(self): + # Use a `lora_alpha` that is different from `rank`. + lora_alpha = 8 + rank = 4 + with tempfile.TemporaryDirectory() as tmpdir: + test_args = f""" + {self.script_path} + --pretrained_model_name_or_path {self.pretrained_model_name_or_path} + --instance_data_dir {self.instance_data_dir} + --instance_prompt {self.instance_prompt} + --resolution 64 + --train_batch_size 1 + --gradient_accumulation_steps 1 + --max_train_steps 2 + --lora_alpha={lora_alpha} + --rank={rank} + --learning_rate 5.0e-04 + --scale_lr + --lr_scheduler constant + --lr_warmup_steps 0 + --output_dir {tmpdir} + """.split() + + run_command(self._launch_args + test_args) + # save_pretrained smoke test + state_dict_file = os.path.join(tmpdir, "pytorch_lora_weights.safetensors") + assert os.path.isfile(state_dict_file) + + # Check if the metadata was properly serialized. + with safetensors.torch.safe_open(state_dict_file, framework="pt", device="cpu") as f: + metadata = f.metadata() or {} + + metadata.pop("format", None) + raw = metadata.get(LORA_ADAPTER_METADATA_KEY) + if raw: + raw = json.loads(raw) + + loaded_lora_alpha = raw["transformer.lora_alpha"] + assert loaded_lora_alpha == lora_alpha + loaded_lora_rank = raw["transformer.r"] + assert loaded_lora_rank == rank + + def test_dreambooth_lora_qwenimage21_aspect_ratio_buckets(self): + # Both latent dimensions have to be even for this model, so bucket sizes must divide by 32. That is a + # model constraint rather than a preference, which is why this runs rather than being skipped. + with tempfile.TemporaryDirectory() as tmpdir: + test_args = f""" + {self.script_path} + --pretrained_model_name_or_path {self.pretrained_model_name_or_path} + --instance_data_dir {self.instance_data_dir} + --instance_prompt {self.instance_prompt} + --use_aspect_ratio_buckets + --aspect_ratio_buckets 64,64;64,128 + --cache_latents + --train_batch_size 1 + --gradient_accumulation_steps 1 + --max_train_steps 2 + --learning_rate 5.0e-04 + --lr_scheduler constant + --lr_warmup_steps 0 + --output_dir {tmpdir} + """.split() + + run_command(self._launch_args + test_args) + assert os.path.isfile(os.path.join(tmpdir, "pytorch_lora_weights.safetensors")) + lora_state_dict = safetensors.torch.load_file(os.path.join(tmpdir, "pytorch_lora_weights.safetensors")) + is_lora = all("lora" in k for k in lora_state_dict.keys()) + assert is_lora + starts_with_transformer = all(key.startswith("transformer") for key in lora_state_dict.keys()) + assert starts_with_transformer + + @pytest.mark.skip(reason="Caption dropout is opt-in and not widely used yet; re-enable when it is.") + def test_dreambooth_lora_qwenimage21_caption_dropout(self): + with tempfile.TemporaryDirectory() as tmpdir: + test_args = f""" + {self.script_path} + --pretrained_model_name_or_path {self.pretrained_model_name_or_path} + --instance_data_dir {self.instance_data_dir} + --instance_prompt {self.instance_prompt} + --resolution 64 + --caption_dropout 1.0 + --train_batch_size 1 + --gradient_accumulation_steps 1 + --max_train_steps 2 + --learning_rate 5.0e-04 + --lr_scheduler constant + --lr_warmup_steps 0 + --output_dir {tmpdir} + """.split() + + run_command(self._launch_args + test_args) + assert os.path.isfile(os.path.join(tmpdir, "pytorch_lora_weights.safetensors")) + lora_state_dict = safetensors.torch.load_file(os.path.join(tmpdir, "pytorch_lora_weights.safetensors")) + is_lora = all("lora" in k for k in lora_state_dict.keys()) + assert is_lora diff --git a/examples/dreambooth/test_dreambooth_lora_qwenimage21_img2img.py b/examples/dreambooth/test_dreambooth_lora_qwenimage21_img2img.py new file mode 100644 index 000000000000..998192d647cb --- /dev/null +++ b/examples/dreambooth/test_dreambooth_lora_qwenimage21_img2img.py @@ -0,0 +1,150 @@ +# coding=utf-8 +# Copyright 2026 HuggingFace Inc. +# +# Licensed under the Apache License, Version 2.0 (the "License"); +# you may not use this file except in compliance with the License. +# You may obtain a copy of the License at +# +# http://www.apache.org/licenses/LICENSE-2.0 +# +# Unless required by applicable law or agreed to in writing, software +# distributed under the License is distributed on an "AS IS" BASIS, +# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. +# See the License for the specific language governing permissions and +# limitations under the License. + +import json +import logging +import os +import sys +import tempfile + +import numpy as np +import safetensors +from PIL import Image + +from diffusers.loaders.lora_base import LORA_ADAPTER_METADATA_KEY + + +sys.path.append("..") +from test_examples_utils import ExamplesTestsAccelerate, run_command # noqa: E402 + + +logging.basicConfig(level=logging.DEBUG) + +logger = logging.getLogger() +stream_handler = logging.StreamHandler(sys.stdout) +logger.addHandler(stream_handler) + + +class TestDreamBoothLoRAQwenImage21Img2Img(ExamplesTestsAccelerate): + instance_prompt = "photo" + pretrained_model_name_or_path = "hf-internal-testing/tiny-qwenimage21-pipe" + script_path = "examples/dreambooth/train_dreambooth_lora_qwenimage21_img2img.py" + transformer_layer_type = "transformer_blocks.0.attn.to_k" + # 256 rather than the 64 the text-to-image tests use: the vision-language processor upsamples images below its + # minimum pixel count, and a condition image it resizes produces more vision tokens than the transformer has + # slots for. 256 is the smallest size where the two line up. + resolution = 256 + + def _paired_dataset(self, directory): + """Write a two-row dataset with a target image, a condition image and a caption.""" + from datasets import Dataset, Features, Value + from datasets import Image as ImageFeature + + rng = np.random.default_rng(0) + + def image(): + return Image.fromarray(rng.integers(0, 255, (self.resolution, self.resolution, 3), dtype=np.uint8)) + + rows = { + "image": [image() for _ in range(2)], + "cond_image": [image() for _ in range(2)], + "caption": [self.instance_prompt] * 2, + } + dataset = Dataset.from_dict( + rows, + features=Features({"image": ImageFeature(), "cond_image": ImageFeature(), "caption": Value("string")}), + ) + path = os.path.join(directory, "dataset") + os.makedirs(path, exist_ok=True) + dataset.to_parquet(os.path.join(path, "data.parquet")) + return path + + def _base_args(self, dataset_dir, tmpdir): + return f""" + {self.script_path} + --pretrained_model_name_or_path {self.pretrained_model_name_or_path} + --dataset_name {dataset_dir} + --cond_image_column cond_image + --caption_column caption + --instance_prompt {self.instance_prompt} + --resolution {self.resolution} + --train_batch_size 1 + --gradient_accumulation_steps 1 + --max_train_steps 2 + --learning_rate 5.0e-04 + --lr_scheduler constant + --lr_warmup_steps 0 + --output_dir {tmpdir} + """ + + def test_dreambooth_lora_qwenimage21_img2img(self): + with tempfile.TemporaryDirectory() as tmpdir: + dataset_dir = self._paired_dataset(tmpdir) + run_command(self._launch_args + self._base_args(dataset_dir, tmpdir).split()) + + # save_pretrained smoke test + assert os.path.isfile(os.path.join(tmpdir, "pytorch_lora_weights.safetensors")) + + # make sure the state_dict has the correct naming in the parameters. + lora_state_dict = safetensors.torch.load_file(os.path.join(tmpdir, "pytorch_lora_weights.safetensors")) + is_lora = all("lora" in k for k in lora_state_dict.keys()) + assert is_lora + + # when not training the text encoder, all the parameters in the state dict should start + # with `"transformer"` in their names. + starts_with_transformer = all(key.startswith("transformer") for key in lora_state_dict.keys()) + assert starts_with_transformer + + def test_dreambooth_lora_qwenimage21_img2img_latent_caching(self): + with tempfile.TemporaryDirectory() as tmpdir: + dataset_dir = self._paired_dataset(tmpdir) + test_args = self._base_args(dataset_dir, tmpdir).split() + ["--cache_latents"] + run_command(self._launch_args + test_args) + + assert os.path.isfile(os.path.join(tmpdir, "pytorch_lora_weights.safetensors")) + lora_state_dict = safetensors.torch.load_file(os.path.join(tmpdir, "pytorch_lora_weights.safetensors")) + is_lora = all("lora" in k for k in lora_state_dict.keys()) + assert is_lora + starts_with_transformer = all(key.startswith("transformer") for key in lora_state_dict.keys()) + assert starts_with_transformer + + def test_dreambooth_lora_qwenimage21_img2img_with_metadata(self): + # Use a `lora_alpha` that is different from `rank`. + lora_alpha = 8 + rank = 4 + with tempfile.TemporaryDirectory() as tmpdir: + dataset_dir = self._paired_dataset(tmpdir) + test_args = self._base_args(dataset_dir, tmpdir).split() + [ + f"--lora_alpha={lora_alpha}", + f"--rank={rank}", + ] + run_command(self._launch_args + test_args) + + state_dict_file = os.path.join(tmpdir, "pytorch_lora_weights.safetensors") + assert os.path.isfile(state_dict_file) + + # Check if the metadata was properly serialized. + with safetensors.torch.safe_open(state_dict_file, framework="pt", device="cpu") as f: + metadata = f.metadata() or {} + + metadata.pop("format", None) + raw = metadata.get(LORA_ADAPTER_METADATA_KEY) + if raw: + raw = json.loads(raw) + + loaded_lora_alpha = raw["transformer.lora_alpha"] + assert loaded_lora_alpha == lora_alpha + loaded_lora_rank = raw["transformer.r"] + assert loaded_lora_rank == rank diff --git a/examples/dreambooth/train_dreambooth_lora_qwenimage21.py b/examples/dreambooth/train_dreambooth_lora_qwenimage21.py new file mode 100644 index 000000000000..502f288f4308 --- /dev/null +++ b/examples/dreambooth/train_dreambooth_lora_qwenimage21.py @@ -0,0 +1,2033 @@ +#!/usr/bin/env python +# coding=utf-8 +# Copyright 2025 The HuggingFace Inc. team. All rights reserved. +# +# Licensed under the Apache License, Version 2.0 (the "License"); +# you may not use this file except in compliance with the License. +# You may obtain a copy of the License at +# +# http://www.apache.org/licenses/LICENSE-2.0 +# +# Unless required by applicable law or agreed to in writing, software +# distributed under the License is distributed on an "AS IS" BASIS, +# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. +# See the License for the specific language governing permissions and + +# /// script +# dependencies = [ +# "diffusers @ git+https://github.com/huggingface/diffusers.git", +# "torch>=2.0.0", +# "accelerate>=0.31.0", +# "transformers>=4.41.2", +# "ftfy", +# "tensorboard", +# "Jinja2", +# "peft>=0.11.1", +# "sentencepiece", +# "torchvision", +# "datasets", +# "bitsandbytes", +# "prodigyopt", +# ] +# /// + +import argparse +import copy +import itertools +import json +import logging +import math +import os +import random +import shutil +import warnings +from contextlib import nullcontext +from pathlib import Path + +import numpy as np +import torch +import transformers +from accelerate import Accelerator, DistributedType +from accelerate.logging import get_logger +from accelerate.utils import DistributedDataParallelKwargs, ProjectConfiguration, set_seed +from huggingface_hub import create_repo, upload_folder +from huggingface_hub.utils import insecure_hashlib +from peft import LoraConfig, prepare_model_for_kbit_training, set_peft_model_state_dict +from peft.utils import get_peft_model_state_dict +from PIL import Image +from PIL.ImageOps import exif_transpose +from torch.utils.data import BatchSampler, Dataset +from torchvision import transforms +from torchvision.transforms import functional as TF +from tqdm.auto import tqdm +from transformers import Qwen3VLForConditionalGeneration, Qwen3VLProcessor + +import diffusers +from diffusers import ( + AutoencoderKLQwenImage21, + BitsAndBytesConfig, + FlowMatchEulerDiscreteScheduler, + QwenImage21Pipeline, + QwenImage21Transformer2DModel, +) +from diffusers.optimization import get_scheduler +from diffusers.training_utils import ( + _collate_lora_metadata, + cast_training_params, + compute_density_for_timestep_sampling, + compute_loss_weighting_for_sd3, + find_nearest_bucket, + free_memory, + generate_aspect_ratio_buckets, + offload_models, + parse_buckets_string, +) +from diffusers.utils import ( + check_min_version, + convert_unet_state_dict_to_peft, + is_wandb_available, +) +from diffusers.utils.hub_utils import load_or_create_model_card, populate_model_card +from diffusers.utils.import_utils import is_torch_npu_available +from diffusers.utils.torch_utils import is_compiled_module + + +if is_wandb_available(): + import wandb + +# Will error if the minimal version of diffusers is not installed. Remove at your own risks. +check_min_version("0.41.0.dev0") + +logger = get_logger(__name__) + +# `vae_scale_factor * 2` for this VAE, checked against its config in `main`. +SIZE_MULTIPLE_OF = 32 + +if is_torch_npu_available(): + torch.npu.config.allow_internal_format = False + + +class QwenImage21ValidationPipeline(QwenImage21Pipeline): + """`QwenImage21Pipeline` that takes the image-pad mask alongside precomputed prompt embeddings. + + `__call__` gets that mask from `encode_prompt`, which it skips when handed `prompt_embeds`, so validation hands + back the one it cached before the text encoder was freed. + """ + + cached_image_pad_mask = None + + def encode_prompt(self, *args, **kwargs): + # The mask goes in before the call, not after it. Handed `prompt_embeds` together with a condition image, + # the base `encode_prompt` raises for a missing mask rather than returning `None` in its place, so a + # substitution made on the way out never runs. + if kwargs.get("image_pad_mask") is None and self.cached_image_pad_mask is not None: + kwargs["image_pad_mask"] = self.cached_image_pad_mask + prompt_embeds, prompt_embeds_mask, image_pad_mask = super().encode_prompt(*args, **kwargs) + if image_pad_mask is None: + image_pad_mask = self.cached_image_pad_mask + return prompt_embeds, prompt_embeds_mask, image_pad_mask + + +def save_model_card( + repo_id: str, + images=None, + base_model: str = None, + instance_prompt=None, + validation_prompt=None, + repo_folder=None, +): + widget_dict = [] + if images is not None: + for i, image in enumerate(images): + image.save(os.path.join(repo_folder, f"image_{i}.png")) + widget_dict.append( + {"text": validation_prompt if validation_prompt else " ", "output": {"url": f"image_{i}.png"}} + ) + + model_description = f""" +# Qwen-Image 2.1 DreamBooth LoRA - {repo_id} + + + +## Model description + +These are {repo_id} DreamBooth LoRA weights for {base_model}. + +The weights were trained using [DreamBooth](https://dreambooth.github.io/) with the [Qwen-Image 2.1 diffusers trainer](https://github.com/huggingface/diffusers/blob/main/examples/dreambooth/README_qwenimage21.md). + +## Trigger words + +You should use `{instance_prompt}` to trigger the image generation. + +## Download model + +[Download the *.safetensors LoRA]({repo_id}/tree/main) in the Files & versions tab. + +## Use it with the [🧨 diffusers library](https://github.com/huggingface/diffusers) + +```py + >>> import torch + >>> from diffusers import QwenImage21Pipeline + + >>> pipe = QwenImage21Pipeline.from_pretrained( + ... "Qwen/Qwen-Image-2.1", + ... torch_dtype=torch.bfloat16, + ... ) + >>> pipe.enable_model_cpu_offload() + >>> pipe.load_lora_weights(f"{repo_id}") + >>> image = pipe(f"{instance_prompt}").images[0] + + +``` + +For more details, including weighting, merging and fusing LoRAs, check the [documentation on loading LoRAs in diffusers](https://huggingface.co/docs/diffusers/main/en/using-diffusers/loading_adapters) +""" + model_card = load_or_create_model_card( + repo_id_or_path=repo_id, + from_training=True, + license="apache-2.0", + base_model=base_model, + prompt=instance_prompt, + model_description=model_description, + widget=widget_dict, + ) + tags = [ + "text-to-image", + "diffusers-training", + "diffusers", + "lora", + "qwen-image", + "qwen-image-2.1", + "qwen-image-diffusers", + "template:sd-lora", + ] + + model_card = populate_model_card(model_card, tags=tags) + model_card.save(os.path.join(repo_folder, "README.md")) + + +def log_validation( + pipeline, + args, + accelerator, + pipeline_args, + epoch, + torch_dtype, + is_final_validation=False, +): + args.num_validation_images = args.num_validation_images if args.num_validation_images else 1 + # `--final_validation_prompt` stands in when only the final pass is asked for. + validation_prompt = args.validation_prompt or args.final_validation_prompt + logger.info( + f"Running validation... \n Generating {args.num_validation_images} images with prompt: {validation_prompt}." + ) + pipeline = pipeline.to(accelerator.device, dtype=torch_dtype) + pipeline.set_progress_bar_config(disable=True) + + # run inference + generator = torch.Generator(device=accelerator.device).manual_seed(args.seed) if args.seed is not None else None + autocast_ctx = torch.autocast(accelerator.device.type) if not is_final_validation else nullcontext() + + images = [] + for _ in range(args.num_validation_images): + with autocast_ctx: + image = pipeline( + **pipeline_args, + num_inference_steps=args.validation_num_inference_steps, + # Classifier-free guidance off, as in the model's own sample script: with `true_cfg_scale <= 1` + # a step is a single forward pass through the transformer. + true_cfg_scale=1.0, + output_resolution=args.resolution, + generator=generator, + ).images[0] + images.append(image) + + for tracker in accelerator.trackers: + phase_name = "test" if is_final_validation else "validation" + if tracker.name == "tensorboard": + # The pipeline returns RGBA, this VAE having four channels, and `add_images` asserts three. + np_images = np.stack([np.asarray(img.convert("RGB")) for img in images]) + tracker.writer.add_images(phase_name, np_images, epoch, dataformats="NHWC") + if tracker.name == "wandb": + tracker.log( + { + phase_name: [ + wandb.Image(image, caption=f"{i}: {validation_prompt}") for i, image in enumerate(images) + ] + } + ) + + del pipeline + free_memory() + + return images + + +def parse_args(input_args=None): + parser = argparse.ArgumentParser(description="Simple example of a training script.") + parser.add_argument( + "--pretrained_model_name_or_path", + type=str, + default=None, + required=True, + help="Path to pretrained model or model identifier from huggingface.co/models.", + ) + parser.add_argument( + "--bnb_quantization_config_path", + type=str, + default=None, + help="Quantization config in a JSON file that will be used to define the bitsandbytes quant config of the DiT.", + ) + parser.add_argument( + "--revision", + type=str, + default=None, + required=False, + help="Revision of pretrained model identifier from huggingface.co/models.", + ) + parser.add_argument( + "--variant", + type=str, + default=None, + help="Variant of the model files of the pretrained model identifier from huggingface.co/models, 'e.g.' fp16", + ) + parser.add_argument( + "--dataset_name", + type=str, + default=None, + help=( + "The name of the Dataset (from the HuggingFace hub) containing the training data of instance images (could be your own, possibly private," + " dataset). It can also be a path pointing to a local copy of a dataset in your filesystem," + " or to a folder containing files that 🤗 Datasets can understand." + ), + ) + parser.add_argument( + "--dataset_config_name", + type=str, + default=None, + help="The config of the Dataset, leave as None if there's only one config.", + ) + parser.add_argument( + "--instance_data_dir", + type=str, + default=None, + help=("A folder containing the training data. "), + ) + + parser.add_argument( + "--cache_dir", + type=str, + default=None, + help="The directory where the downloaded models and datasets will be stored.", + ) + + parser.add_argument( + "--image_column", + type=str, + default="image", + help="The column of the dataset containing the target image. By " + "default, the standard Image Dataset maps out 'file_name' " + "to 'image'.", + ) + parser.add_argument( + "--caption_column", + type=str, + default=None, + help="The column of the dataset containing the instance prompt for each image", + ) + + parser.add_argument("--repeats", type=int, default=1, help="How many times to repeat the training data.") + + parser.add_argument( + "--class_data_dir", + type=str, + default=None, + required=False, + help="A folder containing the training data of class images.", + ) + parser.add_argument( + "--instance_prompt", + type=str, + default=None, + required=True, + help="The prompt with identifier specifying the instance, e.g. 'photo of a TOK dog', 'in the style of TOK'", + ) + parser.add_argument( + "--class_prompt", + type=str, + default=None, + help="The prompt to specify images in the same class as provided instance images.", + ) + parser.add_argument( + "--validation_num_inference_steps", + type=int, + default=40, + help="Denoising steps for validation images. 40 is what the model's own sample script uses.", + ) + + parser.add_argument( + "--validation_prompt", + type=str, + default=None, + help="A prompt that is used during validation to verify that the model is learning.", + ) + + parser.add_argument( + "--skip_final_inference", + default=False, + action="store_true", + help="Whether to skip the final inference step with loaded lora weights upon training completion. This will run intermediate validation inference if `validation_prompt` is provided. Specify to reduce memory.", + ) + + parser.add_argument( + "--final_validation_prompt", + type=str, + default=None, + help="A prompt that is used during a final validation to verify that the model is learning. Ignored if `--validation_prompt` is provided.", + ) + parser.add_argument( + "--num_validation_images", + type=int, + default=4, + help="Number of images that should be generated during validation with `validation_prompt`.", + ) + parser.add_argument( + "--validation_epochs", + type=int, + default=50, + help=( + "Run dreambooth validation every X epochs. Dreambooth validation consists of running the prompt" + " `args.validation_prompt` multiple times: `args.num_validation_images`." + ), + ) + parser.add_argument( + "--rank", + type=int, + default=16, + help=("The dimension of the LoRA update matrices."), + ) + parser.add_argument( + "--lora_alpha", + type=int, + default=16, + help="LoRA alpha. The update is scaled by `lora_alpha / rank`, so keep the two in step.", + ) + parser.add_argument("--lora_dropout", type=float, default=0.0, help="Dropout probability for LoRA layers") + + parser.add_argument( + "--with_prior_preservation", + default=False, + action="store_true", + help="Flag to add prior preservation loss.", + ) + parser.add_argument("--prior_loss_weight", type=float, default=1.0, help="The weight of prior preservation loss.") + parser.add_argument( + "--num_class_images", + type=int, + default=100, + help=( + "Minimal class images for prior preservation loss. If there are not enough images already present in" + " class_data_dir, additional images will be sampled with class_prompt." + ), + ) + parser.add_argument( + "--output_dir", + type=str, + default="qwenimage21-dreambooth-lora", + help="The output directory where the model predictions and checkpoints will be written.", + ) + parser.add_argument("--seed", type=int, default=None, help="A seed for reproducible training.") + parser.add_argument( + "--resolution", + type=int, + default=512, + help=( + "The resolution for input images, all the images in the train/validation dataset will be resized to this" + " resolution" + ), + ) + parser.add_argument( + "--aspect_ratio_buckets", + type=str, + default=None, + help=( + "Aspect ratio buckets to use for training. Define as a string of 'h1,w1;h2,w2;...'. " + "e.g. '1024,1024;768,1360;1360,768;880,1168;1168,880;1248,832;832,1248'. " + "Requires --use_aspect_ratio_buckets. Images are resized to cover and cropped to the nearest " + "listed bucket (smaller images are upscaled). When set, --resolution is ignored." + ), + ) + parser.add_argument( + "--use_aspect_ratio_buckets", + action="store_true", + help=( + "Enable aspect-ratio bucketing. Without --aspect_ratio_buckets, the buckets are computed on the " + "fly from --resolution and capped to each image's own resolution, so smaller images are assigned " + "to a smaller bucket instead of being upscaled. Provide --aspect_ratio_buckets to use an explicit list." + ), + ) + parser.add_argument( + "--center_crop", + default=False, + action="store_true", + help=( + "Whether to center crop the input images to the resolution. If not set, the images will be randomly" + " cropped. The images will be resized to the resolution first before cropping." + ), + ) + parser.add_argument( + "--random_flip", + action="store_true", + help="whether to randomly flip images horizontally", + ) + parser.add_argument( + "--caption_dropout", + type=float, + default=0.0, + help=( + "Probability of replacing an instance image's caption with an empty string during training, so that" + " fraction of samples is trained unconditionally. Improves classifier-free guidance. A common value is" + " 0.1. Class/prior-preservation captions are never dropped." + ), + ) + parser.add_argument( + "--train_batch_size", type=int, default=4, help="Batch size (per device) for the training dataloader." + ) + parser.add_argument( + "--sample_batch_size", type=int, default=4, help="Batch size (per device) for sampling images." + ) + parser.add_argument("--num_train_epochs", type=int, default=1) + parser.add_argument( + "--max_train_steps", + type=int, + default=None, + help="Total number of training steps to perform. If provided, overrides num_train_epochs.", + ) + parser.add_argument( + "--checkpointing_steps", + type=int, + default=500, + help=( + "Save a checkpoint of the training state every X updates. These checkpoints can be used both as final" + " checkpoints in case they are better than the last checkpoint, and are also suitable for resuming" + " training using `--resume_from_checkpoint`." + ), + ) + parser.add_argument( + "--checkpoints_total_limit", + type=int, + default=None, + help=("Max number of checkpoints to store."), + ) + parser.add_argument( + "--resume_from_checkpoint", + type=str, + default=None, + help=( + "Whether training should be resumed from a previous checkpoint. Use a path saved by" + ' `--checkpointing_steps`, or `"latest"` to automatically select the last available checkpoint.' + ), + ) + parser.add_argument( + "--gradient_accumulation_steps", + type=int, + default=1, + help="Number of updates steps to accumulate before performing a backward/update pass.", + ) + parser.add_argument( + "--gradient_checkpointing", + action="store_true", + help="Whether or not to use gradient checkpointing to save memory at the expense of slower backward pass.", + ) + parser.add_argument( + "--learning_rate", + type=float, + default=1e-4, + help="Initial learning rate (after the potential warmup period) to use.", + ) + parser.add_argument( + "--scale_lr", + action="store_true", + default=False, + help="Scale the learning rate by the number of GPUs, gradient accumulation steps, and batch size.", + ) + parser.add_argument( + "--lr_scheduler", + type=str, + default="constant", + help=( + 'The scheduler type to use. Choose between ["linear", "cosine", "cosine_with_restarts", "polynomial",' + ' "constant", "constant_with_warmup"]' + ), + ) + parser.add_argument( + "--lr_warmup_steps", type=int, default=500, help="Number of steps for the warmup in the lr scheduler." + ) + parser.add_argument( + "--lr_num_cycles", + type=int, + default=1, + help="Number of hard resets of the lr in cosine_with_restarts scheduler.", + ) + parser.add_argument("--lr_power", type=float, default=1.0, help="Power factor of the polynomial scheduler.") + parser.add_argument( + "--dataloader_num_workers", + type=int, + default=0, + help=( + "Number of subprocesses to use for data loading. 0 means that the data will be loaded in the main process." + ), + ) + parser.add_argument( + "--weighting_scheme", + type=str, + default="none", + choices=["sigma_sqrt", "logit_normal", "mode", "cosmap", "none"], + help=('We default to the "none" weighting scheme for uniform sampling and uniform loss'), + ) + parser.add_argument( + "--logit_mean", type=float, default=0.0, help="mean to use when using the `'logit_normal'` weighting scheme." + ) + parser.add_argument( + "--logit_std", type=float, default=1.0, help="std to use when using the `'logit_normal'` weighting scheme." + ) + parser.add_argument( + "--mode_scale", + type=float, + default=1.29, + help="Scale of mode weighting scheme. Only effective when using the `'mode'` as the `weighting_scheme`.", + ) + parser.add_argument( + "--optimizer", + type=str, + default="AdamW", + help=('The optimizer type to use. Choose between ["AdamW", "prodigy"]'), + ) + + parser.add_argument( + "--use_8bit_adam", + action="store_true", + help="Whether or not to use 8-bit Adam from bitsandbytes. Ignored if optimizer is not set to AdamW", + ) + + parser.add_argument( + "--adam_beta1", type=float, default=0.9, help="The beta1 parameter for the Adam and Prodigy optimizers." + ) + parser.add_argument( + "--adam_beta2", type=float, default=0.999, help="The beta2 parameter for the Adam and Prodigy optimizers." + ) + parser.add_argument( + "--prodigy_beta3", + type=float, + default=None, + help="coefficients for computing the Prodigy stepsize using running averages. If set to None, " + "uses the value of square root of beta2. Ignored if optimizer is adamW", + ) + parser.add_argument("--prodigy_decouple", type=bool, default=True, help="Use AdamW style decoupled weight decay") + parser.add_argument( + "--adam_weight_decay", type=float, default=1e-04, help="Weight decay to use for the LoRA parameters" + ) + parser.add_argument( + "--lora_layers", + type=str, + default=None, + help=( + 'The transformer modules to apply LoRA training on. Please specify the layers in a comma separated. E.g. - "to_k,to_q,to_v" will result in lora training of attention layers only' + ), + ) + + parser.add_argument( + "--adam_epsilon", + type=float, + default=1e-08, + help="Epsilon value for the Adam optimizer and Prodigy optimizers.", + ) + + parser.add_argument( + "--prodigy_use_bias_correction", + type=bool, + default=True, + help="Turn on Adam's bias correction. True by default. Ignored if optimizer is adamW", + ) + parser.add_argument( + "--prodigy_safeguard_warmup", + type=bool, + default=True, + help="Remove lr from the denominator of D estimate to avoid issues during warm-up stage. True by default. " + "Ignored if optimizer is adamW", + ) + parser.add_argument("--max_grad_norm", default=1.0, type=float, help="Max gradient norm.") + parser.add_argument("--push_to_hub", action="store_true", help="Whether or not to push the model to the Hub.") + parser.add_argument("--hub_token", type=str, default=None, help="The token to use to push to the Model Hub.") + parser.add_argument( + "--hub_model_id", + type=str, + default=None, + help="The name of the repository to keep in sync with the local `output_dir`.", + ) + parser.add_argument( + "--logging_dir", + type=str, + default="logs", + help=( + "[TensorBoard](https://www.tensorflow.org/tensorboard) log directory. Will default to" + " *output_dir/runs/**CURRENT_DATETIME_HOSTNAME***." + ), + ) + parser.add_argument( + "--allow_tf32", + action="store_true", + help=( + "Whether or not to allow TF32 on Ampere GPUs. Can be used to speed up training. For more information, see" + " https://pytorch.org/docs/stable/notes/cuda.html#tensorfloat-32-tf32-on-ampere-devices" + ), + ) + parser.add_argument( + "--cache_latents", + action="store_true", + default=False, + help="Cache the VAE latents", + ) + parser.add_argument( + "--report_to", + type=str, + default="tensorboard", + help=( + 'The integration to report the results and logs to. Supported platforms are `"tensorboard"`' + ' (default), `"wandb"` and `"comet_ml"`. Use `"all"` to report to all integrations.' + ), + ) + parser.add_argument( + "--mixed_precision", + type=str, + default=None, + choices=["no", "fp16", "bf16"], + help=( + "Whether to use mixed precision. Choose between fp16 and bf16 (bfloat16). Bf16 requires PyTorch >=" + " 1.10.and an Nvidia Ampere GPU. Default to the value of accelerate config of the current system or the" + " flag passed with the `accelerate.launch` command. Use this argument to override the accelerate config." + ), + ) + parser.add_argument( + "--upcast_before_saving", + action="store_true", + default=False, + help=( + "Whether to upcast the trained transformer layers to float32 before saving (at the end of training). " + "Defaults to precision dtype used for training to save memory" + ), + ) + parser.add_argument( + "--offload", + action="store_true", + help="Whether to offload the VAE and the text encoder to CPU when they are not used.", + ) + parser.add_argument("--local_rank", type=int, default=-1, help="For distributed training: local_rank") + + if input_args is not None: + args = parser.parse_args(input_args) + else: + args = parser.parse_args() + + if args.dataset_name is None and args.instance_data_dir is None: + raise ValueError("Specify either `--dataset_name` or `--instance_data_dir`") + + if args.dataset_name is not None and args.instance_data_dir is not None: + raise ValueError("Specify only one of `--dataset_name` or `--instance_data_dir`") + + env_local_rank = int(os.environ.get("LOCAL_RANK", -1)) + if env_local_rank != -1 and env_local_rank != args.local_rank: + args.local_rank = env_local_rank + + # An error rather than the pipeline's silent resize: a mismatch only shows up later, as a packing error. + if args.resolution % SIZE_MULTIPLE_OF != 0: + raise ValueError(f"--resolution must be a multiple of {SIZE_MULTIPLE_OF}, got {args.resolution}.") + if args.aspect_ratio_buckets is not None: + for height, width in parse_buckets_string(args.aspect_ratio_buckets): + if height % SIZE_MULTIPLE_OF or width % SIZE_MULTIPLE_OF: + raise ValueError( + f"every --aspect_ratio_buckets entry must be a multiple of {SIZE_MULTIPLE_OF}, got " + f"{height}x{width}." + ) + + if args.with_prior_preservation: + if args.class_data_dir is None: + raise ValueError("You must specify a data directory for class images.") + if args.class_prompt is None: + raise ValueError("You must specify prompt for class images.") + else: + # logger is not available yet + if args.class_data_dir is not None: + warnings.warn("You need not use --class_data_dir without --with_prior_preservation.") + if args.class_prompt is not None: + warnings.warn("You need not use --class_prompt without --with_prior_preservation.") + + return args + + +class DreamBoothDataset(Dataset): + """ + A dataset to prepare the instance and class images with the prompts for fine-tuning the model. + It pre-processes the images. + """ + + def __init__( + self, + instance_data_root, + instance_prompt, + class_prompt, + class_data_root=None, + class_num=None, + size=1024, + repeats=1, + center_crop=False, + buckets=None, + use_aspect_ratio_buckets=False, + # 32, not the usual 16: both latent dimensions have to be even to fill 2x2 slots. + bucket_divisibility=SIZE_MULTIPLE_OF, + bucket_base_resolutions=None, + ): + self.size = size + self.resolution = size + self.center_crop = center_crop + + self.instance_prompt = instance_prompt + self.custom_instance_prompts = None + self.class_prompt = class_prompt + + # Explicit user-provided bucket list (or None). The concrete list of buckets actually used is + # built from the data in `self.buckets` during preprocessing below. + self._explicit_buckets = buckets + self.use_aspect_ratio_buckets = use_aspect_ratio_buckets + self.bucket_divisibility = bucket_divisibility + self.bucket_base_resolutions = bucket_base_resolutions + + # if --dataset_name is provided or a metadata jsonl file is provided in the local --instance_data directory, + # we load the training data using load_dataset + if args.dataset_name is not None: + try: + from datasets import load_dataset + except ImportError: + raise ImportError( + "You are trying to load your data using the datasets library. If you wish to train using custom " + "captions please install the datasets library: `pip install datasets`. If you wish to load a " + "local folder containing images only, specify --instance_data_dir instead." + ) + # Downloading and loading a dataset from the hub. + # See more about loading custom images at + # https://huggingface.co/docs/datasets/v2.0.0/en/dataset_script + dataset = load_dataset( + args.dataset_name, + args.dataset_config_name, + cache_dir=args.cache_dir, + ) + # Preprocessing the datasets. + column_names = dataset["train"].column_names + + # 6. Get the column names for input/target. + if args.image_column is None: + image_column = column_names[0] + logger.info(f"image column defaulting to {image_column}") + else: + image_column = args.image_column + if image_column not in column_names: + raise ValueError( + f"`--image_column` value '{args.image_column}' not found in dataset columns. Dataset columns are: {', '.join(column_names)}" + ) + instance_images = dataset["train"][image_column] + + if args.caption_column is None: + logger.info( + "No caption column provided, defaulting to instance_prompt for all images. If your dataset " + "contains captions/prompts for the images, make sure to specify the " + "column as --caption_column" + ) + self.custom_instance_prompts = None + else: + if args.caption_column not in column_names: + raise ValueError( + f"`--caption_column` value '{args.caption_column}' not found in dataset columns. Dataset columns are: {', '.join(column_names)}" + ) + custom_instance_prompts = dataset["train"][args.caption_column] + # create final list of captions according to --repeats + self.custom_instance_prompts = [] + for caption in custom_instance_prompts: + self.custom_instance_prompts.extend(itertools.repeat(caption, repeats)) + else: + self.instance_data_root = Path(instance_data_root) + if not self.instance_data_root.exists(): + raise ValueError("Instance images root doesn't exists.") + + instance_images = [Image.open(path) for path in list(Path(instance_data_root).iterdir())] + self.custom_instance_prompts = None + + self.instance_images = [] + for img in instance_images: + self.instance_images.extend(itertools.repeat(img, repeats)) + + self.pixel_values = [] + self.buckets = [] + bucket_to_idx = {} + for image in self.instance_images: + image = exif_transpose(image) + # RGBA, not RGB: this VAE takes and returns four channels. + if not image.mode == "RGBA": + image = image.convert("RGBA") + + width, height = image.size + + # Assign the image to a bucket. + target = self._bucket_for_image(height, width) + if target not in bucket_to_idx: + bucket_to_idx[target] = len(self.buckets) + self.buckets.append(target) + bucket_idx = bucket_to_idx[target] + + # based on the bucket assignment, define the transformations + image = self.train_transform( + image, + size=target, + center_crop=args.center_crop, + random_flip=args.random_flip, + ) + self.pixel_values.append((image, bucket_idx)) + + self.num_instance_images = len(self.instance_images) + self._length = self.num_instance_images + + if class_data_root is not None: + self.class_data_root = Path(class_data_root) + self.class_data_root.mkdir(parents=True, exist_ok=True) + self.class_images_path = list(self.class_data_root.iterdir()) + if class_num is not None: + self.num_class_images = min(len(self.class_images_path), class_num) + else: + self.num_class_images = len(self.class_images_path) + self._length = max(self.num_class_images, self.num_instance_images) + else: + self.class_data_root = None + + def __len__(self): + return self._length + + def __getitem__(self, index): + example = {} + instance_image, bucket_idx = self.pixel_values[index % self.num_instance_images] + example["index"] = index + example["instance_images"] = instance_image + example["bucket_idx"] = bucket_idx + if self.custom_instance_prompts: + caption = self.custom_instance_prompts[index % self.num_instance_images] + if caption: + example["instance_prompt"] = caption + else: + example["instance_prompt"] = self.instance_prompt + + else: # custom prompts were provided, but length does not match size of image dataset + example["instance_prompt"] = self.instance_prompt + + if self.class_data_root: + class_image = Image.open(self.class_images_path[index % self.num_class_images]) + class_image = exif_transpose(class_image) + + if not class_image.mode == "RGBA": + class_image = class_image.convert("RGBA") + # Match the class image to the paired instance image's bucket so they can be stacked into one batch. + example["class_images"] = self.train_transform( + class_image, size=self.buckets[bucket_idx], center_crop=self.center_crop + ) + example["class_prompt"] = self.class_prompt + + return example + + def _bucket_for_image(self, height, width): + # An explicit bucket list takes priority: pick the nearest, upscaling smaller images to cover it. + if self._explicit_buckets is not None: + return self._explicit_buckets[find_nearest_bucket(height, width, self._explicit_buckets)] + # On-the-fly bucketing: cap the ladder to the image's own resolution so smaller images are + # assigned to a smaller bucket rather than being upscaled (mirrors ostris' bucketing). + if self.use_aspect_ratio_buckets: + resolution = min(self.resolution, round((height * width) ** 0.5)) + ladder = generate_aspect_ratio_buckets( + resolution, + divisibility=self.bucket_divisibility, + base_resolutions=self.bucket_base_resolutions, + ) + return ladder[find_nearest_bucket(height, width, ladder)] + # No bucketing: a single square bucket reproduces the fixed-size resize + crop. + return (self.resolution, self.resolution) + + def train_transform(self, image, size, center_crop=False, random_flip=False): + # Resize preserving aspect ratio so the image covers the bucket, then crop to the bucket size. + target_height, target_width = size + width, height = image.size + scale = max(target_height / height, target_width / width) + new_height, new_width = round(height * scale), round(width * scale) + image = TF.resize(image, [new_height, new_width], interpolation=transforms.InterpolationMode.BILINEAR) + if center_crop: + image = TF.center_crop(image, size) + else: + i, j, h, w = transforms.RandomCrop.get_params(image, output_size=size) + image = TF.crop(image, i, j, h, w) + if random_flip and random.random() < 0.5: + image = TF.hflip(image) + return TF.normalize(TF.to_tensor(image), [0.5], [0.5]) + + +def collate_fn(examples, with_prior_preservation=False): + indices = [example["index"] for example in examples] + pixel_values = [example["instance_images"] for example in examples] + # Keep instance_prompts unchanged for prompt cache precompute; prompts may be extended with class prompts below. + instance_prompts = [example["instance_prompt"] for example in examples] + prompts = [example["instance_prompt"] for example in examples] + + # Concat class and instance examples for prior preservation. + # We do this to avoid doing two forward passes. + if with_prior_preservation: + pixel_values += [example["class_images"] for example in examples] + prompts += [example["class_prompt"] for example in examples] + + pixel_values = torch.stack(pixel_values) + # Qwen expects a `num_frames` dimension too. + if pixel_values.ndim == 4: + pixel_values = pixel_values.unsqueeze(2) + pixel_values = pixel_values.to(memory_format=torch.contiguous_format).float() + + batch = { + "indices": indices, + "pixel_values": pixel_values, + "instance_prompts": instance_prompts, + "prompts": prompts, + } + return batch + + +class BucketBatchSampler(BatchSampler): + def __init__(self, dataset: DreamBoothDataset, batch_size: int, drop_last: bool = False, seed: int = None): + if not isinstance(batch_size, int) or batch_size <= 0: + raise ValueError("batch_size should be a positive integer value, but got batch_size={}".format(batch_size)) + if not isinstance(drop_last, bool): + raise ValueError("drop_last should be a boolean value, but got drop_last={}".format(drop_last)) + + self.dataset = dataset + self.batch_size = batch_size + self.drop_last = drop_last + self.generator = random.Random(seed) if seed is not None else random + + # Group indices by bucket + self.bucket_indices = [[] for _ in range(len(self.dataset.buckets))] + for idx, (_, bucket_idx) in enumerate(self.dataset.pixel_values): + self.bucket_indices[bucket_idx].append(idx) + + self.sampler_len = 0 + for indices_in_bucket in self.bucket_indices: + num_batches, remainder = divmod(len(indices_in_bucket), self.batch_size) + self.sampler_len += num_batches + if remainder > 0 and not self.drop_last: + self.sampler_len += 1 + + def __iter__(self): + batches = [] + for indices_in_bucket in self.bucket_indices: + shuffled_indices = indices_in_bucket.copy() + self.generator.shuffle(shuffled_indices) + for i in range(0, len(shuffled_indices), self.batch_size): + batch = shuffled_indices[i : i + self.batch_size] + if len(batch) < self.batch_size and self.drop_last: + continue + batches.append(batch) + + self.generator.shuffle(batches) + for batch in batches: + yield batch + + def __len__(self): + return self.sampler_len + + +class PromptDataset(Dataset): + "A simple dataset to prepare the prompts to generate class images on multiple GPUs." + + def __init__(self, prompt, num_samples): + self.prompt = prompt + self.num_samples = num_samples + + def __len__(self): + return self.num_samples + + def __getitem__(self, index): + example = {} + example["prompt"] = self.prompt + example["index"] = index + return example + + +# These helpers only matter for prior preservation, where instance and class prompt +# embedding batches are concatenated and may not share the same mask/sequence length. +def _materialize_prompt_embedding_mask( + prompt_embeds: torch.Tensor, prompt_embeds_mask: torch.Tensor | None +) -> torch.Tensor: + """Return a dense mask tensor for a prompt embedding batch.""" + batch_size, seq_len = prompt_embeds.shape[:2] + + if prompt_embeds_mask is None: + return torch.ones((batch_size, seq_len), dtype=torch.long, device=prompt_embeds.device) + + if prompt_embeds_mask.shape != (batch_size, seq_len): + raise ValueError( + f"`prompt_embeds_mask` shape {prompt_embeds_mask.shape} must match prompt embeddings shape " + f"({batch_size}, {seq_len})." + ) + + return prompt_embeds_mask.to(device=prompt_embeds.device) + + +def _pad_prompt_embedding_pair( + prompt_embeds: torch.Tensor, prompt_embeds_mask: torch.Tensor | None, target_seq_len: int +) -> tuple[torch.Tensor, torch.Tensor]: + """Pad one prompt embedding batch and its mask to a shared sequence length.""" + prompt_embeds_mask = _materialize_prompt_embedding_mask(prompt_embeds, prompt_embeds_mask) + pad_width = target_seq_len - prompt_embeds.shape[1] + + if pad_width <= 0: + return prompt_embeds, prompt_embeds_mask + + prompt_embeds = torch.cat( + [prompt_embeds, prompt_embeds.new_zeros(prompt_embeds.shape[0], pad_width, prompt_embeds.shape[2])], dim=1 + ) + prompt_embeds_mask = torch.cat( + [prompt_embeds_mask, prompt_embeds_mask.new_zeros(prompt_embeds_mask.shape[0], pad_width)], dim=1 + ) + + return prompt_embeds, prompt_embeds_mask + + +def concat_prompt_embedding_batches( + *prompt_embedding_pairs: tuple[torch.Tensor, torch.Tensor | None], +) -> tuple[torch.Tensor, torch.Tensor | None]: + """Concatenate prompt embedding batches while handling missing masks and length mismatches.""" + if not prompt_embedding_pairs: + raise ValueError("At least one prompt embedding pair must be provided.") + + target_seq_len = max(prompt_embeds.shape[1] for prompt_embeds, _ in prompt_embedding_pairs) + padded_pairs = [ + _pad_prompt_embedding_pair(prompt_embeds, prompt_embeds_mask, target_seq_len) + for prompt_embeds, prompt_embeds_mask in prompt_embedding_pairs + ] + + merged_prompt_embeds = torch.cat([prompt_embeds for prompt_embeds, _ in padded_pairs], dim=0) + merged_mask = torch.cat([prompt_embeds_mask for _, prompt_embeds_mask in padded_pairs], dim=0) + + if merged_mask.all(): + return merged_prompt_embeds, None + + return merged_prompt_embeds, merged_mask + + +def main(args): + if args.report_to == "wandb" and args.hub_token is not None: + raise ValueError( + "You cannot use both --report_to=wandb and --hub_token due to a security risk of exposing your token." + " Please use `hf auth login` to authenticate with the Hub." + ) + + if torch.backends.mps.is_available() and args.mixed_precision == "bf16": + # due to pytorch#99272, MPS does not yet support bfloat16. + raise ValueError( + "Mixed precision training with bfloat16 is not supported on MPS. Please use fp16 (recommended) or fp32 instead." + ) + + logging_dir = Path(args.output_dir, args.logging_dir) + + accelerator_project_config = ProjectConfiguration(project_dir=args.output_dir, logging_dir=logging_dir) + kwargs = DistributedDataParallelKwargs(find_unused_parameters=True) + accelerator = Accelerator( + gradient_accumulation_steps=args.gradient_accumulation_steps, + mixed_precision=args.mixed_precision, + log_with=args.report_to, + project_config=accelerator_project_config, + kwargs_handlers=[kwargs], + ) + + # Disable AMP for MPS. + if torch.backends.mps.is_available(): + accelerator.native_amp = False + + if args.report_to == "wandb": + if not is_wandb_available(): + raise ImportError("Make sure to install wandb if you want to use it for logging during training.") + + # Make one log on every process with the configuration for debugging. + logging.basicConfig( + format="%(asctime)s - %(levelname)s - %(name)s - %(message)s", + datefmt="%m/%d/%Y %H:%M:%S", + level=logging.INFO, + ) + logger.info(accelerator.state, main_process_only=False) + if accelerator.is_local_main_process: + transformers.utils.logging.set_verbosity_warning() + diffusers.utils.logging.set_verbosity_info() + else: + transformers.utils.logging.set_verbosity_error() + diffusers.utils.logging.set_verbosity_error() + + # If passed along, set the training seed now. + if args.seed is not None: + set_seed(args.seed) + + # Generate class images if prior preservation is enabled. + if args.with_prior_preservation: + class_images_dir = Path(args.class_data_dir) + if not class_images_dir.exists(): + class_images_dir.mkdir(parents=True) + cur_class_images = len(list(class_images_dir.iterdir())) + + if cur_class_images < args.num_class_images: + pipeline = QwenImage21Pipeline.from_pretrained( + args.pretrained_model_name_or_path, + torch_dtype=torch.bfloat16 if args.mixed_precision == "bf16" else torch.float16, + revision=args.revision, + variant=args.variant, + ) + pipeline.set_progress_bar_config(disable=True) + + num_new_images = args.num_class_images - cur_class_images + logger.info(f"Number of class images to sample: {num_new_images}.") + + sample_dataset = PromptDataset(args.class_prompt, num_new_images) + sample_dataloader = torch.utils.data.DataLoader(sample_dataset, batch_size=args.sample_batch_size) + + sample_dataloader = accelerator.prepare(sample_dataloader) + pipeline.to(accelerator.device) + + for example in tqdm( + sample_dataloader, desc="Generating class images", disable=not accelerator.is_local_main_process + ): + images = pipeline(example["prompt"]).images + + for i, image in enumerate(images): + hash_image = insecure_hashlib.sha1(image.tobytes()).hexdigest() + # PNG, not JPEG: the pipeline returns RGBA and JPEG has no alpha channel. + image_filename = class_images_dir / f"{example['index'][i] + cur_class_images}-{hash_image}.png" + image.save(image_filename) + + pipeline.to("cpu") + del pipeline + free_memory() + + # Handle the repository creation + if accelerator.is_main_process: + if args.output_dir is not None: + os.makedirs(args.output_dir, exist_ok=True) + + if args.push_to_hub: + repo_id = create_repo( + repo_id=args.hub_model_id or Path(args.output_dir).name, + exist_ok=True, + ).repo_id + + # A Qwen3-VL processor rather than a tokenizer: it also expands condition images into vision tokens. + processor = Qwen3VLProcessor.from_pretrained( + args.pretrained_model_name_or_path, + subfolder="processor", + revision=args.revision, + ) + + # For mixed precision training we cast all non-trainable weights (vae, text_encoder and transformer) to half-precision + # as these weights are only used for inference, keeping weights in full precision is not required. + weight_dtype = torch.float32 + if accelerator.mixed_precision == "fp16": + weight_dtype = torch.float16 + elif accelerator.mixed_precision == "bf16": + weight_dtype = torch.bfloat16 + + # Load scheduler and models + # Used as shipped: it sets `use_dynamic_shifting`, so the training sigmas stay unshifted. + noise_scheduler = FlowMatchEulerDiscreteScheduler.from_pretrained( + args.pretrained_model_name_or_path, subfolder="scheduler", revision=args.revision + ) + noise_scheduler_copy = copy.deepcopy(noise_scheduler) + vae = AutoencoderKLQwenImage21.from_pretrained( + args.pretrained_model_name_or_path, + subfolder="vae", + revision=args.revision, + variant=args.variant, + ) + # 16 here, so one latent token covers a 16x16 pixel tile. + vae_scale_factor = 2 ** len(vae.temperal_downsample) + if vae_scale_factor * 2 != SIZE_MULTIPLE_OF: + raise ValueError( + f"This checkpoint's VAE scales by {vae_scale_factor}, so sizes must be multiples of " + f"{vae_scale_factor * 2}, but the resolution checks used {SIZE_MULTIPLE_OF}." + ) + latents_mean = (torch.tensor(vae.config.latents_mean).view(1, vae.config.z_dim, 1, 1, 1)).to(accelerator.device) + latents_std = 1.0 / torch.tensor(vae.config.latents_std).view(1, vae.config.z_dim, 1, 1, 1).to(accelerator.device) + text_encoder = Qwen3VLForConditionalGeneration.from_pretrained( + args.pretrained_model_name_or_path, subfolder="text_encoder", revision=args.revision, torch_dtype=weight_dtype + ) + quantization_config = None + if args.bnb_quantization_config_path is not None: + with open(args.bnb_quantization_config_path, "r") as f: + config_kwargs = json.load(f) + if "load_in_4bit" in config_kwargs and config_kwargs["load_in_4bit"]: + config_kwargs["bnb_4bit_compute_dtype"] = weight_dtype + quantization_config = BitsAndBytesConfig(**config_kwargs) + + transformer = QwenImage21Transformer2DModel.from_pretrained( + args.pretrained_model_name_or_path, + subfolder="transformer", + revision=args.revision, + variant=args.variant, + quantization_config=quantization_config, + torch_dtype=weight_dtype, + ) + if args.bnb_quantization_config_path is not None: + transformer = prepare_model_for_kbit_training(transformer, use_gradient_checkpointing=False) + + # We only train the additional adapter LoRA layers + transformer.requires_grad_(False) + vae.requires_grad_(False) + text_encoder.requires_grad_(False) + + if torch.backends.mps.is_available() and weight_dtype == torch.bfloat16: + # due to pytorch#99272, MPS does not yet support bfloat16. + raise ValueError( + "Mixed precision training with bfloat16 is not supported on MPS. Please use fp16 (recommended) or fp32 instead." + ) + + to_kwargs = {"dtype": weight_dtype, "device": accelerator.device} if not args.offload else {"dtype": weight_dtype} + # The VAE is stable in bf16, so load it in weight_dtype to reduce memory. + vae.to(**to_kwargs) + text_encoder.to(**to_kwargs) + # we never offload the transformer to CPU, so we can just use the accelerator device + transformer_to_kwargs = ( + {"device": accelerator.device} + if args.bnb_quantization_config_path is not None + else {"device": accelerator.device, "dtype": weight_dtype} + ) + transformer.to(**transformer_to_kwargs) + + # Initialize a text encoding pipeline and keep it to CPU for now. + text_encoding_pipeline = QwenImage21Pipeline.from_pretrained( + args.pretrained_model_name_or_path, + vae=None, + transformer=None, + processor=processor, + text_encoder=text_encoder, + scheduler=None, + ) + + if args.gradient_checkpointing: + transformer.enable_gradient_checkpointing() + + if args.lora_layers is not None: + target_modules = [layer.strip() for layer in args.lora_layers.split(",")] + else: + target_modules = ["to_k", "to_q", "to_v", "to_out.0"] + + # now we will add new LoRA weights the transformer layers + transformer_lora_config = LoraConfig( + r=args.rank, + lora_alpha=args.lora_alpha, + lora_dropout=args.lora_dropout, + init_lora_weights="gaussian", + target_modules=target_modules, + ) + transformer.add_adapter(transformer_lora_config) + + def unwrap_model(model): + model = accelerator.unwrap_model(model) + model = model._orig_mod if is_compiled_module(model) else model + return model + + # create custom saving & loading hooks so that `accelerator.save_state(...)` serializes in a nice format + def save_model_hook(models, weights, output_dir): + if accelerator.is_main_process: + transformer_lora_layers_to_save = None + modules_to_save = {} + + for model in models: + if isinstance(unwrap_model(model), type(unwrap_model(transformer))): + model = unwrap_model(model) + transformer_lora_layers_to_save = get_peft_model_state_dict(model) + modules_to_save["transformer"] = model + else: + raise ValueError(f"unexpected save model: {model.__class__}") + + # make sure to pop weight so that corresponding model is not saved again + if weights: + weights.pop() + + QwenImage21Pipeline.save_lora_weights( + output_dir, + transformer_lora_layers=transformer_lora_layers_to_save, + **_collate_lora_metadata(modules_to_save), + ) + + def load_model_hook(models, input_dir): + transformer_ = None + + if not accelerator.distributed_type == DistributedType.DEEPSPEED: + while len(models) > 0: + model = models.pop() + + if isinstance(unwrap_model(model), type(unwrap_model(transformer))): + model = unwrap_model(model) + transformer_ = model + else: + raise ValueError(f"unexpected save model: {model.__class__}") + else: + transformer_ = QwenImage21Transformer2DModel.from_pretrained( + args.pretrained_model_name_or_path, subfolder="transformer" + ) + transformer_.add_adapter(transformer_lora_config) + + lora_state_dict = QwenImage21Pipeline.lora_state_dict(input_dir) + + transformer_state_dict = { + f"{k.replace('transformer.', '')}": v for k, v in lora_state_dict.items() if k.startswith("transformer.") + } + transformer_state_dict = convert_unet_state_dict_to_peft(transformer_state_dict) + incompatible_keys = set_peft_model_state_dict(transformer_, transformer_state_dict, adapter_name="default") + if incompatible_keys is not None: + # check only for unexpected keys + unexpected_keys = getattr(incompatible_keys, "unexpected_keys", None) + if unexpected_keys: + logger.warning( + f"Loading adapter weights from state_dict led to unexpected keys not found in the model: " + f" {unexpected_keys}. " + ) + + # Make sure the trainable params are in float32. This is again needed since the base models + # are in `weight_dtype`. More details: + # https://github.com/huggingface/diffusers/pull/6514#discussion_r1449796804 + if args.mixed_precision == "fp16": + models = [transformer_] + # only upcast trainable parameters (LoRA) into fp32 + cast_training_params(models) + + accelerator.register_save_state_pre_hook(save_model_hook) + accelerator.register_load_state_pre_hook(load_model_hook) + + # Enable TF32 for faster training on Ampere GPUs, + # cf https://pytorch.org/docs/stable/notes/cuda.html#tensorfloat-32-tf32-on-ampere-devices + if args.allow_tf32 and torch.cuda.is_available(): + torch.backends.cuda.matmul.allow_tf32 = True + + if args.scale_lr: + args.learning_rate = ( + args.learning_rate * args.gradient_accumulation_steps * args.train_batch_size * accelerator.num_processes + ) + + # Make sure the trainable params are in float32. + if args.mixed_precision == "fp16": + models = [transformer] + # only upcast trainable parameters (LoRA) into fp32 + cast_training_params(models, dtype=torch.float32) + + transformer_lora_parameters = list(filter(lambda p: p.requires_grad, transformer.parameters())) + + # Optimization parameters + transformer_parameters_with_lr = {"params": transformer_lora_parameters, "lr": args.learning_rate} + params_to_optimize = [transformer_parameters_with_lr] + + # Optimizer creation + if not (args.optimizer.lower() == "prodigy" or args.optimizer.lower() == "adamw"): + logger.warning( + f"Unsupported choice of optimizer: {args.optimizer}.Supported optimizers include [adamW, prodigy]." + "Defaulting to adamW" + ) + args.optimizer = "adamw" + + if args.use_8bit_adam and not args.optimizer.lower() == "adamw": + logger.warning( + f"use_8bit_adam is ignored when optimizer is not set to 'AdamW'. Optimizer was " + f"set to {args.optimizer.lower()}" + ) + + if args.optimizer.lower() == "adamw": + if args.use_8bit_adam: + try: + import bitsandbytes as bnb + except ImportError: + raise ImportError( + "To use 8-bit Adam, please install the bitsandbytes library: `pip install bitsandbytes`." + ) + + optimizer_class = bnb.optim.AdamW8bit + else: + optimizer_class = torch.optim.AdamW + + optimizer = optimizer_class( + params_to_optimize, + betas=(args.adam_beta1, args.adam_beta2), + weight_decay=args.adam_weight_decay, + eps=args.adam_epsilon, + ) + + if args.optimizer.lower() == "prodigy": + try: + import prodigyopt + except ImportError: + raise ImportError("To use Prodigy, please install the prodigyopt library: `pip install prodigyopt`") + + optimizer_class = prodigyopt.Prodigy + + if args.learning_rate <= 0.1: + logger.warning( + "Learning rate is too low. When using prodigy, it's generally better to set learning rate around 1.0" + ) + + optimizer = optimizer_class( + params_to_optimize, + betas=(args.adam_beta1, args.adam_beta2), + beta3=args.prodigy_beta3, + weight_decay=args.adam_weight_decay, + eps=args.adam_epsilon, + decouple=args.prodigy_decouple, + use_bias_correction=args.prodigy_use_bias_correction, + safeguard_warmup=args.prodigy_safeguard_warmup, + ) + + # Resolve the bucketing mode. Bucketing must be enabled explicitly with --use_aspect_ratio_buckets; + # a bucket list without that flag is an error. With the flag, an explicit --aspect_ratio_buckets list + # drives assignment, otherwise buckets are computed on the fly inside the dataset. Without the flag a + # single square bucket reproduces the fixed-size resize + crop. + if args.aspect_ratio_buckets is not None and not args.use_aspect_ratio_buckets: + raise ValueError("--aspect_ratio_buckets requires --use_aspect_ratio_buckets to be set.") + if args.aspect_ratio_buckets is not None: + buckets = parse_buckets_string(args.aspect_ratio_buckets) + use_aspect_ratio_buckets = False + logger.info(f"Using explicit aspect ratio buckets: {buckets}") + elif args.use_aspect_ratio_buckets: + buckets = None + use_aspect_ratio_buckets = True + logger.info( + "No --aspect_ratio_buckets provided; auto-computing aspect ratio buckets on the fly from --resolution." + ) + else: + buckets = [(args.resolution, args.resolution)] + use_aspect_ratio_buckets = False + + # Dataset and DataLoaders creation: + train_dataset = DreamBoothDataset( + instance_data_root=args.instance_data_dir, + instance_prompt=args.instance_prompt, + class_prompt=args.class_prompt, + class_data_root=args.class_data_dir if args.with_prior_preservation else None, + class_num=args.num_class_images, + size=args.resolution, + repeats=args.repeats, + center_crop=args.center_crop, + buckets=buckets, + use_aspect_ratio_buckets=use_aspect_ratio_buckets, + ) + precompute_latents = args.cache_latents or train_dataset.custom_instance_prompts + batch_sampler = BucketBatchSampler(train_dataset, batch_size=args.train_batch_size, drop_last=True, seed=args.seed) + train_dataloader = torch.utils.data.DataLoader( + train_dataset, + batch_sampler=batch_sampler, + collate_fn=lambda examples: collate_fn(examples, args.with_prior_preservation), + num_workers=args.dataloader_num_workers, + ) + + def compute_text_embeddings(prompt, text_encoding_pipeline): + with torch.no_grad(): + # The image-pad mask marks where condition-image tokens sit; the pipeline needs it alongside the embeddings. + prompt_embeds, prompt_embeds_mask, image_pad_mask = text_encoding_pipeline.encode_prompt(prompt=prompt) + return prompt_embeds, prompt_embeds_mask, image_pad_mask + + # If no type of tuning is done on the text_encoder and custom instance prompts are NOT + # provided (i.e. the --instance_prompt is used for all images), we encode the instance prompt once to avoid + # the redundant encoding. + if not train_dataset.custom_instance_prompts: + with offload_models(text_encoding_pipeline, device=accelerator.device, offload=args.offload): + instance_prompt_embeds, instance_prompt_embeds_mask, _ = compute_text_embeddings( + args.instance_prompt, text_encoding_pipeline + ) + + # Handle class prompt for prior-preservation. + if args.with_prior_preservation: + with offload_models(text_encoding_pipeline, device=accelerator.device, offload=args.offload): + class_prompt_embeds, class_prompt_embeds_mask, _ = compute_text_embeddings( + args.class_prompt, text_encoding_pipeline + ) + + # When caption dropout is enabled, we precompute the empty ("") prompt embedding once and swap it in + # for randomly selected instance samples at training time (see the training loop below). + if args.caption_dropout > 0: + with offload_models(text_encoding_pipeline, device=accelerator.device, offload=args.offload): + empty_prompt_embeds, empty_prompt_embeds_mask, _ = compute_text_embeddings("", text_encoding_pipeline) + + validation_pipeline_args = {} + validation_image_pad_mask = None + # The final pass runs on `--final_validation_prompt` when `--validation_prompt` is absent, so the embeddings + # have to be built for whichever one is set - the text encoder is freed before that pass reaches the pipeline. + effective_validation_prompt = args.validation_prompt or args.final_validation_prompt + if effective_validation_prompt is not None: + with offload_models(text_encoding_pipeline, device=accelerator.device, offload=args.offload): + embeds, embeds_mask, image_pad_mask = compute_text_embeddings( + effective_validation_prompt, text_encoding_pipeline + ) + validation_pipeline_args = {"prompt_embeds": embeds, "prompt_embeds_mask": embeds_mask} + validation_image_pad_mask = image_pad_mask + + # if cache_latents is set to True, we encode images to latents and store them. + # Similar to pre-encoding in the case of a single instance prompt, if custom prompts are provided + # we encode them in advance as well. Caches are keyed by dataset index so they stay correct under + # aspect-ratio bucketing, where the batch composition differs between the caching pass and training. + if args.cache_latents: + instance_latents_cache = [None] * train_dataset.num_instance_images + class_latents_cache = [None] * train_dataset.num_instance_images if args.with_prior_preservation else None + if train_dataset.custom_instance_prompts: + prompt_embeds_cache = [None] * train_dataset.num_instance_images + prompt_embeds_mask_cache = [None] * train_dataset.num_instance_images + if precompute_latents: + cache_batch_sampler = BucketBatchSampler( + train_dataset, batch_size=args.train_batch_size, drop_last=False, seed=args.seed + ) + cache_dataloader = torch.utils.data.DataLoader( + train_dataset, + batch_sampler=cache_batch_sampler, + collate_fn=lambda examples: collate_fn(examples, args.with_prior_preservation), + num_workers=args.dataloader_num_workers, + ) + for batch in tqdm(cache_dataloader, desc="Caching latents"): + with torch.no_grad(): + sample_indices = batch["indices"] + if args.cache_latents: + with offload_models(vae, device=accelerator.device, offload=args.offload): + batch["pixel_values"] = batch["pixel_values"].to( + accelerator.device, non_blocking=True, dtype=vae.dtype + ) + latents = vae.encode(batch["pixel_values"]).latent_dist.sample() + if args.with_prior_preservation: + instance_latents, class_latents = torch.chunk(latents, 2, dim=0) + else: + instance_latents = latents + for i, idx in enumerate(sample_indices): + instance_latents_cache[idx] = instance_latents[i : i + 1] + if args.with_prior_preservation: + class_latents_cache[idx] = class_latents[i : i + 1] + if train_dataset.custom_instance_prompts: + with offload_models(text_encoding_pipeline, device=accelerator.device, offload=args.offload): + prompt_embeds, prompt_embeds_mask, _ = compute_text_embeddings( + batch["instance_prompts"], text_encoding_pipeline + ) + # `encode_prompt` returns no mask when nothing in the batch is padded, which is the + # common case here since a bucket's captions often tokenize to the same length. The + # cache is read back per sample, so store a dense mask rather than a `None` to slice. + prompt_embeds_mask = _materialize_prompt_embedding_mask(prompt_embeds, prompt_embeds_mask) + for i, idx in enumerate(sample_indices): + prompt_embeds_cache[idx] = prompt_embeds[i : i + 1] + prompt_embeds_mask_cache[idx] = prompt_embeds_mask[i : i + 1] + + if args.cache_latents: + assert all(latents is not None for latents in instance_latents_cache), "Latent cache has unfilled entries." + if args.with_prior_preservation: + assert all(latents is not None for latents in class_latents_cache), ( + "Class latent cache has unfilled entries." + ) + if train_dataset.custom_instance_prompts: + assert all(embeds is not None for embeds in prompt_embeds_cache), ( + "Prompt embedding cache has unfilled entries." + ) + + # move back to cpu before deleting to ensure memory is freed see: https://github.com/huggingface/diffusers/issues/11376#issue-3008144624 + if args.cache_latents: + vae = vae.to("cpu") + del vae + + # move back to cpu before deleting to ensure memory is freed see: https://github.com/huggingface/diffusers/issues/11376#issue-3008144624 + text_encoding_pipeline = text_encoding_pipeline.to("cpu") + # The processor stays: it holds no weights, and the pipeline cannot be constructed without one. + del text_encoder + free_memory() + + # Scheduler and math around the number of training steps. + overrode_max_train_steps = False + num_update_steps_per_epoch = math.ceil(len(train_dataloader) / args.gradient_accumulation_steps) + if args.max_train_steps is None: + args.max_train_steps = args.num_train_epochs * num_update_steps_per_epoch + overrode_max_train_steps = True + + lr_scheduler = get_scheduler( + args.lr_scheduler, + optimizer=optimizer, + num_warmup_steps=args.lr_warmup_steps * accelerator.num_processes, + num_training_steps=args.max_train_steps * accelerator.num_processes, + num_cycles=args.lr_num_cycles, + power=args.lr_power, + ) + + # Prepare everything with our `accelerator`. + transformer, optimizer, train_dataloader, lr_scheduler = accelerator.prepare( + transformer, optimizer, train_dataloader, lr_scheduler + ) + + # We need to recalculate our total training steps as the size of the training dataloader may have changed. + num_update_steps_per_epoch = math.ceil(len(train_dataloader) / args.gradient_accumulation_steps) + if overrode_max_train_steps: + args.max_train_steps = args.num_train_epochs * num_update_steps_per_epoch + # Afterwards we recalculate our number of training epochs + args.num_train_epochs = math.ceil(args.max_train_steps / num_update_steps_per_epoch) + + # We need to initialize the trackers we use, and also store our configuration. + # The trackers initializes automatically on the main process. + if accelerator.is_main_process: + tracker_name = "dreambooth-qwen-image-lora" + accelerator.init_trackers(tracker_name, config=vars(args)) + + # Train! + total_batch_size = args.train_batch_size * accelerator.num_processes * args.gradient_accumulation_steps + + logger.info("***** Running training *****") + logger.info(f" Num examples = {len(train_dataset)}") + logger.info(f" Num batches each epoch = {len(train_dataloader)}") + logger.info(f" Num Epochs = {args.num_train_epochs}") + logger.info(f" Instantaneous batch size per device = {args.train_batch_size}") + logger.info(f" Total train batch size (w. parallel, distributed & accumulation) = {total_batch_size}") + logger.info(f" Gradient Accumulation steps = {args.gradient_accumulation_steps}") + logger.info(f" Total optimization steps = {args.max_train_steps}") + global_step = 0 + first_epoch = 0 + + # Potentially load in the weights and states from a previous save + if args.resume_from_checkpoint: + if args.resume_from_checkpoint != "latest": + path = os.path.basename(args.resume_from_checkpoint) + else: + # Get the mos recent checkpoint + dirs = os.listdir(args.output_dir) + dirs = [d for d in dirs if d.startswith("checkpoint")] + dirs = sorted(dirs, key=lambda x: int(x.split("-")[1])) + path = dirs[-1] if len(dirs) > 0 else None + + if path is None: + accelerator.print( + f"Checkpoint '{args.resume_from_checkpoint}' does not exist. Starting a new training run." + ) + args.resume_from_checkpoint = None + initial_global_step = 0 + else: + accelerator.print(f"Resuming from checkpoint {path}") + accelerator.load_state(os.path.join(args.output_dir, path)) + global_step = int(path.split("-")[1]) + + initial_global_step = global_step + first_epoch = global_step // num_update_steps_per_epoch + + else: + initial_global_step = 0 + + progress_bar = tqdm( + range(0, args.max_train_steps), + initial=initial_global_step, + desc="Steps", + # Only show the progress bar once on each machine. + disable=not accelerator.is_local_main_process, + ) + + def get_sigmas(timesteps, n_dim=4, dtype=torch.float32): + sigmas = noise_scheduler_copy.sigmas.to(device=accelerator.device, dtype=dtype) + schedule_timesteps = noise_scheduler_copy.timesteps.to(accelerator.device) + timesteps = timesteps.to(accelerator.device) + step_indices = [(schedule_timesteps == t).nonzero().item() for t in timesteps] + + sigma = sigmas[step_indices].flatten() + while len(sigma.shape) < n_dim: + sigma = sigma.unsqueeze(-1) + return sigma + + for epoch in range(first_epoch, args.num_train_epochs): + transformer.train() + + for batch in train_dataloader: + models_to_accumulate = [transformer] + sample_indices = batch["indices"] + n_inst = len(sample_indices) + + with accelerator.accumulate(models_to_accumulate): + # Assemble this batch's instance prompt embeddings as one (embeds, mask) pair per sample, + # gathered by dataset index so they stay aligned with the latents under aspect-ratio bucketing. + if train_dataset.custom_instance_prompts: + instance_pairs = [ + (prompt_embeds_cache[idx], prompt_embeds_mask_cache[idx]) for idx in sample_indices + ] + else: + instance_pairs = [(instance_prompt_embeds, instance_prompt_embeds_mask)] * n_inst + + # Caption dropout: replace a sample's caption embedding with the empty-prompt embedding so it + # trains unconditionally. Only instance captions are dropped, never class/prior captions. + if args.caption_dropout > 0: + instance_pairs = [ + (empty_prompt_embeds, empty_prompt_embeds_mask) + if random.random() < args.caption_dropout + else pair + for pair in instance_pairs + ] + + # collate_fn orders batches as [instance..., class...]; keep the prompt embeddings in the same order. + prompt_pairs = instance_pairs + if args.with_prior_preservation: + prompt_pairs = prompt_pairs + [(class_prompt_embeds, class_prompt_embeds_mask)] * n_inst + prompt_embeds, prompt_embeds_mask = concat_prompt_embedding_batches(*prompt_pairs) + + # Convert images to latent space + if args.cache_latents: + model_input = torch.cat([instance_latents_cache[idx] for idx in sample_indices], dim=0) + if args.with_prior_preservation: + model_input = torch.cat( + [model_input, torch.cat([class_latents_cache[idx] for idx in sample_indices], dim=0)], + dim=0, + ) + else: + # `vae.encode` belongs inside the context manager: on the way out it puts the VAE back + # on the CPU, and encoding a batch the prepared dataloader has already placed on the + # accelerator would then raise. + with offload_models(vae, device=accelerator.device, offload=args.offload): + pixel_values = batch["pixel_values"].to(dtype=vae.dtype) + model_input = vae.encode(pixel_values).latent_dist.sample() + + model_input = (model_input - latents_mean) * latents_std + model_input = model_input.to(dtype=weight_dtype) + + # Sample noise that we'll add to the latents + noise = torch.randn_like(model_input) + bsz = model_input.shape[0] + + # Sample a random timestep for each image + # for weighting schemes where we sample timesteps non-uniformly + u = compute_density_for_timestep_sampling( + weighting_scheme=args.weighting_scheme, + batch_size=bsz, + logit_mean=args.logit_mean, + logit_std=args.logit_std, + mode_scale=args.mode_scale, + ) + indices = (u * noise_scheduler_copy.config.num_train_timesteps).long() + timesteps = noise_scheduler_copy.timesteps[indices].to(device=model_input.device) + + # Add noise according to flow matching. + # zt = (1 - texp) * x + texp * z1 + sigmas = get_sigmas(timesteps, n_dim=model_input.ndim, dtype=model_input.dtype) + noisy_model_input = (1.0 - sigmas) * model_input + sigmas * noise + + # Predict the noise residual. A batch is single-bucket, so the latent height/width are shared + # across the batch; derive them from the latents to support aspect-ratio buckets. + latent_height, latent_width = model_input.shape[3], model_input.shape[4] + # One shape per image in the sequence; text to image has only the target. + img_shapes = [[(1, latent_height, latent_width)]] * bsz + # Latents are consumed unpatched, so packing is a plain spatial flatten. + packed_noisy_model_input = QwenImage21Pipeline._pack_latents( + noisy_model_input, + batch_size=model_input.shape[0], + num_channels_latents=model_input.shape[1], + height=latent_height, + width=latent_width, + ) + # `img_mask` marks which positions over [prompt tokens, target slots] stand for image latents, one + # slot per 2x2 group of latents. The prompt half is all-False without condition images. + target_slots = (latent_height * latent_width) // 4 + img_mask = torch.cat( + [ + torch.zeros(bsz, prompt_embeds.shape[1], dtype=torch.bool, device=accelerator.device), + torch.ones(bsz, target_slots, dtype=torch.bool, device=accelerator.device), + ], + dim=1, + ) + model_pred = transformer( + hidden_states=packed_noisy_model_input, + encoder_hidden_states=prompt_embeds, + encoder_hidden_states_mask=prompt_embeds_mask, + timestep=timesteps / 1000, + img_shapes=img_shapes, + img_mask=img_mask, + return_dict=False, + )[0] + # The prediction spans the joint sequence, so keep the target's tail, as `__call__` does. + model_pred = model_pred[:, -packed_noisy_model_input.shape[1] :] + model_pred = QwenImage21Pipeline._unpack_latents( + model_pred, latent_height * vae_scale_factor, latent_width * vae_scale_factor, vae_scale_factor + ) + + # these weighting schemes use a uniform timestep sampling + # and instead post-weight the loss + weighting = compute_loss_weighting_for_sd3(weighting_scheme=args.weighting_scheme, sigmas=sigmas) + + target = noise - model_input + if args.with_prior_preservation: + # Chunk the noise and model_pred into two parts and compute the loss on each part separately. + model_pred, model_pred_prior = torch.chunk(model_pred, 2, dim=0) + target, target_prior = torch.chunk(target, 2, dim=0) + weighting, weighting_prior = torch.chunk(weighting, 2, dim=0) + + # Compute prior loss + prior_loss = torch.mean( + (weighting_prior.float() * (model_pred_prior.float() - target_prior.float()) ** 2).reshape( + target_prior.shape[0], -1 + ), + 1, + ) + prior_loss = prior_loss.mean() + + # Compute regular loss. + loss = torch.mean( + (weighting.float() * (model_pred.float() - target.float()) ** 2).reshape(target.shape[0], -1), + 1, + ) + loss = loss.mean() + + if args.with_prior_preservation: + # Add the prior loss to the instance loss. + loss = loss + args.prior_loss_weight * prior_loss + + accelerator.backward(loss) + if accelerator.sync_gradients: + params_to_clip = transformer.parameters() + accelerator.clip_grad_norm_(params_to_clip, args.max_grad_norm) + + optimizer.step() + lr_scheduler.step() + optimizer.zero_grad() + + # Checks if the accelerator has performed an optimization step behind the scenes + if accelerator.sync_gradients: + progress_bar.update(1) + global_step += 1 + + if accelerator.is_main_process or accelerator.distributed_type == DistributedType.DEEPSPEED: + if global_step % args.checkpointing_steps == 0: + # _before_ saving state, check if this save would set us over the `checkpoints_total_limit` + if args.checkpoints_total_limit is not None: + checkpoints = os.listdir(args.output_dir) + checkpoints = [d for d in checkpoints if d.startswith("checkpoint")] + checkpoints = sorted(checkpoints, key=lambda x: int(x.split("-")[1])) + + # before we save the new checkpoint, we need to have at _most_ `checkpoints_total_limit - 1` checkpoints + if len(checkpoints) >= args.checkpoints_total_limit: + num_to_remove = len(checkpoints) - args.checkpoints_total_limit + 1 + removing_checkpoints = checkpoints[0:num_to_remove] + + logger.info( + f"{len(checkpoints)} checkpoints already exist, removing {len(removing_checkpoints)} checkpoints" + ) + logger.info(f"removing checkpoints: {', '.join(removing_checkpoints)}") + + for removing_checkpoint in removing_checkpoints: + removing_checkpoint = os.path.join(args.output_dir, removing_checkpoint) + shutil.rmtree(removing_checkpoint) + + save_path = os.path.join(args.output_dir, f"checkpoint-{global_step}") + accelerator.save_state(save_path) + logger.info(f"Saved state to {save_path}") + + logs = {"loss": loss.detach().item(), "lr": lr_scheduler.get_last_lr()[0]} + progress_bar.set_postfix(**logs) + accelerator.log(logs, step=global_step) + + if global_step >= args.max_train_steps: + break + + if accelerator.is_main_process: + if args.validation_prompt is not None and epoch % args.validation_epochs == 0: + # create pipeline. The prompt is supplied as embeddings, so no text encoder is loaded. + pipeline = QwenImage21ValidationPipeline.from_pretrained( + args.pretrained_model_name_or_path, + text_encoder=None, + processor=processor, + transformer=accelerator.unwrap_model(transformer), + revision=args.revision, + variant=args.variant, + torch_dtype=weight_dtype, + ) + pipeline.cached_image_pad_mask = validation_image_pad_mask + images = log_validation( + pipeline=pipeline, + args=args, + accelerator=accelerator, + pipeline_args=validation_pipeline_args, + torch_dtype=weight_dtype, + epoch=epoch, + ) + del pipeline + images = None + free_memory() + + # Save the lora layers + accelerator.wait_for_everyone() + if accelerator.is_main_process: + modules_to_save = {} + transformer = unwrap_model(transformer) + if args.bnb_quantization_config_path is None: + if args.upcast_before_saving: + transformer.to(torch.float32) + else: + transformer = transformer.to(weight_dtype) + transformer_lora_layers = get_peft_model_state_dict(transformer) + modules_to_save["transformer"] = transformer + + QwenImage21Pipeline.save_lora_weights( + save_directory=args.output_dir, + transformer_lora_layers=transformer_lora_layers, + **_collate_lora_metadata(modules_to_save), + ) + + images = [] + run_validation = (args.validation_prompt and args.num_validation_images > 0) or (args.final_validation_prompt) + should_run_final_inference = not args.skip_final_inference and run_validation + if should_run_final_inference: + # Final inference + # Load previous pipeline + # The transformer is reloaded, so this exercises the adapter that was written to disk. + pipeline = QwenImage21ValidationPipeline.from_pretrained( + args.pretrained_model_name_or_path, + text_encoder=None, + processor=processor, + revision=args.revision, + variant=args.variant, + torch_dtype=weight_dtype, + ) + # load attention processors + pipeline.load_lora_weights(args.output_dir) + pipeline.cached_image_pad_mask = validation_image_pad_mask + + # run inference + images = log_validation( + pipeline=pipeline, + args=args, + accelerator=accelerator, + pipeline_args=validation_pipeline_args, + epoch=epoch, + is_final_validation=True, + torch_dtype=weight_dtype, + ) + del pipeline + free_memory() + + validation_prompt = args.validation_prompt if args.validation_prompt else args.final_validation_prompt + save_model_card( + (args.hub_model_id or Path(args.output_dir).name) if not args.push_to_hub else repo_id, + images=images, + base_model=args.pretrained_model_name_or_path, + instance_prompt=args.instance_prompt, + validation_prompt=validation_prompt, + repo_folder=args.output_dir, + ) + + if args.push_to_hub: + upload_folder( + repo_id=repo_id, + folder_path=args.output_dir, + commit_message="End of training", + ignore_patterns=["step_*", "epoch_*"], + ) + + images = None + + accelerator.end_training() + + +if __name__ == "__main__": + args = parse_args() + main(args) diff --git a/examples/dreambooth/train_dreambooth_lora_qwenimage21_img2img.py b/examples/dreambooth/train_dreambooth_lora_qwenimage21_img2img.py new file mode 100644 index 000000000000..530f95914f66 --- /dev/null +++ b/examples/dreambooth/train_dreambooth_lora_qwenimage21_img2img.py @@ -0,0 +1,2129 @@ +#!/usr/bin/env python +# coding=utf-8 +# Copyright 2025 The HuggingFace Inc. team. All rights reserved. +# +# Licensed under the Apache License, Version 2.0 (the "License"); +# you may not use this file except in compliance with the License. +# You may obtain a copy of the License at +# +# http://www.apache.org/licenses/LICENSE-2.0 +# +# Unless required by applicable law or agreed to in writing, software +# distributed under the License is distributed on an "AS IS" BASIS, +# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. +# See the License for the specific language governing permissions and + +# /// script +# dependencies = [ +# "diffusers @ git+https://github.com/huggingface/diffusers.git", +# "torch>=2.0.0", +# "accelerate>=0.31.0", +# "transformers>=4.41.2", +# "ftfy", +# "tensorboard", +# "Jinja2", +# "peft>=0.11.1", +# "sentencepiece", +# "torchvision", +# "datasets", +# "bitsandbytes", +# "prodigyopt", +# ] +# /// + +import argparse +import copy +import itertools +import json +import logging +import math +import os +import random +import shutil +import warnings +from contextlib import nullcontext +from pathlib import Path + +import numpy as np +import torch +import transformers +from accelerate import Accelerator, DistributedType +from accelerate.logging import get_logger +from accelerate.utils import DistributedDataParallelKwargs, ProjectConfiguration, set_seed +from huggingface_hub import create_repo, upload_folder +from huggingface_hub.utils import insecure_hashlib +from peft import LoraConfig, prepare_model_for_kbit_training, set_peft_model_state_dict +from peft.utils import get_peft_model_state_dict +from PIL import Image +from PIL.ImageOps import exif_transpose +from torch.utils.data import BatchSampler, Dataset +from torchvision import transforms +from torchvision.transforms import functional as TF +from tqdm.auto import tqdm +from transformers import Qwen3VLForConditionalGeneration, Qwen3VLProcessor + +import diffusers +from diffusers import ( + AutoencoderKLQwenImage21, + BitsAndBytesConfig, + FlowMatchEulerDiscreteScheduler, + QwenImage21Pipeline, + QwenImage21Transformer2DModel, +) +from diffusers.optimization import get_scheduler +from diffusers.pipelines.qwenimage21.pipeline_qwenimage21 import calculate_dimensions +from diffusers.training_utils import ( + _collate_lora_metadata, + cast_training_params, + compute_density_for_timestep_sampling, + compute_loss_weighting_for_sd3, + find_nearest_bucket, + free_memory, + generate_aspect_ratio_buckets, + offload_models, + parse_buckets_string, +) +from diffusers.utils import ( + check_min_version, + convert_unet_state_dict_to_peft, + is_wandb_available, + load_image, +) +from diffusers.utils.hub_utils import load_or_create_model_card, populate_model_card +from diffusers.utils.import_utils import is_torch_npu_available +from diffusers.utils.torch_utils import is_compiled_module + + +if is_wandb_available(): + import wandb + +# Will error if the minimal version of diffusers is not installed. Remove at your own risks. +check_min_version("0.41.0.dev0") + +logger = get_logger(__name__) + +# `vae_scale_factor * 2` for this VAE, checked against its config in `main`. +SIZE_MULTIPLE_OF = 32 + +if is_torch_npu_available(): + torch.npu.config.allow_internal_format = False + + +class QwenImage21ValidationPipeline(QwenImage21Pipeline): + """`QwenImage21Pipeline` that takes the image-pad mask alongside precomputed prompt embeddings. + + `__call__` gets that mask from `encode_prompt`, which it skips when handed `prompt_embeds`, so validation hands + back the one it cached before the text encoder was freed. + """ + + cached_image_pad_mask = None + + def encode_prompt(self, *args, **kwargs): + # The mask goes in before the call, not after it. Handed `prompt_embeds` together with a condition image, + # the base `encode_prompt` raises for a missing mask rather than returning `None` in its place, so a + # substitution made on the way out never runs. + if kwargs.get("image_pad_mask") is None and self.cached_image_pad_mask is not None: + kwargs["image_pad_mask"] = self.cached_image_pad_mask + prompt_embeds, prompt_embeds_mask, image_pad_mask = super().encode_prompt(*args, **kwargs) + if image_pad_mask is None: + image_pad_mask = self.cached_image_pad_mask + return prompt_embeds, prompt_embeds_mask, image_pad_mask + + +def save_model_card( + repo_id: str, + images=None, + base_model: str = None, + instance_prompt=None, + validation_prompt=None, + repo_folder=None, +): + widget_dict = [] + if images is not None: + for i, image in enumerate(images): + image.save(os.path.join(repo_folder, f"image_{i}.png")) + widget_dict.append( + {"text": validation_prompt if validation_prompt else " ", "output": {"url": f"image_{i}.png"}} + ) + + model_description = f""" +# Qwen-Image 2.1 image-to-image DreamBooth LoRA - {repo_id} + + + +## Model description + +These are {repo_id} DreamBooth LoRA weights for {base_model}. + +The weights were trained using [DreamBooth](https://dreambooth.github.io/) with the [Qwen-Image 2.1 image-to-image diffusers trainer](https://github.com/huggingface/diffusers/blob/main/examples/dreambooth/README_qwenimage21.md#image-to-image-editing). + +## Trigger words + +You should use `{instance_prompt}` as the edit instruction. + +## Download model + +[Download the *.safetensors LoRA]({repo_id}/tree/main) in the Files & versions tab. + +## Use it with the [🧨 diffusers library](https://github.com/huggingface/diffusers) + +```py + >>> import torch + >>> from diffusers import QwenImage21Pipeline + >>> from diffusers.utils import load_image + + >>> pipe = QwenImage21Pipeline.from_pretrained( + ... "Qwen/Qwen-Image-2.1", + ... torch_dtype=torch.bfloat16, + ... ) + >>> pipe.enable_model_cpu_offload() + >>> pipe.load_lora_weights(f"{repo_id}") + >>> condition = load_image("https://huggingface.co/datasets/huggingface/documentation-images/resolve/main/diffusers/cat.png") + >>> image = pipe(f"{instance_prompt}", image=condition).images[0] + + +``` + +For more details, including weighting, merging and fusing LoRAs, check the [documentation on loading LoRAs in diffusers](https://huggingface.co/docs/diffusers/main/en/using-diffusers/loading_adapters) +""" + model_card = load_or_create_model_card( + repo_id_or_path=repo_id, + from_training=True, + license="apache-2.0", + base_model=base_model, + prompt=instance_prompt, + model_description=model_description, + widget=widget_dict, + ) + tags = [ + "image-to-image", + "diffusers-training", + "diffusers", + "lora", + "qwen-image", + "qwen-image-2.1", + "qwen-image-diffusers", + "template:sd-lora", + ] + + model_card = populate_model_card(model_card, tags=tags) + model_card.save(os.path.join(repo_folder, "README.md")) + + +def log_validation( + pipeline, + args, + accelerator, + pipeline_args, + epoch, + torch_dtype, + is_final_validation=False, +): + args.num_validation_images = args.num_validation_images if args.num_validation_images else 1 + # `--final_validation_prompt` stands in when only the final pass is asked for. + validation_prompt = args.validation_prompt or args.final_validation_prompt + logger.info( + f"Running validation... \n Generating {args.num_validation_images} images with prompt: {validation_prompt}." + ) + pipeline = pipeline.to(accelerator.device, dtype=torch_dtype) + pipeline.set_progress_bar_config(disable=True) + + # run inference + generator = torch.Generator(device=accelerator.device).manual_seed(args.seed) if args.seed is not None else None + autocast_ctx = torch.autocast(accelerator.device.type) if not is_final_validation else nullcontext() + + images = [] + for _ in range(args.num_validation_images): + with autocast_ctx: + image = pipeline( + **pipeline_args, + num_inference_steps=args.validation_num_inference_steps, + # Classifier-free guidance off, as in the model's own sample script: with `true_cfg_scale <= 1` + # a step is a single forward pass through the transformer. + true_cfg_scale=1.0, + output_resolution=args.resolution, + generator=generator, + ).images[0] + images.append(image) + + for tracker in accelerator.trackers: + phase_name = "test" if is_final_validation else "validation" + if tracker.name == "tensorboard": + # The pipeline returns RGBA, this VAE having four channels, and `add_images` asserts three. + np_images = np.stack([np.asarray(img.convert("RGB")) for img in images]) + tracker.writer.add_images(phase_name, np_images, epoch, dataformats="NHWC") + if tracker.name == "wandb": + tracker.log( + { + phase_name: [ + wandb.Image(image, caption=f"{i}: {validation_prompt}") for i, image in enumerate(images) + ] + } + ) + + del pipeline + free_memory() + + return images + + +def parse_args(input_args=None): + parser = argparse.ArgumentParser(description="Simple example of a training script.") + parser.add_argument( + "--pretrained_model_name_or_path", + type=str, + default=None, + required=True, + help="Path to pretrained model or model identifier from huggingface.co/models.", + ) + parser.add_argument( + "--bnb_quantization_config_path", + type=str, + default=None, + help="Quantization config in a JSON file that will be used to define the bitsandbytes quant config of the DiT.", + ) + parser.add_argument( + "--revision", + type=str, + default=None, + required=False, + help="Revision of pretrained model identifier from huggingface.co/models.", + ) + parser.add_argument( + "--variant", + type=str, + default=None, + help="Variant of the model files of the pretrained model identifier from huggingface.co/models, 'e.g.' fp16", + ) + parser.add_argument( + "--dataset_name", + type=str, + default=None, + help=( + "The name of the Dataset (from the HuggingFace hub) containing the training data of instance images (could be your own, possibly private," + " dataset). It can also be a path pointing to a local copy of a dataset in your filesystem," + " or to a folder containing files that 🤗 Datasets can understand." + ), + ) + parser.add_argument( + "--dataset_config_name", + type=str, + default=None, + help="The config of the Dataset, leave as None if there's only one config.", + ) + parser.add_argument( + "--instance_data_dir", + type=str, + default=None, + help=("A folder containing the training data. "), + ) + + parser.add_argument( + "--cache_dir", + type=str, + default=None, + help="The directory where the downloaded models and datasets will be stored.", + ) + + parser.add_argument( + "--image_column", + type=str, + default="image", + help="The column of the dataset containing the target image. By " + "default, the standard Image Dataset maps out 'file_name' " + "to 'image'.", + ) + parser.add_argument( + "--cond_image_column", + type=str, + default=None, + help="Column in the dataset containing the condition image the edit is applied to. Required here.", + ) + parser.add_argument( + "--caption_column", + type=str, + default=None, + help="The column of the dataset containing the instance prompt for each image", + ) + + parser.add_argument("--repeats", type=int, default=1, help="How many times to repeat the training data.") + + parser.add_argument( + "--class_data_dir", + type=str, + default=None, + required=False, + help="A folder containing the training data of class images.", + ) + parser.add_argument( + "--instance_prompt", + type=str, + default=None, + required=True, + help="The prompt with identifier specifying the instance, e.g. 'photo of a TOK dog', 'in the style of TOK'", + ) + parser.add_argument( + "--class_prompt", + type=str, + default=None, + help="The prompt to specify images in the same class as provided instance images.", + ) + parser.add_argument( + "--validation_image", + type=str, + default=None, + help="Path or URL of the condition image to edit during validation.", + ) + parser.add_argument( + "--validation_num_inference_steps", + type=int, + default=40, + help="Denoising steps for validation images. 40 is what the model's own sample script uses.", + ) + + parser.add_argument( + "--validation_prompt", + type=str, + default=None, + help="A prompt that is used during validation to verify that the model is learning.", + ) + + parser.add_argument( + "--skip_final_inference", + default=False, + action="store_true", + help="Whether to skip the final inference step with loaded lora weights upon training completion. This will run intermediate validation inference if `validation_prompt` is provided. Specify to reduce memory.", + ) + + parser.add_argument( + "--final_validation_prompt", + type=str, + default=None, + help="A prompt that is used during a final validation to verify that the model is learning. Ignored if `--validation_prompt` is provided.", + ) + parser.add_argument( + "--num_validation_images", + type=int, + default=4, + help="Number of images that should be generated during validation with `validation_prompt`.", + ) + parser.add_argument( + "--validation_epochs", + type=int, + default=50, + help=( + "Run dreambooth validation every X epochs. Dreambooth validation consists of running the prompt" + " `args.validation_prompt` multiple times: `args.num_validation_images`." + ), + ) + parser.add_argument( + "--rank", + type=int, + default=16, + help=("The dimension of the LoRA update matrices."), + ) + parser.add_argument( + "--lora_alpha", + type=int, + default=16, + help="LoRA alpha. The update is scaled by `lora_alpha / rank`, so keep the two in step.", + ) + parser.add_argument("--lora_dropout", type=float, default=0.0, help="Dropout probability for LoRA layers") + + parser.add_argument( + "--with_prior_preservation", + default=False, + action="store_true", + help="Flag to add prior preservation loss.", + ) + parser.add_argument("--prior_loss_weight", type=float, default=1.0, help="The weight of prior preservation loss.") + parser.add_argument( + "--num_class_images", + type=int, + default=100, + help=( + "Minimal class images for prior preservation loss. If there are not enough images already present in" + " class_data_dir, additional images will be sampled with class_prompt." + ), + ) + parser.add_argument( + "--output_dir", + type=str, + default="qwenimage21-img2img-dreambooth-lora", + help="The output directory where the model predictions and checkpoints will be written.", + ) + parser.add_argument("--seed", type=int, default=None, help="A seed for reproducible training.") + parser.add_argument( + "--resolution", + type=int, + default=512, + help=( + "The resolution for input images, all the images in the train/validation dataset will be resized to this" + " resolution" + ), + ) + parser.add_argument( + "--aspect_ratio_buckets", + type=str, + default=None, + help=( + "Aspect ratio buckets to use for training. Define as a string of 'h1,w1;h2,w2;...'. " + "e.g. '1024,1024;768,1360;1360,768;880,1168;1168,880;1248,832;832,1248'. " + "Requires --use_aspect_ratio_buckets. Images are resized to cover and cropped to the nearest " + "listed bucket (smaller images are upscaled). When set, --resolution is ignored." + ), + ) + parser.add_argument( + "--use_aspect_ratio_buckets", + action="store_true", + help=( + "Enable aspect-ratio bucketing. Without --aspect_ratio_buckets, the buckets are computed on the " + "fly from --resolution and capped to each image's own resolution, so smaller images are assigned " + "to a smaller bucket instead of being upscaled. Provide --aspect_ratio_buckets to use an explicit list." + ), + ) + parser.add_argument( + "--center_crop", + default=False, + action="store_true", + help=( + "Whether to center crop the input images to the resolution. If not set, the images will be randomly" + " cropped. The images will be resized to the resolution first before cropping." + ), + ) + parser.add_argument( + "--random_flip", + action="store_true", + help="whether to randomly flip images horizontally", + ) + parser.add_argument( + "--caption_dropout", + type=float, + default=0.0, + help=( + "Probability of replacing an instance image's caption with an empty string during training, so that" + " fraction of samples is trained unconditionally. Improves classifier-free guidance. A common value is" + " 0.1. Class/prior-preservation captions are never dropped." + ), + ) + parser.add_argument( + "--train_batch_size", type=int, default=4, help="Batch size (per device) for the training dataloader." + ) + parser.add_argument( + "--sample_batch_size", type=int, default=4, help="Batch size (per device) for sampling images." + ) + parser.add_argument("--num_train_epochs", type=int, default=1) + parser.add_argument( + "--max_train_steps", + type=int, + default=None, + help="Total number of training steps to perform. If provided, overrides num_train_epochs.", + ) + parser.add_argument( + "--checkpointing_steps", + type=int, + default=500, + help=( + "Save a checkpoint of the training state every X updates. These checkpoints can be used both as final" + " checkpoints in case they are better than the last checkpoint, and are also suitable for resuming" + " training using `--resume_from_checkpoint`." + ), + ) + parser.add_argument( + "--checkpoints_total_limit", + type=int, + default=None, + help=("Max number of checkpoints to store."), + ) + parser.add_argument( + "--resume_from_checkpoint", + type=str, + default=None, + help=( + "Whether training should be resumed from a previous checkpoint. Use a path saved by" + ' `--checkpointing_steps`, or `"latest"` to automatically select the last available checkpoint.' + ), + ) + parser.add_argument( + "--gradient_accumulation_steps", + type=int, + default=1, + help="Number of updates steps to accumulate before performing a backward/update pass.", + ) + parser.add_argument( + "--gradient_checkpointing", + action="store_true", + help="Whether or not to use gradient checkpointing to save memory at the expense of slower backward pass.", + ) + parser.add_argument( + "--learning_rate", + type=float, + default=1e-4, + help="Initial learning rate (after the potential warmup period) to use.", + ) + parser.add_argument( + "--scale_lr", + action="store_true", + default=False, + help="Scale the learning rate by the number of GPUs, gradient accumulation steps, and batch size.", + ) + parser.add_argument( + "--lr_scheduler", + type=str, + default="constant", + help=( + 'The scheduler type to use. Choose between ["linear", "cosine", "cosine_with_restarts", "polynomial",' + ' "constant", "constant_with_warmup"]' + ), + ) + parser.add_argument( + "--lr_warmup_steps", type=int, default=500, help="Number of steps for the warmup in the lr scheduler." + ) + parser.add_argument( + "--lr_num_cycles", + type=int, + default=1, + help="Number of hard resets of the lr in cosine_with_restarts scheduler.", + ) + parser.add_argument("--lr_power", type=float, default=1.0, help="Power factor of the polynomial scheduler.") + parser.add_argument( + "--dataloader_num_workers", + type=int, + default=0, + help=( + "Number of subprocesses to use for data loading. 0 means that the data will be loaded in the main process." + ), + ) + parser.add_argument( + "--weighting_scheme", + type=str, + default="none", + choices=["sigma_sqrt", "logit_normal", "mode", "cosmap", "none"], + help=('We default to the "none" weighting scheme for uniform sampling and uniform loss'), + ) + parser.add_argument( + "--logit_mean", type=float, default=0.0, help="mean to use when using the `'logit_normal'` weighting scheme." + ) + parser.add_argument( + "--logit_std", type=float, default=1.0, help="std to use when using the `'logit_normal'` weighting scheme." + ) + parser.add_argument( + "--mode_scale", + type=float, + default=1.29, + help="Scale of mode weighting scheme. Only effective when using the `'mode'` as the `weighting_scheme`.", + ) + parser.add_argument( + "--optimizer", + type=str, + default="AdamW", + help=('The optimizer type to use. Choose between ["AdamW", "prodigy"]'), + ) + + parser.add_argument( + "--use_8bit_adam", + action="store_true", + help="Whether or not to use 8-bit Adam from bitsandbytes. Ignored if optimizer is not set to AdamW", + ) + + parser.add_argument( + "--adam_beta1", type=float, default=0.9, help="The beta1 parameter for the Adam and Prodigy optimizers." + ) + parser.add_argument( + "--adam_beta2", type=float, default=0.999, help="The beta2 parameter for the Adam and Prodigy optimizers." + ) + parser.add_argument( + "--prodigy_beta3", + type=float, + default=None, + help="coefficients for computing the Prodigy stepsize using running averages. If set to None, " + "uses the value of square root of beta2. Ignored if optimizer is adamW", + ) + parser.add_argument("--prodigy_decouple", type=bool, default=True, help="Use AdamW style decoupled weight decay") + parser.add_argument( + "--adam_weight_decay", type=float, default=1e-04, help="Weight decay to use for the LoRA parameters" + ) + parser.add_argument( + "--lora_layers", + type=str, + default=None, + help=( + 'The transformer modules to apply LoRA training on. Please specify the layers in a comma separated. E.g. - "to_k,to_q,to_v" will result in lora training of attention layers only' + ), + ) + + parser.add_argument( + "--adam_epsilon", + type=float, + default=1e-08, + help="Epsilon value for the Adam optimizer and Prodigy optimizers.", + ) + + parser.add_argument( + "--prodigy_use_bias_correction", + type=bool, + default=True, + help="Turn on Adam's bias correction. True by default. Ignored if optimizer is adamW", + ) + parser.add_argument( + "--prodigy_safeguard_warmup", + type=bool, + default=True, + help="Remove lr from the denominator of D estimate to avoid issues during warm-up stage. True by default. " + "Ignored if optimizer is adamW", + ) + parser.add_argument("--max_grad_norm", default=1.0, type=float, help="Max gradient norm.") + parser.add_argument("--push_to_hub", action="store_true", help="Whether or not to push the model to the Hub.") + parser.add_argument("--hub_token", type=str, default=None, help="The token to use to push to the Model Hub.") + parser.add_argument( + "--hub_model_id", + type=str, + default=None, + help="The name of the repository to keep in sync with the local `output_dir`.", + ) + parser.add_argument( + "--logging_dir", + type=str, + default="logs", + help=( + "[TensorBoard](https://www.tensorflow.org/tensorboard) log directory. Will default to" + " *output_dir/runs/**CURRENT_DATETIME_HOSTNAME***." + ), + ) + parser.add_argument( + "--allow_tf32", + action="store_true", + help=( + "Whether or not to allow TF32 on Ampere GPUs. Can be used to speed up training. For more information, see" + " https://pytorch.org/docs/stable/notes/cuda.html#tensorfloat-32-tf32-on-ampere-devices" + ), + ) + parser.add_argument( + "--cache_latents", + action="store_true", + default=False, + help="Cache the VAE latents", + ) + parser.add_argument( + "--report_to", + type=str, + default="tensorboard", + help=( + 'The integration to report the results and logs to. Supported platforms are `"tensorboard"`' + ' (default), `"wandb"` and `"comet_ml"`. Use `"all"` to report to all integrations.' + ), + ) + parser.add_argument( + "--mixed_precision", + type=str, + default=None, + choices=["no", "fp16", "bf16"], + help=( + "Whether to use mixed precision. Choose between fp16 and bf16 (bfloat16). Bf16 requires PyTorch >=" + " 1.10.and an Nvidia Ampere GPU. Default to the value of accelerate config of the current system or the" + " flag passed with the `accelerate.launch` command. Use this argument to override the accelerate config." + ), + ) + parser.add_argument( + "--upcast_before_saving", + action="store_true", + default=False, + help=( + "Whether to upcast the trained transformer layers to float32 before saving (at the end of training). " + "Defaults to precision dtype used for training to save memory" + ), + ) + parser.add_argument( + "--offload", + action="store_true", + help="Whether to offload the VAE and the text encoder to CPU when they are not used.", + ) + parser.add_argument("--local_rank", type=int, default=-1, help="For distributed training: local_rank") + + if input_args is not None: + args = parser.parse_args(input_args) + else: + args = parser.parse_args() + + if args.dataset_name is None and args.instance_data_dir is None: + raise ValueError("Specify either `--dataset_name` or `--instance_data_dir`") + + if args.dataset_name is not None and args.instance_data_dir is not None: + raise ValueError("Specify only one of `--dataset_name` or `--instance_data_dir`") + + env_local_rank = int(os.environ.get("LOCAL_RANK", -1)) + if env_local_rank != -1 and env_local_rank != args.local_rank: + args.local_rank = env_local_rank + + if args.dataset_name is None or args.cond_image_column is None: + raise ValueError( + "Image-to-image training pairs each target image with a condition image, so it needs a dataset that " + "holds both: pass `--dataset_name` together with `--cond_image_column`. For text-to-image training use " + "train_dreambooth_lora_qwenimage21.py." + ) + + if args.with_prior_preservation: + raise ValueError( + "`--with_prior_preservation` is not supported here. Class images carry a different prompt, and every " + "sample in a batch has to share one image-pad layout, which a second prompt breaks." + ) + + if args.caption_dropout > 0: + raise ValueError( + "`--caption_dropout` is not supported here. Dropping a caption changes where the condition image's " + "tokens land in the prompt, and the batch has to share one image-pad layout." + ) + + # An error rather than the pipeline's silent resize: a mismatch only shows up later, as a packing error. + if args.resolution % SIZE_MULTIPLE_OF != 0: + raise ValueError(f"--resolution must be a multiple of {SIZE_MULTIPLE_OF}, got {args.resolution}.") + if args.aspect_ratio_buckets is not None: + for height, width in parse_buckets_string(args.aspect_ratio_buckets): + if height % SIZE_MULTIPLE_OF or width % SIZE_MULTIPLE_OF: + raise ValueError( + f"every --aspect_ratio_buckets entry must be a multiple of {SIZE_MULTIPLE_OF}, got " + f"{height}x{width}." + ) + + if args.with_prior_preservation: + if args.class_data_dir is None: + raise ValueError("You must specify a data directory for class images.") + if args.class_prompt is None: + raise ValueError("You must specify prompt for class images.") + else: + # logger is not available yet + if args.class_data_dir is not None: + warnings.warn("You need not use --class_data_dir without --with_prior_preservation.") + if args.class_prompt is not None: + warnings.warn("You need not use --class_prompt without --with_prior_preservation.") + + return args + + +class DreamBoothDataset(Dataset): + """ + A dataset to prepare the instance and class images with the prompts for fine-tuning the model. + It pre-processes the images. + """ + + def __init__( + self, + instance_data_root, + instance_prompt, + class_prompt, + class_data_root=None, + class_num=None, + size=1024, + repeats=1, + center_crop=False, + buckets=None, + use_aspect_ratio_buckets=False, + # 32, not the usual 16: both latent dimensions have to be even to fill 2x2 slots. + bucket_divisibility=SIZE_MULTIPLE_OF, + bucket_base_resolutions=None, + ): + self.size = size + self.resolution = size + self.center_crop = center_crop + + self.instance_prompt = instance_prompt + self.custom_instance_prompts = None + self.class_prompt = class_prompt + + # Explicit user-provided bucket list (or None). The concrete list of buckets actually used is + # built from the data in `self.buckets` during preprocessing below. + self._explicit_buckets = buckets + self.use_aspect_ratio_buckets = use_aspect_ratio_buckets + self.bucket_divisibility = bucket_divisibility + self.bucket_base_resolutions = bucket_base_resolutions + + # if --dataset_name is provided or a metadata jsonl file is provided in the local --instance_data directory, + # we load the training data using load_dataset + if args.dataset_name is not None: + try: + from datasets import load_dataset + except ImportError: + raise ImportError( + "You are trying to load your data using the datasets library. If you wish to train using custom " + "captions please install the datasets library: `pip install datasets`. If you wish to load a " + "local folder containing images only, specify --instance_data_dir instead." + ) + # Downloading and loading a dataset from the hub. + # See more about loading custom images at + # https://huggingface.co/docs/datasets/v2.0.0/en/dataset_script + dataset = load_dataset( + args.dataset_name, + args.dataset_config_name, + cache_dir=args.cache_dir, + ) + # Preprocessing the datasets. + column_names = dataset["train"].column_names + + # 6. Get the column names for input/target. + if args.image_column is None: + image_column = column_names[0] + logger.info(f"image column defaulting to {image_column}") + else: + image_column = args.image_column + if image_column not in column_names: + raise ValueError( + f"`--image_column` value '{args.image_column}' not found in dataset columns. Dataset columns are: {', '.join(column_names)}" + ) + instance_images = dataset["train"][image_column] + + if args.cond_image_column not in column_names: + raise ValueError( + f"`--cond_image_column` value '{args.cond_image_column}' not found in dataset columns. Dataset " + f"columns are: {', '.join(column_names)}" + ) + cond_images = dataset["train"][args.cond_image_column] + if len(cond_images) != len(instance_images): + raise ValueError( + f"The dataset has {len(instance_images)} target images but {len(cond_images)} condition images." + ) + + if args.caption_column is None: + logger.info( + "No caption column provided, defaulting to instance_prompt for all images. If your dataset " + "contains captions/prompts for the images, make sure to specify the " + "column as --caption_column" + ) + self.custom_instance_prompts = None + else: + if args.caption_column not in column_names: + raise ValueError( + f"`--caption_column` value '{args.caption_column}' not found in dataset columns. Dataset columns are: {', '.join(column_names)}" + ) + custom_instance_prompts = dataset["train"][args.caption_column] + # create final list of captions according to --repeats + self.custom_instance_prompts = [] + for caption in custom_instance_prompts: + self.custom_instance_prompts.extend(itertools.repeat(caption, repeats)) + else: + self.instance_data_root = Path(instance_data_root) + if not self.instance_data_root.exists(): + raise ValueError("Instance images root doesn't exists.") + + instance_images = [Image.open(path) for path in list(Path(instance_data_root).iterdir())] + self.custom_instance_prompts = None + + self.instance_images = [] + self.cond_images = [] + for img, cond_img in zip(instance_images, cond_images): + self.instance_images.extend(itertools.repeat(img, repeats)) + self.cond_images.extend(itertools.repeat(cond_img, repeats)) + + self.pixel_values = [] + self.cond_pixel_values = [] + self.cond_pil_images = [] + self.buckets = [] + bucket_to_idx = {} + for image, cond_image in zip(self.instance_images, self.cond_images): + image = exif_transpose(image) + cond_image = exif_transpose(cond_image) + # RGBA, not RGB: this VAE takes and returns four channels. + if not image.mode == "RGBA": + image = image.convert("RGBA") + if not cond_image.mode == "RGBA": + cond_image = cond_image.convert("RGBA") + + width, height = image.size + + # Assign the image to a bucket. + target = self._bucket_for_image(height, width) + if target not in bucket_to_idx: + bucket_to_idx[target] = len(self.buckets) + self.buckets.append(target) + bucket_idx = bucket_to_idx[target] + + # based on the bucket assignment, define the transformations + image, cond_tensor, cond_pil = self.train_transform_pair( + image, + cond_image, + size=target, + center_crop=args.center_crop, + random_flip=args.random_flip, + ) + self.pixel_values.append((image, bucket_idx)) + self.cond_pixel_values.append((cond_tensor, bucket_idx)) + # The vision-language model reads the condition image as pixels, so the cropped frame is kept as well. + self.cond_pil_images.append(cond_pil) + + self.num_instance_images = len(self.instance_images) + self._length = self.num_instance_images + + if class_data_root is not None: + self.class_data_root = Path(class_data_root) + self.class_data_root.mkdir(parents=True, exist_ok=True) + self.class_images_path = list(self.class_data_root.iterdir()) + if class_num is not None: + self.num_class_images = min(len(self.class_images_path), class_num) + else: + self.num_class_images = len(self.class_images_path) + self._length = max(self.num_class_images, self.num_instance_images) + else: + self.class_data_root = None + + def __len__(self): + return self._length + + def __getitem__(self, index): + example = {} + instance_image, bucket_idx = self.pixel_values[index % self.num_instance_images] + example["index"] = index + example["instance_images"] = instance_image + example["cond_images"] = self.cond_pixel_values[index % self.num_instance_images][0] + example["cond_pil_images"] = self.cond_pil_images[index % self.num_instance_images] + example["bucket_idx"] = bucket_idx + if self.custom_instance_prompts: + caption = self.custom_instance_prompts[index % self.num_instance_images] + if caption: + example["instance_prompt"] = caption + else: + example["instance_prompt"] = self.instance_prompt + + else: # custom prompts were provided, but length does not match size of image dataset + example["instance_prompt"] = self.instance_prompt + + if self.class_data_root: + class_image = Image.open(self.class_images_path[index % self.num_class_images]) + class_image = exif_transpose(class_image) + + if not class_image.mode == "RGBA": + class_image = class_image.convert("RGBA") + # Match the class image to the paired instance image's bucket so they can be stacked into one batch. + example["class_images"] = self.train_transform( + class_image, size=self.buckets[bucket_idx], center_crop=self.center_crop + ) + example["class_prompt"] = self.class_prompt + + return example + + def _bucket_for_image(self, height, width): + # An explicit bucket list takes priority: pick the nearest, upscaling smaller images to cover it. + if self._explicit_buckets is not None: + return self._explicit_buckets[find_nearest_bucket(height, width, self._explicit_buckets)] + # On-the-fly bucketing: cap the ladder to the image's own resolution so smaller images are + # assigned to a smaller bucket rather than being upscaled (mirrors ostris' bucketing). + if self.use_aspect_ratio_buckets: + resolution = min(self.resolution, round((height * width) ** 0.5)) + ladder = generate_aspect_ratio_buckets( + resolution, + divisibility=self.bucket_divisibility, + base_resolutions=self.bucket_base_resolutions, + ) + return ladder[find_nearest_bucket(height, width, ladder)] + # No bucketing: a single square bucket reproduces the fixed-size resize + crop. + return (self.resolution, self.resolution) + + def train_transform_pair(self, image, cond_image, size, center_crop=False, random_flip=False): + """Resize and crop a target image and its condition image through one geometry, so an aligned pair stays + aligned. Returns both tensors and the cropped condition image, which the prompt encoder needs.""" + target_height, target_width = size + width, height = image.size + scale = max(target_height / height, target_width / width) + new_height, new_width = round(height * scale), round(width * scale) + resize = lambda img: TF.resize( # noqa: E731 + img, [new_height, new_width], interpolation=transforms.InterpolationMode.BILINEAR + ) + image, cond_image = resize(image), resize(cond_image) + if center_crop: + image, cond_image = TF.center_crop(image, size), TF.center_crop(cond_image, size) + else: + i, j, h, w = transforms.RandomCrop.get_params(image, output_size=size) + image, cond_image = TF.crop(image, i, j, h, w), TF.crop(cond_image, i, j, h, w) + if random_flip and random.random() < 0.5: + image, cond_image = TF.hflip(image), TF.hflip(cond_image) + return ( + TF.normalize(TF.to_tensor(image), [0.5], [0.5]), + TF.normalize(TF.to_tensor(cond_image), [0.5], [0.5]), + cond_image, + ) + + def train_transform(self, image, size, center_crop=False, random_flip=False): + # Resize preserving aspect ratio so the image covers the bucket, then crop to the bucket size. + target_height, target_width = size + width, height = image.size + scale = max(target_height / height, target_width / width) + new_height, new_width = round(height * scale), round(width * scale) + image = TF.resize(image, [new_height, new_width], interpolation=transforms.InterpolationMode.BILINEAR) + if center_crop: + image = TF.center_crop(image, size) + else: + i, j, h, w = transforms.RandomCrop.get_params(image, output_size=size) + image = TF.crop(image, i, j, h, w) + if random_flip and random.random() < 0.5: + image = TF.hflip(image) + return TF.normalize(TF.to_tensor(image), [0.5], [0.5]) + + +def collate_fn(examples, with_prior_preservation=False): + indices = [example["index"] for example in examples] + pixel_values = [example["instance_images"] for example in examples] + # Keep instance_prompts unchanged for prompt cache precompute; prompts may be extended with class prompts below. + instance_prompts = [example["instance_prompt"] for example in examples] + prompts = [example["instance_prompt"] for example in examples] + + # Concat class and instance examples for prior preservation. + # We do this to avoid doing two forward passes. + if with_prior_preservation: + pixel_values += [example["class_images"] for example in examples] + prompts += [example["class_prompt"] for example in examples] + + pixel_values = torch.stack(pixel_values) + # Qwen expects a `num_frames` dimension too. + if pixel_values.ndim == 4: + pixel_values = pixel_values.unsqueeze(2) + pixel_values = pixel_values.to(memory_format=torch.contiguous_format).float() + + cond_pixel_values = torch.stack([example["cond_images"] for example in examples]) + if cond_pixel_values.ndim == 4: + cond_pixel_values = cond_pixel_values.unsqueeze(2) + cond_pixel_values = cond_pixel_values.to(memory_format=torch.contiguous_format).float() + + batch = { + "indices": indices, + "pixel_values": pixel_values, + "cond_pixel_values": cond_pixel_values, + # Images rather than tensors: the prompt encoder feeds the condition image to the vision-language model. + "cond_pil_images": [example["cond_pil_images"] for example in examples], + "instance_prompts": instance_prompts, + "prompts": prompts, + } + return batch + + +class BucketBatchSampler(BatchSampler): + def __init__(self, dataset: DreamBoothDataset, batch_size: int, drop_last: bool = False, seed: int = None): + if not isinstance(batch_size, int) or batch_size <= 0: + raise ValueError("batch_size should be a positive integer value, but got batch_size={}".format(batch_size)) + if not isinstance(drop_last, bool): + raise ValueError("drop_last should be a boolean value, but got drop_last={}".format(drop_last)) + + self.dataset = dataset + self.batch_size = batch_size + self.drop_last = drop_last + self.generator = random.Random(seed) if seed is not None else random + + # Group indices by bucket + self.bucket_indices = [[] for _ in range(len(self.dataset.buckets))] + for idx, (_, bucket_idx) in enumerate(self.dataset.pixel_values): + self.bucket_indices[bucket_idx].append(idx) + + self.sampler_len = 0 + for indices_in_bucket in self.bucket_indices: + num_batches, remainder = divmod(len(indices_in_bucket), self.batch_size) + self.sampler_len += num_batches + if remainder > 0 and not self.drop_last: + self.sampler_len += 1 + + def __iter__(self): + batches = [] + for indices_in_bucket in self.bucket_indices: + shuffled_indices = indices_in_bucket.copy() + self.generator.shuffle(shuffled_indices) + for i in range(0, len(shuffled_indices), self.batch_size): + batch = shuffled_indices[i : i + self.batch_size] + if len(batch) < self.batch_size and self.drop_last: + continue + batches.append(batch) + + self.generator.shuffle(batches) + for batch in batches: + yield batch + + def __len__(self): + return self.sampler_len + + +class PromptDataset(Dataset): + "A simple dataset to prepare the prompts to generate class images on multiple GPUs." + + def __init__(self, prompt, num_samples): + self.prompt = prompt + self.num_samples = num_samples + + def __len__(self): + return self.num_samples + + def __getitem__(self, index): + example = {} + example["prompt"] = self.prompt + example["index"] = index + return example + + +# These helpers only matter for prior preservation, where instance and class prompt +# embedding batches are concatenated and may not share the same mask/sequence length. +def _materialize_prompt_embedding_mask( + prompt_embeds: torch.Tensor, prompt_embeds_mask: torch.Tensor | None +) -> torch.Tensor: + """Return a dense mask tensor for a prompt embedding batch.""" + batch_size, seq_len = prompt_embeds.shape[:2] + + if prompt_embeds_mask is None: + return torch.ones((batch_size, seq_len), dtype=torch.long, device=prompt_embeds.device) + + if prompt_embeds_mask.shape != (batch_size, seq_len): + raise ValueError( + f"`prompt_embeds_mask` shape {prompt_embeds_mask.shape} must match prompt embeddings shape " + f"({batch_size}, {seq_len})." + ) + + return prompt_embeds_mask.to(device=prompt_embeds.device) + + +def _pad_prompt_embedding_pair( + prompt_embeds: torch.Tensor, prompt_embeds_mask: torch.Tensor | None, target_seq_len: int +) -> tuple[torch.Tensor, torch.Tensor]: + """Pad one prompt embedding batch and its mask to a shared sequence length.""" + prompt_embeds_mask = _materialize_prompt_embedding_mask(prompt_embeds, prompt_embeds_mask) + pad_width = target_seq_len - prompt_embeds.shape[1] + + if pad_width <= 0: + return prompt_embeds, prompt_embeds_mask + + prompt_embeds = torch.cat( + [prompt_embeds, prompt_embeds.new_zeros(prompt_embeds.shape[0], pad_width, prompt_embeds.shape[2])], dim=1 + ) + prompt_embeds_mask = torch.cat( + [prompt_embeds_mask, prompt_embeds_mask.new_zeros(prompt_embeds_mask.shape[0], pad_width)], dim=1 + ) + + return prompt_embeds, prompt_embeds_mask + + +def concat_prompt_embedding_batches( + *prompt_embedding_pairs: tuple[torch.Tensor, torch.Tensor | None], +) -> tuple[torch.Tensor, torch.Tensor | None]: + """Concatenate prompt embedding batches while handling missing masks and length mismatches.""" + if not prompt_embedding_pairs: + raise ValueError("At least one prompt embedding pair must be provided.") + + target_seq_len = max(prompt_embeds.shape[1] for prompt_embeds, _ in prompt_embedding_pairs) + padded_pairs = [ + _pad_prompt_embedding_pair(prompt_embeds, prompt_embeds_mask, target_seq_len) + for prompt_embeds, prompt_embeds_mask in prompt_embedding_pairs + ] + + merged_prompt_embeds = torch.cat([prompt_embeds for prompt_embeds, _ in padded_pairs], dim=0) + merged_mask = torch.cat([prompt_embeds_mask for _, prompt_embeds_mask in padded_pairs], dim=0) + + if merged_mask.all(): + return merged_prompt_embeds, None + + return merged_prompt_embeds, merged_mask + + +def main(args): + if args.report_to == "wandb" and args.hub_token is not None: + raise ValueError( + "You cannot use both --report_to=wandb and --hub_token due to a security risk of exposing your token." + " Please use `hf auth login` to authenticate with the Hub." + ) + + if torch.backends.mps.is_available() and args.mixed_precision == "bf16": + # due to pytorch#99272, MPS does not yet support bfloat16. + raise ValueError( + "Mixed precision training with bfloat16 is not supported on MPS. Please use fp16 (recommended) or fp32 instead." + ) + + logging_dir = Path(args.output_dir, args.logging_dir) + + accelerator_project_config = ProjectConfiguration(project_dir=args.output_dir, logging_dir=logging_dir) + kwargs = DistributedDataParallelKwargs(find_unused_parameters=True) + accelerator = Accelerator( + gradient_accumulation_steps=args.gradient_accumulation_steps, + mixed_precision=args.mixed_precision, + log_with=args.report_to, + project_config=accelerator_project_config, + kwargs_handlers=[kwargs], + ) + + # Disable AMP for MPS. + if torch.backends.mps.is_available(): + accelerator.native_amp = False + + if args.report_to == "wandb": + if not is_wandb_available(): + raise ImportError("Make sure to install wandb if you want to use it for logging during training.") + + # Make one log on every process with the configuration for debugging. + logging.basicConfig( + format="%(asctime)s - %(levelname)s - %(name)s - %(message)s", + datefmt="%m/%d/%Y %H:%M:%S", + level=logging.INFO, + ) + logger.info(accelerator.state, main_process_only=False) + if accelerator.is_local_main_process: + transformers.utils.logging.set_verbosity_warning() + diffusers.utils.logging.set_verbosity_info() + else: + transformers.utils.logging.set_verbosity_error() + diffusers.utils.logging.set_verbosity_error() + + # If passed along, set the training seed now. + if args.seed is not None: + set_seed(args.seed) + + # Generate class images if prior preservation is enabled. + if args.with_prior_preservation: + class_images_dir = Path(args.class_data_dir) + if not class_images_dir.exists(): + class_images_dir.mkdir(parents=True) + cur_class_images = len(list(class_images_dir.iterdir())) + + if cur_class_images < args.num_class_images: + pipeline = QwenImage21Pipeline.from_pretrained( + args.pretrained_model_name_or_path, + torch_dtype=torch.bfloat16 if args.mixed_precision == "bf16" else torch.float16, + revision=args.revision, + variant=args.variant, + ) + pipeline.set_progress_bar_config(disable=True) + + num_new_images = args.num_class_images - cur_class_images + logger.info(f"Number of class images to sample: {num_new_images}.") + + sample_dataset = PromptDataset(args.class_prompt, num_new_images) + sample_dataloader = torch.utils.data.DataLoader(sample_dataset, batch_size=args.sample_batch_size) + + sample_dataloader = accelerator.prepare(sample_dataloader) + pipeline.to(accelerator.device) + + for example in tqdm( + sample_dataloader, desc="Generating class images", disable=not accelerator.is_local_main_process + ): + images = pipeline(example["prompt"]).images + + for i, image in enumerate(images): + hash_image = insecure_hashlib.sha1(image.tobytes()).hexdigest() + # PNG, not JPEG: the pipeline returns RGBA and JPEG has no alpha channel. + image_filename = class_images_dir / f"{example['index'][i] + cur_class_images}-{hash_image}.png" + image.save(image_filename) + + pipeline.to("cpu") + del pipeline + free_memory() + + # Handle the repository creation + if accelerator.is_main_process: + if args.output_dir is not None: + os.makedirs(args.output_dir, exist_ok=True) + + if args.push_to_hub: + repo_id = create_repo( + repo_id=args.hub_model_id or Path(args.output_dir).name, + exist_ok=True, + ).repo_id + + # A Qwen3-VL processor rather than a tokenizer: it also expands condition images into vision tokens. + processor = Qwen3VLProcessor.from_pretrained( + args.pretrained_model_name_or_path, + subfolder="processor", + revision=args.revision, + ) + + # For mixed precision training we cast all non-trainable weights (vae, text_encoder and transformer) to half-precision + # as these weights are only used for inference, keeping weights in full precision is not required. + weight_dtype = torch.float32 + if accelerator.mixed_precision == "fp16": + weight_dtype = torch.float16 + elif accelerator.mixed_precision == "bf16": + weight_dtype = torch.bfloat16 + + # Load scheduler and models + # Used as shipped: it sets `use_dynamic_shifting`, so the training sigmas stay unshifted. + noise_scheduler = FlowMatchEulerDiscreteScheduler.from_pretrained( + args.pretrained_model_name_or_path, subfolder="scheduler", revision=args.revision + ) + noise_scheduler_copy = copy.deepcopy(noise_scheduler) + vae = AutoencoderKLQwenImage21.from_pretrained( + args.pretrained_model_name_or_path, + subfolder="vae", + revision=args.revision, + variant=args.variant, + ) + # 16 here, so one latent token covers a 16x16 pixel tile. + vae_scale_factor = 2 ** len(vae.temperal_downsample) + if vae_scale_factor * 2 != SIZE_MULTIPLE_OF: + raise ValueError( + f"This checkpoint's VAE scales by {vae_scale_factor}, so sizes must be multiples of " + f"{vae_scale_factor * 2}, but the resolution checks used {SIZE_MULTIPLE_OF}." + ) + latents_mean = (torch.tensor(vae.config.latents_mean).view(1, vae.config.z_dim, 1, 1, 1)).to(accelerator.device) + latents_std = 1.0 / torch.tensor(vae.config.latents_std).view(1, vae.config.z_dim, 1, 1, 1).to(accelerator.device) + text_encoder = Qwen3VLForConditionalGeneration.from_pretrained( + args.pretrained_model_name_or_path, subfolder="text_encoder", revision=args.revision, torch_dtype=weight_dtype + ) + quantization_config = None + if args.bnb_quantization_config_path is not None: + with open(args.bnb_quantization_config_path, "r") as f: + config_kwargs = json.load(f) + if "load_in_4bit" in config_kwargs and config_kwargs["load_in_4bit"]: + config_kwargs["bnb_4bit_compute_dtype"] = weight_dtype + quantization_config = BitsAndBytesConfig(**config_kwargs) + + transformer = QwenImage21Transformer2DModel.from_pretrained( + args.pretrained_model_name_or_path, + subfolder="transformer", + revision=args.revision, + variant=args.variant, + quantization_config=quantization_config, + torch_dtype=weight_dtype, + ) + if args.bnb_quantization_config_path is not None: + transformer = prepare_model_for_kbit_training(transformer, use_gradient_checkpointing=False) + + # We only train the additional adapter LoRA layers + transformer.requires_grad_(False) + vae.requires_grad_(False) + text_encoder.requires_grad_(False) + + if torch.backends.mps.is_available() and weight_dtype == torch.bfloat16: + # due to pytorch#99272, MPS does not yet support bfloat16. + raise ValueError( + "Mixed precision training with bfloat16 is not supported on MPS. Please use fp16 (recommended) or fp32 instead." + ) + + to_kwargs = {"dtype": weight_dtype, "device": accelerator.device} if not args.offload else {"dtype": weight_dtype} + # The VAE is stable in bf16, so load it in weight_dtype to reduce memory. + vae.to(**to_kwargs) + text_encoder.to(**to_kwargs) + # we never offload the transformer to CPU, so we can just use the accelerator device + transformer_to_kwargs = ( + {"device": accelerator.device} + if args.bnb_quantization_config_path is not None + else {"device": accelerator.device, "dtype": weight_dtype} + ) + transformer.to(**transformer_to_kwargs) + + # Initialize a text encoding pipeline and keep it to CPU for now. + text_encoding_pipeline = QwenImage21Pipeline.from_pretrained( + args.pretrained_model_name_or_path, + vae=None, + transformer=None, + processor=processor, + text_encoder=text_encoder, + scheduler=None, + ) + + if args.gradient_checkpointing: + transformer.enable_gradient_checkpointing() + + if args.lora_layers is not None: + target_modules = [layer.strip() for layer in args.lora_layers.split(",")] + else: + target_modules = ["to_k", "to_q", "to_v", "to_out.0"] + + # now we will add new LoRA weights the transformer layers + transformer_lora_config = LoraConfig( + r=args.rank, + lora_alpha=args.lora_alpha, + lora_dropout=args.lora_dropout, + init_lora_weights="gaussian", + target_modules=target_modules, + ) + transformer.add_adapter(transformer_lora_config) + + def unwrap_model(model): + model = accelerator.unwrap_model(model) + model = model._orig_mod if is_compiled_module(model) else model + return model + + # create custom saving & loading hooks so that `accelerator.save_state(...)` serializes in a nice format + def save_model_hook(models, weights, output_dir): + if accelerator.is_main_process: + transformer_lora_layers_to_save = None + modules_to_save = {} + + for model in models: + if isinstance(unwrap_model(model), type(unwrap_model(transformer))): + model = unwrap_model(model) + transformer_lora_layers_to_save = get_peft_model_state_dict(model) + modules_to_save["transformer"] = model + else: + raise ValueError(f"unexpected save model: {model.__class__}") + + # make sure to pop weight so that corresponding model is not saved again + if weights: + weights.pop() + + QwenImage21Pipeline.save_lora_weights( + output_dir, + transformer_lora_layers=transformer_lora_layers_to_save, + **_collate_lora_metadata(modules_to_save), + ) + + def load_model_hook(models, input_dir): + transformer_ = None + + if not accelerator.distributed_type == DistributedType.DEEPSPEED: + while len(models) > 0: + model = models.pop() + + if isinstance(unwrap_model(model), type(unwrap_model(transformer))): + model = unwrap_model(model) + transformer_ = model + else: + raise ValueError(f"unexpected save model: {model.__class__}") + else: + transformer_ = QwenImage21Transformer2DModel.from_pretrained( + args.pretrained_model_name_or_path, subfolder="transformer" + ) + transformer_.add_adapter(transformer_lora_config) + + lora_state_dict = QwenImage21Pipeline.lora_state_dict(input_dir) + + transformer_state_dict = { + f"{k.replace('transformer.', '')}": v for k, v in lora_state_dict.items() if k.startswith("transformer.") + } + transformer_state_dict = convert_unet_state_dict_to_peft(transformer_state_dict) + incompatible_keys = set_peft_model_state_dict(transformer_, transformer_state_dict, adapter_name="default") + if incompatible_keys is not None: + # check only for unexpected keys + unexpected_keys = getattr(incompatible_keys, "unexpected_keys", None) + if unexpected_keys: + logger.warning( + f"Loading adapter weights from state_dict led to unexpected keys not found in the model: " + f" {unexpected_keys}. " + ) + + # Make sure the trainable params are in float32. This is again needed since the base models + # are in `weight_dtype`. More details: + # https://github.com/huggingface/diffusers/pull/6514#discussion_r1449796804 + if args.mixed_precision == "fp16": + models = [transformer_] + # only upcast trainable parameters (LoRA) into fp32 + cast_training_params(models) + + accelerator.register_save_state_pre_hook(save_model_hook) + accelerator.register_load_state_pre_hook(load_model_hook) + + # Enable TF32 for faster training on Ampere GPUs, + # cf https://pytorch.org/docs/stable/notes/cuda.html#tensorfloat-32-tf32-on-ampere-devices + if args.allow_tf32 and torch.cuda.is_available(): + torch.backends.cuda.matmul.allow_tf32 = True + + if args.scale_lr: + args.learning_rate = ( + args.learning_rate * args.gradient_accumulation_steps * args.train_batch_size * accelerator.num_processes + ) + + # Make sure the trainable params are in float32. + if args.mixed_precision == "fp16": + models = [transformer] + # only upcast trainable parameters (LoRA) into fp32 + cast_training_params(models, dtype=torch.float32) + + transformer_lora_parameters = list(filter(lambda p: p.requires_grad, transformer.parameters())) + + # Optimization parameters + transformer_parameters_with_lr = {"params": transformer_lora_parameters, "lr": args.learning_rate} + params_to_optimize = [transformer_parameters_with_lr] + + # Optimizer creation + if not (args.optimizer.lower() == "prodigy" or args.optimizer.lower() == "adamw"): + logger.warning( + f"Unsupported choice of optimizer: {args.optimizer}.Supported optimizers include [adamW, prodigy]." + "Defaulting to adamW" + ) + args.optimizer = "adamw" + + if args.use_8bit_adam and not args.optimizer.lower() == "adamw": + logger.warning( + f"use_8bit_adam is ignored when optimizer is not set to 'AdamW'. Optimizer was " + f"set to {args.optimizer.lower()}" + ) + + if args.optimizer.lower() == "adamw": + if args.use_8bit_adam: + try: + import bitsandbytes as bnb + except ImportError: + raise ImportError( + "To use 8-bit Adam, please install the bitsandbytes library: `pip install bitsandbytes`." + ) + + optimizer_class = bnb.optim.AdamW8bit + else: + optimizer_class = torch.optim.AdamW + + optimizer = optimizer_class( + params_to_optimize, + betas=(args.adam_beta1, args.adam_beta2), + weight_decay=args.adam_weight_decay, + eps=args.adam_epsilon, + ) + + if args.optimizer.lower() == "prodigy": + try: + import prodigyopt + except ImportError: + raise ImportError("To use Prodigy, please install the prodigyopt library: `pip install prodigyopt`") + + optimizer_class = prodigyopt.Prodigy + + if args.learning_rate <= 0.1: + logger.warning( + "Learning rate is too low. When using prodigy, it's generally better to set learning rate around 1.0" + ) + + optimizer = optimizer_class( + params_to_optimize, + betas=(args.adam_beta1, args.adam_beta2), + beta3=args.prodigy_beta3, + weight_decay=args.adam_weight_decay, + eps=args.adam_epsilon, + decouple=args.prodigy_decouple, + use_bias_correction=args.prodigy_use_bias_correction, + safeguard_warmup=args.prodigy_safeguard_warmup, + ) + + # Resolve the bucketing mode. Bucketing must be enabled explicitly with --use_aspect_ratio_buckets; + # a bucket list without that flag is an error. With the flag, an explicit --aspect_ratio_buckets list + # drives assignment, otherwise buckets are computed on the fly inside the dataset. Without the flag a + # single square bucket reproduces the fixed-size resize + crop. + if args.aspect_ratio_buckets is not None and not args.use_aspect_ratio_buckets: + raise ValueError("--aspect_ratio_buckets requires --use_aspect_ratio_buckets to be set.") + if args.aspect_ratio_buckets is not None: + buckets = parse_buckets_string(args.aspect_ratio_buckets) + use_aspect_ratio_buckets = False + logger.info(f"Using explicit aspect ratio buckets: {buckets}") + elif args.use_aspect_ratio_buckets: + buckets = None + use_aspect_ratio_buckets = True + logger.info( + "No --aspect_ratio_buckets provided; auto-computing aspect ratio buckets on the fly from --resolution." + ) + else: + buckets = [(args.resolution, args.resolution)] + use_aspect_ratio_buckets = False + + # Dataset and DataLoaders creation: + train_dataset = DreamBoothDataset( + instance_data_root=args.instance_data_dir, + instance_prompt=args.instance_prompt, + class_prompt=args.class_prompt, + class_data_root=args.class_data_dir if args.with_prior_preservation else None, + class_num=args.num_class_images, + size=args.resolution, + repeats=args.repeats, + center_crop=args.center_crop, + buckets=buckets, + use_aspect_ratio_buckets=use_aspect_ratio_buckets, + ) + # Prompt embeddings depend on the sample's condition image, so they are always precomputed per sample. + precompute_latents = True + batch_sampler = BucketBatchSampler(train_dataset, batch_size=args.train_batch_size, drop_last=True, seed=args.seed) + train_dataloader = torch.utils.data.DataLoader( + train_dataset, + batch_sampler=batch_sampler, + collate_fn=lambda examples: collate_fn(examples, args.with_prior_preservation), + num_workers=args.dataloader_num_workers, + ) + + def compute_text_embeddings(prompt, text_encoding_pipeline, cond_image=None): + with torch.no_grad(): + # One sample at a time: a list of images applies to every prompt in the call, not pairwise. + prompt_embeds, prompt_embeds_mask, image_pad_mask = text_encoding_pipeline.encode_prompt( + prompt=prompt, image=None if cond_image is None else [cond_image] + ) + return prompt_embeds, prompt_embeds_mask, image_pad_mask + + # Handle class prompt for prior-preservation. + if args.with_prior_preservation: + with offload_models(text_encoding_pipeline, device=accelerator.device, offload=args.offload): + class_prompt_embeds, class_prompt_embeds_mask, _ = compute_text_embeddings( + args.class_prompt, text_encoding_pipeline + ) + + # When caption dropout is enabled, we precompute the empty ("") prompt embedding once and swap it in + # for randomly selected instance samples at training time (see the training loop below). + if args.caption_dropout > 0: + with offload_models(text_encoding_pipeline, device=accelerator.device, offload=args.offload): + empty_prompt_embeds, empty_prompt_embeds_mask, _ = compute_text_embeddings("", text_encoding_pipeline) + + validation_pipeline_args = {} + validation_image_pad_mask = None + # The final pass runs on `--final_validation_prompt` when `--validation_prompt` is absent, so the embeddings + # have to be built for whichever one is set - the text encoder is freed before that pass reaches the pipeline. + effective_validation_prompt = args.validation_prompt or args.final_validation_prompt + if effective_validation_prompt is not None: + if args.validation_image is None: + raise ValueError( + "A validation prompt needs `--validation_image`, the image the edit is applied to. Pass it " + "alongside `--validation_prompt` or `--final_validation_prompt`." + ) + validation_image = load_image(args.validation_image) + # Encoded at the size the pipeline will resize to, so vision tokens and latents line up. + width, height, _ = calculate_dimensions( + args.resolution * args.resolution, validation_image.size[0] / validation_image.size[1] + ) + resized_validation_image = validation_image.resize((width, height)) + with offload_models(text_encoding_pipeline, device=accelerator.device, offload=args.offload): + embeds, embeds_mask, image_pad_mask = compute_text_embeddings( + effective_validation_prompt, text_encoding_pipeline, resized_validation_image + ) + validation_pipeline_args = { + "prompt_embeds": embeds, + "prompt_embeds_mask": embeds_mask, + "image": validation_image, + } + validation_image_pad_mask = image_pad_mask + + # if cache_latents is set to True, we encode images to latents and store them. + # Similar to pre-encoding in the case of a single instance prompt, if custom prompts are provided + # we encode them in advance as well. Caches are keyed by dataset index so they stay correct under + # aspect-ratio bucketing, where the batch composition differs between the caching pass and training. + if args.cache_latents: + instance_latents_cache = [None] * train_dataset.num_instance_images + cond_latents_cache = [None] * train_dataset.num_instance_images + prompt_embeds_cache = [None] * train_dataset.num_instance_images + prompt_embeds_mask_cache = [None] * train_dataset.num_instance_images + image_pad_mask_cache = [None] * train_dataset.num_instance_images + if precompute_latents: + cache_batch_sampler = BucketBatchSampler( + train_dataset, batch_size=args.train_batch_size, drop_last=False, seed=args.seed + ) + cache_dataloader = torch.utils.data.DataLoader( + train_dataset, + batch_sampler=cache_batch_sampler, + collate_fn=lambda examples: collate_fn(examples, args.with_prior_preservation), + num_workers=args.dataloader_num_workers, + ) + for batch in tqdm(cache_dataloader, desc="Caching latents"): + with torch.no_grad(): + sample_indices = batch["indices"] + if args.cache_latents: + with offload_models(vae, device=accelerator.device, offload=args.offload): + batch["pixel_values"] = batch["pixel_values"].to( + accelerator.device, non_blocking=True, dtype=vae.dtype + ) + instance_latents = vae.encode(batch["pixel_values"]).latent_dist.sample() + # Taken at the mode, not sampled: the condition image is not noised. + cond_latents = vae.encode( + batch["cond_pixel_values"].to(accelerator.device, non_blocking=True, dtype=vae.dtype) + ).latent_dist.mode() + for i, idx in enumerate(sample_indices): + instance_latents_cache[idx] = instance_latents[i : i + 1] + cond_latents_cache[idx] = cond_latents[i : i + 1] + with offload_models(text_encoding_pipeline, device=accelerator.device, offload=args.offload): + for i, idx in enumerate(sample_indices): + prompt_embeds, prompt_embeds_mask, image_pad_mask = compute_text_embeddings( + batch["instance_prompts"][i], text_encoding_pipeline, batch["cond_pil_images"][i] + ) + prompt_embeds_cache[idx] = prompt_embeds + prompt_embeds_mask_cache[idx] = prompt_embeds_mask + image_pad_mask_cache[idx] = image_pad_mask + + if args.cache_latents: + assert all(latents is not None for latents in instance_latents_cache), "Latent cache has unfilled entries." + assert all(latents is not None for latents in cond_latents_cache), ( + "Condition latent cache has unfilled entries." + ) + assert all(embeds is not None for embeds in prompt_embeds_cache), ( + "Prompt embedding cache has unfilled entries." + ) + + # move back to cpu before deleting to ensure memory is freed see: https://github.com/huggingface/diffusers/issues/11376#issue-3008144624 + if args.cache_latents: + vae = vae.to("cpu") + del vae + + # move back to cpu before deleting to ensure memory is freed see: https://github.com/huggingface/diffusers/issues/11376#issue-3008144624 + text_encoding_pipeline = text_encoding_pipeline.to("cpu") + # The processor stays: it holds no weights, and the pipeline cannot be constructed without one. + del text_encoder + free_memory() + + # Scheduler and math around the number of training steps. + overrode_max_train_steps = False + num_update_steps_per_epoch = math.ceil(len(train_dataloader) / args.gradient_accumulation_steps) + if args.max_train_steps is None: + args.max_train_steps = args.num_train_epochs * num_update_steps_per_epoch + overrode_max_train_steps = True + + lr_scheduler = get_scheduler( + args.lr_scheduler, + optimizer=optimizer, + num_warmup_steps=args.lr_warmup_steps * accelerator.num_processes, + num_training_steps=args.max_train_steps * accelerator.num_processes, + num_cycles=args.lr_num_cycles, + power=args.lr_power, + ) + + # Prepare everything with our `accelerator`. + transformer, optimizer, train_dataloader, lr_scheduler = accelerator.prepare( + transformer, optimizer, train_dataloader, lr_scheduler + ) + + # We need to recalculate our total training steps as the size of the training dataloader may have changed. + num_update_steps_per_epoch = math.ceil(len(train_dataloader) / args.gradient_accumulation_steps) + if overrode_max_train_steps: + args.max_train_steps = args.num_train_epochs * num_update_steps_per_epoch + # Afterwards we recalculate our number of training epochs + args.num_train_epochs = math.ceil(args.max_train_steps / num_update_steps_per_epoch) + + # We need to initialize the trackers we use, and also store our configuration. + # The trackers initializes automatically on the main process. + if accelerator.is_main_process: + tracker_name = "dreambooth-qwen-image-lora" + accelerator.init_trackers(tracker_name, config=vars(args)) + + # Train! + total_batch_size = args.train_batch_size * accelerator.num_processes * args.gradient_accumulation_steps + + logger.info("***** Running training *****") + logger.info(f" Num examples = {len(train_dataset)}") + logger.info(f" Num batches each epoch = {len(train_dataloader)}") + logger.info(f" Num Epochs = {args.num_train_epochs}") + logger.info(f" Instantaneous batch size per device = {args.train_batch_size}") + logger.info(f" Total train batch size (w. parallel, distributed & accumulation) = {total_batch_size}") + logger.info(f" Gradient Accumulation steps = {args.gradient_accumulation_steps}") + logger.info(f" Total optimization steps = {args.max_train_steps}") + global_step = 0 + first_epoch = 0 + + # Potentially load in the weights and states from a previous save + if args.resume_from_checkpoint: + if args.resume_from_checkpoint != "latest": + path = os.path.basename(args.resume_from_checkpoint) + else: + # Get the mos recent checkpoint + dirs = os.listdir(args.output_dir) + dirs = [d for d in dirs if d.startswith("checkpoint")] + dirs = sorted(dirs, key=lambda x: int(x.split("-")[1])) + path = dirs[-1] if len(dirs) > 0 else None + + if path is None: + accelerator.print( + f"Checkpoint '{args.resume_from_checkpoint}' does not exist. Starting a new training run." + ) + args.resume_from_checkpoint = None + initial_global_step = 0 + else: + accelerator.print(f"Resuming from checkpoint {path}") + accelerator.load_state(os.path.join(args.output_dir, path)) + global_step = int(path.split("-")[1]) + + initial_global_step = global_step + first_epoch = global_step // num_update_steps_per_epoch + + else: + initial_global_step = 0 + + progress_bar = tqdm( + range(0, args.max_train_steps), + initial=initial_global_step, + desc="Steps", + # Only show the progress bar once on each machine. + disable=not accelerator.is_local_main_process, + ) + + def get_sigmas(timesteps, n_dim=4, dtype=torch.float32): + sigmas = noise_scheduler_copy.sigmas.to(device=accelerator.device, dtype=dtype) + schedule_timesteps = noise_scheduler_copy.timesteps.to(accelerator.device) + timesteps = timesteps.to(accelerator.device) + step_indices = [(schedule_timesteps == t).nonzero().item() for t in timesteps] + + sigma = sigmas[step_indices].flatten() + while len(sigma.shape) < n_dim: + sigma = sigma.unsqueeze(-1) + return sigma + + for epoch in range(first_epoch, args.num_train_epochs): + transformer.train() + + for batch in train_dataloader: + models_to_accumulate = [transformer] + sample_indices = batch["indices"] + + with accelerator.accumulate(models_to_accumulate): + # Each sample's embeddings were encoded with its own condition image, gathered by dataset index. + prompt_pairs = [(prompt_embeds_cache[idx], prompt_embeds_mask_cache[idx]) for idx in sample_indices] + prompt_embeds, prompt_embeds_mask = concat_prompt_embedding_batches(*prompt_pairs) + + # The transformer reads one image-pad layout for the whole batch, so the samples have to agree. + batch_image_pad_masks = [image_pad_mask_cache[idx] for idx in sample_indices] + if any(not torch.equal(mask, batch_image_pad_masks[0]) for mask in batch_image_pad_masks[1:]): + raise ValueError( + "The samples in this batch place the condition image's tokens differently, which happens " + "when their prompts differ in length. Train with `--train_batch_size 1`, or give the " + "samples in a batch the same prompt." + ) + image_pad_mask = batch_image_pad_masks[0].repeat(len(sample_indices), 1).to(accelerator.device) + + # Convert images to latent space + if args.cache_latents: + model_input = torch.cat([instance_latents_cache[idx] for idx in sample_indices], dim=0) + cond_model_input = torch.cat([cond_latents_cache[idx] for idx in sample_indices], dim=0) + else: + with offload_models(vae, device=accelerator.device, offload=args.offload): + pixel_values = batch["pixel_values"].to(dtype=vae.dtype) + cond_pixel_values = batch["cond_pixel_values"].to(device=accelerator.device, dtype=vae.dtype) + model_input = vae.encode(pixel_values).latent_dist.sample() + cond_model_input = vae.encode(cond_pixel_values).latent_dist.mode() + + model_input = (model_input - latents_mean) * latents_std + model_input = model_input.to(dtype=weight_dtype) + # Clean, at the same normalisation as the target. + cond_model_input = ((cond_model_input - latents_mean) * latents_std).to(dtype=weight_dtype) + + # Sample noise that we'll add to the latents + noise = torch.randn_like(model_input) + bsz = model_input.shape[0] + + # Sample a random timestep for each image + # for weighting schemes where we sample timesteps non-uniformly + u = compute_density_for_timestep_sampling( + weighting_scheme=args.weighting_scheme, + batch_size=bsz, + logit_mean=args.logit_mean, + logit_std=args.logit_std, + mode_scale=args.mode_scale, + ) + indices = (u * noise_scheduler_copy.config.num_train_timesteps).long() + timesteps = noise_scheduler_copy.timesteps[indices].to(device=model_input.device) + + # Add noise according to flow matching. + # zt = (1 - texp) * x + texp * z1 + sigmas = get_sigmas(timesteps, n_dim=model_input.ndim, dtype=model_input.dtype) + noisy_model_input = (1.0 - sigmas) * model_input + sigmas * noise + + # Predict the noise residual. A batch is single-bucket, so the latent height/width are shared + # across the batch; derive them from the latents to support aspect-ratio buckets. + latent_height, latent_width = model_input.shape[3], model_input.shape[4] + cond_latent_height, cond_latent_width = cond_model_input.shape[3], cond_model_input.shape[4] + # One shape per image in the sequence, ordered [condition, target]. + img_shapes = [[(1, cond_latent_height, cond_latent_width), (1, latent_height, latent_width)]] * bsz + # Latents are consumed unpatched, so packing is a plain spatial flatten. + packed_noisy_model_input = QwenImage21Pipeline._pack_latents( + noisy_model_input, + batch_size=model_input.shape[0], + num_channels_latents=model_input.shape[1], + height=latent_height, + width=latent_width, + ) + packed_cond_model_input = QwenImage21Pipeline._pack_latents( + cond_model_input, + batch_size=cond_model_input.shape[0], + num_channels_latents=cond_model_input.shape[1], + height=cond_latent_height, + width=cond_latent_width, + ) + # Condition latents lead, noisy target follows, the order `img_shapes` declares. + packed_input = torch.cat([packed_cond_model_input, packed_noisy_model_input], dim=1) + # `img_mask` marks which positions over [prompt tokens, target slots] stand for image latents, one + # slot per 2x2 group. The prompt half comes from the encoder, where the vision tokens sit. + target_slots = (latent_height * latent_width) // 4 + cond_slots = (cond_latent_height * cond_latent_width) // 4 + if int(image_pad_mask[0].sum()) != cond_slots: + raise ValueError( + f"The prompt carries {int(image_pad_mask[0].sum())} condition-image slots but the condition " + f"latents need {cond_slots}. The vision-language processor resizes images below its minimum " + f"pixel count, so a condition image this small ({cond_latent_height * 16}x" + f"{cond_latent_width * 16}) does not line up. Train at a larger resolution." + ) + img_mask = torch.cat( + [ + image_pad_mask, + torch.ones(bsz, target_slots, dtype=image_pad_mask.dtype, device=accelerator.device), + ], + dim=1, + ) + model_pred = transformer( + hidden_states=packed_input, + encoder_hidden_states=prompt_embeds, + encoder_hidden_states_mask=prompt_embeds_mask, + timestep=timesteps / 1000, + img_shapes=img_shapes, + img_mask=img_mask, + return_dict=False, + )[0] + # The prediction spans the joint sequence, so keep the target's tail, as `__call__` does. + model_pred = model_pred[:, -packed_noisy_model_input.shape[1] :] + model_pred = QwenImage21Pipeline._unpack_latents( + model_pred, latent_height * vae_scale_factor, latent_width * vae_scale_factor, vae_scale_factor + ) + + # these weighting schemes use a uniform timestep sampling + # and instead post-weight the loss + weighting = compute_loss_weighting_for_sd3(weighting_scheme=args.weighting_scheme, sigmas=sigmas) + + target = noise - model_input + if args.with_prior_preservation: + # Chunk the noise and model_pred into two parts and compute the loss on each part separately. + model_pred, model_pred_prior = torch.chunk(model_pred, 2, dim=0) + target, target_prior = torch.chunk(target, 2, dim=0) + weighting, weighting_prior = torch.chunk(weighting, 2, dim=0) + + # Compute prior loss + prior_loss = torch.mean( + (weighting_prior.float() * (model_pred_prior.float() - target_prior.float()) ** 2).reshape( + target_prior.shape[0], -1 + ), + 1, + ) + prior_loss = prior_loss.mean() + + # Compute regular loss. + loss = torch.mean( + (weighting.float() * (model_pred.float() - target.float()) ** 2).reshape(target.shape[0], -1), + 1, + ) + loss = loss.mean() + + if args.with_prior_preservation: + # Add the prior loss to the instance loss. + loss = loss + args.prior_loss_weight * prior_loss + + accelerator.backward(loss) + if accelerator.sync_gradients: + params_to_clip = transformer.parameters() + accelerator.clip_grad_norm_(params_to_clip, args.max_grad_norm) + + optimizer.step() + lr_scheduler.step() + optimizer.zero_grad() + + # Checks if the accelerator has performed an optimization step behind the scenes + if accelerator.sync_gradients: + progress_bar.update(1) + global_step += 1 + + if accelerator.is_main_process or accelerator.distributed_type == DistributedType.DEEPSPEED: + if global_step % args.checkpointing_steps == 0: + # _before_ saving state, check if this save would set us over the `checkpoints_total_limit` + if args.checkpoints_total_limit is not None: + checkpoints = os.listdir(args.output_dir) + checkpoints = [d for d in checkpoints if d.startswith("checkpoint")] + checkpoints = sorted(checkpoints, key=lambda x: int(x.split("-")[1])) + + # before we save the new checkpoint, we need to have at _most_ `checkpoints_total_limit - 1` checkpoints + if len(checkpoints) >= args.checkpoints_total_limit: + num_to_remove = len(checkpoints) - args.checkpoints_total_limit + 1 + removing_checkpoints = checkpoints[0:num_to_remove] + + logger.info( + f"{len(checkpoints)} checkpoints already exist, removing {len(removing_checkpoints)} checkpoints" + ) + logger.info(f"removing checkpoints: {', '.join(removing_checkpoints)}") + + for removing_checkpoint in removing_checkpoints: + removing_checkpoint = os.path.join(args.output_dir, removing_checkpoint) + shutil.rmtree(removing_checkpoint) + + save_path = os.path.join(args.output_dir, f"checkpoint-{global_step}") + accelerator.save_state(save_path) + logger.info(f"Saved state to {save_path}") + + logs = {"loss": loss.detach().item(), "lr": lr_scheduler.get_last_lr()[0]} + progress_bar.set_postfix(**logs) + accelerator.log(logs, step=global_step) + + if global_step >= args.max_train_steps: + break + + if accelerator.is_main_process: + if args.validation_prompt is not None and epoch % args.validation_epochs == 0: + # create pipeline. The prompt is supplied as embeddings, so no text encoder is loaded. + pipeline = QwenImage21ValidationPipeline.from_pretrained( + args.pretrained_model_name_or_path, + text_encoder=None, + processor=processor, + transformer=accelerator.unwrap_model(transformer), + revision=args.revision, + variant=args.variant, + torch_dtype=weight_dtype, + ) + pipeline.cached_image_pad_mask = validation_image_pad_mask + images = log_validation( + pipeline=pipeline, + args=args, + accelerator=accelerator, + pipeline_args=validation_pipeline_args, + torch_dtype=weight_dtype, + epoch=epoch, + ) + del pipeline + images = None + free_memory() + + # Save the lora layers + accelerator.wait_for_everyone() + if accelerator.is_main_process: + modules_to_save = {} + transformer = unwrap_model(transformer) + if args.bnb_quantization_config_path is None: + if args.upcast_before_saving: + transformer.to(torch.float32) + else: + transformer = transformer.to(weight_dtype) + transformer_lora_layers = get_peft_model_state_dict(transformer) + modules_to_save["transformer"] = transformer + + QwenImage21Pipeline.save_lora_weights( + save_directory=args.output_dir, + transformer_lora_layers=transformer_lora_layers, + **_collate_lora_metadata(modules_to_save), + ) + + images = [] + run_validation = (args.validation_prompt and args.num_validation_images > 0) or (args.final_validation_prompt) + should_run_final_inference = not args.skip_final_inference and run_validation + if should_run_final_inference: + # Final inference + # Load previous pipeline + # The transformer is reloaded, so this exercises the adapter that was written to disk. + pipeline = QwenImage21ValidationPipeline.from_pretrained( + args.pretrained_model_name_or_path, + text_encoder=None, + processor=processor, + revision=args.revision, + variant=args.variant, + torch_dtype=weight_dtype, + ) + # load attention processors + pipeline.load_lora_weights(args.output_dir) + pipeline.cached_image_pad_mask = validation_image_pad_mask + + # run inference + images = log_validation( + pipeline=pipeline, + args=args, + accelerator=accelerator, + pipeline_args=validation_pipeline_args, + epoch=epoch, + is_final_validation=True, + torch_dtype=weight_dtype, + ) + del pipeline + free_memory() + + validation_prompt = args.validation_prompt if args.validation_prompt else args.final_validation_prompt + save_model_card( + (args.hub_model_id or Path(args.output_dir).name) if not args.push_to_hub else repo_id, + images=images, + base_model=args.pretrained_model_name_or_path, + instance_prompt=args.instance_prompt, + validation_prompt=validation_prompt, + repo_folder=args.output_dir, + ) + + if args.push_to_hub: + upload_folder( + repo_id=repo_id, + folder_path=args.output_dir, + commit_message="End of training", + ignore_patterns=["step_*", "epoch_*"], + ) + + images = None + + accelerator.end_training() + + +if __name__ == "__main__": + args = parse_args() + main(args) diff --git a/tests/pipelines/qwenimage21/test_qwenimage21.py b/tests/pipelines/qwenimage21/test_qwenimage21.py index 1d79334cdfc9..e009f8694388 100644 --- a/tests/pipelines/qwenimage21/test_qwenimage21.py +++ b/tests/pipelines/qwenimage21/test_qwenimage21.py @@ -36,6 +36,8 @@ from ...testing_utils import assert_tensors_close from ..testing_utils import ( BasePipelineTesterConfig, + LoraMemoryTesterMixin, + LoraTesterMixin, MemoryTesterMixin, PipelineTesterMixin, ) @@ -260,3 +262,11 @@ def test_inference_with_condition_image(self): class TestQwenImage21PipelineMemory(QwenImage21PipelineTesterConfig, MemoryTesterMixin): pass + + +class TestQwenImage21PipelineLoRA(QwenImage21PipelineTesterConfig, LoraTesterMixin): + """LoRA tests for the Qwen-Image 2.1 pipeline.""" + + +class TestQwenImage21PipelineLoRAMemory(QwenImage21PipelineTesterConfig, LoraMemoryTesterMixin): + """LoRA x memory-optimization tests for the Qwen-Image 2.1 pipeline."""