Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
53 changes: 53 additions & 0 deletions examples/dreambooth/test_dreambooth_lora_qwenimage.py
Original file line number Diff line number Diff line change
Expand Up @@ -19,6 +19,7 @@
import sys
import tempfile

import pytest
import safetensors

from diffusers.loaders.lora_base import LORA_ADAPTER_METADATA_KEY
Expand Down Expand Up @@ -243,3 +244,55 @@ def test_dreambooth_lora_with_metadata(self):
assert loaded_lora_alpha == lora_alpha
loaded_lora_rank = raw["transformer.r"]
assert loaded_lora_rank == rank

@pytest.mark.skip(reason="Aspect-ratio bucketing is opt-in and not widely used yet; re-enable when it is.")
def test_dreambooth_lora_qwen_aspect_ratio_buckets(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}
--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_qwen_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
Loading
Loading