diff --git a/examples/ace_step/model_training/special/split_training/acestep-v15-xl-sft.sh b/examples/ace_step/model_training/special/split_training/acestep-v15-xl-sft.sh new file mode 100644 index 000000000..457c0b8a7 --- /dev/null +++ b/examples/ace_step/model_training/special/split_training/acestep-v15-xl-sft.sh @@ -0,0 +1,42 @@ +modelscope download --dataset DiffSynth-Studio/diffsynth_example_dataset --include "ace_step/acestep-v15-xl-sft/*" --local_dir ./data/diffsynth_example_dataset + +# Stage 1: cache deterministic preprocessing outputs. +accelerate launch examples/ace_step/model_training/train.py \ + --learning_rate 1e-4 \ + --num_epochs 20 \ + --lora_rank 32 \ + --use_gradient_checkpointing \ + --find_unused_parameters \ + --dataset_base_path ./data/diffsynth_example_dataset/ace_step/acestep-v15-xl-sft \ + --dataset_metadata_path ./data/diffsynth_example_dataset/ace_step/acestep-v15-xl-sft/metadata.json \ + --model_id_with_origin_paths 'ACE-Step/acestep-v15-xl-sft:model-*.safetensors,ACE-Step/Ace-Step1.5:Qwen3-Embedding-0.6B/model.safetensors,ACE-Step/Ace-Step1.5:vae/diffusion_pytorch_model.safetensors' \ + --tokenizer_path ACE-Step/Ace-Step1.5:Qwen3-Embedding-0.6B/ \ + --silence_latent_path ACE-Step/Ace-Step1.5:acestep-v15-turbo/silence_latent.pt \ + --lora_base_model dit \ + --remove_prefix_in_ckpt pipe.dit. \ + --dataset_repeat 1 \ + --output_path ./models/train/acestep-v15-xl-sft_split_cache \ + --lora_target_modules q_proj,k_proj,v_proj,o_proj,gate_proj,up_proj,down_proj \ + --data_file_keys audio \ + --offload_models "" \ + --task sft:data_process + +# Stage 2: train LoRA from the cached dataset. +accelerate launch examples/ace_step/model_training/train.py \ + --learning_rate 1e-4 \ + --num_epochs 20 \ + --lora_rank 32 \ + --use_gradient_checkpointing \ + --find_unused_parameters \ + --dataset_base_path ./models/train/acestep-v15-xl-sft_split_cache \ + --model_id_with_origin_paths 'ACE-Step/acestep-v15-xl-sft:model-*.safetensors,ACE-Step/Ace-Step1.5:Qwen3-Embedding-0.6B/model.safetensors,ACE-Step/Ace-Step1.5:vae/diffusion_pytorch_model.safetensors' \ + --tokenizer_path ACE-Step/Ace-Step1.5:Qwen3-Embedding-0.6B/ \ + --silence_latent_path ACE-Step/Ace-Step1.5:acestep-v15-turbo/silence_latent.pt \ + --lora_base_model dit \ + --remove_prefix_in_ckpt pipe.dit. \ + --dataset_repeat 50 \ + --output_path ./models/train/acestep-v15-xl-sft_split \ + --lora_target_modules q_proj,k_proj,v_proj,o_proj,gate_proj,up_proj,down_proj \ + --data_file_keys audio \ + --offload_models "" \ + --task sft:train \ No newline at end of file diff --git a/examples/ace_step/model_training/special/split_training/validate.py b/examples/ace_step/model_training/special/split_training/validate.py new file mode 100644 index 000000000..91acf4705 --- /dev/null +++ b/examples/ace_step/model_training/special/split_training/validate.py @@ -0,0 +1,43 @@ +from pathlib import Path + + +def latest_checkpoint(directory): + checkpoints = list(Path(directory).glob("epoch-*.safetensors")) + if not checkpoints: + raise FileNotFoundError(f"No checkpoint found in {directory}") + return max(checkpoints, key=lambda path: int(path.stem.rsplit("-", 1)[-1])) + + +from diffsynth.pipelines.ace_step import AceStepPipeline, ModelConfig +from diffsynth.utils.data.audio import save_audio +import torch + + +pipe = AceStepPipeline.from_pretrained( + torch_dtype=torch.bfloat16, + device="cuda", + model_configs=[ + ModelConfig(model_id="ACE-Step/acestep-v15-xl-sft", origin_file_pattern="model-*.safetensors"), + ModelConfig(model_id="ACE-Step/Ace-Step1.5", origin_file_pattern="Qwen3-Embedding-0.6B/model.safetensors"), + ModelConfig(model_id="ACE-Step/Ace-Step1.5", origin_file_pattern="vae/diffusion_pytorch_model.safetensors"), + ], + text_tokenizer_config=ModelConfig(model_id="ACE-Step/Ace-Step1.5", origin_file_pattern="Qwen3-Embedding-0.6B/"), + silence_latent_config=ModelConfig(model_id="ACE-Step/Ace-Step1.5", origin_file_pattern="acestep-v15-turbo/silence_latent.pt"), +) +pipe.load_lora(pipe.dit, str(latest_checkpoint('./models/train/acestep-v15-xl-sft_split')), alpha=1) + +prompt = "An explosive, high-energy pop-rock track with a strong anime theme song feel. The song kicks off with a catchy, synthesized brass fanfare over a driving rock beat with punchy drums and a solid bassline. A powerful, clear male vocal enters with a theatrical and energetic delivery, soaring through the verses and hitting powerful high notes in the chorus. The arrangement is dense and dynamic, featuring rhythmic electric guitar chords, brief instrumental breaks with synth flourishes, and a consistent, danceable groove throughout. The overall mood is triumphant, adventurous, and exhilarating." +lyrics = '[Intro - Synth Brass Fanfare]\n\n[Verse 1]\n黑夜里的风吹过耳畔\n甜蜜时光转瞬即万\n脚步飘摇在星光上\n心追节奏心跳狂乱\n耳边传来电吉他呼唤\n手指轻触碰点流点燃\n梦在云端任它蔓延\n疯狂跳跃自由无间\n\n[Chorus]\n心电感应在震动间\n拥抱未来勇敢冒险\n那旋律在心中无限\n世界变得如此耀眼\n\n[Instrumental Break - Synth Brass Melody]\n\n[Verse 2]\n鼓点撞击黑夜的底端\n跳动节拍连接你我俩\n在这里让灵魂发光\n燃尽所有不留遗憾\n\n[Instrumental Break - Synth Brass Melody]\n\n[Bridge]\n光影交错彼此的视线\n霓虹之下夜空的蔚蓝\n月光洒下温热心田\n追逐梦想它不会遥远\n\n[Chorus]\n心电感应在震动间\n拥抱未来勇敢冒险\n那旋律在心中无限\n世界变得如此耀眼\n\n[Outro - Instrumental with Synth Brass Melody]\n[Song ends abruptly]' +audio = pipe( + prompt=prompt, + lyrics=lyrics, + duration=160, + bpm=100, + keyscale="B minor", + timesignature="4", + vocal_language="zh", + seed=1, + num_inference_steps=50, + cfg_scale=4.0, +) +save_audio(audio, pipe.vae.sampling_rate, 'split_training_acestep-v15-xl-sft.wav') diff --git a/examples/anima/model_training/special/split_training/anima-preview.sh b/examples/anima/model_training/special/split_training/anima-preview.sh new file mode 100644 index 000000000..bdd2c405d --- /dev/null +++ b/examples/anima/model_training/special/split_training/anima-preview.sh @@ -0,0 +1,42 @@ +set -e + +modelscope download --dataset DiffSynth-Studio/diffsynth_example_dataset --include "anima/anima-preview/*" --local_dir ./data/diffsynth_example_dataset + +# Stage 1: cache deterministic preprocessing outputs. +accelerate launch examples/anima/model_training/train.py \ + --dataset_base_path data/diffsynth_example_dataset/anima/anima-preview \ + --dataset_metadata_path data/diffsynth_example_dataset/anima/anima-preview/metadata.csv \ + --max_pixels 1048576 \ + --dataset_repeat 1 \ + --model_id_with_origin_paths circlestone-labs/Anima:split_files/diffusion_models/anima-preview.safetensors,circlestone-labs/Anima:split_files/text_encoders/qwen_3_06b_base.safetensors,circlestone-labs/Anima:split_files/vae/qwen_image_vae.safetensors \ + --tokenizer_path Qwen/Qwen3-0.6B:./ \ + --tokenizer_t5xxl_path stabilityai/stable-diffusion-3.5-large:tokenizer_3/ \ + --learning_rate 1e-4 \ + --num_epochs 5 \ + --remove_prefix_in_ckpt pipe.dit. \ + --output_path ./models/train/anima-preview_split_cache \ + --lora_base_model dit \ + --lora_target_modules '' \ + --lora_rank 32 \ + --use_gradient_checkpointing \ + --offload_models circlestone-labs/Anima:split_files/diffusion_models/anima-preview.safetensors \ + --task sft:data_process + +# Stage 2: train LoRA from the cached dataset. +accelerate launch examples/anima/model_training/train.py \ + --dataset_base_path ./models/train/anima-preview_split_cache \ + --max_pixels 1048576 \ + --dataset_repeat 50 \ + --model_id_with_origin_paths circlestone-labs/Anima:split_files/diffusion_models/anima-preview.safetensors,circlestone-labs/Anima:split_files/text_encoders/qwen_3_06b_base.safetensors,circlestone-labs/Anima:split_files/vae/qwen_image_vae.safetensors \ + --tokenizer_path Qwen/Qwen3-0.6B:./ \ + --tokenizer_t5xxl_path stabilityai/stable-diffusion-3.5-large:tokenizer_3/ \ + --learning_rate 1e-4 \ + --num_epochs 5 \ + --remove_prefix_in_ckpt pipe.dit. \ + --output_path ./models/train/anima-preview_split \ + --lora_base_model dit \ + --lora_target_modules '' \ + --lora_rank 32 \ + --use_gradient_checkpointing \ + --offload_models circlestone-labs/Anima:split_files/text_encoders/qwen_3_06b_base.safetensors,circlestone-labs/Anima:split_files/vae/qwen_image_vae.safetensors \ + --task sft:train diff --git a/examples/anima/model_training/special/split_training/validate.py b/examples/anima/model_training/special/split_training/validate.py new file mode 100644 index 000000000..1435e9e65 --- /dev/null +++ b/examples/anima/model_training/special/split_training/validate.py @@ -0,0 +1,29 @@ +from pathlib import Path + + +def latest_checkpoint(directory): + checkpoints = list(Path(directory).glob("epoch-*.safetensors")) + if not checkpoints: + raise FileNotFoundError(f"No checkpoint found in {directory}") + return max(checkpoints, key=lambda path: int(path.stem.rsplit("-", 1)[-1])) + + +from diffsynth.pipelines.anima_image import AnimaImagePipeline, ModelConfig +import torch + + +pipe = AnimaImagePipeline.from_pretrained( + torch_dtype=torch.bfloat16, + device="cuda", + model_configs=[ + ModelConfig(model_id="circlestone-labs/Anima", origin_file_pattern="split_files/diffusion_models/anima-preview.safetensors"), + ModelConfig(model_id="circlestone-labs/Anima", origin_file_pattern="split_files/text_encoders/qwen_3_06b_base.safetensors"), + ModelConfig(model_id="circlestone-labs/Anima", origin_file_pattern="split_files/vae/qwen_image_vae.safetensors"), + ], + tokenizer_config=ModelConfig(model_id="Qwen/Qwen3-0.6B", origin_file_pattern="./"), + tokenizer_t5xxl_config=ModelConfig(model_id="stabilityai/stable-diffusion-3.5-large", origin_file_pattern="tokenizer_3/") +) +pipe.load_lora(pipe.dit, str(latest_checkpoint('./models/train/anima-preview_split'))) +prompt = "a dog" +image = pipe(prompt=prompt, seed=0) +image.save('split_training_anima-preview.jpg') \ No newline at end of file diff --git a/examples/boogu_image/model_training/special/split_training/Boogu-Image-0.1-Base.sh b/examples/boogu_image/model_training/special/split_training/Boogu-Image-0.1-Base.sh new file mode 100644 index 000000000..c048a690f --- /dev/null +++ b/examples/boogu_image/model_training/special/split_training/Boogu-Image-0.1-Base.sh @@ -0,0 +1,44 @@ +set -e + +modelscope download --dataset DiffSynth-Studio/diffsynth_example_dataset --include "boogu_image/Boogu-Image-0.1-Base/*" --local_dir ./data/diffsynth_example_dataset + +# Stage 1: cache deterministic preprocessing outputs. +accelerate launch examples/boogu_image/model_training/train.py \ + --dataset_base_path ./data/diffsynth_example_dataset/boogu_image/Boogu-Image-0.1-Base \ + --dataset_metadata_path ./data/diffsynth_example_dataset/boogu_image/Boogu-Image-0.1-Base/metadata.csv \ + --max_pixels 1048576 \ + --dataset_repeat 1 \ + --model_id_with_origin_paths 'Boogu/Boogu-Image-0.1-Base:transformer/*.safetensors,Boogu/Boogu-Image-0.1-Base:mllm/*.safetensors,Boogu/Boogu-Image-0.1-Base:vae/*.safetensors' \ + --learning_rate 1e-4 \ + --num_epochs 5 \ + --remove_prefix_in_ckpt pipe.dit. \ + --output_path ./models/train/Boogu-Image-0.1-Base_split_cache \ + --lora_base_model dit \ + --lora_target_modules to_q,to_k,to_v,to_out.0,img_to_q,img_to_k,img_to_v,img_out,instruct_to_q,instruct_to_k,instruct_to_v,instruct_out \ + --lora_rank 32 \ + --use_gradient_checkpointing \ + --dataset_num_workers 8 \ + --find_unused_parameters \ + --data_file_keys image \ + --offload_models 'Boogu/Boogu-Image-0.1-Base:transformer/*.safetensors' \ + --task sft:data_process + +# Stage 2: train LoRA from the cached dataset. +accelerate launch examples/boogu_image/model_training/train.py \ + --dataset_base_path ./models/train/Boogu-Image-0.1-Base_split_cache \ + --max_pixels 1048576 \ + --dataset_repeat 50 \ + --model_id_with_origin_paths 'Boogu/Boogu-Image-0.1-Base:transformer/*.safetensors,Boogu/Boogu-Image-0.1-Base:mllm/*.safetensors,Boogu/Boogu-Image-0.1-Base:vae/*.safetensors' \ + --learning_rate 1e-4 \ + --num_epochs 5 \ + --remove_prefix_in_ckpt pipe.dit. \ + --output_path ./models/train/Boogu-Image-0.1-Base_split \ + --lora_base_model dit \ + --lora_target_modules to_q,to_k,to_v,to_out.0,img_to_q,img_to_k,img_to_v,img_out,instruct_to_q,instruct_to_k,instruct_to_v,instruct_out \ + --lora_rank 32 \ + --use_gradient_checkpointing \ + --dataset_num_workers 8 \ + --find_unused_parameters \ + --data_file_keys image \ + --offload_models 'Boogu/Boogu-Image-0.1-Base:mllm/*.safetensors,Boogu/Boogu-Image-0.1-Base:vae/*.safetensors' \ + --task sft:train diff --git a/examples/boogu_image/model_training/special/split_training/validate.py b/examples/boogu_image/model_training/special/split_training/validate.py new file mode 100644 index 000000000..eea3a9945 --- /dev/null +++ b/examples/boogu_image/model_training/special/split_training/validate.py @@ -0,0 +1,38 @@ +from pathlib import Path + + +def latest_checkpoint(directory): + checkpoints = list(Path(directory).glob("epoch-*.safetensors")) + if not checkpoints: + raise FileNotFoundError(f"No checkpoint found in {directory}") + return max(checkpoints, key=lambda path: int(path.stem.rsplit("-", 1)[-1])) + + +import torch +from diffsynth.pipelines.boogu_image import BooguImagePipeline, ModelConfig + +pipe = BooguImagePipeline.from_pretrained( + torch_dtype=torch.bfloat16, + device="cuda", + model_configs=[ + ModelConfig(model_id="Boogu/Boogu-Image-0.1-Base", origin_file_pattern="transformer/*.safetensors"), + ModelConfig(model_id="Boogu/Boogu-Image-0.1-Base", origin_file_pattern="mllm/*.safetensors"), + ModelConfig(model_id="Boogu/Boogu-Image-0.1-Base", origin_file_pattern="vae/*.safetensors"), + ], + processor_config=ModelConfig(model_id="Boogu/Boogu-Image-0.1-Base", origin_file_pattern="mllm/"), +) + +pipe.load_lora(pipe.dit, str(latest_checkpoint('./models/train/Boogu-Image-0.1-Base_split'))) + +prompt = "dog,white and brown dog, sitting on wall, under pink flowers" + +output = pipe( + prompt=prompt, + negative_prompt="", + height=1024, + width=1024, + seed=42, + num_inference_steps=50, + cfg_scale=4.0, +) +output.save('split_training_Boogu-Image-0.1-Base.jpg') diff --git a/examples/ernie_image/model_training/special/split_training/ERNIE-Image.sh b/examples/ernie_image/model_training/special/split_training/ERNIE-Image.sh new file mode 100644 index 000000000..9fc5639e8 --- /dev/null +++ b/examples/ernie_image/model_training/special/split_training/ERNIE-Image.sh @@ -0,0 +1,42 @@ +set -e + +modelscope download --dataset DiffSynth-Studio/diffsynth_example_dataset --include "ernie_image/Ernie-Image-T2I/*" --local_dir ./data/diffsynth_example_dataset + +# Stage 1: cache deterministic preprocessing outputs. +accelerate launch examples/ernie_image/model_training/train.py \ + --dataset_base_path data/diffsynth_example_dataset/ernie_image/Ernie-Image-T2I \ + --dataset_metadata_path data/diffsynth_example_dataset/ernie_image/Ernie-Image-T2I/metadata.csv \ + --max_pixels 1048576 \ + --dataset_repeat 1 \ + --model_id_with_origin_paths 'PaddlePaddle/ERNIE-Image:transformer/diffusion_pytorch_model*.safetensors,PaddlePaddle/ERNIE-Image:text_encoder/model.safetensors,PaddlePaddle/ERNIE-Image:vae/diffusion_pytorch_model.safetensors' \ + --learning_rate 1e-4 \ + --num_epochs 5 \ + --remove_prefix_in_ckpt pipe.dit. \ + --output_path ./models/train/Ernie-Image-T2I_split_cache \ + --lora_base_model dit \ + --lora_target_modules to_q,to_k,to_v,to_out.0 \ + --lora_rank 32 \ + --use_gradient_checkpointing \ + --dataset_num_workers 8 \ + --find_unused_parameters \ + --offload_models 'PaddlePaddle/ERNIE-Image:transformer/diffusion_pytorch_model*.safetensors' \ + --task sft:data_process + +# Stage 2: train LoRA from the cached dataset. +accelerate launch examples/ernie_image/model_training/train.py \ + --dataset_base_path ./models/train/Ernie-Image-T2I_split_cache \ + --max_pixels 1048576 \ + --dataset_repeat 50 \ + --model_id_with_origin_paths 'PaddlePaddle/ERNIE-Image:transformer/diffusion_pytorch_model*.safetensors,PaddlePaddle/ERNIE-Image:text_encoder/model.safetensors,PaddlePaddle/ERNIE-Image:vae/diffusion_pytorch_model.safetensors' \ + --learning_rate 1e-4 \ + --num_epochs 5 \ + --remove_prefix_in_ckpt pipe.dit. \ + --output_path ./models/train/Ernie-Image-T2I_split \ + --lora_base_model dit \ + --lora_target_modules to_q,to_k,to_v,to_out.0 \ + --lora_rank 32 \ + --use_gradient_checkpointing \ + --dataset_num_workers 8 \ + --find_unused_parameters \ + --offload_models PaddlePaddle/ERNIE-Image:text_encoder/model.safetensors,PaddlePaddle/ERNIE-Image:vae/diffusion_pytorch_model.safetensors \ + --task sft:train diff --git a/examples/ernie_image/model_training/special/split_training/validate.py b/examples/ernie_image/model_training/special/split_training/validate.py new file mode 100644 index 000000000..b06ecc267 --- /dev/null +++ b/examples/ernie_image/model_training/special/split_training/validate.py @@ -0,0 +1,35 @@ +from pathlib import Path + + +def latest_checkpoint(directory): + checkpoints = list(Path(directory).glob("epoch-*.safetensors")) + if not checkpoints: + raise FileNotFoundError(f"No checkpoint found in {directory}") + return max(checkpoints, key=lambda path: int(path.stem.rsplit("-", 1)[-1])) + + +import torch +from diffsynth.pipelines.ernie_image import ErnieImagePipeline, ModelConfig +from diffsynth.core.loader.file import load_state_dict + +pipe = ErnieImagePipeline.from_pretrained( + torch_dtype=torch.bfloat16, + device="cuda", + model_configs=[ + ModelConfig(model_id="PaddlePaddle/ERNIE-Image", origin_file_pattern="transformer/diffusion_pytorch_model*.safetensors"), + ModelConfig(model_id="PaddlePaddle/ERNIE-Image", origin_file_pattern="text_encoder/model.safetensors"), + ModelConfig(model_id="PaddlePaddle/ERNIE-Image", origin_file_pattern="vae/diffusion_pytorch_model.safetensors"), + ], +) + +lora_state_dict = load_state_dict(str(latest_checkpoint('./models/train/Ernie-Image-T2I_split')), torch_dtype=torch.bfloat16, device="cuda") +pipe.load_lora(pipe.dit, state_dict=lora_state_dict, alpha=1.0) + +image = pipe( + prompt="a professional photo of a cute dog", + seed=0, + num_inference_steps=50, + cfg_scale=4.0, +) +image.save('split_training_ERNIE-Image.jpg') +print("LoRA validation image saved to image_lora.jpg") diff --git a/examples/flux/model_training/special/split_training/FLUX.1-dev.sh b/examples/flux/model_training/special/split_training/FLUX.1-dev.sh new file mode 100644 index 000000000..a844e38e1 --- /dev/null +++ b/examples/flux/model_training/special/split_training/FLUX.1-dev.sh @@ -0,0 +1,40 @@ +set -e + +modelscope download --dataset DiffSynth-Studio/diffsynth_example_dataset --include "flux/FLUX.1-dev/*" --local_dir ./data/diffsynth_example_dataset + +# Stage 1: cache deterministic preprocessing outputs. +accelerate launch examples/flux/model_training/train.py \ + --dataset_base_path data/diffsynth_example_dataset/flux/FLUX.1-dev \ + --dataset_metadata_path data/diffsynth_example_dataset/flux/FLUX.1-dev/metadata.csv \ + --max_pixels 1048576 \ + --dataset_repeat 1 \ + --model_id_with_origin_paths 'black-forest-labs/FLUX.1-dev:flux1-dev.safetensors,black-forest-labs/FLUX.1-dev:text_encoder/model.safetensors,black-forest-labs/FLUX.1-dev:text_encoder_2/*.safetensors,black-forest-labs/FLUX.1-dev:ae.safetensors' \ + --learning_rate 1e-4 \ + --num_epochs 5 \ + --remove_prefix_in_ckpt pipe.dit. \ + --output_path ./models/train/FLUX.1-dev_split_cache \ + --lora_base_model dit \ + --lora_target_modules a_to_qkv,b_to_qkv,ff_a.0,ff_a.2,ff_b.0,ff_b.2,a_to_out,b_to_out,proj_out,norm.linear,norm1_a.linear,norm1_b.linear,to_qkv_mlp \ + --lora_rank 32 \ + --align_to_opensource_format \ + --use_gradient_checkpointing \ + --offload_models black-forest-labs/FLUX.1-dev:flux1-dev.safetensors \ + --task sft:data_process + +# Stage 2: train LoRA from the cached dataset. +accelerate launch examples/flux/model_training/train.py \ + --dataset_base_path ./models/train/FLUX.1-dev_split_cache \ + --max_pixels 1048576 \ + --dataset_repeat 50 \ + --model_id_with_origin_paths 'black-forest-labs/FLUX.1-dev:flux1-dev.safetensors,black-forest-labs/FLUX.1-dev:text_encoder/model.safetensors,black-forest-labs/FLUX.1-dev:text_encoder_2/*.safetensors,black-forest-labs/FLUX.1-dev:ae.safetensors' \ + --learning_rate 1e-4 \ + --num_epochs 5 \ + --remove_prefix_in_ckpt pipe.dit. \ + --output_path ./models/train/FLUX.1-dev_split \ + --lora_base_model dit \ + --lora_target_modules a_to_qkv,b_to_qkv,ff_a.0,ff_a.2,ff_b.0,ff_b.2,a_to_out,b_to_out,proj_out,norm.linear,norm1_a.linear,norm1_b.linear,to_qkv_mlp \ + --lora_rank 32 \ + --align_to_opensource_format \ + --use_gradient_checkpointing \ + --offload_models 'black-forest-labs/FLUX.1-dev:text_encoder_2/*.safetensors,black-forest-labs/FLUX.1-dev:ae.safetensors' \ + --task sft:train diff --git a/examples/flux/model_training/special/split_training/validate.py b/examples/flux/model_training/special/split_training/validate.py new file mode 100644 index 000000000..896064a7d --- /dev/null +++ b/examples/flux/model_training/special/split_training/validate.py @@ -0,0 +1,28 @@ +from pathlib import Path + + +def latest_checkpoint(directory): + checkpoints = list(Path(directory).glob("epoch-*.safetensors")) + if not checkpoints: + raise FileNotFoundError(f"No checkpoint found in {directory}") + return max(checkpoints, key=lambda path: int(path.stem.rsplit("-", 1)[-1])) + + +import torch +from diffsynth.pipelines.flux_image import FluxImagePipeline, ModelConfig + + +pipe = FluxImagePipeline.from_pretrained( + torch_dtype=torch.bfloat16, + device="cuda", + model_configs=[ + ModelConfig(model_id="black-forest-labs/FLUX.1-dev", origin_file_pattern="flux1-dev.safetensors"), + ModelConfig(model_id="black-forest-labs/FLUX.1-dev", origin_file_pattern="text_encoder/model.safetensors"), + ModelConfig(model_id="black-forest-labs/FLUX.1-dev", origin_file_pattern="text_encoder_2/*.safetensors"), + ModelConfig(model_id="black-forest-labs/FLUX.1-dev", origin_file_pattern="ae.safetensors"), + ], +) +pipe.load_lora(pipe.dit, str(latest_checkpoint('./models/train/FLUX.1-dev_split')), alpha=1) + +image = pipe(prompt="a dog", seed=0) +image.save('split_training_FLUX.1-dev.jpg') diff --git a/examples/flux2/model_training/special/split_training/FLUX.2-klein-base-4B_lora.sh b/examples/flux2/model_training/special/split_training/FLUX.2-klein-base-4B_lora.sh index 25751d1a0..5bde53ab7 100644 --- a/examples/flux2/model_training/special/split_training/FLUX.2-klein-base-4B_lora.sh +++ b/examples/flux2/model_training/special/split_training/FLUX.2-klein-base-4B_lora.sh @@ -1,34 +1,40 @@ +set -e + modelscope download --dataset DiffSynth-Studio/diffsynth_example_dataset --include "flux2/FLUX.2-klein-base-4B/*" --local_dir ./data/diffsynth_example_dataset +# Stage 1: cache deterministic preprocessing outputs. accelerate launch examples/flux2/model_training/train.py \ - --dataset_base_path data/example_image_dataset \ - --dataset_metadata_path data/example_image_dataset/metadata.csv \ + --dataset_base_path data/diffsynth_example_dataset/flux2/FLUX.2-klein-base-4B \ + --dataset_metadata_path data/diffsynth_example_dataset/flux2/FLUX.2-klein-base-4B/metadata.csv \ --max_pixels 1048576 \ --dataset_repeat 1 \ - --model_id_with_origin_paths "black-forest-labs/FLUX.2-klein-4B:text_encoder/*.safetensors,black-forest-labs/FLUX.2-klein-4B:vae/diffusion_pytorch_model.safetensors" \ - --tokenizer_path "black-forest-labs/FLUX.2-klein-4B:tokenizer/" \ + --model_id_with_origin_paths 'black-forest-labs/FLUX.2-klein-4B:text_encoder/*.safetensors,black-forest-labs/FLUX.2-klein-base-4B:transformer/*.safetensors,black-forest-labs/FLUX.2-klein-4B:vae/diffusion_pytorch_model.safetensors' \ + --tokenizer_path black-forest-labs/FLUX.2-klein-4B:tokenizer/ \ --learning_rate 1e-4 \ --num_epochs 5 \ - --remove_prefix_in_ckpt "pipe.dit." \ - --output_path "./models/train/FLUX.2-klein-base-4B_lora_cache" \ - --lora_base_model "dit" \ - --lora_target_modules "to_q,to_k,to_v,to_out.0,add_q_proj,add_k_proj,add_v_proj,to_add_out,linear_in,linear_out,to_qkv_mlp_proj,single_transformer_blocks.0.attn.to_out,single_transformer_blocks.1.attn.to_out,single_transformer_blocks.2.attn.to_out,single_transformer_blocks.3.attn.to_out,single_transformer_blocks.4.attn.to_out,single_transformer_blocks.5.attn.to_out,single_transformer_blocks.6.attn.to_out,single_transformer_blocks.7.attn.to_out,single_transformer_blocks.8.attn.to_out,single_transformer_blocks.9.attn.to_out,single_transformer_blocks.10.attn.to_out,single_transformer_blocks.11.attn.to_out,single_transformer_blocks.12.attn.to_out,single_transformer_blocks.13.attn.to_out,single_transformer_blocks.14.attn.to_out,single_transformer_blocks.15.attn.to_out,single_transformer_blocks.16.attn.to_out,single_transformer_blocks.17.attn.to_out,single_transformer_blocks.18.attn.to_out,single_transformer_blocks.19.attn.to_out" \ + --remove_prefix_in_ckpt pipe.dit. \ + --output_path ./models/train/FLUX.2-klein-base-4B_lora_cache \ + --lora_base_model dit \ + --lora_target_modules to_q,to_k,to_v,to_out.0,add_q_proj,add_k_proj,add_v_proj,to_add_out,linear_in,linear_out,to_qkv_mlp_proj,single_transformer_blocks.0.attn.to_out,single_transformer_blocks.1.attn.to_out,single_transformer_blocks.2.attn.to_out,single_transformer_blocks.3.attn.to_out,single_transformer_blocks.4.attn.to_out,single_transformer_blocks.5.attn.to_out,single_transformer_blocks.6.attn.to_out,single_transformer_blocks.7.attn.to_out,single_transformer_blocks.8.attn.to_out,single_transformer_blocks.9.attn.to_out,single_transformer_blocks.10.attn.to_out,single_transformer_blocks.11.attn.to_out,single_transformer_blocks.12.attn.to_out,single_transformer_blocks.13.attn.to_out,single_transformer_blocks.14.attn.to_out,single_transformer_blocks.15.attn.to_out,single_transformer_blocks.16.attn.to_out,single_transformer_blocks.17.attn.to_out,single_transformer_blocks.18.attn.to_out,single_transformer_blocks.19.attn.to_out \ --lora_rank 32 \ --use_gradient_checkpointing \ - --task "sft:data_process" + --offload_models 'black-forest-labs/FLUX.2-klein-base-4B:transformer/*.safetensors' \ + --task sft:data_process +# Stage 2: train LoRA from the cached dataset. accelerate launch examples/flux2/model_training/train.py \ - --dataset_base_path "./models/train/FLUX.2-klein-base-4B_lora_cache" \ + --dataset_base_path ./models/train/FLUX.2-klein-base-4B_lora_cache \ --max_pixels 1048576 \ --dataset_repeat 50 \ - --model_id_with_origin_paths "black-forest-labs/FLUX.2-klein-base-4B:transformer/*.safetensors" \ - --tokenizer_path "black-forest-labs/FLUX.2-klein-4B:tokenizer/" \ + --model_id_with_origin_paths 'black-forest-labs/FLUX.2-klein-4B:text_encoder/*.safetensors,black-forest-labs/FLUX.2-klein-base-4B:transformer/*.safetensors,black-forest-labs/FLUX.2-klein-4B:vae/diffusion_pytorch_model.safetensors' \ + --tokenizer_path black-forest-labs/FLUX.2-klein-4B:tokenizer/ \ --learning_rate 1e-4 \ --num_epochs 5 \ - --remove_prefix_in_ckpt "pipe.dit." \ - --output_path "./models/train/FLUX.2-klein-base-4B_lora" \ - --lora_base_model "dit" \ - --lora_target_modules "to_q,to_k,to_v,to_out.0,add_q_proj,add_k_proj,add_v_proj,to_add_out,linear_in,linear_out,to_qkv_mlp_proj,single_transformer_blocks.0.attn.to_out,single_transformer_blocks.1.attn.to_out,single_transformer_blocks.2.attn.to_out,single_transformer_blocks.3.attn.to_out,single_transformer_blocks.4.attn.to_out,single_transformer_blocks.5.attn.to_out,single_transformer_blocks.6.attn.to_out,single_transformer_blocks.7.attn.to_out,single_transformer_blocks.8.attn.to_out,single_transformer_blocks.9.attn.to_out,single_transformer_blocks.10.attn.to_out,single_transformer_blocks.11.attn.to_out,single_transformer_blocks.12.attn.to_out,single_transformer_blocks.13.attn.to_out,single_transformer_blocks.14.attn.to_out,single_transformer_blocks.15.attn.to_out,single_transformer_blocks.16.attn.to_out,single_transformer_blocks.17.attn.to_out,single_transformer_blocks.18.attn.to_out,single_transformer_blocks.19.attn.to_out" \ + --remove_prefix_in_ckpt pipe.dit. \ + --output_path ./models/train/FLUX.2-klein-base-4B_lora \ + --lora_base_model dit \ + --lora_target_modules to_q,to_k,to_v,to_out.0,add_q_proj,add_k_proj,add_v_proj,to_add_out,linear_in,linear_out,to_qkv_mlp_proj,single_transformer_blocks.0.attn.to_out,single_transformer_blocks.1.attn.to_out,single_transformer_blocks.2.attn.to_out,single_transformer_blocks.3.attn.to_out,single_transformer_blocks.4.attn.to_out,single_transformer_blocks.5.attn.to_out,single_transformer_blocks.6.attn.to_out,single_transformer_blocks.7.attn.to_out,single_transformer_blocks.8.attn.to_out,single_transformer_blocks.9.attn.to_out,single_transformer_blocks.10.attn.to_out,single_transformer_blocks.11.attn.to_out,single_transformer_blocks.12.attn.to_out,single_transformer_blocks.13.attn.to_out,single_transformer_blocks.14.attn.to_out,single_transformer_blocks.15.attn.to_out,single_transformer_blocks.16.attn.to_out,single_transformer_blocks.17.attn.to_out,single_transformer_blocks.18.attn.to_out,single_transformer_blocks.19.attn.to_out \ --lora_rank 32 \ --use_gradient_checkpointing \ - --task "sft:train" + --offload_models 'black-forest-labs/FLUX.2-klein-4B:text_encoder/*.safetensors,black-forest-labs/FLUX.2-klein-4B:vae/diffusion_pytorch_model.safetensors' \ + --task sft:train diff --git a/examples/flux2/model_training/special/split_training/Template-KleinBase4B-Brightness.sh b/examples/flux2/model_training/special/split_training/Template-KleinBase4B-Brightness.sh index b214595e6..ed522a7b6 100644 --- a/examples/flux2/model_training/special/split_training/Template-KleinBase4B-Brightness.sh +++ b/examples/flux2/model_training/special/split_training/Template-KleinBase4B-Brightness.sh @@ -6,7 +6,8 @@ accelerate launch examples/flux2/model_training/train.py \ --extra_inputs "template_inputs" \ --max_pixels 1048576 \ --dataset_repeat 1 \ - --model_id_with_origin_paths "black-forest-labs/FLUX.2-klein-4B:text_encoder/*.safetensors,black-forest-labs/FLUX.2-klein-4B:vae/diffusion_pytorch_model.safetensors" \ + --model_id_with_origin_paths "black-forest-labs/FLUX.2-klein-4B:text_encoder/*.safetensors,black-forest-labs/FLUX.2-klein-base-4B:transformer/*.safetensors,black-forest-labs/FLUX.2-klein-4B:vae/diffusion_pytorch_model.safetensors" \ + --offload_models "black-forest-labs/FLUX.2-klein-base-4B:transformer/*.safetensors" \ --template_model_id_or_path "DiffSynth-Studio/Template-KleinBase4B-Brightness:" \ --tokenizer_path "black-forest-labs/FLUX.2-klein-4B:tokenizer/" \ --learning_rate 1e-4 \ @@ -23,7 +24,8 @@ accelerate launch examples/flux2/model_training/train.py \ --extra_inputs "template_inputs" \ --max_pixels 1048576 \ --dataset_repeat 50 \ - --model_id_with_origin_paths "black-forest-labs/FLUX.2-klein-base-4B:transformer/*.safetensors" \ + --model_id_with_origin_paths "black-forest-labs/FLUX.2-klein-4B:text_encoder/*.safetensors,black-forest-labs/FLUX.2-klein-base-4B:transformer/*.safetensors,black-forest-labs/FLUX.2-klein-4B:vae/diffusion_pytorch_model.safetensors" \ + --offload_models "black-forest-labs/FLUX.2-klein-4B:text_encoder/*.safetensors,black-forest-labs/FLUX.2-klein-4B:vae/diffusion_pytorch_model.safetensors" \ --template_model_id_or_path "DiffSynth-Studio/Template-KleinBase4B-Brightness:" \ --tokenizer_path "black-forest-labs/FLUX.2-klein-4B:tokenizer/" \ --learning_rate 1e-4 \ diff --git a/examples/flux2/model_training/special/split_training/validate.py b/examples/flux2/model_training/special/split_training/validate.py new file mode 100644 index 000000000..44461e6b2 --- /dev/null +++ b/examples/flux2/model_training/special/split_training/validate.py @@ -0,0 +1,28 @@ +from pathlib import Path + + +def latest_checkpoint(directory): + checkpoints = list(Path(directory).glob("epoch-*.safetensors")) + if not checkpoints: + raise FileNotFoundError(f"No checkpoint found in {directory}") + return max(checkpoints, key=lambda path: int(path.stem.rsplit("-", 1)[-1])) + + +from diffsynth.pipelines.flux2_image import Flux2ImagePipeline, ModelConfig +import torch + + +pipe = Flux2ImagePipeline.from_pretrained( + torch_dtype=torch.bfloat16, + device="cuda", + model_configs=[ + ModelConfig(model_id="black-forest-labs/FLUX.2-klein-4B", origin_file_pattern="text_encoder/*.safetensors"), + ModelConfig(model_id="black-forest-labs/FLUX.2-klein-base-4B", origin_file_pattern="transformer/*.safetensors"), + ModelConfig(model_id="black-forest-labs/FLUX.2-klein-4B", origin_file_pattern="vae/diffusion_pytorch_model.safetensors"), + ], + tokenizer_config=ModelConfig(model_id="black-forest-labs/FLUX.2-klein-base-4B", origin_file_pattern="tokenizer/"), +) +pipe.load_lora(pipe.dit, str(latest_checkpoint('./models/train/FLUX.2-klein-base-4B_lora'))) +prompt = "a dog" +image = pipe(prompt=prompt, seed=0, num_inference_steps=40, cfg_scale=4, height=768, width=768) +image.save('split_training_FLUX.2-klein-base-4B.jpg') diff --git a/examples/hidream_o1_image/model_training/special/split_training/HiDream-O1-Image.sh b/examples/hidream_o1_image/model_training/special/split_training/HiDream-O1-Image.sh new file mode 100644 index 000000000..87c1a9192 --- /dev/null +++ b/examples/hidream_o1_image/model_training/special/split_training/HiDream-O1-Image.sh @@ -0,0 +1,41 @@ +set -e + +modelscope download --dataset DiffSynth-Studio/diffsynth_example_dataset --include "hidream_o1_image/HiDream-O1-Image/*" --local_dir ./data/diffsynth_example_dataset + +# Stage 1: cache deterministic preprocessing outputs. +accelerate launch examples/hidream_o1_image/model_training/train.py \ + --dataset_base_path data/diffsynth_example_dataset/hidream_o1_image/HiDream-O1-Image \ + --dataset_metadata_path data/diffsynth_example_dataset/hidream_o1_image/HiDream-O1-Image/metadata.csv \ + --max_pixels 4194304 \ + --dataset_repeat 1 \ + --model_id_with_origin_paths 'HiDream-ai/HiDream-O1-Image:model-*.safetensors' \ + --processor_config HiDream-ai/HiDream-O1-Image:./ \ + --learning_rate 1e-4 \ + --num_epochs 5 \ + --lora_rank 32 \ + --remove_prefix_in_ckpt pipe.dit. \ + --output_path ./models/train/HiDream-O1-Image_split_cache \ + --lora_base_model dit \ + --lora_target_modules q_proj,k_proj,v_proj,o_proj,gate_proj,up_proj,down_proj,attn.qkv,attn.proj,mlp.linear_fc1,mlp.linear_fc2 \ + --use_gradient_checkpointing \ + --noise_scale 8.0 \ + --offload_models 'HiDream-ai/HiDream-O1-Image:model-*.safetensors' \ + --task sft:data_process + +# Stage 2: train LoRA from the cached dataset. +accelerate launch examples/hidream_o1_image/model_training/train.py \ + --dataset_base_path ./models/train/HiDream-O1-Image_split_cache \ + --max_pixels 4194304 \ + --dataset_repeat 50 \ + --model_id_with_origin_paths 'HiDream-ai/HiDream-O1-Image:model-*.safetensors' \ + --processor_config HiDream-ai/HiDream-O1-Image:./ \ + --learning_rate 1e-4 \ + --num_epochs 5 \ + --lora_rank 32 \ + --remove_prefix_in_ckpt pipe.dit. \ + --output_path ./models/train/HiDream-O1-Image_split \ + --lora_base_model dit \ + --lora_target_modules q_proj,k_proj,v_proj,o_proj,gate_proj,up_proj,down_proj,attn.qkv,attn.proj,mlp.linear_fc1,mlp.linear_fc2 \ + --use_gradient_checkpointing \ + --noise_scale 8.0 \ + --task sft:train diff --git a/examples/hidream_o1_image/model_training/special/split_training/validate.py b/examples/hidream_o1_image/model_training/special/split_training/validate.py new file mode 100644 index 000000000..5131a27f1 --- /dev/null +++ b/examples/hidream_o1_image/model_training/special/split_training/validate.py @@ -0,0 +1,33 @@ +from pathlib import Path + + +def latest_checkpoint(directory): + checkpoints = list(Path(directory).glob("epoch-*.safetensors")) + if not checkpoints: + raise FileNotFoundError(f"No checkpoint found in {directory}") + return max(checkpoints, key=lambda path: int(path.stem.rsplit("-", 1)[-1])) + + +import torch +from diffsynth.pipelines.hidream_o1_image import HiDreamO1ImagePipeline, ModelConfig + + +pipe = HiDreamO1ImagePipeline.from_pretrained( + torch_dtype=torch.bfloat16, + device="cuda", + model_configs=[ + ModelConfig(model_id="HiDream-ai/HiDream-O1-Image", origin_file_pattern="model-*.safetensors"), + ], + processor_config=ModelConfig(model_id="HiDream-ai/HiDream-O1-Image", origin_file_pattern="./"), +) +pipe.load_lora(pipe.dit, str(latest_checkpoint('./models/train/HiDream-O1-Image_split'))) +image = pipe( + prompt="dog,white and brown dog, sitting on wall, under pink flowers", + negative_prompt=" ", + cfg_scale=4.0, + height=2048, + width=2048, + seed=42, + num_inference_steps=50, +) +image.save('split_training_HiDream-O1-Image.jpg') diff --git a/examples/ideogram4/model_training/special/split_training/Ideogram-4-bf16-repackage.sh b/examples/ideogram4/model_training/special/split_training/Ideogram-4-bf16-repackage.sh new file mode 100644 index 000000000..937d622fd --- /dev/null +++ b/examples/ideogram4/model_training/special/split_training/Ideogram-4-bf16-repackage.sh @@ -0,0 +1,42 @@ +set -e + +modelscope download --dataset DiffSynth-Studio/diffsynth_example_dataset --include "ideogram4/Ideogram-4-bf16-repackage/*" --local_dir ./data/diffsynth_example_dataset + +# Stage 1: cache deterministic preprocessing outputs. +accelerate launch examples/ideogram4/model_training/train.py \ + --learning_rate 1e-4 \ + --num_epochs 5 \ + --lora_rank 32 \ + --use_gradient_checkpointing \ + --find_unused_parameters \ + --dataset_base_path ./data/diffsynth_example_dataset/ideogram4/Ideogram-4-bf16-repackage \ + --dataset_metadata_path ./data/diffsynth_example_dataset/ideogram4/Ideogram-4-bf16-repackage/metadata.json \ + --model_id_with_origin_paths DiffSynth-Studio/ideogram-4-bf16-repackage:transformer/diffusion_pytorch_model.safetensors,DiffSynth-Studio/ideogram-4-bf16-repackage:text_encoder/model.safetensors,DiffSynth-Studio/ideogram-4-bf16-repackage:vae/diffusion_pytorch_model.safetensors \ + --lora_base_model dit \ + --remove_prefix_in_ckpt pipe.dit. \ + --max_pixels 1048576 \ + --dataset_repeat 1 \ + --output_path ./models/train/Ideogram-4-bf16-repackage_split_cache \ + --lora_target_modules attention.qkv,attention.o,feed_forward.w1,feed_forward.w2,feed_forward.w3,adaln_modulation \ + --data_file_keys image \ + --offload_models DiffSynth-Studio/ideogram-4-bf16-repackage:transformer/diffusion_pytorch_model.safetensors \ + --task sft:data_process + +# Stage 2: train LoRA from the cached dataset. +accelerate launch examples/ideogram4/model_training/train.py \ + --learning_rate 1e-4 \ + --num_epochs 5 \ + --lora_rank 32 \ + --use_gradient_checkpointing \ + --find_unused_parameters \ + --dataset_base_path ./models/train/Ideogram-4-bf16-repackage_split_cache \ + --model_id_with_origin_paths DiffSynth-Studio/ideogram-4-bf16-repackage:transformer/diffusion_pytorch_model.safetensors,DiffSynth-Studio/ideogram-4-bf16-repackage:text_encoder/model.safetensors,DiffSynth-Studio/ideogram-4-bf16-repackage:vae/diffusion_pytorch_model.safetensors \ + --lora_base_model dit \ + --remove_prefix_in_ckpt pipe.dit. \ + --max_pixels 1048576 \ + --dataset_repeat 100 \ + --output_path ./models/train/Ideogram-4-bf16-repackage_split \ + --lora_target_modules attention.qkv,attention.o,feed_forward.w1,feed_forward.w2,feed_forward.w3,adaln_modulation \ + --data_file_keys image \ + --offload_models DiffSynth-Studio/ideogram-4-bf16-repackage:text_encoder/model.safetensors,DiffSynth-Studio/ideogram-4-bf16-repackage:vae/diffusion_pytorch_model.safetensors \ + --task sft:train diff --git a/examples/ideogram4/model_training/special/split_training/validate.py b/examples/ideogram4/model_training/special/split_training/validate.py new file mode 100644 index 000000000..9eebb45c8 --- /dev/null +++ b/examples/ideogram4/model_training/special/split_training/validate.py @@ -0,0 +1,31 @@ +from pathlib import Path + + +def latest_checkpoint(directory): + checkpoints = list(Path(directory).glob("epoch-*.safetensors")) + if not checkpoints: + raise FileNotFoundError(f"No checkpoint found in {directory}") + return max(checkpoints, key=lambda path: int(path.stem.rsplit("-", 1)[-1])) + + +from diffsynth.pipelines.ideogram4 import Ideogram4Pipeline +from diffsynth.core import ModelConfig +import torch + +pipe = Ideogram4Pipeline.from_pretrained( + torch_dtype=torch.bfloat16, + device="cuda", + model_configs=[ + ModelConfig(model_id="DiffSynth-Studio/ideogram-4-bf16-repackage", origin_file_pattern="transformer/diffusion_pytorch_model.safetensors"), + ModelConfig(model_id="DiffSynth-Studio/ideogram-4-bf16-repackage", origin_file_pattern="unconditional_transformer/diffusion_pytorch_model.safetensors"), + ModelConfig(model_id="DiffSynth-Studio/ideogram-4-bf16-repackage", origin_file_pattern="text_encoder/model.safetensors"), + ModelConfig(model_id="DiffSynth-Studio/ideogram-4-bf16-repackage", origin_file_pattern="vae/diffusion_pytorch_model.safetensors"), + ], + tokenizer_config=ModelConfig(model_id="ideogram-ai/ideogram-4-fp8", origin_file_pattern="tokenizer/"), +) +pipe.load_lora(pipe.dit, str(latest_checkpoint('./models/train/Ideogram-4-bf16-repackage_split')), alpha=1) +# pipe.load_lora(pipe.dit_uncond, str(latest_checkpoint('./models/train/Ideogram-4-bf16-repackage_split')), alpha=1) + +prompt = "{\"high_level_description\":\"A close-up photograph of a happy Pembroke Welsh Corgi sitting on a concrete wall, panting with its tongue out, set against a backdrop of blurred pink cherry blossoms and blue sky.\",\"style_description\":{\"aesthetics\":\"joyful, vibrant, spring-like, cute, energetic\",\"lighting\":\"bright natural daylight, soft diffuse sunlight, shallow depth of field\",\"photo\":\"85mm lens, f/2.0, bokeh background, sharp focus on dog face\",\"medium\":\"photograph\",\"color_palette\":[\"#E2725B\",\"#FFFFFF\",\"#FFB7C5\",\"#87CEEB\",\"#A9A9A9\"]},\"compositional_deconstruction\":{\"background\":\"Softly blurred background of pink cherry blossom branches against a pale blue sky. The bokeh effect creates a dreamy spring atmosphere. The background is out of focus to highlight the sharp details of the dog in the foreground.\",\"elements\":[{\"type\":\"obj\",\"bbox\":[150,200,900,850],\"desc\":\"A Pembroke Welsh Corgi with fluffy orange and white fur. Its mouth is open, panting with a pink tongue hanging out, expression is happy and excited. Ears are perked up. Sharp focus on the face and eyes.\"},{\"type\":\"obj\",\"bbox\":[850,0,1000,1000],\"desc\":\"A grey concrete wall or ledge at the bottom of the frame. The dog's front paws are resting near the edge. Rough texture.\"}]}}" +image = pipe(prompt=prompt, height=1024, width=1024, num_inference_steps=48, cfg_scale=7.0, seed=0) +image.save('split_training_Ideogram-4.jpg') diff --git a/examples/joyai_image/model_training/special/split_training/JoyAI-Image-Edit.sh b/examples/joyai_image/model_training/special/split_training/JoyAI-Image-Edit.sh new file mode 100644 index 000000000..a8796486b --- /dev/null +++ b/examples/joyai_image/model_training/special/split_training/JoyAI-Image-Edit.sh @@ -0,0 +1,44 @@ +set -e + +modelscope download --dataset DiffSynth-Studio/diffsynth_example_dataset --include "joyai_image/JoyAI-Image-Edit/*" --local_dir ./data/diffsynth_example_dataset + +# Stage 1: cache deterministic preprocessing outputs. +accelerate launch examples/joyai_image/model_training/train.py \ + --dataset_base_path ./data/diffsynth_example_dataset/joyai_image/JoyAI-Image-Edit \ + --dataset_metadata_path ./data/diffsynth_example_dataset/joyai_image/JoyAI-Image-Edit/metadata.csv \ + --max_pixels 1048576 \ + --dataset_repeat 1 \ + --model_id_with_origin_paths 'jd-opensource/JoyAI-Image-Edit:transformer/transformer.pth,jd-opensource/JoyAI-Image-Edit:JoyAI-Image-Und/model*.safetensors,jd-opensource/JoyAI-Image-Edit:vae/Wan2.1_VAE.pth' \ + --learning_rate 1e-4 \ + --num_epochs 5 \ + --remove_prefix_in_ckpt pipe.dit. \ + --output_path ./models/train/JoyAI-Image-Edit-split-cache \ + --lora_base_model dit \ + --lora_target_modules img_attn_qkv,txt_attn_qkv,img_attn_proj,txt_attn_proj \ + --lora_rank 32 \ + --use_gradient_checkpointing \ + --find_unused_parameters \ + --data_file_keys image,edit_image \ + --extra_inputs edit_image \ + --task sft:data_process \ + --offload_models jd-opensource/JoyAI-Image-Edit:transformer/transformer.pth + +# Stage 2: train LoRA from the cached dataset. +accelerate launch examples/joyai_image/model_training/train.py \ + --dataset_base_path ./models/train/JoyAI-Image-Edit-split-cache \ + --max_pixels 1048576 \ + --dataset_repeat 50 \ + --model_id_with_origin_paths 'jd-opensource/JoyAI-Image-Edit:transformer/transformer.pth,jd-opensource/JoyAI-Image-Edit:JoyAI-Image-Und/model*.safetensors,jd-opensource/JoyAI-Image-Edit:vae/Wan2.1_VAE.pth' \ + --learning_rate 1e-4 \ + --num_epochs 5 \ + --remove_prefix_in_ckpt pipe.dit. \ + --output_path ./models/train/JoyAI-Image-Edit-split \ + --lora_base_model dit \ + --lora_target_modules img_attn_qkv,txt_attn_qkv,img_attn_proj,txt_attn_proj \ + --lora_rank 32 \ + --use_gradient_checkpointing \ + --find_unused_parameters \ + --data_file_keys image,edit_image \ + --extra_inputs edit_image \ + --task sft:train \ + --offload_models 'jd-opensource/JoyAI-Image-Edit:JoyAI-Image-Und/model*.safetensors,jd-opensource/JoyAI-Image-Edit:vae/Wan2.1_VAE.pth' diff --git a/examples/joyai_image/model_training/special/split_training/validate.py b/examples/joyai_image/model_training/special/split_training/validate.py new file mode 100644 index 000000000..0319e3d79 --- /dev/null +++ b/examples/joyai_image/model_training/special/split_training/validate.py @@ -0,0 +1,40 @@ +from pathlib import Path + + +def latest_checkpoint(directory): + checkpoints = list(Path(directory).glob("epoch-*.safetensors")) + if not checkpoints: + raise FileNotFoundError(f"No checkpoint found in {directory}") + return max(checkpoints, key=lambda path: int(path.stem.rsplit("-", 1)[-1])) + + +import torch +from PIL import Image +from diffsynth.pipelines.joyai_image import JoyAIImagePipeline, ModelConfig + +pipe = JoyAIImagePipeline.from_pretrained( + torch_dtype=torch.bfloat16, + device="cuda", + model_configs=[ + ModelConfig(model_id="jd-opensource/JoyAI-Image-Edit", origin_file_pattern="transformer/transformer.pth"), + ModelConfig(model_id="jd-opensource/JoyAI-Image-Edit", origin_file_pattern="JoyAI-Image-Und/model*.safetensors"), + ModelConfig(model_id="jd-opensource/JoyAI-Image-Edit", origin_file_pattern="vae/Wan2.1_VAE.pth"), + ], + processor_config=ModelConfig(model_id="jd-opensource/JoyAI-Image-Edit", origin_file_pattern="JoyAI-Image-Und/"), +) + +pipe.load_lora(pipe.dit, str(latest_checkpoint('./models/train/JoyAI-Image-Edit-split'))) + +prompt = "将裙子改为粉色" +edit_image = Image.open("data/diffsynth_example_dataset/joyai_image/JoyAI-Image-Edit/edit/image1.jpg").convert("RGB") + +image = pipe( + prompt=prompt, + edit_image=edit_image, + height=1024, + width=1024, + seed=0, + num_inference_steps=30, + cfg_scale=5.0, +) +image.save('split_training_JoyAI-Image-Edit.jpg') diff --git a/examples/krea2/model_training/special/split_training/Krea-2-Raw.sh b/examples/krea2/model_training/special/split_training/Krea-2-Raw.sh new file mode 100644 index 000000000..90a0a501f --- /dev/null +++ b/examples/krea2/model_training/special/split_training/Krea-2-Raw.sh @@ -0,0 +1,44 @@ +set -e + +modelscope download --dataset DiffSynth-Studio/diffsynth_example_dataset --include "krea2/Krea-2-Raw/*" --local_dir ./data/diffsynth_example_dataset + +# Stage 1: cache deterministic preprocessing outputs. +accelerate launch examples/krea2/model_training/train.py \ + --dataset_base_path data/diffsynth_example_dataset/krea2/Krea-2-Raw \ + --dataset_metadata_path data/diffsynth_example_dataset/krea2/Krea-2-Raw/metadata.csv \ + --max_pixels 1048576 \ + --dataset_repeat 1 \ + --model_id_with_origin_paths 'krea/Krea-2-Raw:raw.safetensors,Qwen/Qwen3-VL-4B-Instruct:*.safetensors,Qwen/Qwen-Image:vae/diffusion_pytorch_model.safetensors' \ + --tokenizer_path Qwen/Qwen3-VL-4B-Instruct: \ + --learning_rate 1e-4 \ + --num_epochs 5 \ + --remove_prefix_in_ckpt pipe.dit. \ + --output_path ./models/train/Krea-2-Raw_split_cache \ + --lora_base_model dit \ + --lora_target_modules wq,wk,wv,gate,wo,gate,up,down,first,tmlp.0,tmlp.2,projector,txtmlp.1,txtmlp.3,last.linear,tproj.1 \ + --lora_rank 32 \ + --use_gradient_checkpointing \ + --find_unused_parameters \ + --align_to_opensource_format \ + --offload_models krea/Krea-2-Raw:raw.safetensors \ + --task sft:data_process + +# Stage 2: train LoRA from the cached dataset. +accelerate launch examples/krea2/model_training/train.py \ + --dataset_base_path ./models/train/Krea-2-Raw_split_cache \ + --max_pixels 1048576 \ + --dataset_repeat 50 \ + --model_id_with_origin_paths 'krea/Krea-2-Raw:raw.safetensors,Qwen/Qwen3-VL-4B-Instruct:*.safetensors,Qwen/Qwen-Image:vae/diffusion_pytorch_model.safetensors' \ + --tokenizer_path Qwen/Qwen3-VL-4B-Instruct: \ + --learning_rate 1e-4 \ + --num_epochs 5 \ + --remove_prefix_in_ckpt pipe.dit. \ + --output_path ./models/train/Krea-2-Raw_split \ + --lora_base_model dit \ + --lora_target_modules wq,wk,wv,gate,wo,gate,up,down,first,tmlp.0,tmlp.2,projector,txtmlp.1,txtmlp.3,last.linear,tproj.1 \ + --lora_rank 32 \ + --use_gradient_checkpointing \ + --find_unused_parameters \ + --align_to_opensource_format \ + --offload_models 'Qwen/Qwen3-VL-4B-Instruct:*.safetensors,Qwen/Qwen-Image:vae/diffusion_pytorch_model.safetensors' \ + --task sft:train diff --git a/examples/krea2/model_training/special/split_training/validate.py b/examples/krea2/model_training/special/split_training/validate.py new file mode 100644 index 000000000..4c20347ba --- /dev/null +++ b/examples/krea2/model_training/special/split_training/validate.py @@ -0,0 +1,29 @@ +from pathlib import Path + + +def latest_checkpoint(directory): + checkpoints = list(Path(directory).glob("epoch-*.safetensors")) + if not checkpoints: + raise FileNotFoundError(f"No checkpoint found in {directory}") + return max(checkpoints, key=lambda path: int(path.stem.rsplit("-", 1)[-1])) + + +from diffsynth.pipelines.krea2 import Krea2Pipeline, ModelConfig +import torch + + +pipe = Krea2Pipeline.from_pretrained( + torch_dtype=torch.bfloat16, + device="cuda", + model_configs=[ + # For LoRA models trained on Krea-2-Raw, we recommend using them on Krea-2-Turbo. + ModelConfig(model_id="krea/Krea-2-Raw", origin_file_pattern="raw.safetensors"), + ModelConfig(model_id="Qwen/Qwen3-VL-4B-Instruct", origin_file_pattern="*.safetensors"), + ModelConfig(model_id="Qwen/Qwen-Image", origin_file_pattern="vae/diffusion_pytorch_model.safetensors"), + ], + tokenizer_config=ModelConfig(model_id="Qwen/Qwen3-VL-4B-Instruct", origin_file_pattern=""), +) +pipe.load_lora(pipe.dit, str(latest_checkpoint('./models/train/Krea-2-Raw_split'))) +prompt = "A dog" +image = pipe(prompt, seed=0, num_inference_steps=52, cfg_scale=4.5) +image.save('split_training_Krea-2-Raw.jpg') diff --git a/examples/lingbot_video/model_training/special/split_training/lingbot-video-dense-1.3b_t2v.sh b/examples/lingbot_video/model_training/special/split_training/lingbot-video-dense-1.3b_t2v.sh new file mode 100644 index 000000000..53a665f57 --- /dev/null +++ b/examples/lingbot_video/model_training/special/split_training/lingbot-video-dense-1.3b_t2v.sh @@ -0,0 +1,46 @@ +set -e + +modelscope download --dataset DiffSynth-Studio/diffsynth_example_dataset --include "lingbot_video/lingbot-video-dense-1.3b_t2v/*" --local_dir ./data/diffsynth_example_dataset + +# Stage 1: cache deterministic preprocessing outputs. +accelerate launch examples/lingbot_video/model_training/train.py \ + --dataset_base_path data/diffsynth_example_dataset/lingbot_video/lingbot-video-dense-1.3b_t2v \ + --dataset_metadata_path data/diffsynth_example_dataset/lingbot_video/lingbot-video-dense-1.3b_t2v/metadata.json \ + --data_file_keys video \ + --height 480 \ + --width 832 \ + --num_frames 81 \ + --dataset_repeat 1 \ + --model_id_with_origin_paths 'Robbyant/lingbot-video-dense-1.3b:transformer/diffusion_pytorch_model.safetensors,Qwen/Qwen3-VL-4B-Instruct:*.safetensors,Robbyant/lingbot-video-dense-1.3b:vae/diffusion_pytorch_model.safetensors' \ + --processor_path Qwen/Qwen3-VL-4B-Instruct: \ + --learning_rate 1e-4 \ + --num_epochs 5 \ + --remove_prefix_in_ckpt pipe.dit. \ + --output_path ./models/train/lingbot-video-dense-1.3b_t2v_split_cache \ + --lora_base_model dit \ + --lora_target_modules to_q,to_k,to_v,to_out \ + --lora_rank 32 \ + --use_gradient_checkpointing \ + --offload_models Robbyant/lingbot-video-dense-1.3b:transformer/diffusion_pytorch_model.safetensors \ + --task sft:data_process + +# Stage 2: train LoRA from the cached dataset. +accelerate launch examples/lingbot_video/model_training/train.py \ + --dataset_base_path ./models/train/lingbot-video-dense-1.3b_t2v_split_cache \ + --data_file_keys video \ + --height 480 \ + --width 832 \ + --num_frames 81 \ + --dataset_repeat 50 \ + --model_id_with_origin_paths 'Robbyant/lingbot-video-dense-1.3b:transformer/diffusion_pytorch_model.safetensors,Qwen/Qwen3-VL-4B-Instruct:*.safetensors,Robbyant/lingbot-video-dense-1.3b:vae/diffusion_pytorch_model.safetensors' \ + --processor_path Qwen/Qwen3-VL-4B-Instruct: \ + --learning_rate 1e-4 \ + --num_epochs 5 \ + --remove_prefix_in_ckpt pipe.dit. \ + --output_path ./models/train/lingbot-video-dense-1.3b_t2v_split \ + --lora_base_model dit \ + --lora_target_modules to_q,to_k,to_v,to_out \ + --lora_rank 32 \ + --use_gradient_checkpointing \ + --offload_models 'Qwen/Qwen3-VL-4B-Instruct:*.safetensors,Robbyant/lingbot-video-dense-1.3b:vae/diffusion_pytorch_model.safetensors' \ + --task sft:train diff --git a/examples/lingbot_video/model_training/special/split_training/validate.py b/examples/lingbot_video/model_training/special/split_training/validate.py new file mode 100644 index 000000000..ed5cb2bc8 --- /dev/null +++ b/examples/lingbot_video/model_training/special/split_training/validate.py @@ -0,0 +1,44 @@ +from pathlib import Path + + +def latest_checkpoint(directory): + checkpoints = list(Path(directory).glob("epoch-*.safetensors")) + if not checkpoints: + raise FileNotFoundError(f"No checkpoint found in {directory}") + return max(checkpoints, key=lambda path: int(path.stem.rsplit("-", 1)[-1])) + + +import torch +import json +from diffsynth.utils.data import save_video +from diffsynth.pipelines.lingbot_video import LingBotVideoPipeline, ModelConfig +from modelscope import dataset_snapshot_download + + +pipe = LingBotVideoPipeline.from_pretrained( + torch_dtype=torch.bfloat16, + device="cuda", + model_configs=[ + ModelConfig(model_id="Robbyant/lingbot-video-dense-1.3b", origin_file_pattern="transformer/diffusion_pytorch_model.safetensors"), + ModelConfig(model_id="Qwen/Qwen3-VL-4B-Instruct", origin_file_pattern="*.safetensors"), + ModelConfig(model_id="Robbyant/lingbot-video-dense-1.3b", origin_file_pattern="vae/diffusion_pytorch_model.safetensors"), + ], + processor_config=ModelConfig(model_id="Qwen/Qwen3-VL-4B-Instruct", origin_file_pattern=""), +) +pipe.load_lora(pipe.dit, str(latest_checkpoint('./models/train/lingbot-video-dense-1.3b_t2v_split')), alpha=1) +dataset_snapshot_download( + dataset_id="DiffSynth-Studio/diffsynth_example_dataset", + local_dir="data/diffsynth_example_dataset", + allow_file_pattern="lingbot_video/lingbot-video-dense-1.3b_t2v/*", +) +with open("data/diffsynth_example_dataset/lingbot_video/lingbot-video-dense-1.3b_t2v/t2v_example_1.json", "r", encoding="utf-8") as f: + caption = json.load(f) + +video = pipe( + prompt=caption, + negative_prompt=pipe.default_negative_prompt, + height=480, width=832, num_frames=81, + num_inference_steps=40, cfg_scale=3.0, + seed=0, +) +save_video(video, 'split_training_lingbot-video-dense-1.3b_t2v.mp4', fps=15, quality=10) diff --git a/examples/ltx2/model_training/special/split_training/LTX-2-T2AV-noaudio.sh b/examples/ltx2/model_training/special/split_training/LTX-2-T2AV-noaudio.sh new file mode 100644 index 000000000..76c599a18 --- /dev/null +++ b/examples/ltx2/model_training/special/split_training/LTX-2-T2AV-noaudio.sh @@ -0,0 +1,44 @@ +set -e + +modelscope download --dataset DiffSynth-Studio/diffsynth_example_dataset --include "ltx2/LTX-2-T2AV-noaudio/*" --local_dir ./data/diffsynth_example_dataset + +# Stage 1: cache deterministic preprocessing outputs. +accelerate launch examples/ltx2/model_training/train.py \ + --dataset_base_path data/diffsynth_example_dataset/ltx2/LTX-2-T2AV-noaudio \ + --dataset_metadata_path data/diffsynth_example_dataset/ltx2/LTX-2-T2AV-noaudio/metadata.csv \ + --height 256 \ + --width 384 \ + --num_frames 25 \ + --dataset_repeat 1 \ + --model_id_with_origin_paths 'DiffSynth-Studio/LTX-2-Repackage:transformer.safetensors,DiffSynth-Studio/LTX-2-Repackage:text_encoder_post_modules.safetensors,DiffSynth-Studio/LTX-2-Repackage:video_vae_encoder.safetensors,DiffSynth-Studio/LTX-2-Repackage:audio_vae_encoder.safetensors,google/gemma-3-12b-it-qat-q4_0-unquantized:model-*.safetensors' \ + --learning_rate 1e-4 \ + --num_epochs 5 \ + --remove_prefix_in_ckpt pipe.dit. \ + --output_path ./models/train/LTX2-T2AV-noaudio_split_cache \ + --lora_base_model dit \ + --lora_target_modules to_k,to_q,to_v,to_out.0 \ + --lora_rank 32 \ + --use_gradient_checkpointing \ + --find_unused_parameters \ + --offload_models DiffSynth-Studio/LTX-2-Repackage:transformer.safetensors \ + --task sft:data_process + +# Stage 2: train LoRA from the cached dataset. +accelerate launch examples/ltx2/model_training/train.py \ + --dataset_base_path ./models/train/LTX2-T2AV-noaudio_split_cache \ + --height 256 \ + --width 384 \ + --num_frames 25 \ + --dataset_repeat 100 \ + --model_id_with_origin_paths 'DiffSynth-Studio/LTX-2-Repackage:transformer.safetensors,DiffSynth-Studio/LTX-2-Repackage:text_encoder_post_modules.safetensors,DiffSynth-Studio/LTX-2-Repackage:video_vae_encoder.safetensors,DiffSynth-Studio/LTX-2-Repackage:audio_vae_encoder.safetensors,google/gemma-3-12b-it-qat-q4_0-unquantized:model-*.safetensors' \ + --learning_rate 1e-4 \ + --num_epochs 5 \ + --remove_prefix_in_ckpt pipe.dit. \ + --output_path ./models/train/LTX2-T2AV-noaudio_split \ + --lora_base_model dit \ + --lora_target_modules to_k,to_q,to_v,to_out.0 \ + --lora_rank 32 \ + --use_gradient_checkpointing \ + --find_unused_parameters \ + --offload_models 'DiffSynth-Studio/LTX-2-Repackage:text_encoder_post_modules.safetensors,DiffSynth-Studio/LTX-2-Repackage:video_vae_encoder.safetensors,DiffSynth-Studio/LTX-2-Repackage:audio_vae_encoder.safetensors,google/gemma-3-12b-it-qat-q4_0-unquantized:model-*.safetensors' \ + --task sft:train diff --git a/examples/ltx2/model_training/special/split_training/validate.py b/examples/ltx2/model_training/special/split_training/validate.py new file mode 100644 index 000000000..a1b7c8dad --- /dev/null +++ b/examples/ltx2/model_training/special/split_training/validate.py @@ -0,0 +1,58 @@ +from pathlib import Path + + +def latest_checkpoint(directory): + checkpoints = list(Path(directory).glob("epoch-*.safetensors")) + if not checkpoints: + raise FileNotFoundError(f"No checkpoint found in {directory}") + return max(checkpoints, key=lambda path: int(path.stem.rsplit("-", 1)[-1])) + + +import torch +from diffsynth.pipelines.ltx2_audio_video import LTX2AudioVideoPipeline, ModelConfig +from diffsynth.utils.data.media_io_ltx2 import write_video_audio_ltx2 + +vram_config = { + "offload_dtype": torch.bfloat16, + "offload_device": "cpu", + "onload_dtype": torch.bfloat16, + "onload_device": "cuda", + "preparing_dtype": torch.bfloat16, + "preparing_device": "cuda", + "computation_dtype": torch.bfloat16, + "computation_device": "cuda", +} +pipe = LTX2AudioVideoPipeline.from_pretrained( + torch_dtype=torch.bfloat16, + device="cuda", + model_configs=[ + ModelConfig(model_id="google/gemma-3-12b-it-qat-q4_0-unquantized", origin_file_pattern="model-*.safetensors", **vram_config), + ModelConfig(model_id="DiffSynth-Studio/LTX-2-Repackage", origin_file_pattern="transformer.safetensors", **vram_config), + ModelConfig(model_id="DiffSynth-Studio/LTX-2-Repackage", origin_file_pattern="text_encoder_post_modules.safetensors", **vram_config), + ModelConfig(model_id="DiffSynth-Studio/LTX-2-Repackage", origin_file_pattern="video_vae_decoder.safetensors", **vram_config), + ModelConfig(model_id="DiffSynth-Studio/LTX-2-Repackage", origin_file_pattern="audio_vae_decoder.safetensors", **vram_config), + ModelConfig(model_id="DiffSynth-Studio/LTX-2-Repackage", origin_file_pattern="audio_vocoder.safetensors", **vram_config), + ], + tokenizer_config=ModelConfig(model_id="google/gemma-3-12b-it-qat-q4_0-unquantized"), +) +pipe.load_lora(pipe.dit, str(latest_checkpoint('./models/train/LTX2-T2AV-noaudio_split'))) +prompt = "A beautiful sunset over the ocean." +negative_prompt = "blurry, out of focus, overexposed, underexposed, low contrast, washed out colors, excessive noise, grainy texture, poor lighting, flickering, motion blur, distorted proportions, unnatural skin tones, deformed facial features, asymmetrical face, missing facial features, extra limbs, disfigured hands, wrong hand count, artifacts around text, inconsistent perspective, camera shake, incorrect depth of field, background too sharp, background clutter, distracting reflections, harsh shadows, inconsistent lighting direction, color banding, cartoonish rendering, 3D CGI look, unrealistic materials, uncanny valley effect, incorrect ethnicity, wrong gender, exaggerated expressions, wrong gaze direction, mismatched lip sync, silent or muted audio, distorted voice, robotic voice, echo, background noise, off-sync audio, incorrect dialogue, added dialogue, repetitive speech, jittery movement, awkward pauses, incorrect timing, unnatural transitions, inconsistent framing, tilted camera, flat lighting, inconsistent tone, cinematic oversaturation, stylized filters, or AI artifacts." +height, width, num_frames = 512, 768, 121 +video, audio = pipe( + prompt=prompt, + negative_prompt=negative_prompt, + seed=43, + height=height, + width=width, + num_frames=num_frames, + tiled=True, + cfg_scale=4.0 +) +write_video_audio_ltx2( + video=video, + audio=audio, + output_path='split_training_LTX-2-T2AV-noaudio.mp4', + fps=24, + audio_sample_rate=24000, +) diff --git a/examples/minimax_h3/model_training/special/split_training/MiniMax-H3-FL2VA.sh b/examples/minimax_h3/model_training/special/split_training/MiniMax-H3-FL2VA.sh new file mode 100644 index 000000000..5d0926898 --- /dev/null +++ b/examples/minimax_h3/model_training/special/split_training/MiniMax-H3-FL2VA.sh @@ -0,0 +1,46 @@ +set -e + +modelscope download --dataset DiffSynth-Studio/diffsynth_example_dataset --include "minimax_h3/MiniMax-H3-FL2VA/*" --local_dir ./data/diffsynth_example_dataset + +# Stage 1: cache deterministic preprocessing outputs. +accelerate launch examples/minimax_h3/model_training/train.py \ + --dataset_base_path data/diffsynth_example_dataset/minimax_h3/MiniMax-H3-FL2VA \ + --dataset_metadata_path data/diffsynth_example_dataset/minimax_h3/MiniMax-H3-FL2VA/metadata.csv \ + --data_file_keys video,input_audio \ + --extra_inputs input_audio,input_image,end_image \ + --height 480 \ + --width 832 \ + --num_frames 124 \ + --dataset_repeat 1 \ + --model_id_with_origin_paths 'MiniMax/MiniMax-H3:FL2VA/transformer/model*.safetensors,MiniMax/MiniMax-H3:FL2VA/text_encoder/model*.safetensors,MiniMax/MiniMax-H3:FL2VA/video_vae/source/model.safetensors,MiniMax/MiniMax-H3:FL2VA/audio_vae/model.safetensors' \ + --learning_rate 1e-4 \ + --num_epochs 1 \ + --remove_prefix_in_ckpt pipe.dit. \ + --output_path ./models/train/MiniMax-H3-FL2VA-split-cache \ + --lora_base_model dit \ + --lora_target_modules attn.qkv_proj,attn.out_proj,mlp.fc1,mlp.fc2 \ + --lora_rank 32 \ + --use_gradient_checkpointing \ + --task sft:data_process \ + --offload_models 'MiniMax/MiniMax-H3:FL2VA/transformer/model*.safetensors' + +# Stage 2: train LoRA from the cached dataset. +accelerate launch examples/minimax_h3/model_training/train.py \ + --dataset_base_path ./models/train/MiniMax-H3-FL2VA-split-cache \ + --data_file_keys video,input_audio \ + --extra_inputs input_audio,input_image,end_image \ + --height 480 \ + --width 832 \ + --num_frames 124 \ + --dataset_repeat 100 \ + --model_id_with_origin_paths 'MiniMax/MiniMax-H3:FL2VA/transformer/model*.safetensors,MiniMax/MiniMax-H3:FL2VA/text_encoder/model*.safetensors,MiniMax/MiniMax-H3:FL2VA/video_vae/source/model.safetensors,MiniMax/MiniMax-H3:FL2VA/audio_vae/model.safetensors' \ + --learning_rate 1e-4 \ + --num_epochs 1 \ + --remove_prefix_in_ckpt pipe.dit. \ + --output_path ./models/train/MiniMax-H3-FL2VA-split \ + --lora_base_model dit \ + --lora_target_modules attn.qkv_proj,attn.out_proj,mlp.fc1,mlp.fc2 \ + --lora_rank 32 \ + --use_gradient_checkpointing \ + --task sft:train \ + --offload_models 'MiniMax/MiniMax-H3:FL2VA/text_encoder/model*.safetensors,MiniMax/MiniMax-H3:FL2VA/video_vae/source/model.safetensors,MiniMax/MiniMax-H3:FL2VA/audio_vae/model.safetensors' diff --git a/examples/minimax_h3/model_training/special/split_training/validate.py b/examples/minimax_h3/model_training/special/split_training/validate.py new file mode 100644 index 000000000..70c19349d --- /dev/null +++ b/examples/minimax_h3/model_training/special/split_training/validate.py @@ -0,0 +1,63 @@ +from pathlib import Path + + +def latest_checkpoint(directory): + checkpoints = list(Path(directory).glob("epoch-*.safetensors")) + if not checkpoints: + raise FileNotFoundError(f"No checkpoint found in {directory}") + return max(checkpoints, key=lambda path: int(path.stem.rsplit("-", 1)[-1])) + + +import torch +from diffsynth.pipelines.minimax_h3_audio_video import MiniMaxH3Pipeline, ModelConfig +from diffsynth.utils.data.audio_video import write_video_audio +from diffsynth.utils.data import VideoData +from modelscope import dataset_snapshot_download + +vram_config = { + "offload_dtype": torch.bfloat16, + "offload_device": "cpu", + "onload_dtype": torch.bfloat16, + "onload_device": "cpu", + "preparing_dtype": torch.bfloat16, + "preparing_device": "cuda", + "computation_dtype": torch.bfloat16, + "computation_device": "cuda", +} +pipe = MiniMaxH3Pipeline.from_pretrained( + torch_dtype=torch.bfloat16, + device="cuda", + model_configs=[ + ModelConfig(model_id="MiniMax/MiniMax-H3", origin_file_pattern="FL2VA/text_encoder/model*.safetensors", **vram_config), + ModelConfig(model_id="MiniMax/MiniMax-H3", origin_file_pattern="FL2VA/transformer/model*.safetensors", **vram_config), + ModelConfig(model_id="MiniMax/MiniMax-H3", origin_file_pattern="FL2VA/video_vae/source/model.safetensors", **vram_config), + ModelConfig(model_id="MiniMax/MiniMax-H3", origin_file_pattern="FL2VA/audio_vae/model.safetensors", **vram_config), + ], + vram_limit=torch.cuda.mem_get_info("cuda")[1] / (1024 ** 3) - 2, +) + +dataset_snapshot_download( + dataset_id="DiffSynth-Studio/diffsynth_example_dataset", + local_dir="data/diffsynth_example_dataset", + allow_file_pattern="minimax_h3/MiniMax-H3-FL2VA/*", +) +dataset_base_path = "data/diffsynth_example_dataset/minimax_h3/MiniMax-H3-FL2VA" +height, width, num_frames = 480, 832, 124 + +prompt = "A girl is very happy, she is speaking in english: “I enjoy working with Diffsynth-Studio, it's a perfect framework.”" + +frames = VideoData(f"{dataset_base_path}/video.mp4", height=height, width=width).raw_data() +pipe.clear_lora() +pipe.load_lora(pipe.dit, str(latest_checkpoint('./models/train/MiniMax-H3-FL2VA-split'))) +video, audio = pipe( + prompt=prompt, + height=height, width=width, num_frames=num_frames, + num_inference_steps=50, seed=0, + keyframes=[frames[0], frames[num_frames - 1]], + keyframe_indices=[0, -1], +) +write_video_audio( + video=video, audio=audio, output_path='split_training_MiniMax-H3-FL2VA.mp4', + fps=24, audio_sample_rate=pipe.audio_vae.sample_rate, +) +print("saved minimax_h3_fl2va_lora.mp4", "frames:", len(video), "audio:", tuple(audio.shape)) diff --git a/examples/mova/model_training/special/split_training/MOVA-360P-I2AV.sh b/examples/mova/model_training/special/split_training/MOVA-360P-I2AV.sh new file mode 100644 index 000000000..2e86f59eb --- /dev/null +++ b/examples/mova/model_training/special/split_training/MOVA-360P-I2AV.sh @@ -0,0 +1,50 @@ +set -e + +modelscope download --dataset DiffSynth-Studio/diffsynth_example_dataset --include "mova/MOVA-360P-I2AV/*" --local_dir ./data/diffsynth_example_dataset + +# Stage 1: cache deterministic preprocessing outputs. +accelerate launch examples/mova/model_training/train.py \ + --dataset_base_path data/diffsynth_example_dataset/mova/MOVA-360P-I2AV \ + --dataset_metadata_path data/diffsynth_example_dataset/mova/MOVA-360P-I2AV/metadata.csv \ + --data_file_keys video,input_audio \ + --extra_inputs input_audio,input_image \ + --height 352 \ + --width 640 \ + --num_frames 121 \ + --dataset_repeat 1 \ + --model_id_with_origin_paths 'openmoss/MOVA-360p:video_dit/diffusion_pytorch_model-*.safetensors,openmoss/MOVA-360p:audio_dit/diffusion_pytorch_model.safetensors,openmoss/MOVA-360p:dual_tower_bridge/diffusion_pytorch_model.safetensors,openmoss/MOVA-720p:audio_vae/diffusion_pytorch_model.safetensors,DiffSynth-Studio/Wan-Series-Converted-Safetensors:Wan2.1_VAE.safetensors,DiffSynth-Studio/Wan-Series-Converted-Safetensors:models_t5_umt5-xxl-enc-bf16.safetensors' \ + --learning_rate 1e-4 \ + --num_epochs 5 \ + --remove_prefix_in_ckpt pipe.video_dit. \ + --output_path ./models/train/MOVA-360p-I2AV_high_noise_split_cache \ + --lora_base_model video_dit \ + --lora_target_modules q,k,v,o,ffn.0,ffn.2 \ + --lora_rank 32 \ + --max_timestep_boundary 0.358 \ + --min_timestep_boundary 0 \ + --use_gradient_checkpointing \ + --offload_models 'openmoss/MOVA-360p:video_dit/diffusion_pytorch_model-*.safetensors,openmoss/MOVA-360p:audio_dit/diffusion_pytorch_model.safetensors,openmoss/MOVA-360p:dual_tower_bridge/diffusion_pytorch_model.safetensors' \ + --task sft:data_process + +# Stage 2: train LoRA from the cached dataset. +accelerate launch examples/mova/model_training/train.py \ + --dataset_base_path ./models/train/MOVA-360p-I2AV_high_noise_split_cache \ + --data_file_keys video,input_audio \ + --extra_inputs input_audio,input_image \ + --height 352 \ + --width 640 \ + --num_frames 121 \ + --dataset_repeat 100 \ + --model_id_with_origin_paths 'openmoss/MOVA-360p:video_dit/diffusion_pytorch_model-*.safetensors,openmoss/MOVA-360p:audio_dit/diffusion_pytorch_model.safetensors,openmoss/MOVA-360p:dual_tower_bridge/diffusion_pytorch_model.safetensors,openmoss/MOVA-720p:audio_vae/diffusion_pytorch_model.safetensors,DiffSynth-Studio/Wan-Series-Converted-Safetensors:Wan2.1_VAE.safetensors,DiffSynth-Studio/Wan-Series-Converted-Safetensors:models_t5_umt5-xxl-enc-bf16.safetensors' \ + --learning_rate 1e-4 \ + --num_epochs 5 \ + --remove_prefix_in_ckpt pipe.video_dit. \ + --output_path ./models/train/MOVA-360p-I2AV_high_noise_split \ + --lora_base_model video_dit \ + --lora_target_modules q,k,v,o,ffn.0,ffn.2 \ + --lora_rank 32 \ + --max_timestep_boundary 0.358 \ + --min_timestep_boundary 0 \ + --use_gradient_checkpointing \ + --offload_models openmoss/MOVA-720p:audio_vae/diffusion_pytorch_model.safetensors,DiffSynth-Studio/Wan-Series-Converted-Safetensors:Wan2.1_VAE.safetensors,DiffSynth-Studio/Wan-Series-Converted-Safetensors:models_t5_umt5-xxl-enc-bf16.safetensors \ + --task sft:train diff --git a/examples/mova/model_training/special/split_training/validate.py b/examples/mova/model_training/special/split_training/validate.py new file mode 100644 index 000000000..95fd1e714 --- /dev/null +++ b/examples/mova/model_training/special/split_training/validate.py @@ -0,0 +1,64 @@ +from pathlib import Path + + +def latest_checkpoint(directory): + checkpoints = list(Path(directory).glob("epoch-*.safetensors")) + if not checkpoints: + raise FileNotFoundError(f"No checkpoint found in {directory}") + return max(checkpoints, key=lambda path: int(path.stem.rsplit("-", 1)[-1])) + + +import torch +from PIL import Image +from diffsynth.pipelines.mova_audio_video import ModelConfig, MovaAudioVideoPipeline +from diffsynth.utils.data.audio_video import write_video_audio +from diffsynth.utils.data import VideoData + + +vram_config = { + "offload_dtype": torch.bfloat16, + "offload_device": "cpu", + "onload_dtype": torch.bfloat16, + "onload_device": "cuda", + "preparing_dtype": torch.bfloat16, + "preparing_device": "cuda", + "computation_dtype": torch.bfloat16, + "computation_device": "cuda", +} +pipe = MovaAudioVideoPipeline.from_pretrained( + torch_dtype=torch.bfloat16, + device="cuda", + model_configs=[ + ModelConfig(model_id="openmoss/MOVA-360p", origin_file_pattern="video_dit/diffusion_pytorch_model-*.safetensors", **vram_config), + ModelConfig(model_id="openmoss/MOVA-360p", origin_file_pattern="video_dit_2/diffusion_pytorch_model-*.safetensors", **vram_config), + ModelConfig(model_id="openmoss/MOVA-360p", origin_file_pattern="audio_dit/diffusion_pytorch_model.safetensors", **vram_config), + ModelConfig(model_id="openmoss/MOVA-360p", origin_file_pattern="dual_tower_bridge/diffusion_pytorch_model.safetensors", **vram_config), + ModelConfig(model_id="openmoss/MOVA-720p", origin_file_pattern="audio_vae/diffusion_pytorch_model.safetensors", **vram_config), + ModelConfig(model_id="DiffSynth-Studio/Wan-Series-Converted-Safetensors", origin_file_pattern="Wan2.1_VAE.safetensors", **vram_config), + ModelConfig(model_id="DiffSynth-Studio/Wan-Series-Converted-Safetensors", origin_file_pattern="models_t5_umt5-xxl-enc-bf16.safetensors", **vram_config), + ], + tokenizer_config=ModelConfig(model_id="openmoss/MOVA-720p", origin_file_pattern="tokenizer/"), +) +pipe.load_lora(pipe.video_dit, str(latest_checkpoint('./models/train/MOVA-360p-I2AV_high_noise_split'))) +negative_prompt = ( + "色调艳丽,过曝,静态,细节模糊不清,字幕,风格,作品,画作,画面,静止," + "整体发灰,最差质量,低质量,JPEG压缩残留,丑陋的,残缺的,多余的手指" +) +prompt = "A beautiful sunset over the ocean." +height, width, num_frames = 352, 640, 121 +frame_rate = 24 +input_image = VideoData("data/example_video_dataset/ltx2/video.mp4", height=height, width=width)[0] +# Image-to-video +video, audio = pipe( + prompt=prompt, + negative_prompt=negative_prompt, + height=height, + width=width, + num_frames=num_frames, + input_image=input_image, + num_inference_steps=50, + seed=0, + tiled=True, + frame_rate=frame_rate, +) +write_video_audio(video, audio, 'split_training_MOVA-360p-I2AV.mp4', fps=24, audio_sample_rate=pipe.audio_vae.sample_rate) diff --git a/examples/qwen_image/model_training/special/split_training/Qwen-Image-LoRA.sh b/examples/qwen_image/model_training/special/split_training/Qwen-Image-LoRA.sh index edb10f74d..cb17606d2 100644 --- a/examples/qwen_image/model_training/special/split_training/Qwen-Image-LoRA.sh +++ b/examples/qwen_image/model_training/special/split_training/Qwen-Image-LoRA.sh @@ -1,38 +1,42 @@ +set -e + modelscope download --dataset DiffSynth-Studio/diffsynth_example_dataset --include "qwen_image/Qwen-Image/*" --local_dir ./data/diffsynth_example_dataset +# Stage 1: cache deterministic preprocessing outputs. accelerate launch examples/qwen_image/model_training/train.py \ --dataset_base_path data/diffsynth_example_dataset/qwen_image/Qwen-Image \ --dataset_metadata_path data/diffsynth_example_dataset/qwen_image/Qwen-Image/metadata.csv \ --max_pixels 1048576 \ --dataset_repeat 1 \ - --model_id_with_origin_paths "Qwen/Qwen-Image:text_encoder/model*.safetensors,Qwen/Qwen-Image:vae/diffusion_pytorch_model.safetensors,Qwen/Qwen-Image:transformer/diffusion_pytorch_model*.safetensors" \ - --offload_models "Qwen/Qwen-Image:transformer/diffusion_pytorch_model*.safetensors" \ + --model_id_with_origin_paths 'Qwen/Qwen-Image:transformer/diffusion_pytorch_model*.safetensors,Qwen/Qwen-Image:text_encoder/model*.safetensors,Qwen/Qwen-Image:vae/diffusion_pytorch_model.safetensors' \ --learning_rate 1e-4 \ --num_epochs 5 \ - --remove_prefix_in_ckpt "pipe.dit." \ - --output_path "./models/train/Qwen-Image-LoRA-splited-cache" \ - --lora_base_model "dit" \ - --lora_target_modules "to_q,to_k,to_v,add_q_proj,add_k_proj,add_v_proj,to_out.0,to_add_out,img_mlp.net.2,img_mod.1,txt_mlp.net.2,txt_mod.1" \ + --remove_prefix_in_ckpt pipe.dit. \ + --output_path ./models/train/Qwen-Image-LoRA-splited-cache \ + --lora_base_model dit \ + --lora_target_modules to_q,to_k,to_v,add_q_proj,add_k_proj,add_v_proj,to_out.0,to_add_out,img_mlp.net.2,img_mod.1,txt_mlp.net.2,txt_mod.1 \ --lora_rank 32 \ --use_gradient_checkpointing \ --dataset_num_workers 8 \ --find_unused_parameters \ - --task "sft:data_process" + --offload_models 'Qwen/Qwen-Image:transformer/diffusion_pytorch_model*.safetensors' \ + --task sft:data_process +# Stage 2: train LoRA from the cached dataset. accelerate launch examples/qwen_image/model_training/train.py \ - --dataset_base_path "./models/train/Qwen-Image-LoRA-splited-cache" \ + --dataset_base_path ./models/train/Qwen-Image-LoRA-splited-cache \ --max_pixels 1048576 \ --dataset_repeat 50 \ - --model_id_with_origin_paths "Qwen/Qwen-Image:text_encoder/model*.safetensors,Qwen/Qwen-Image:vae/diffusion_pytorch_model.safetensors,Qwen/Qwen-Image:transformer/diffusion_pytorch_model*.safetensors" \ - --offload_models "Qwen/Qwen-Image:text_encoder/model*.safetensors,Qwen/Qwen-Image:vae/diffusion_pytorch_model.safetensors" \ + --model_id_with_origin_paths 'Qwen/Qwen-Image:transformer/diffusion_pytorch_model*.safetensors,Qwen/Qwen-Image:text_encoder/model*.safetensors,Qwen/Qwen-Image:vae/diffusion_pytorch_model.safetensors' \ --learning_rate 1e-4 \ --num_epochs 5 \ - --remove_prefix_in_ckpt "pipe.dit." \ - --output_path "./models/train/Qwen-Image-LoRA-splited" \ - --lora_base_model "dit" \ - --lora_target_modules "to_q,to_k,to_v,add_q_proj,add_k_proj,add_v_proj,to_out.0,to_add_out,img_mlp.net.2,img_mod.1,txt_mlp.net.2,txt_mod.1" \ + --remove_prefix_in_ckpt pipe.dit. \ + --output_path ./models/train/Qwen-Image-LoRA-splited \ + --lora_base_model dit \ + --lora_target_modules to_q,to_k,to_v,add_q_proj,add_k_proj,add_v_proj,to_out.0,to_add_out,img_mlp.net.2,img_mod.1,txt_mlp.net.2,txt_mod.1 \ --lora_rank 32 \ --use_gradient_checkpointing \ --dataset_num_workers 8 \ --find_unused_parameters \ - --task "sft:train" + --offload_models 'Qwen/Qwen-Image:text_encoder/model*.safetensors,Qwen/Qwen-Image:vae/diffusion_pytorch_model.safetensors' \ + --task sft:train diff --git a/examples/qwen_image/model_training/special/split_training/validate.py b/examples/qwen_image/model_training/special/split_training/validate.py index a2f9e07a9..ab3516c46 100644 --- a/examples/qwen_image/model_training/special/split_training/validate.py +++ b/examples/qwen_image/model_training/special/split_training/validate.py @@ -1,3 +1,13 @@ +from pathlib import Path + + +def latest_checkpoint(directory): + checkpoints = list(Path(directory).glob("epoch-*.safetensors")) + if not checkpoints: + raise FileNotFoundError(f"No checkpoint found in {directory}") + return max(checkpoints, key=lambda path: int(path.stem.rsplit("-", 1)[-1])) + + from diffsynth.pipelines.qwen_image import QwenImagePipeline, ModelConfig import torch @@ -12,7 +22,7 @@ ], tokenizer_config=ModelConfig(model_id="Qwen/Qwen-Image", origin_file_pattern="tokenizer/"), ) -pipe.load_lora(pipe.dit, "models/train/Qwen-Image-LoRA-splited/epoch-4.safetensors") +pipe.load_lora(pipe.dit, str(latest_checkpoint('./models/train/Qwen-Image-LoRA-splited'))) prompt = "a dog" image = pipe(prompt, seed=0) -image.save("image.jpg") +image.save('split_training_Qwen-Image.jpg') diff --git a/examples/stable_diffusion/model_training/special/split_training/stable-diffusion-v1-5.sh b/examples/stable_diffusion/model_training/special/split_training/stable-diffusion-v1-5.sh new file mode 100644 index 000000000..3efdf00b6 --- /dev/null +++ b/examples/stable_diffusion/model_training/special/split_training/stable-diffusion-v1-5.sh @@ -0,0 +1,40 @@ +set -e + +modelscope download --dataset DiffSynth-Studio/diffsynth_example_dataset --include "stable_diffusion/stable-diffusion-v1-5/*" --local_dir ./data/diffsynth_example_dataset + +# Stage 1: cache deterministic preprocessing outputs. +accelerate launch examples/stable_diffusion/model_training/train.py \ + --dataset_base_path data/diffsynth_example_dataset/stable_diffusion/stable-diffusion-v1-5 \ + --dataset_metadata_path data/diffsynth_example_dataset/stable_diffusion/stable-diffusion-v1-5/metadata.csv \ + --height 512 \ + --width 512 \ + --dataset_repeat 1 \ + --model_id_with_origin_paths AI-ModelScope/stable-diffusion-v1-5:text_encoder/model.safetensors,AI-ModelScope/stable-diffusion-v1-5:unet/diffusion_pytorch_model.safetensors,AI-ModelScope/stable-diffusion-v1-5:vae/diffusion_pytorch_model.safetensors \ + --learning_rate 1e-4 \ + --num_epochs 5 \ + --remove_prefix_in_ckpt pipe.unet. \ + --output_path ./models/train/stable-diffusion-v1-5_split_cache \ + --lora_base_model unet \ + --lora_target_modules '' \ + --lora_rank 32 \ + --use_gradient_checkpointing \ + --offload_models AI-ModelScope/stable-diffusion-v1-5:unet/diffusion_pytorch_model.safetensors \ + --task sft:data_process + +# Stage 2: train LoRA from the cached dataset. +accelerate launch examples/stable_diffusion/model_training/train.py \ + --dataset_base_path ./models/train/stable-diffusion-v1-5_split_cache \ + --height 512 \ + --width 512 \ + --dataset_repeat 50 \ + --model_id_with_origin_paths AI-ModelScope/stable-diffusion-v1-5:text_encoder/model.safetensors,AI-ModelScope/stable-diffusion-v1-5:unet/diffusion_pytorch_model.safetensors,AI-ModelScope/stable-diffusion-v1-5:vae/diffusion_pytorch_model.safetensors \ + --learning_rate 1e-4 \ + --num_epochs 5 \ + --remove_prefix_in_ckpt pipe.unet. \ + --output_path ./models/train/stable-diffusion-v1-5_split \ + --lora_base_model unet \ + --lora_target_modules '' \ + --lora_rank 32 \ + --use_gradient_checkpointing \ + --offload_models AI-ModelScope/stable-diffusion-v1-5:text_encoder/model.safetensors,AI-ModelScope/stable-diffusion-v1-5:vae/diffusion_pytorch_model.safetensors \ + --task sft:train diff --git a/examples/stable_diffusion/model_training/special/split_training/validate.py b/examples/stable_diffusion/model_training/special/split_training/validate.py new file mode 100644 index 000000000..501f3bb5c --- /dev/null +++ b/examples/stable_diffusion/model_training/special/split_training/validate.py @@ -0,0 +1,36 @@ +from pathlib import Path + + +def latest_checkpoint(directory): + checkpoints = list(Path(directory).glob("epoch-*.safetensors")) + if not checkpoints: + raise FileNotFoundError(f"No checkpoint found in {directory}") + return max(checkpoints, key=lambda path: int(path.stem.rsplit("-", 1)[-1])) + + +import torch +from diffsynth.core import ModelConfig +from diffsynth.pipelines.stable_diffusion import StableDiffusionPipeline + +pipe = StableDiffusionPipeline.from_pretrained( + torch_dtype=torch.float32, + model_configs=[ + ModelConfig(model_id="AI-ModelScope/stable-diffusion-v1-5", origin_file_pattern="text_encoder/model.safetensors"), + ModelConfig(model_id="AI-ModelScope/stable-diffusion-v1-5", origin_file_pattern="unet/diffusion_pytorch_model.safetensors"), + ModelConfig(model_id="AI-ModelScope/stable-diffusion-v1-5", origin_file_pattern="vae/diffusion_pytorch_model.safetensors"), + ], + tokenizer_config=ModelConfig(model_id="AI-ModelScope/stable-diffusion-v1-5", origin_file_pattern="tokenizer/"), +) +pipe.load_lora(pipe.unet, str(latest_checkpoint('./models/train/stable-diffusion-v1-5_split'))) + +image = pipe( + prompt="a dog", + negative_prompt="blurry, low quality, deformed", + cfg_scale=7.5, + height=512, + width=512, + seed=42, + rand_device="cuda", + num_inference_steps=50, +) +image.save('split_training_stable-diffusion-v1-5.jpg') diff --git a/examples/stable_diffusion_xl/model_training/special/split_training/stable-diffusion-xl-base-1.0.sh b/examples/stable_diffusion_xl/model_training/special/split_training/stable-diffusion-xl-base-1.0.sh new file mode 100644 index 000000000..65129fdc5 --- /dev/null +++ b/examples/stable_diffusion_xl/model_training/special/split_training/stable-diffusion-xl-base-1.0.sh @@ -0,0 +1,42 @@ +set -e + +modelscope download --dataset DiffSynth-Studio/diffsynth_example_dataset --include "stable_diffusion_xl/stable-diffusion-xl-base-1.0/*" --local_dir ./data/diffsynth_example_dataset + +# Stage 1: cache deterministic preprocessing outputs. +accelerate launch examples/stable_diffusion_xl/model_training/train.py \ + --dataset_base_path data/diffsynth_example_dataset/stable_diffusion_xl/stable-diffusion-xl-base-1.0 \ + --dataset_metadata_path data/diffsynth_example_dataset/stable_diffusion_xl/stable-diffusion-xl-base-1.0/metadata.csv \ + --height 1024 \ + --width 1024 \ + --dataset_repeat 1 \ + --model_id_with_origin_paths stabilityai/stable-diffusion-xl-base-1.0:text_encoder/model.safetensors,stabilityai/stable-diffusion-xl-base-1.0:text_encoder_2/model.safetensors,stabilityai/stable-diffusion-xl-base-1.0:unet/diffusion_pytorch_model.safetensors,stabilityai/stable-diffusion-xl-base-1.0:vae/diffusion_pytorch_model.safetensors \ + --learning_rate 1e-4 \ + --num_epochs 5 \ + --remove_prefix_in_ckpt pipe.unet. \ + --output_path ./models/train/stable-diffusion-xl-base-1.0_split_cache \ + --lora_base_model unet \ + --lora_target_modules mid_block.attentions.0.proj_in,mid_block.attentions.0.proj_out,down_blocks.1.attentions.0.proj_in,down_blocks.1.attentions.0.proj_out,down_blocks.1.attentions.0.transformer_blocks.0.attn1.to_k,down_blocks.1.attentions.0.transformer_blocks.0.attn1.to_out.0,down_blocks.1.attentions.0.transformer_blocks.0.attn1.to_q,down_blocks.1.attentions.0.transformer_blocks.0.attn1.to_v,down_blocks.1.attentions.0.transformer_blocks.0.attn2.to_k,down_blocks.1.attentions.0.transformer_blocks.0.attn2.to_out.0,down_blocks.1.attentions.0.transformer_blocks.0.attn2.to_q,down_blocks.1.attentions.0.transformer_blocks.0.attn2.to_v,down_blocks.1.attentions.0.transformer_blocks.0.ff.net.0.proj,down_blocks.1.attentions.0.transformer_blocks.0.ff.net.2,down_blocks.1.attentions.0.transformer_blocks.1.attn1.to_k,down_blocks.1.attentions.0.transformer_blocks.1.attn1.to_out.0,down_blocks.1.attentions.0.transformer_blocks.1.attn1.to_q,down_blocks.1.attentions.0.transformer_blocks.1.attn1.to_v,down_blocks.1.attentions.0.transformer_blocks.1.attn2.to_k,down_blocks.1.attentions.0.transformer_blocks.1.attn2.to_out.0,down_blocks.1.attentions.0.transformer_blocks.1.attn2.to_q,down_blocks.1.attentions.0.transformer_blocks.1.attn2.to_v,down_blocks.1.attentions.0.transformer_blocks.1.ff.net.0.proj,down_blocks.1.attentions.0.transformer_blocks.1.ff.net.2,down_blocks.1.attentions.1.proj_in,down_blocks.1.attentions.1.proj_out,down_blocks.1.attentions.1.transformer_blocks.0.attn1.to_k,down_blocks.1.attentions.1.transformer_blocks.0.attn1.to_out.0,down_blocks.1.attentions.1.transformer_blocks.0.attn1.to_q,down_blocks.1.attentions.1.transformer_blocks.0.attn1.to_v,down_blocks.1.attentions.1.transformer_blocks.0.attn2.to_k,down_blocks.1.attentions.1.transformer_blocks.0.attn2.to_out.0,down_blocks.1.attentions.1.transformer_blocks.0.attn2.to_q,down_blocks.1.attentions.1.transformer_blocks.0.attn2.to_v,down_blocks.1.attentions.1.transformer_blocks.0.ff.net.0.proj,down_blocks.1.attentions.1.transformer_blocks.0.ff.net.2,down_blocks.1.attentions.1.transformer_blocks.1.attn1.to_k,down_blocks.1.attentions.1.transformer_blocks.1.attn1.to_out.0,down_blocks.1.attentions.1.transformer_blocks.1.attn1.to_q,down_blocks.1.attentions.1.transformer_blocks.1.attn1.to_v,down_blocks.1.attentions.1.transformer_blocks.1.attn2.to_k,down_blocks.1.attentions.1.transformer_blocks.1.attn2.to_out.0,down_blocks.1.attentions.1.transformer_blocks.1.attn2.to_q,down_blocks.1.attentions.1.transformer_blocks.1.attn2.to_v,down_blocks.1.attentions.1.transformer_blocks.1.ff.net.0.proj,down_blocks.1.attentions.1.transformer_blocks.1.ff.net.2,down_blocks.2.attentions.0.proj_in,down_blocks.2.attentions.0.proj_out,down_blocks.2.attentions.0.transformer_blocks.0.attn1.to_k,down_blocks.2.attentions.0.transformer_blocks.0.attn1.to_out.0,down_blocks.2.attentions.0.transformer_blocks.0.attn1.to_q,down_blocks.2.attentions.0.transformer_blocks.0.attn1.to_v,down_blocks.2.attentions.0.transformer_blocks.0.attn2.to_k,down_blocks.2.attentions.0.transformer_blocks.0.attn2.to_out.0,down_blocks.2.attentions.0.transformer_blocks.0.attn2.to_q,down_blocks.2.attentions.0.transformer_blocks.0.attn2.to_v,down_blocks.2.attentions.0.transformer_blocks.0.ff.net.0.proj,down_blocks.2.attentions.0.transformer_blocks.0.ff.net.2,down_blocks.2.attentions.0.transformer_blocks.1.attn1.to_k,down_blocks.2.attentions.0.transformer_blocks.1.attn1.to_out.0,down_blocks.2.attentions.0.transformer_blocks.1.attn1.to_q,down_blocks.2.attentions.0.transformer_blocks.1.attn1.to_v,down_blocks.2.attentions.0.transformer_blocks.1.attn2.to_k,down_blocks.2.attentions.0.transformer_blocks.1.attn2.to_out.0,down_blocks.2.attentions.0.transformer_blocks.1.attn2.to_q,down_blocks.2.attentions.0.transformer_blocks.1.attn2.to_v,down_blocks.2.attentions.0.transformer_blocks.1.ff.net.0.proj,down_blocks.2.attentions.0.transformer_blocks.1.ff.net.2,down_blocks.2.attentions.0.transformer_blocks.2.attn1.to_k,down_blocks.2.attentions.0.transformer_blocks.2.attn1.to_out.0,down_blocks.2.attentions.0.transformer_blocks.2.attn1.to_q,down_blocks.2.attentions.0.transformer_blocks.2.attn1.to_v,down_blocks.2.attentions.0.transformer_blocks.2.attn2.to_k,down_blocks.2.attentions.0.transformer_blocks.2.attn2.to_out.0,down_blocks.2.attentions.0.transformer_blocks.2.attn2.to_q,down_blocks.2.attentions.0.transformer_blocks.2.attn2.to_v,down_blocks.2.attentions.0.transformer_blocks.2.ff.net.0.proj,down_blocks.2.attentions.0.transformer_blocks.2.ff.net.2,down_blocks.2.attentions.0.transformer_blocks.3.attn1.to_k,down_blocks.2.attentions.0.transformer_blocks.3.attn1.to_out.0,down_blocks.2.attentions.0.transformer_blocks.3.attn1.to_q,down_blocks.2.attentions.0.transformer_blocks.3.attn1.to_v,down_blocks.2.attentions.0.transformer_blocks.3.attn2.to_k,down_blocks.2.attentions.0.transformer_blocks.3.attn2.to_out.0,down_blocks.2.attentions.0.transformer_blocks.3.attn2.to_q,down_blocks.2.attentions.0.transformer_blocks.3.attn2.to_v,down_blocks.2.attentions.0.transformer_blocks.3.ff.net.0.proj,down_blocks.2.attentions.0.transformer_blocks.3.ff.net.2,down_blocks.2.attentions.0.transformer_blocks.4.attn1.to_k,down_blocks.2.attentions.0.transformer_blocks.4.attn1.to_out.0,down_blocks.2.attentions.0.transformer_blocks.4.attn1.to_q,down_blocks.2.attentions.0.transformer_blocks.4.attn1.to_v,down_blocks.2.attentions.0.transformer_blocks.4.attn2.to_k,down_blocks.2.attentions.0.transformer_blocks.4.attn2.to_out.0,down_blocks.2.attentions.0.transformer_blocks.4.attn2.to_q,down_blocks.2.attentions.0.transformer_blocks.4.attn2.to_v,down_blocks.2.attentions.0.transformer_blocks.4.ff.net.0.proj,down_blocks.2.attentions.0.transformer_blocks.4.ff.net.2,down_blocks.2.attentions.0.transformer_blocks.5.attn1.to_k,down_blocks.2.attentions.0.transformer_blocks.5.attn1.to_out.0,down_blocks.2.attentions.0.transformer_blocks.5.attn1.to_q,down_blocks.2.attentions.0.transformer_blocks.5.attn1.to_v,down_blocks.2.attentions.0.transformer_blocks.5.attn2.to_k,down_blocks.2.attentions.0.transformer_blocks.5.attn2.to_out.0,down_blocks.2.attentions.0.transformer_blocks.5.attn2.to_q,down_blocks.2.attentions.0.transformer_blocks.5.attn2.to_v,down_blocks.2.attentions.0.transformer_blocks.5.ff.net.0.proj,down_blocks.2.attentions.0.transformer_blocks.5.ff.net.2,down_blocks.2.attentions.0.transformer_blocks.6.attn1.to_k,down_blocks.2.attentions.0.transformer_blocks.6.attn1.to_out.0,down_blocks.2.attentions.0.transformer_blocks.6.attn1.to_q,down_blocks.2.attentions.0.transformer_blocks.6.attn1.to_v,down_blocks.2.attentions.0.transformer_blocks.6.attn2.to_k,down_blocks.2.attentions.0.transformer_blocks.6.attn2.to_out.0,down_blocks.2.attentions.0.transformer_blocks.6.attn2.to_q,down_blocks.2.attentions.0.transformer_blocks.6.attn2.to_v,down_blocks.2.attentions.0.transformer_blocks.6.ff.net.0.proj,down_blocks.2.attentions.0.transformer_blocks.6.ff.net.2,down_blocks.2.attentions.0.transformer_blocks.7.attn1.to_k,down_blocks.2.attentions.0.transformer_blocks.7.attn1.to_out.0,down_blocks.2.attentions.0.transformer_blocks.7.attn1.to_q,down_blocks.2.attentions.0.transformer_blocks.7.attn1.to_v,down_blocks.2.attentions.0.transformer_blocks.7.attn2.to_k,down_blocks.2.attentions.0.transformer_blocks.7.attn2.to_out.0,down_blocks.2.attentions.0.transformer_blocks.7.attn2.to_q,down_blocks.2.attentions.0.transformer_blocks.7.attn2.to_v,down_blocks.2.attentions.0.transformer_blocks.7.ff.net.0.proj,down_blocks.2.attentions.0.transformer_blocks.7.ff.net.2,down_blocks.2.attentions.0.transformer_blocks.8.attn1.to_k,down_blocks.2.attentions.0.transformer_blocks.8.attn1.to_out.0,down_blocks.2.attentions.0.transformer_blocks.8.attn1.to_q,down_blocks.2.attentions.0.transformer_blocks.8.attn1.to_v,down_blocks.2.attentions.0.transformer_blocks.8.attn2.to_k,down_blocks.2.attentions.0.transformer_blocks.8.attn2.to_out.0,down_blocks.2.attentions.0.transformer_blocks.8.attn2.to_q,down_blocks.2.attentions.0.transformer_blocks.8.attn2.to_v,down_blocks.2.attentions.0.transformer_blocks.8.ff.net.0.proj,down_blocks.2.attentions.0.transformer_blocks.8.ff.net.2,down_blocks.2.attentions.0.transformer_blocks.9.attn1.to_k,down_blocks.2.attentions.0.transformer_blocks.9.attn1.to_out.0,down_blocks.2.attentions.0.transformer_blocks.9.attn1.to_q,down_blocks.2.attentions.0.transformer_blocks.9.attn1.to_v,down_blocks.2.attentions.0.transformer_blocks.9.attn2.to_k,down_blocks.2.attentions.0.transformer_blocks.9.attn2.to_out.0,down_blocks.2.attentions.0.transformer_blocks.9.attn2.to_q,down_blocks.2.attentions.0.transformer_blocks.9.attn2.to_v,down_blocks.2.attentions.0.transformer_blocks.9.ff.net.0.proj,down_blocks.2.attentions.0.transformer_blocks.9.ff.net.2,down_blocks.2.attentions.1.proj_in,down_blocks.2.attentions.1.proj_out,down_blocks.2.attentions.1.transformer_blocks.0.attn1.to_k,down_blocks.2.attentions.1.transformer_blocks.0.attn1.to_out.0,down_blocks.2.attentions.1.transformer_blocks.0.attn1.to_q,down_blocks.2.attentions.1.transformer_blocks.0.attn1.to_v,down_blocks.2.attentions.1.transformer_blocks.0.attn2.to_k,down_blocks.2.attentions.1.transformer_blocks.0.attn2.to_out.0,down_blocks.2.attentions.1.transformer_blocks.0.attn2.to_q,down_blocks.2.attentions.1.transformer_blocks.0.attn2.to_v,down_blocks.2.attentions.1.transformer_blocks.0.ff.net.0.proj,down_blocks.2.attentions.1.transformer_blocks.0.ff.net.2,down_blocks.2.attentions.1.transformer_blocks.1.attn1.to_k,down_blocks.2.attentions.1.transformer_blocks.1.attn1.to_out.0,down_blocks.2.attentions.1.transformer_blocks.1.attn1.to_q,down_blocks.2.attentions.1.transformer_blocks.1.attn1.to_v,down_blocks.2.attentions.1.transformer_blocks.1.attn2.to_k,down_blocks.2.attentions.1.transformer_blocks.1.attn2.to_out.0,down_blocks.2.attentions.1.transformer_blocks.1.attn2.to_q,down_blocks.2.attentions.1.transformer_blocks.1.attn2.to_v,down_blocks.2.attentions.1.transformer_blocks.1.ff.net.0.proj,down_blocks.2.attentions.1.transformer_blocks.1.ff.net.2,down_blocks.2.attentions.1.transformer_blocks.2.attn1.to_k,down_blocks.2.attentions.1.transformer_blocks.2.attn1.to_out.0,down_blocks.2.attentions.1.transformer_blocks.2.attn1.to_q,down_blocks.2.attentions.1.transformer_blocks.2.attn1.to_v,down_blocks.2.attentions.1.transformer_blocks.2.attn2.to_k,down_blocks.2.attentions.1.transformer_blocks.2.attn2.to_out.0,down_blocks.2.attentions.1.transformer_blocks.2.attn2.to_q,down_blocks.2.attentions.1.transformer_blocks.2.attn2.to_v,down_blocks.2.attentions.1.transformer_blocks.2.ff.net.0.proj,down_blocks.2.attentions.1.transformer_blocks.2.ff.net.2,down_blocks.2.attentions.1.transformer_blocks.3.attn1.to_k,down_blocks.2.attentions.1.transformer_blocks.3.attn1.to_out.0,down_blocks.2.attentions.1.transformer_blocks.3.attn1.to_q,down_blocks.2.attentions.1.transformer_blocks.3.attn1.to_v,down_blocks.2.attentions.1.transformer_blocks.3.attn2.to_k,down_blocks.2.attentions.1.transformer_blocks.3.attn2.to_out.0,down_blocks.2.attentions.1.transformer_blocks.3.attn2.to_q,down_blocks.2.attentions.1.transformer_blocks.3.attn2.to_v,down_blocks.2.attentions.1.transformer_blocks.3.ff.net.0.proj,down_blocks.2.attentions.1.transformer_blocks.3.ff.net.2,down_blocks.2.attentions.1.transformer_blocks.4.attn1.to_k,down_blocks.2.attentions.1.transformer_blocks.4.attn1.to_out.0,down_blocks.2.attentions.1.transformer_blocks.4.attn1.to_q,down_blocks.2.attentions.1.transformer_blocks.4.attn1.to_v,down_blocks.2.attentions.1.transformer_blocks.4.attn2.to_k,down_blocks.2.attentions.1.transformer_blocks.4.attn2.to_out.0,down_blocks.2.attentions.1.transformer_blocks.4.attn2.to_q,down_blocks.2.attentions.1.transformer_blocks.4.attn2.to_v,down_blocks.2.attentions.1.transformer_blocks.4.ff.net.0.proj,down_blocks.2.attentions.1.transformer_blocks.4.ff.net.2,down_blocks.2.attentions.1.transformer_blocks.5.attn1.to_k,down_blocks.2.attentions.1.transformer_blocks.5.attn1.to_out.0,down_blocks.2.attentions.1.transformer_blocks.5.attn1.to_q,down_blocks.2.attentions.1.transformer_blocks.5.attn1.to_v,down_blocks.2.attentions.1.transformer_blocks.5.attn2.to_k,down_blocks.2.attentions.1.transformer_blocks.5.attn2.to_out.0,down_blocks.2.attentions.1.transformer_blocks.5.attn2.to_q,down_blocks.2.attentions.1.transformer_blocks.5.attn2.to_v,down_blocks.2.attentions.1.transformer_blocks.5.ff.net.0.proj,down_blocks.2.attentions.1.transformer_blocks.5.ff.net.2,down_blocks.2.attentions.1.transformer_blocks.6.attn1.to_k,down_blocks.2.attentions.1.transformer_blocks.6.attn1.to_out.0,down_blocks.2.attentions.1.transformer_blocks.6.attn1.to_q,down_blocks.2.attentions.1.transformer_blocks.6.attn1.to_v,down_blocks.2.attentions.1.transformer_blocks.6.attn2.to_k,down_blocks.2.attentions.1.transformer_blocks.6.attn2.to_out.0,down_blocks.2.attentions.1.transformer_blocks.6.attn2.to_q,down_blocks.2.attentions.1.transformer_blocks.6.attn2.to_v,down_blocks.2.attentions.1.transformer_blocks.6.ff.net.0.proj,down_blocks.2.attentions.1.transformer_blocks.6.ff.net.2,down_blocks.2.attentions.1.transformer_blocks.7.attn1.to_k,down_blocks.2.attentions.1.transformer_blocks.7.attn1.to_out.0,down_blocks.2.attentions.1.transformer_blocks.7.attn1.to_q,down_blocks.2.attentions.1.transformer_blocks.7.attn1.to_v,down_blocks.2.attentions.1.transformer_blocks.7.attn2.to_k,down_blocks.2.attentions.1.transformer_blocks.7.attn2.to_out.0,down_blocks.2.attentions.1.transformer_blocks.7.attn2.to_q,down_blocks.2.attentions.1.transformer_blocks.7.attn2.to_v,down_blocks.2.attentions.1.transformer_blocks.7.ff.net.0.proj,down_blocks.2.attentions.1.transformer_blocks.7.ff.net.2,down_blocks.2.attentions.1.transformer_blocks.8.attn1.to_k,down_blocks.2.attentions.1.transformer_blocks.8.attn1.to_out.0,down_blocks.2.attentions.1.transformer_blocks.8.attn1.to_q,down_blocks.2.attentions.1.transformer_blocks.8.attn1.to_v,down_blocks.2.attentions.1.transformer_blocks.8.attn2.to_k,down_blocks.2.attentions.1.transformer_blocks.8.attn2.to_out.0,down_blocks.2.attentions.1.transformer_blocks.8.attn2.to_q,down_blocks.2.attentions.1.transformer_blocks.8.attn2.to_v,down_blocks.2.attentions.1.transformer_blocks.8.ff.net.0.proj,down_blocks.2.attentions.1.transformer_blocks.8.ff.net.2,down_blocks.2.attentions.1.transformer_blocks.9.attn1.to_k,down_blocks.2.attentions.1.transformer_blocks.9.attn1.to_out.0,down_blocks.2.attentions.1.transformer_blocks.9.attn1.to_q,down_blocks.2.attentions.1.transformer_blocks.9.attn1.to_v,down_blocks.2.attentions.1.transformer_blocks.9.attn2.to_k,down_blocks.2.attentions.1.transformer_blocks.9.attn2.to_out.0,down_blocks.2.attentions.1.transformer_blocks.9.attn2.to_q,down_blocks.2.attentions.1.transformer_blocks.9.attn2.to_v,down_blocks.2.attentions.1.transformer_blocks.9.ff.net.0.proj,down_blocks.2.attentions.1.transformer_blocks.9.ff.net.2,mid_block.attentions.0.transformer_blocks.0.attn1.to_k,mid_block.attentions.0.transformer_blocks.0.attn1.to_out.0,mid_block.attentions.0.transformer_blocks.0.attn1.to_q,mid_block.attentions.0.transformer_blocks.0.attn1.to_v,mid_block.attentions.0.transformer_blocks.0.attn2.to_k,mid_block.attentions.0.transformer_blocks.0.attn2.to_out.0,mid_block.attentions.0.transformer_blocks.0.attn2.to_q,mid_block.attentions.0.transformer_blocks.0.attn2.to_v,mid_block.attentions.0.transformer_blocks.0.ff.net.0.proj,mid_block.attentions.0.transformer_blocks.0.ff.net.2,mid_block.attentions.0.transformer_blocks.1.attn1.to_k,mid_block.attentions.0.transformer_blocks.1.attn1.to_out.0,mid_block.attentions.0.transformer_blocks.1.attn1.to_q,mid_block.attentions.0.transformer_blocks.1.attn1.to_v,mid_block.attentions.0.transformer_blocks.1.attn2.to_k,mid_block.attentions.0.transformer_blocks.1.attn2.to_out.0,mid_block.attentions.0.transformer_blocks.1.attn2.to_q,mid_block.attentions.0.transformer_blocks.1.attn2.to_v,mid_block.attentions.0.transformer_blocks.1.ff.net.0.proj,mid_block.attentions.0.transformer_blocks.1.ff.net.2,mid_block.attentions.0.transformer_blocks.2.attn1.to_k,mid_block.attentions.0.transformer_blocks.2.attn1.to_out.0,mid_block.attentions.0.transformer_blocks.2.attn1.to_q,mid_block.attentions.0.transformer_blocks.2.attn1.to_v,mid_block.attentions.0.transformer_blocks.2.attn2.to_k,mid_block.attentions.0.transformer_blocks.2.attn2.to_out.0,mid_block.attentions.0.transformer_blocks.2.attn2.to_q,mid_block.attentions.0.transformer_blocks.2.attn2.to_v,mid_block.attentions.0.transformer_blocks.2.ff.net.0.proj,mid_block.attentions.0.transformer_blocks.2.ff.net.2,mid_block.attentions.0.transformer_blocks.3.attn1.to_k,mid_block.attentions.0.transformer_blocks.3.attn1.to_out.0,mid_block.attentions.0.transformer_blocks.3.attn1.to_q,mid_block.attentions.0.transformer_blocks.3.attn1.to_v,mid_block.attentions.0.transformer_blocks.3.attn2.to_k,mid_block.attentions.0.transformer_blocks.3.attn2.to_out.0,mid_block.attentions.0.transformer_blocks.3.attn2.to_q,mid_block.attentions.0.transformer_blocks.3.attn2.to_v,mid_block.attentions.0.transformer_blocks.3.ff.net.0.proj,mid_block.attentions.0.transformer_blocks.3.ff.net.2,mid_block.attentions.0.transformer_blocks.4.attn1.to_k,mid_block.attentions.0.transformer_blocks.4.attn1.to_out.0,mid_block.attentions.0.transformer_blocks.4.attn1.to_q,mid_block.attentions.0.transformer_blocks.4.attn1.to_v,mid_block.attentions.0.transformer_blocks.4.attn2.to_k,mid_block.attentions.0.transformer_blocks.4.attn2.to_out.0,mid_block.attentions.0.transformer_blocks.4.attn2.to_q,mid_block.attentions.0.transformer_blocks.4.attn2.to_v,mid_block.attentions.0.transformer_blocks.4.ff.net.0.proj,mid_block.attentions.0.transformer_blocks.4.ff.net.2,mid_block.attentions.0.transformer_blocks.5.attn1.to_k,mid_block.attentions.0.transformer_blocks.5.attn1.to_out.0,mid_block.attentions.0.transformer_blocks.5.attn1.to_q,mid_block.attentions.0.transformer_blocks.5.attn1.to_v,mid_block.attentions.0.transformer_blocks.5.attn2.to_k,mid_block.attentions.0.transformer_blocks.5.attn2.to_out.0,mid_block.attentions.0.transformer_blocks.5.attn2.to_q,mid_block.attentions.0.transformer_blocks.5.attn2.to_v,mid_block.attentions.0.transformer_blocks.5.ff.net.0.proj,mid_block.attentions.0.transformer_blocks.5.ff.net.2,mid_block.attentions.0.transformer_blocks.6.attn1.to_k,mid_block.attentions.0.transformer_blocks.6.attn1.to_out.0,mid_block.attentions.0.transformer_blocks.6.attn1.to_q,mid_block.attentions.0.transformer_blocks.6.attn1.to_v,mid_block.attentions.0.transformer_blocks.6.attn2.to_k,mid_block.attentions.0.transformer_blocks.6.attn2.to_out.0,mid_block.attentions.0.transformer_blocks.6.attn2.to_q,mid_block.attentions.0.transformer_blocks.6.attn2.to_v,mid_block.attentions.0.transformer_blocks.6.ff.net.0.proj,mid_block.attentions.0.transformer_blocks.6.ff.net.2,mid_block.attentions.0.transformer_blocks.7.attn1.to_k,mid_block.attentions.0.transformer_blocks.7.attn1.to_out.0,mid_block.attentions.0.transformer_blocks.7.attn1.to_q,mid_block.attentions.0.transformer_blocks.7.attn1.to_v,mid_block.attentions.0.transformer_blocks.7.attn2.to_k,mid_block.attentions.0.transformer_blocks.7.attn2.to_out.0,mid_block.attentions.0.transformer_blocks.7.attn2.to_q,mid_block.attentions.0.transformer_blocks.7.attn2.to_v,mid_block.attentions.0.transformer_blocks.7.ff.net.0.proj,mid_block.attentions.0.transformer_blocks.7.ff.net.2,mid_block.attentions.0.transformer_blocks.8.attn1.to_k,mid_block.attentions.0.transformer_blocks.8.attn1.to_out.0,mid_block.attentions.0.transformer_blocks.8.attn1.to_q,mid_block.attentions.0.transformer_blocks.8.attn1.to_v,mid_block.attentions.0.transformer_blocks.8.attn2.to_k,mid_block.attentions.0.transformer_blocks.8.attn2.to_out.0,mid_block.attentions.0.transformer_blocks.8.attn2.to_q,mid_block.attentions.0.transformer_blocks.8.attn2.to_v,mid_block.attentions.0.transformer_blocks.8.ff.net.0.proj,mid_block.attentions.0.transformer_blocks.8.ff.net.2,mid_block.attentions.0.transformer_blocks.9.attn1.to_k,mid_block.attentions.0.transformer_blocks.9.attn1.to_out.0,mid_block.attentions.0.transformer_blocks.9.attn1.to_q,mid_block.attentions.0.transformer_blocks.9.attn1.to_v,mid_block.attentions.0.transformer_blocks.9.attn2.to_k,mid_block.attentions.0.transformer_blocks.9.attn2.to_out.0,mid_block.attentions.0.transformer_blocks.9.attn2.to_q,mid_block.attentions.0.transformer_blocks.9.attn2.to_v,mid_block.attentions.0.transformer_blocks.9.ff.net.0.proj,mid_block.attentions.0.transformer_blocks.9.ff.net.2,up_blocks.0.attentions.0.proj_in,up_blocks.0.attentions.0.proj_out,up_blocks.0.attentions.0.transformer_blocks.0.attn1.to_k,up_blocks.0.attentions.0.transformer_blocks.0.attn1.to_out.0,up_blocks.0.attentions.0.transformer_blocks.0.attn1.to_q,up_blocks.0.attentions.0.transformer_blocks.0.attn1.to_v,up_blocks.0.attentions.0.transformer_blocks.0.attn2.to_k,up_blocks.0.attentions.0.transformer_blocks.0.attn2.to_out.0,up_blocks.0.attentions.0.transformer_blocks.0.attn2.to_q,up_blocks.0.attentions.0.transformer_blocks.0.attn2.to_v,up_blocks.0.attentions.0.transformer_blocks.0.ff.net.0.proj,up_blocks.0.attentions.0.transformer_blocks.0.ff.net.2,up_blocks.0.attentions.0.transformer_blocks.1.attn1.to_k,up_blocks.0.attentions.0.transformer_blocks.1.attn1.to_out.0,up_blocks.0.attentions.0.transformer_blocks.1.attn1.to_q,up_blocks.0.attentions.0.transformer_blocks.1.attn1.to_v,up_blocks.0.attentions.0.transformer_blocks.1.attn2.to_k,up_blocks.0.attentions.0.transformer_blocks.1.attn2.to_out.0,up_blocks.0.attentions.0.transformer_blocks.1.attn2.to_q,up_blocks.0.attentions.0.transformer_blocks.1.attn2.to_v,up_blocks.0.attentions.0.transformer_blocks.1.ff.net.0.proj,up_blocks.0.attentions.0.transformer_blocks.1.ff.net.2,up_blocks.0.attentions.0.transformer_blocks.2.attn1.to_k,up_blocks.0.attentions.0.transformer_blocks.2.attn1.to_out.0,up_blocks.0.attentions.0.transformer_blocks.2.attn1.to_q,up_blocks.0.attentions.0.transformer_blocks.2.attn1.to_v,up_blocks.0.attentions.0.transformer_blocks.2.attn2.to_k,up_blocks.0.attentions.0.transformer_blocks.2.attn2.to_out.0,up_blocks.0.attentions.0.transformer_blocks.2.attn2.to_q,up_blocks.0.attentions.0.transformer_blocks.2.attn2.to_v,up_blocks.0.attentions.0.transformer_blocks.2.ff.net.0.proj,up_blocks.0.attentions.0.transformer_blocks.2.ff.net.2,up_blocks.0.attentions.0.transformer_blocks.3.attn1.to_k,up_blocks.0.attentions.0.transformer_blocks.3.attn1.to_out.0,up_blocks.0.attentions.0.transformer_blocks.3.attn1.to_q,up_blocks.0.attentions.0.transformer_blocks.3.attn1.to_v,up_blocks.0.attentions.0.transformer_blocks.3.attn2.to_k,up_blocks.0.attentions.0.transformer_blocks.3.attn2.to_out.0,up_blocks.0.attentions.0.transformer_blocks.3.attn2.to_q,up_blocks.0.attentions.0.transformer_blocks.3.attn2.to_v,up_blocks.0.attentions.0.transformer_blocks.3.ff.net.0.proj,up_blocks.0.attentions.0.transformer_blocks.3.ff.net.2,up_blocks.0.attentions.0.transformer_blocks.4.attn1.to_k,up_blocks.0.attentions.0.transformer_blocks.4.attn1.to_out.0,up_blocks.0.attentions.0.transformer_blocks.4.attn1.to_q,up_blocks.0.attentions.0.transformer_blocks.4.attn1.to_v,up_blocks.0.attentions.0.transformer_blocks.4.attn2.to_k,up_blocks.0.attentions.0.transformer_blocks.4.attn2.to_out.0,up_blocks.0.attentions.0.transformer_blocks.4.attn2.to_q,up_blocks.0.attentions.0.transformer_blocks.4.attn2.to_v,up_blocks.0.attentions.0.transformer_blocks.4.ff.net.0.proj,up_blocks.0.attentions.0.transformer_blocks.4.ff.net.2,up_blocks.0.attentions.0.transformer_blocks.5.attn1.to_k,up_blocks.0.attentions.0.transformer_blocks.5.attn1.to_out.0,up_blocks.0.attentions.0.transformer_blocks.5.attn1.to_q,up_blocks.0.attentions.0.transformer_blocks.5.attn1.to_v,up_blocks.0.attentions.0.transformer_blocks.5.attn2.to_k,up_blocks.0.attentions.0.transformer_blocks.5.attn2.to_out.0,up_blocks.0.attentions.0.transformer_blocks.5.attn2.to_q,up_blocks.0.attentions.0.transformer_blocks.5.attn2.to_v,up_blocks.0.attentions.0.transformer_blocks.5.ff.net.0.proj,up_blocks.0.attentions.0.transformer_blocks.5.ff.net.2,up_blocks.0.attentions.0.transformer_blocks.6.attn1.to_k,up_blocks.0.attentions.0.transformer_blocks.6.attn1.to_out.0,up_blocks.0.attentions.0.transformer_blocks.6.attn1.to_q,up_blocks.0.attentions.0.transformer_blocks.6.attn1.to_v,up_blocks.0.attentions.0.transformer_blocks.6.attn2.to_k,up_blocks.0.attentions.0.transformer_blocks.6.attn2.to_out.0,up_blocks.0.attentions.0.transformer_blocks.6.attn2.to_q,up_blocks.0.attentions.0.transformer_blocks.6.attn2.to_v,up_blocks.0.attentions.0.transformer_blocks.6.ff.net.0.proj,up_blocks.0.attentions.0.transformer_blocks.6.ff.net.2,up_blocks.0.attentions.0.transformer_blocks.7.attn1.to_k,up_blocks.0.attentions.0.transformer_blocks.7.attn1.to_out.0,up_blocks.0.attentions.0.transformer_blocks.7.attn1.to_q,up_blocks.0.attentions.0.transformer_blocks.7.attn1.to_v,up_blocks.0.attentions.0.transformer_blocks.7.attn2.to_k,up_blocks.0.attentions.0.transformer_blocks.7.attn2.to_out.0,up_blocks.0.attentions.0.transformer_blocks.7.attn2.to_q,up_blocks.0.attentions.0.transformer_blocks.7.attn2.to_v,up_blocks.0.attentions.0.transformer_blocks.7.ff.net.0.proj,up_blocks.0.attentions.0.transformer_blocks.7.ff.net.2,up_blocks.0.attentions.0.transformer_blocks.8.attn1.to_k,up_blocks.0.attentions.0.transformer_blocks.8.attn1.to_out.0,up_blocks.0.attentions.0.transformer_blocks.8.attn1.to_q,up_blocks.0.attentions.0.transformer_blocks.8.attn1.to_v,up_blocks.0.attentions.0.transformer_blocks.8.attn2.to_k,up_blocks.0.attentions.0.transformer_blocks.8.attn2.to_out.0,up_blocks.0.attentions.0.transformer_blocks.8.attn2.to_q,up_blocks.0.attentions.0.transformer_blocks.8.attn2.to_v,up_blocks.0.attentions.0.transformer_blocks.8.ff.net.0.proj,up_blocks.0.attentions.0.transformer_blocks.8.ff.net.2,up_blocks.0.attentions.0.transformer_blocks.9.attn1.to_k,up_blocks.0.attentions.0.transformer_blocks.9.attn1.to_out.0,up_blocks.0.attentions.0.transformer_blocks.9.attn1.to_q,up_blocks.0.attentions.0.transformer_blocks.9.attn1.to_v,up_blocks.0.attentions.0.transformer_blocks.9.attn2.to_k,up_blocks.0.attentions.0.transformer_blocks.9.attn2.to_out.0,up_blocks.0.attentions.0.transformer_blocks.9.attn2.to_q,up_blocks.0.attentions.0.transformer_blocks.9.attn2.to_v,up_blocks.0.attentions.0.transformer_blocks.9.ff.net.0.proj,up_blocks.0.attentions.0.transformer_blocks.9.ff.net.2,up_blocks.0.attentions.1.proj_in,up_blocks.0.attentions.1.proj_out,up_blocks.0.attentions.1.transformer_blocks.0.attn1.to_k,up_blocks.0.attentions.1.transformer_blocks.0.attn1.to_out.0,up_blocks.0.attentions.1.transformer_blocks.0.attn1.to_q,up_blocks.0.attentions.1.transformer_blocks.0.attn1.to_v,up_blocks.0.attentions.1.transformer_blocks.0.attn2.to_k,up_blocks.0.attentions.1.transformer_blocks.0.attn2.to_out.0,up_blocks.0.attentions.1.transformer_blocks.0.attn2.to_q,up_blocks.0.attentions.1.transformer_blocks.0.attn2.to_v,up_blocks.0.attentions.1.transformer_blocks.0.ff.net.0.proj,up_blocks.0.attentions.1.transformer_blocks.0.ff.net.2,up_blocks.0.attentions.1.transformer_blocks.1.attn1.to_k,up_blocks.0.attentions.1.transformer_blocks.1.attn1.to_out.0,up_blocks.0.attentions.1.transformer_blocks.1.attn1.to_q,up_blocks.0.attentions.1.transformer_blocks.1.attn1.to_v,up_blocks.0.attentions.1.transformer_blocks.1.attn2.to_k,up_blocks.0.attentions.1.transformer_blocks.1.attn2.to_out.0,up_blocks.0.attentions.1.transformer_blocks.1.attn2.to_q,up_blocks.0.attentions.1.transformer_blocks.1.attn2.to_v,up_blocks.0.attentions.1.transformer_blocks.1.ff.net.0.proj,up_blocks.0.attentions.1.transformer_blocks.1.ff.net.2,up_blocks.0.attentions.1.transformer_blocks.2.attn1.to_k,up_blocks.0.attentions.1.transformer_blocks.2.attn1.to_out.0,up_blocks.0.attentions.1.transformer_blocks.2.attn1.to_q,up_blocks.0.attentions.1.transformer_blocks.2.attn1.to_v,up_blocks.0.attentions.1.transformer_blocks.2.attn2.to_k,up_blocks.0.attentions.1.transformer_blocks.2.attn2.to_out.0,up_blocks.0.attentions.1.transformer_blocks.2.attn2.to_q,up_blocks.0.attentions.1.transformer_blocks.2.attn2.to_v,up_blocks.0.attentions.1.transformer_blocks.2.ff.net.0.proj,up_blocks.0.attentions.1.transformer_blocks.2.ff.net.2,up_blocks.0.attentions.1.transformer_blocks.3.attn1.to_k,up_blocks.0.attentions.1.transformer_blocks.3.attn1.to_out.0,up_blocks.0.attentions.1.transformer_blocks.3.attn1.to_q,up_blocks.0.attentions.1.transformer_blocks.3.attn1.to_v,up_blocks.0.attentions.1.transformer_blocks.3.attn2.to_k,up_blocks.0.attentions.1.transformer_blocks.3.attn2.to_out.0,up_blocks.0.attentions.1.transformer_blocks.3.attn2.to_q,up_blocks.0.attentions.1.transformer_blocks.3.attn2.to_v,up_blocks.0.attentions.1.transformer_blocks.3.ff.net.0.proj,up_blocks.0.attentions.1.transformer_blocks.3.ff.net.2,up_blocks.0.attentions.1.transformer_blocks.4.attn1.to_k,up_blocks.0.attentions.1.transformer_blocks.4.attn1.to_out.0,up_blocks.0.attentions.1.transformer_blocks.4.attn1.to_q,up_blocks.0.attentions.1.transformer_blocks.4.attn1.to_v,up_blocks.0.attentions.1.transformer_blocks.4.attn2.to_k,up_blocks.0.attentions.1.transformer_blocks.4.attn2.to_out.0,up_blocks.0.attentions.1.transformer_blocks.4.attn2.to_q,up_blocks.0.attentions.1.transformer_blocks.4.attn2.to_v,up_blocks.0.attentions.1.transformer_blocks.4.ff.net.0.proj,up_blocks.0.attentions.1.transformer_blocks.4.ff.net.2,up_blocks.0.attentions.1.transformer_blocks.5.attn1.to_k,up_blocks.0.attentions.1.transformer_blocks.5.attn1.to_out.0,up_blocks.0.attentions.1.transformer_blocks.5.attn1.to_q,up_blocks.0.attentions.1.transformer_blocks.5.attn1.to_v,up_blocks.0.attentions.1.transformer_blocks.5.attn2.to_k,up_blocks.0.attentions.1.transformer_blocks.5.attn2.to_out.0,up_blocks.0.attentions.1.transformer_blocks.5.attn2.to_q,up_blocks.0.attentions.1.transformer_blocks.5.attn2.to_v,up_blocks.0.attentions.1.transformer_blocks.5.ff.net.0.proj,up_blocks.0.attentions.1.transformer_blocks.5.ff.net.2,up_blocks.0.attentions.1.transformer_blocks.6.attn1.to_k,up_blocks.0.attentions.1.transformer_blocks.6.attn1.to_out.0,up_blocks.0.attentions.1.transformer_blocks.6.attn1.to_q,up_blocks.0.attentions.1.transformer_blocks.6.attn1.to_v,up_blocks.0.attentions.1.transformer_blocks.6.attn2.to_k,up_blocks.0.attentions.1.transformer_blocks.6.attn2.to_out.0,up_blocks.0.attentions.1.transformer_blocks.6.attn2.to_q,up_blocks.0.attentions.1.transformer_blocks.6.attn2.to_v,up_blocks.0.attentions.1.transformer_blocks.6.ff.net.0.proj,up_blocks.0.attentions.1.transformer_blocks.6.ff.net.2,up_blocks.0.attentions.1.transformer_blocks.7.attn1.to_k,up_blocks.0.attentions.1.transformer_blocks.7.attn1.to_out.0,up_blocks.0.attentions.1.transformer_blocks.7.attn1.to_q,up_blocks.0.attentions.1.transformer_blocks.7.attn1.to_v,up_blocks.0.attentions.1.transformer_blocks.7.attn2.to_k,up_blocks.0.attentions.1.transformer_blocks.7.attn2.to_out.0,up_blocks.0.attentions.1.transformer_blocks.7.attn2.to_q,up_blocks.0.attentions.1.transformer_blocks.7.attn2.to_v,up_blocks.0.attentions.1.transformer_blocks.7.ff.net.0.proj,up_blocks.0.attentions.1.transformer_blocks.7.ff.net.2,up_blocks.0.attentions.1.transformer_blocks.8.attn1.to_k,up_blocks.0.attentions.1.transformer_blocks.8.attn1.to_out.0,up_blocks.0.attentions.1.transformer_blocks.8.attn1.to_q,up_blocks.0.attentions.1.transformer_blocks.8.attn1.to_v,up_blocks.0.attentions.1.transformer_blocks.8.attn2.to_k,up_blocks.0.attentions.1.transformer_blocks.8.attn2.to_out.0,up_blocks.0.attentions.1.transformer_blocks.8.attn2.to_q,up_blocks.0.attentions.1.transformer_blocks.8.attn2.to_v,up_blocks.0.attentions.1.transformer_blocks.8.ff.net.0.proj,up_blocks.0.attentions.1.transformer_blocks.8.ff.net.2,up_blocks.0.attentions.1.transformer_blocks.9.attn1.to_k,up_blocks.0.attentions.1.transformer_blocks.9.attn1.to_out.0,up_blocks.0.attentions.1.transformer_blocks.9.attn1.to_q,up_blocks.0.attentions.1.transformer_blocks.9.attn1.to_v,up_blocks.0.attentions.1.transformer_blocks.9.attn2.to_k,up_blocks.0.attentions.1.transformer_blocks.9.attn2.to_out.0,up_blocks.0.attentions.1.transformer_blocks.9.attn2.to_q,up_blocks.0.attentions.1.transformer_blocks.9.attn2.to_v,up_blocks.0.attentions.1.transformer_blocks.9.ff.net.0.proj,up_blocks.0.attentions.1.transformer_blocks.9.ff.net.2,up_blocks.0.attentions.2.proj_in,up_blocks.0.attentions.2.proj_out,up_blocks.0.attentions.2.transformer_blocks.0.attn1.to_k,up_blocks.0.attentions.2.transformer_blocks.0.attn1.to_out.0,up_blocks.0.attentions.2.transformer_blocks.0.attn1.to_q,up_blocks.0.attentions.2.transformer_blocks.0.attn1.to_v,up_blocks.0.attentions.2.transformer_blocks.0.attn2.to_k,up_blocks.0.attentions.2.transformer_blocks.0.attn2.to_out.0,up_blocks.0.attentions.2.transformer_blocks.0.attn2.to_q,up_blocks.0.attentions.2.transformer_blocks.0.attn2.to_v,up_blocks.0.attentions.2.transformer_blocks.0.ff.net.0.proj,up_blocks.0.attentions.2.transformer_blocks.0.ff.net.2,up_blocks.0.attentions.2.transformer_blocks.1.attn1.to_k,up_blocks.0.attentions.2.transformer_blocks.1.attn1.to_out.0,up_blocks.0.attentions.2.transformer_blocks.1.attn1.to_q,up_blocks.0.attentions.2.transformer_blocks.1.attn1.to_v,up_blocks.0.attentions.2.transformer_blocks.1.attn2.to_k,up_blocks.0.attentions.2.transformer_blocks.1.attn2.to_out.0,up_blocks.0.attentions.2.transformer_blocks.1.attn2.to_q,up_blocks.0.attentions.2.transformer_blocks.1.attn2.to_v,up_blocks.0.attentions.2.transformer_blocks.1.ff.net.0.proj,up_blocks.0.attentions.2.transformer_blocks.1.ff.net.2,up_blocks.0.attentions.2.transformer_blocks.2.attn1.to_k,up_blocks.0.attentions.2.transformer_blocks.2.attn1.to_out.0,up_blocks.0.attentions.2.transformer_blocks.2.attn1.to_q,up_blocks.0.attentions.2.transformer_blocks.2.attn1.to_v,up_blocks.0.attentions.2.transformer_blocks.2.attn2.to_k,up_blocks.0.attentions.2.transformer_blocks.2.attn2.to_out.0,up_blocks.0.attentions.2.transformer_blocks.2.attn2.to_q,up_blocks.0.attentions.2.transformer_blocks.2.attn2.to_v,up_blocks.0.attentions.2.transformer_blocks.2.ff.net.0.proj,up_blocks.0.attentions.2.transformer_blocks.2.ff.net.2,up_blocks.0.attentions.2.transformer_blocks.3.attn1.to_k,up_blocks.0.attentions.2.transformer_blocks.3.attn1.to_out.0,up_blocks.0.attentions.2.transformer_blocks.3.attn1.to_q,up_blocks.0.attentions.2.transformer_blocks.3.attn1.to_v,up_blocks.0.attentions.2.transformer_blocks.3.attn2.to_k,up_blocks.0.attentions.2.transformer_blocks.3.attn2.to_out.0,up_blocks.0.attentions.2.transformer_blocks.3.attn2.to_q,up_blocks.0.attentions.2.transformer_blocks.3.attn2.to_v,up_blocks.0.attentions.2.transformer_blocks.3.ff.net.0.proj,up_blocks.0.attentions.2.transformer_blocks.3.ff.net.2,up_blocks.0.attentions.2.transformer_blocks.4.attn1.to_k,up_blocks.0.attentions.2.transformer_blocks.4.attn1.to_out.0,up_blocks.0.attentions.2.transformer_blocks.4.attn1.to_q,up_blocks.0.attentions.2.transformer_blocks.4.attn1.to_v,up_blocks.0.attentions.2.transformer_blocks.4.attn2.to_k,up_blocks.0.attentions.2.transformer_blocks.4.attn2.to_out.0,up_blocks.0.attentions.2.transformer_blocks.4.attn2.to_q,up_blocks.0.attentions.2.transformer_blocks.4.attn2.to_v,up_blocks.0.attentions.2.transformer_blocks.4.ff.net.0.proj,up_blocks.0.attentions.2.transformer_blocks.4.ff.net.2,up_blocks.0.attentions.2.transformer_blocks.5.attn1.to_k,up_blocks.0.attentions.2.transformer_blocks.5.attn1.to_out.0,up_blocks.0.attentions.2.transformer_blocks.5.attn1.to_q,up_blocks.0.attentions.2.transformer_blocks.5.attn1.to_v,up_blocks.0.attentions.2.transformer_blocks.5.attn2.to_k,up_blocks.0.attentions.2.transformer_blocks.5.attn2.to_out.0,up_blocks.0.attentions.2.transformer_blocks.5.attn2.to_q,up_blocks.0.attentions.2.transformer_blocks.5.attn2.to_v,up_blocks.0.attentions.2.transformer_blocks.5.ff.net.0.proj,up_blocks.0.attentions.2.transformer_blocks.5.ff.net.2,up_blocks.0.attentions.2.transformer_blocks.6.attn1.to_k,up_blocks.0.attentions.2.transformer_blocks.6.attn1.to_out.0,up_blocks.0.attentions.2.transformer_blocks.6.attn1.to_q,up_blocks.0.attentions.2.transformer_blocks.6.attn1.to_v,up_blocks.0.attentions.2.transformer_blocks.6.attn2.to_k,up_blocks.0.attentions.2.transformer_blocks.6.attn2.to_out.0,up_blocks.0.attentions.2.transformer_blocks.6.attn2.to_q,up_blocks.0.attentions.2.transformer_blocks.6.attn2.to_v,up_blocks.0.attentions.2.transformer_blocks.6.ff.net.0.proj,up_blocks.0.attentions.2.transformer_blocks.6.ff.net.2,up_blocks.0.attentions.2.transformer_blocks.7.attn1.to_k,up_blocks.0.attentions.2.transformer_blocks.7.attn1.to_out.0,up_blocks.0.attentions.2.transformer_blocks.7.attn1.to_q,up_blocks.0.attentions.2.transformer_blocks.7.attn1.to_v,up_blocks.0.attentions.2.transformer_blocks.7.attn2.to_k,up_blocks.0.attentions.2.transformer_blocks.7.attn2.to_out.0,up_blocks.0.attentions.2.transformer_blocks.7.attn2.to_q,up_blocks.0.attentions.2.transformer_blocks.7.attn2.to_v,up_blocks.0.attentions.2.transformer_blocks.7.ff.net.0.proj,up_blocks.0.attentions.2.transformer_blocks.7.ff.net.2,up_blocks.0.attentions.2.transformer_blocks.8.attn1.to_k,up_blocks.0.attentions.2.transformer_blocks.8.attn1.to_out.0,up_blocks.0.attentions.2.transformer_blocks.8.attn1.to_q,up_blocks.0.attentions.2.transformer_blocks.8.attn1.to_v,up_blocks.0.attentions.2.transformer_blocks.8.attn2.to_k,up_blocks.0.attentions.2.transformer_blocks.8.attn2.to_out.0,up_blocks.0.attentions.2.transformer_blocks.8.attn2.to_q,up_blocks.0.attentions.2.transformer_blocks.8.attn2.to_v,up_blocks.0.attentions.2.transformer_blocks.8.ff.net.0.proj,up_blocks.0.attentions.2.transformer_blocks.8.ff.net.2,up_blocks.0.attentions.2.transformer_blocks.9.attn1.to_k,up_blocks.0.attentions.2.transformer_blocks.9.attn1.to_out.0,up_blocks.0.attentions.2.transformer_blocks.9.attn1.to_q,up_blocks.0.attentions.2.transformer_blocks.9.attn1.to_v,up_blocks.0.attentions.2.transformer_blocks.9.attn2.to_k,up_blocks.0.attentions.2.transformer_blocks.9.attn2.to_out.0,up_blocks.0.attentions.2.transformer_blocks.9.attn2.to_q,up_blocks.0.attentions.2.transformer_blocks.9.attn2.to_v,up_blocks.0.attentions.2.transformer_blocks.9.ff.net.0.proj,up_blocks.0.attentions.2.transformer_blocks.9.ff.net.2,up_blocks.1.attentions.0.proj_in,up_blocks.1.attentions.0.proj_out,up_blocks.1.attentions.0.transformer_blocks.0.attn1.to_k,up_blocks.1.attentions.0.transformer_blocks.0.attn1.to_out.0,up_blocks.1.attentions.0.transformer_blocks.0.attn1.to_q,up_blocks.1.attentions.0.transformer_blocks.0.attn1.to_v,up_blocks.1.attentions.0.transformer_blocks.0.attn2.to_k,up_blocks.1.attentions.0.transformer_blocks.0.attn2.to_out.0,up_blocks.1.attentions.0.transformer_blocks.0.attn2.to_q,up_blocks.1.attentions.0.transformer_blocks.0.attn2.to_v,up_blocks.1.attentions.0.transformer_blocks.0.ff.net.0.proj,up_blocks.1.attentions.0.transformer_blocks.0.ff.net.2,up_blocks.1.attentions.0.transformer_blocks.1.attn1.to_k,up_blocks.1.attentions.0.transformer_blocks.1.attn1.to_out.0,up_blocks.1.attentions.0.transformer_blocks.1.attn1.to_q,up_blocks.1.attentions.0.transformer_blocks.1.attn1.to_v,up_blocks.1.attentions.0.transformer_blocks.1.attn2.to_k,up_blocks.1.attentions.0.transformer_blocks.1.attn2.to_out.0,up_blocks.1.attentions.0.transformer_blocks.1.attn2.to_q,up_blocks.1.attentions.0.transformer_blocks.1.attn2.to_v,up_blocks.1.attentions.0.transformer_blocks.1.ff.net.0.proj,up_blocks.1.attentions.0.transformer_blocks.1.ff.net.2,up_blocks.1.attentions.1.proj_in,up_blocks.1.attentions.1.proj_out,up_blocks.1.attentions.1.transformer_blocks.0.attn1.to_k,up_blocks.1.attentions.1.transformer_blocks.0.attn1.to_out.0,up_blocks.1.attentions.1.transformer_blocks.0.attn1.to_q,up_blocks.1.attentions.1.transformer_blocks.0.attn1.to_v,up_blocks.1.attentions.1.transformer_blocks.0.attn2.to_k,up_blocks.1.attentions.1.transformer_blocks.0.attn2.to_out.0,up_blocks.1.attentions.1.transformer_blocks.0.attn2.to_q,up_blocks.1.attentions.1.transformer_blocks.0.attn2.to_v,up_blocks.1.attentions.1.transformer_blocks.0.ff.net.0.proj,up_blocks.1.attentions.1.transformer_blocks.0.ff.net.2,up_blocks.1.attentions.1.transformer_blocks.1.attn1.to_k,up_blocks.1.attentions.1.transformer_blocks.1.attn1.to_out.0,up_blocks.1.attentions.1.transformer_blocks.1.attn1.to_q,up_blocks.1.attentions.1.transformer_blocks.1.attn1.to_v,up_blocks.1.attentions.1.transformer_blocks.1.attn2.to_k,up_blocks.1.attentions.1.transformer_blocks.1.attn2.to_out.0,up_blocks.1.attentions.1.transformer_blocks.1.attn2.to_q,up_blocks.1.attentions.1.transformer_blocks.1.attn2.to_v,up_blocks.1.attentions.1.transformer_blocks.1.ff.net.0.proj,up_blocks.1.attentions.1.transformer_blocks.1.ff.net.2,up_blocks.1.attentions.2.proj_in,up_blocks.1.attentions.2.proj_out,up_blocks.1.attentions.2.transformer_blocks.0.attn1.to_k,up_blocks.1.attentions.2.transformer_blocks.0.attn1.to_out.0,up_blocks.1.attentions.2.transformer_blocks.0.attn1.to_q,up_blocks.1.attentions.2.transformer_blocks.0.attn1.to_v,up_blocks.1.attentions.2.transformer_blocks.0.attn2.to_k,up_blocks.1.attentions.2.transformer_blocks.0.attn2.to_out.0,up_blocks.1.attentions.2.transformer_blocks.0.attn2.to_q,up_blocks.1.attentions.2.transformer_blocks.0.attn2.to_v,up_blocks.1.attentions.2.transformer_blocks.0.ff.net.0.proj,up_blocks.1.attentions.2.transformer_blocks.0.ff.net.2,up_blocks.1.attentions.2.transformer_blocks.1.attn1.to_k,up_blocks.1.attentions.2.transformer_blocks.1.attn1.to_out.0,up_blocks.1.attentions.2.transformer_blocks.1.attn1.to_q,up_blocks.1.attentions.2.transformer_blocks.1.attn1.to_v,up_blocks.1.attentions.2.transformer_blocks.1.attn2.to_k,up_blocks.1.attentions.2.transformer_blocks.1.attn2.to_out.0,up_blocks.1.attentions.2.transformer_blocks.1.attn2.to_q,up_blocks.1.attentions.2.transformer_blocks.1.attn2.to_v,up_blocks.1.attentions.2.transformer_blocks.1.ff.net.0.proj,up_blocks.1.attentions.2.transformer_blocks.1.ff.net.2 \ + --lora_rank 32 \ + --use_gradient_checkpointing \ + --align_to_opensource_format \ + --offload_models stabilityai/stable-diffusion-xl-base-1.0:unet/diffusion_pytorch_model.safetensors \ + --task sft:data_process + +# Stage 2: train LoRA from the cached dataset. +accelerate launch examples/stable_diffusion_xl/model_training/train.py \ + --dataset_base_path ./models/train/stable-diffusion-xl-base-1.0_split_cache \ + --height 1024 \ + --width 1024 \ + --dataset_repeat 10 \ + --model_id_with_origin_paths stabilityai/stable-diffusion-xl-base-1.0:text_encoder/model.safetensors,stabilityai/stable-diffusion-xl-base-1.0:text_encoder_2/model.safetensors,stabilityai/stable-diffusion-xl-base-1.0:unet/diffusion_pytorch_model.safetensors,stabilityai/stable-diffusion-xl-base-1.0:vae/diffusion_pytorch_model.safetensors \ + --learning_rate 1e-4 \ + --num_epochs 5 \ + --remove_prefix_in_ckpt pipe.unet. \ + --output_path ./models/train/stable-diffusion-xl-base-1.0_split \ + --lora_base_model unet \ + --lora_target_modules mid_block.attentions.0.proj_in,mid_block.attentions.0.proj_out,down_blocks.1.attentions.0.proj_in,down_blocks.1.attentions.0.proj_out,down_blocks.1.attentions.0.transformer_blocks.0.attn1.to_k,down_blocks.1.attentions.0.transformer_blocks.0.attn1.to_out.0,down_blocks.1.attentions.0.transformer_blocks.0.attn1.to_q,down_blocks.1.attentions.0.transformer_blocks.0.attn1.to_v,down_blocks.1.attentions.0.transformer_blocks.0.attn2.to_k,down_blocks.1.attentions.0.transformer_blocks.0.attn2.to_out.0,down_blocks.1.attentions.0.transformer_blocks.0.attn2.to_q,down_blocks.1.attentions.0.transformer_blocks.0.attn2.to_v,down_blocks.1.attentions.0.transformer_blocks.0.ff.net.0.proj,down_blocks.1.attentions.0.transformer_blocks.0.ff.net.2,down_blocks.1.attentions.0.transformer_blocks.1.attn1.to_k,down_blocks.1.attentions.0.transformer_blocks.1.attn1.to_out.0,down_blocks.1.attentions.0.transformer_blocks.1.attn1.to_q,down_blocks.1.attentions.0.transformer_blocks.1.attn1.to_v,down_blocks.1.attentions.0.transformer_blocks.1.attn2.to_k,down_blocks.1.attentions.0.transformer_blocks.1.attn2.to_out.0,down_blocks.1.attentions.0.transformer_blocks.1.attn2.to_q,down_blocks.1.attentions.0.transformer_blocks.1.attn2.to_v,down_blocks.1.attentions.0.transformer_blocks.1.ff.net.0.proj,down_blocks.1.attentions.0.transformer_blocks.1.ff.net.2,down_blocks.1.attentions.1.proj_in,down_blocks.1.attentions.1.proj_out,down_blocks.1.attentions.1.transformer_blocks.0.attn1.to_k,down_blocks.1.attentions.1.transformer_blocks.0.attn1.to_out.0,down_blocks.1.attentions.1.transformer_blocks.0.attn1.to_q,down_blocks.1.attentions.1.transformer_blocks.0.attn1.to_v,down_blocks.1.attentions.1.transformer_blocks.0.attn2.to_k,down_blocks.1.attentions.1.transformer_blocks.0.attn2.to_out.0,down_blocks.1.attentions.1.transformer_blocks.0.attn2.to_q,down_blocks.1.attentions.1.transformer_blocks.0.attn2.to_v,down_blocks.1.attentions.1.transformer_blocks.0.ff.net.0.proj,down_blocks.1.attentions.1.transformer_blocks.0.ff.net.2,down_blocks.1.attentions.1.transformer_blocks.1.attn1.to_k,down_blocks.1.attentions.1.transformer_blocks.1.attn1.to_out.0,down_blocks.1.attentions.1.transformer_blocks.1.attn1.to_q,down_blocks.1.attentions.1.transformer_blocks.1.attn1.to_v,down_blocks.1.attentions.1.transformer_blocks.1.attn2.to_k,down_blocks.1.attentions.1.transformer_blocks.1.attn2.to_out.0,down_blocks.1.attentions.1.transformer_blocks.1.attn2.to_q,down_blocks.1.attentions.1.transformer_blocks.1.attn2.to_v,down_blocks.1.attentions.1.transformer_blocks.1.ff.net.0.proj,down_blocks.1.attentions.1.transformer_blocks.1.ff.net.2,down_blocks.2.attentions.0.proj_in,down_blocks.2.attentions.0.proj_out,down_blocks.2.attentions.0.transformer_blocks.0.attn1.to_k,down_blocks.2.attentions.0.transformer_blocks.0.attn1.to_out.0,down_blocks.2.attentions.0.transformer_blocks.0.attn1.to_q,down_blocks.2.attentions.0.transformer_blocks.0.attn1.to_v,down_blocks.2.attentions.0.transformer_blocks.0.attn2.to_k,down_blocks.2.attentions.0.transformer_blocks.0.attn2.to_out.0,down_blocks.2.attentions.0.transformer_blocks.0.attn2.to_q,down_blocks.2.attentions.0.transformer_blocks.0.attn2.to_v,down_blocks.2.attentions.0.transformer_blocks.0.ff.net.0.proj,down_blocks.2.attentions.0.transformer_blocks.0.ff.net.2,down_blocks.2.attentions.0.transformer_blocks.1.attn1.to_k,down_blocks.2.attentions.0.transformer_blocks.1.attn1.to_out.0,down_blocks.2.attentions.0.transformer_blocks.1.attn1.to_q,down_blocks.2.attentions.0.transformer_blocks.1.attn1.to_v,down_blocks.2.attentions.0.transformer_blocks.1.attn2.to_k,down_blocks.2.attentions.0.transformer_blocks.1.attn2.to_out.0,down_blocks.2.attentions.0.transformer_blocks.1.attn2.to_q,down_blocks.2.attentions.0.transformer_blocks.1.attn2.to_v,down_blocks.2.attentions.0.transformer_blocks.1.ff.net.0.proj,down_blocks.2.attentions.0.transformer_blocks.1.ff.net.2,down_blocks.2.attentions.0.transformer_blocks.2.attn1.to_k,down_blocks.2.attentions.0.transformer_blocks.2.attn1.to_out.0,down_blocks.2.attentions.0.transformer_blocks.2.attn1.to_q,down_blocks.2.attentions.0.transformer_blocks.2.attn1.to_v,down_blocks.2.attentions.0.transformer_blocks.2.attn2.to_k,down_blocks.2.attentions.0.transformer_blocks.2.attn2.to_out.0,down_blocks.2.attentions.0.transformer_blocks.2.attn2.to_q,down_blocks.2.attentions.0.transformer_blocks.2.attn2.to_v,down_blocks.2.attentions.0.transformer_blocks.2.ff.net.0.proj,down_blocks.2.attentions.0.transformer_blocks.2.ff.net.2,down_blocks.2.attentions.0.transformer_blocks.3.attn1.to_k,down_blocks.2.attentions.0.transformer_blocks.3.attn1.to_out.0,down_blocks.2.attentions.0.transformer_blocks.3.attn1.to_q,down_blocks.2.attentions.0.transformer_blocks.3.attn1.to_v,down_blocks.2.attentions.0.transformer_blocks.3.attn2.to_k,down_blocks.2.attentions.0.transformer_blocks.3.attn2.to_out.0,down_blocks.2.attentions.0.transformer_blocks.3.attn2.to_q,down_blocks.2.attentions.0.transformer_blocks.3.attn2.to_v,down_blocks.2.attentions.0.transformer_blocks.3.ff.net.0.proj,down_blocks.2.attentions.0.transformer_blocks.3.ff.net.2,down_blocks.2.attentions.0.transformer_blocks.4.attn1.to_k,down_blocks.2.attentions.0.transformer_blocks.4.attn1.to_out.0,down_blocks.2.attentions.0.transformer_blocks.4.attn1.to_q,down_blocks.2.attentions.0.transformer_blocks.4.attn1.to_v,down_blocks.2.attentions.0.transformer_blocks.4.attn2.to_k,down_blocks.2.attentions.0.transformer_blocks.4.attn2.to_out.0,down_blocks.2.attentions.0.transformer_blocks.4.attn2.to_q,down_blocks.2.attentions.0.transformer_blocks.4.attn2.to_v,down_blocks.2.attentions.0.transformer_blocks.4.ff.net.0.proj,down_blocks.2.attentions.0.transformer_blocks.4.ff.net.2,down_blocks.2.attentions.0.transformer_blocks.5.attn1.to_k,down_blocks.2.attentions.0.transformer_blocks.5.attn1.to_out.0,down_blocks.2.attentions.0.transformer_blocks.5.attn1.to_q,down_blocks.2.attentions.0.transformer_blocks.5.attn1.to_v,down_blocks.2.attentions.0.transformer_blocks.5.attn2.to_k,down_blocks.2.attentions.0.transformer_blocks.5.attn2.to_out.0,down_blocks.2.attentions.0.transformer_blocks.5.attn2.to_q,down_blocks.2.attentions.0.transformer_blocks.5.attn2.to_v,down_blocks.2.attentions.0.transformer_blocks.5.ff.net.0.proj,down_blocks.2.attentions.0.transformer_blocks.5.ff.net.2,down_blocks.2.attentions.0.transformer_blocks.6.attn1.to_k,down_blocks.2.attentions.0.transformer_blocks.6.attn1.to_out.0,down_blocks.2.attentions.0.transformer_blocks.6.attn1.to_q,down_blocks.2.attentions.0.transformer_blocks.6.attn1.to_v,down_blocks.2.attentions.0.transformer_blocks.6.attn2.to_k,down_blocks.2.attentions.0.transformer_blocks.6.attn2.to_out.0,down_blocks.2.attentions.0.transformer_blocks.6.attn2.to_q,down_blocks.2.attentions.0.transformer_blocks.6.attn2.to_v,down_blocks.2.attentions.0.transformer_blocks.6.ff.net.0.proj,down_blocks.2.attentions.0.transformer_blocks.6.ff.net.2,down_blocks.2.attentions.0.transformer_blocks.7.attn1.to_k,down_blocks.2.attentions.0.transformer_blocks.7.attn1.to_out.0,down_blocks.2.attentions.0.transformer_blocks.7.attn1.to_q,down_blocks.2.attentions.0.transformer_blocks.7.attn1.to_v,down_blocks.2.attentions.0.transformer_blocks.7.attn2.to_k,down_blocks.2.attentions.0.transformer_blocks.7.attn2.to_out.0,down_blocks.2.attentions.0.transformer_blocks.7.attn2.to_q,down_blocks.2.attentions.0.transformer_blocks.7.attn2.to_v,down_blocks.2.attentions.0.transformer_blocks.7.ff.net.0.proj,down_blocks.2.attentions.0.transformer_blocks.7.ff.net.2,down_blocks.2.attentions.0.transformer_blocks.8.attn1.to_k,down_blocks.2.attentions.0.transformer_blocks.8.attn1.to_out.0,down_blocks.2.attentions.0.transformer_blocks.8.attn1.to_q,down_blocks.2.attentions.0.transformer_blocks.8.attn1.to_v,down_blocks.2.attentions.0.transformer_blocks.8.attn2.to_k,down_blocks.2.attentions.0.transformer_blocks.8.attn2.to_out.0,down_blocks.2.attentions.0.transformer_blocks.8.attn2.to_q,down_blocks.2.attentions.0.transformer_blocks.8.attn2.to_v,down_blocks.2.attentions.0.transformer_blocks.8.ff.net.0.proj,down_blocks.2.attentions.0.transformer_blocks.8.ff.net.2,down_blocks.2.attentions.0.transformer_blocks.9.attn1.to_k,down_blocks.2.attentions.0.transformer_blocks.9.attn1.to_out.0,down_blocks.2.attentions.0.transformer_blocks.9.attn1.to_q,down_blocks.2.attentions.0.transformer_blocks.9.attn1.to_v,down_blocks.2.attentions.0.transformer_blocks.9.attn2.to_k,down_blocks.2.attentions.0.transformer_blocks.9.attn2.to_out.0,down_blocks.2.attentions.0.transformer_blocks.9.attn2.to_q,down_blocks.2.attentions.0.transformer_blocks.9.attn2.to_v,down_blocks.2.attentions.0.transformer_blocks.9.ff.net.0.proj,down_blocks.2.attentions.0.transformer_blocks.9.ff.net.2,down_blocks.2.attentions.1.proj_in,down_blocks.2.attentions.1.proj_out,down_blocks.2.attentions.1.transformer_blocks.0.attn1.to_k,down_blocks.2.attentions.1.transformer_blocks.0.attn1.to_out.0,down_blocks.2.attentions.1.transformer_blocks.0.attn1.to_q,down_blocks.2.attentions.1.transformer_blocks.0.attn1.to_v,down_blocks.2.attentions.1.transformer_blocks.0.attn2.to_k,down_blocks.2.attentions.1.transformer_blocks.0.attn2.to_out.0,down_blocks.2.attentions.1.transformer_blocks.0.attn2.to_q,down_blocks.2.attentions.1.transformer_blocks.0.attn2.to_v,down_blocks.2.attentions.1.transformer_blocks.0.ff.net.0.proj,down_blocks.2.attentions.1.transformer_blocks.0.ff.net.2,down_blocks.2.attentions.1.transformer_blocks.1.attn1.to_k,down_blocks.2.attentions.1.transformer_blocks.1.attn1.to_out.0,down_blocks.2.attentions.1.transformer_blocks.1.attn1.to_q,down_blocks.2.attentions.1.transformer_blocks.1.attn1.to_v,down_blocks.2.attentions.1.transformer_blocks.1.attn2.to_k,down_blocks.2.attentions.1.transformer_blocks.1.attn2.to_out.0,down_blocks.2.attentions.1.transformer_blocks.1.attn2.to_q,down_blocks.2.attentions.1.transformer_blocks.1.attn2.to_v,down_blocks.2.attentions.1.transformer_blocks.1.ff.net.0.proj,down_blocks.2.attentions.1.transformer_blocks.1.ff.net.2,down_blocks.2.attentions.1.transformer_blocks.2.attn1.to_k,down_blocks.2.attentions.1.transformer_blocks.2.attn1.to_out.0,down_blocks.2.attentions.1.transformer_blocks.2.attn1.to_q,down_blocks.2.attentions.1.transformer_blocks.2.attn1.to_v,down_blocks.2.attentions.1.transformer_blocks.2.attn2.to_k,down_blocks.2.attentions.1.transformer_blocks.2.attn2.to_out.0,down_blocks.2.attentions.1.transformer_blocks.2.attn2.to_q,down_blocks.2.attentions.1.transformer_blocks.2.attn2.to_v,down_blocks.2.attentions.1.transformer_blocks.2.ff.net.0.proj,down_blocks.2.attentions.1.transformer_blocks.2.ff.net.2,down_blocks.2.attentions.1.transformer_blocks.3.attn1.to_k,down_blocks.2.attentions.1.transformer_blocks.3.attn1.to_out.0,down_blocks.2.attentions.1.transformer_blocks.3.attn1.to_q,down_blocks.2.attentions.1.transformer_blocks.3.attn1.to_v,down_blocks.2.attentions.1.transformer_blocks.3.attn2.to_k,down_blocks.2.attentions.1.transformer_blocks.3.attn2.to_out.0,down_blocks.2.attentions.1.transformer_blocks.3.attn2.to_q,down_blocks.2.attentions.1.transformer_blocks.3.attn2.to_v,down_blocks.2.attentions.1.transformer_blocks.3.ff.net.0.proj,down_blocks.2.attentions.1.transformer_blocks.3.ff.net.2,down_blocks.2.attentions.1.transformer_blocks.4.attn1.to_k,down_blocks.2.attentions.1.transformer_blocks.4.attn1.to_out.0,down_blocks.2.attentions.1.transformer_blocks.4.attn1.to_q,down_blocks.2.attentions.1.transformer_blocks.4.attn1.to_v,down_blocks.2.attentions.1.transformer_blocks.4.attn2.to_k,down_blocks.2.attentions.1.transformer_blocks.4.attn2.to_out.0,down_blocks.2.attentions.1.transformer_blocks.4.attn2.to_q,down_blocks.2.attentions.1.transformer_blocks.4.attn2.to_v,down_blocks.2.attentions.1.transformer_blocks.4.ff.net.0.proj,down_blocks.2.attentions.1.transformer_blocks.4.ff.net.2,down_blocks.2.attentions.1.transformer_blocks.5.attn1.to_k,down_blocks.2.attentions.1.transformer_blocks.5.attn1.to_out.0,down_blocks.2.attentions.1.transformer_blocks.5.attn1.to_q,down_blocks.2.attentions.1.transformer_blocks.5.attn1.to_v,down_blocks.2.attentions.1.transformer_blocks.5.attn2.to_k,down_blocks.2.attentions.1.transformer_blocks.5.attn2.to_out.0,down_blocks.2.attentions.1.transformer_blocks.5.attn2.to_q,down_blocks.2.attentions.1.transformer_blocks.5.attn2.to_v,down_blocks.2.attentions.1.transformer_blocks.5.ff.net.0.proj,down_blocks.2.attentions.1.transformer_blocks.5.ff.net.2,down_blocks.2.attentions.1.transformer_blocks.6.attn1.to_k,down_blocks.2.attentions.1.transformer_blocks.6.attn1.to_out.0,down_blocks.2.attentions.1.transformer_blocks.6.attn1.to_q,down_blocks.2.attentions.1.transformer_blocks.6.attn1.to_v,down_blocks.2.attentions.1.transformer_blocks.6.attn2.to_k,down_blocks.2.attentions.1.transformer_blocks.6.attn2.to_out.0,down_blocks.2.attentions.1.transformer_blocks.6.attn2.to_q,down_blocks.2.attentions.1.transformer_blocks.6.attn2.to_v,down_blocks.2.attentions.1.transformer_blocks.6.ff.net.0.proj,down_blocks.2.attentions.1.transformer_blocks.6.ff.net.2,down_blocks.2.attentions.1.transformer_blocks.7.attn1.to_k,down_blocks.2.attentions.1.transformer_blocks.7.attn1.to_out.0,down_blocks.2.attentions.1.transformer_blocks.7.attn1.to_q,down_blocks.2.attentions.1.transformer_blocks.7.attn1.to_v,down_blocks.2.attentions.1.transformer_blocks.7.attn2.to_k,down_blocks.2.attentions.1.transformer_blocks.7.attn2.to_out.0,down_blocks.2.attentions.1.transformer_blocks.7.attn2.to_q,down_blocks.2.attentions.1.transformer_blocks.7.attn2.to_v,down_blocks.2.attentions.1.transformer_blocks.7.ff.net.0.proj,down_blocks.2.attentions.1.transformer_blocks.7.ff.net.2,down_blocks.2.attentions.1.transformer_blocks.8.attn1.to_k,down_blocks.2.attentions.1.transformer_blocks.8.attn1.to_out.0,down_blocks.2.attentions.1.transformer_blocks.8.attn1.to_q,down_blocks.2.attentions.1.transformer_blocks.8.attn1.to_v,down_blocks.2.attentions.1.transformer_blocks.8.attn2.to_k,down_blocks.2.attentions.1.transformer_blocks.8.attn2.to_out.0,down_blocks.2.attentions.1.transformer_blocks.8.attn2.to_q,down_blocks.2.attentions.1.transformer_blocks.8.attn2.to_v,down_blocks.2.attentions.1.transformer_blocks.8.ff.net.0.proj,down_blocks.2.attentions.1.transformer_blocks.8.ff.net.2,down_blocks.2.attentions.1.transformer_blocks.9.attn1.to_k,down_blocks.2.attentions.1.transformer_blocks.9.attn1.to_out.0,down_blocks.2.attentions.1.transformer_blocks.9.attn1.to_q,down_blocks.2.attentions.1.transformer_blocks.9.attn1.to_v,down_blocks.2.attentions.1.transformer_blocks.9.attn2.to_k,down_blocks.2.attentions.1.transformer_blocks.9.attn2.to_out.0,down_blocks.2.attentions.1.transformer_blocks.9.attn2.to_q,down_blocks.2.attentions.1.transformer_blocks.9.attn2.to_v,down_blocks.2.attentions.1.transformer_blocks.9.ff.net.0.proj,down_blocks.2.attentions.1.transformer_blocks.9.ff.net.2,mid_block.attentions.0.transformer_blocks.0.attn1.to_k,mid_block.attentions.0.transformer_blocks.0.attn1.to_out.0,mid_block.attentions.0.transformer_blocks.0.attn1.to_q,mid_block.attentions.0.transformer_blocks.0.attn1.to_v,mid_block.attentions.0.transformer_blocks.0.attn2.to_k,mid_block.attentions.0.transformer_blocks.0.attn2.to_out.0,mid_block.attentions.0.transformer_blocks.0.attn2.to_q,mid_block.attentions.0.transformer_blocks.0.attn2.to_v,mid_block.attentions.0.transformer_blocks.0.ff.net.0.proj,mid_block.attentions.0.transformer_blocks.0.ff.net.2,mid_block.attentions.0.transformer_blocks.1.attn1.to_k,mid_block.attentions.0.transformer_blocks.1.attn1.to_out.0,mid_block.attentions.0.transformer_blocks.1.attn1.to_q,mid_block.attentions.0.transformer_blocks.1.attn1.to_v,mid_block.attentions.0.transformer_blocks.1.attn2.to_k,mid_block.attentions.0.transformer_blocks.1.attn2.to_out.0,mid_block.attentions.0.transformer_blocks.1.attn2.to_q,mid_block.attentions.0.transformer_blocks.1.attn2.to_v,mid_block.attentions.0.transformer_blocks.1.ff.net.0.proj,mid_block.attentions.0.transformer_blocks.1.ff.net.2,mid_block.attentions.0.transformer_blocks.2.attn1.to_k,mid_block.attentions.0.transformer_blocks.2.attn1.to_out.0,mid_block.attentions.0.transformer_blocks.2.attn1.to_q,mid_block.attentions.0.transformer_blocks.2.attn1.to_v,mid_block.attentions.0.transformer_blocks.2.attn2.to_k,mid_block.attentions.0.transformer_blocks.2.attn2.to_out.0,mid_block.attentions.0.transformer_blocks.2.attn2.to_q,mid_block.attentions.0.transformer_blocks.2.attn2.to_v,mid_block.attentions.0.transformer_blocks.2.ff.net.0.proj,mid_block.attentions.0.transformer_blocks.2.ff.net.2,mid_block.attentions.0.transformer_blocks.3.attn1.to_k,mid_block.attentions.0.transformer_blocks.3.attn1.to_out.0,mid_block.attentions.0.transformer_blocks.3.attn1.to_q,mid_block.attentions.0.transformer_blocks.3.attn1.to_v,mid_block.attentions.0.transformer_blocks.3.attn2.to_k,mid_block.attentions.0.transformer_blocks.3.attn2.to_out.0,mid_block.attentions.0.transformer_blocks.3.attn2.to_q,mid_block.attentions.0.transformer_blocks.3.attn2.to_v,mid_block.attentions.0.transformer_blocks.3.ff.net.0.proj,mid_block.attentions.0.transformer_blocks.3.ff.net.2,mid_block.attentions.0.transformer_blocks.4.attn1.to_k,mid_block.attentions.0.transformer_blocks.4.attn1.to_out.0,mid_block.attentions.0.transformer_blocks.4.attn1.to_q,mid_block.attentions.0.transformer_blocks.4.attn1.to_v,mid_block.attentions.0.transformer_blocks.4.attn2.to_k,mid_block.attentions.0.transformer_blocks.4.attn2.to_out.0,mid_block.attentions.0.transformer_blocks.4.attn2.to_q,mid_block.attentions.0.transformer_blocks.4.attn2.to_v,mid_block.attentions.0.transformer_blocks.4.ff.net.0.proj,mid_block.attentions.0.transformer_blocks.4.ff.net.2,mid_block.attentions.0.transformer_blocks.5.attn1.to_k,mid_block.attentions.0.transformer_blocks.5.attn1.to_out.0,mid_block.attentions.0.transformer_blocks.5.attn1.to_q,mid_block.attentions.0.transformer_blocks.5.attn1.to_v,mid_block.attentions.0.transformer_blocks.5.attn2.to_k,mid_block.attentions.0.transformer_blocks.5.attn2.to_out.0,mid_block.attentions.0.transformer_blocks.5.attn2.to_q,mid_block.attentions.0.transformer_blocks.5.attn2.to_v,mid_block.attentions.0.transformer_blocks.5.ff.net.0.proj,mid_block.attentions.0.transformer_blocks.5.ff.net.2,mid_block.attentions.0.transformer_blocks.6.attn1.to_k,mid_block.attentions.0.transformer_blocks.6.attn1.to_out.0,mid_block.attentions.0.transformer_blocks.6.attn1.to_q,mid_block.attentions.0.transformer_blocks.6.attn1.to_v,mid_block.attentions.0.transformer_blocks.6.attn2.to_k,mid_block.attentions.0.transformer_blocks.6.attn2.to_out.0,mid_block.attentions.0.transformer_blocks.6.attn2.to_q,mid_block.attentions.0.transformer_blocks.6.attn2.to_v,mid_block.attentions.0.transformer_blocks.6.ff.net.0.proj,mid_block.attentions.0.transformer_blocks.6.ff.net.2,mid_block.attentions.0.transformer_blocks.7.attn1.to_k,mid_block.attentions.0.transformer_blocks.7.attn1.to_out.0,mid_block.attentions.0.transformer_blocks.7.attn1.to_q,mid_block.attentions.0.transformer_blocks.7.attn1.to_v,mid_block.attentions.0.transformer_blocks.7.attn2.to_k,mid_block.attentions.0.transformer_blocks.7.attn2.to_out.0,mid_block.attentions.0.transformer_blocks.7.attn2.to_q,mid_block.attentions.0.transformer_blocks.7.attn2.to_v,mid_block.attentions.0.transformer_blocks.7.ff.net.0.proj,mid_block.attentions.0.transformer_blocks.7.ff.net.2,mid_block.attentions.0.transformer_blocks.8.attn1.to_k,mid_block.attentions.0.transformer_blocks.8.attn1.to_out.0,mid_block.attentions.0.transformer_blocks.8.attn1.to_q,mid_block.attentions.0.transformer_blocks.8.attn1.to_v,mid_block.attentions.0.transformer_blocks.8.attn2.to_k,mid_block.attentions.0.transformer_blocks.8.attn2.to_out.0,mid_block.attentions.0.transformer_blocks.8.attn2.to_q,mid_block.attentions.0.transformer_blocks.8.attn2.to_v,mid_block.attentions.0.transformer_blocks.8.ff.net.0.proj,mid_block.attentions.0.transformer_blocks.8.ff.net.2,mid_block.attentions.0.transformer_blocks.9.attn1.to_k,mid_block.attentions.0.transformer_blocks.9.attn1.to_out.0,mid_block.attentions.0.transformer_blocks.9.attn1.to_q,mid_block.attentions.0.transformer_blocks.9.attn1.to_v,mid_block.attentions.0.transformer_blocks.9.attn2.to_k,mid_block.attentions.0.transformer_blocks.9.attn2.to_out.0,mid_block.attentions.0.transformer_blocks.9.attn2.to_q,mid_block.attentions.0.transformer_blocks.9.attn2.to_v,mid_block.attentions.0.transformer_blocks.9.ff.net.0.proj,mid_block.attentions.0.transformer_blocks.9.ff.net.2,up_blocks.0.attentions.0.proj_in,up_blocks.0.attentions.0.proj_out,up_blocks.0.attentions.0.transformer_blocks.0.attn1.to_k,up_blocks.0.attentions.0.transformer_blocks.0.attn1.to_out.0,up_blocks.0.attentions.0.transformer_blocks.0.attn1.to_q,up_blocks.0.attentions.0.transformer_blocks.0.attn1.to_v,up_blocks.0.attentions.0.transformer_blocks.0.attn2.to_k,up_blocks.0.attentions.0.transformer_blocks.0.attn2.to_out.0,up_blocks.0.attentions.0.transformer_blocks.0.attn2.to_q,up_blocks.0.attentions.0.transformer_blocks.0.attn2.to_v,up_blocks.0.attentions.0.transformer_blocks.0.ff.net.0.proj,up_blocks.0.attentions.0.transformer_blocks.0.ff.net.2,up_blocks.0.attentions.0.transformer_blocks.1.attn1.to_k,up_blocks.0.attentions.0.transformer_blocks.1.attn1.to_out.0,up_blocks.0.attentions.0.transformer_blocks.1.attn1.to_q,up_blocks.0.attentions.0.transformer_blocks.1.attn1.to_v,up_blocks.0.attentions.0.transformer_blocks.1.attn2.to_k,up_blocks.0.attentions.0.transformer_blocks.1.attn2.to_out.0,up_blocks.0.attentions.0.transformer_blocks.1.attn2.to_q,up_blocks.0.attentions.0.transformer_blocks.1.attn2.to_v,up_blocks.0.attentions.0.transformer_blocks.1.ff.net.0.proj,up_blocks.0.attentions.0.transformer_blocks.1.ff.net.2,up_blocks.0.attentions.0.transformer_blocks.2.attn1.to_k,up_blocks.0.attentions.0.transformer_blocks.2.attn1.to_out.0,up_blocks.0.attentions.0.transformer_blocks.2.attn1.to_q,up_blocks.0.attentions.0.transformer_blocks.2.attn1.to_v,up_blocks.0.attentions.0.transformer_blocks.2.attn2.to_k,up_blocks.0.attentions.0.transformer_blocks.2.attn2.to_out.0,up_blocks.0.attentions.0.transformer_blocks.2.attn2.to_q,up_blocks.0.attentions.0.transformer_blocks.2.attn2.to_v,up_blocks.0.attentions.0.transformer_blocks.2.ff.net.0.proj,up_blocks.0.attentions.0.transformer_blocks.2.ff.net.2,up_blocks.0.attentions.0.transformer_blocks.3.attn1.to_k,up_blocks.0.attentions.0.transformer_blocks.3.attn1.to_out.0,up_blocks.0.attentions.0.transformer_blocks.3.attn1.to_q,up_blocks.0.attentions.0.transformer_blocks.3.attn1.to_v,up_blocks.0.attentions.0.transformer_blocks.3.attn2.to_k,up_blocks.0.attentions.0.transformer_blocks.3.attn2.to_out.0,up_blocks.0.attentions.0.transformer_blocks.3.attn2.to_q,up_blocks.0.attentions.0.transformer_blocks.3.attn2.to_v,up_blocks.0.attentions.0.transformer_blocks.3.ff.net.0.proj,up_blocks.0.attentions.0.transformer_blocks.3.ff.net.2,up_blocks.0.attentions.0.transformer_blocks.4.attn1.to_k,up_blocks.0.attentions.0.transformer_blocks.4.attn1.to_out.0,up_blocks.0.attentions.0.transformer_blocks.4.attn1.to_q,up_blocks.0.attentions.0.transformer_blocks.4.attn1.to_v,up_blocks.0.attentions.0.transformer_blocks.4.attn2.to_k,up_blocks.0.attentions.0.transformer_blocks.4.attn2.to_out.0,up_blocks.0.attentions.0.transformer_blocks.4.attn2.to_q,up_blocks.0.attentions.0.transformer_blocks.4.attn2.to_v,up_blocks.0.attentions.0.transformer_blocks.4.ff.net.0.proj,up_blocks.0.attentions.0.transformer_blocks.4.ff.net.2,up_blocks.0.attentions.0.transformer_blocks.5.attn1.to_k,up_blocks.0.attentions.0.transformer_blocks.5.attn1.to_out.0,up_blocks.0.attentions.0.transformer_blocks.5.attn1.to_q,up_blocks.0.attentions.0.transformer_blocks.5.attn1.to_v,up_blocks.0.attentions.0.transformer_blocks.5.attn2.to_k,up_blocks.0.attentions.0.transformer_blocks.5.attn2.to_out.0,up_blocks.0.attentions.0.transformer_blocks.5.attn2.to_q,up_blocks.0.attentions.0.transformer_blocks.5.attn2.to_v,up_blocks.0.attentions.0.transformer_blocks.5.ff.net.0.proj,up_blocks.0.attentions.0.transformer_blocks.5.ff.net.2,up_blocks.0.attentions.0.transformer_blocks.6.attn1.to_k,up_blocks.0.attentions.0.transformer_blocks.6.attn1.to_out.0,up_blocks.0.attentions.0.transformer_blocks.6.attn1.to_q,up_blocks.0.attentions.0.transformer_blocks.6.attn1.to_v,up_blocks.0.attentions.0.transformer_blocks.6.attn2.to_k,up_blocks.0.attentions.0.transformer_blocks.6.attn2.to_out.0,up_blocks.0.attentions.0.transformer_blocks.6.attn2.to_q,up_blocks.0.attentions.0.transformer_blocks.6.attn2.to_v,up_blocks.0.attentions.0.transformer_blocks.6.ff.net.0.proj,up_blocks.0.attentions.0.transformer_blocks.6.ff.net.2,up_blocks.0.attentions.0.transformer_blocks.7.attn1.to_k,up_blocks.0.attentions.0.transformer_blocks.7.attn1.to_out.0,up_blocks.0.attentions.0.transformer_blocks.7.attn1.to_q,up_blocks.0.attentions.0.transformer_blocks.7.attn1.to_v,up_blocks.0.attentions.0.transformer_blocks.7.attn2.to_k,up_blocks.0.attentions.0.transformer_blocks.7.attn2.to_out.0,up_blocks.0.attentions.0.transformer_blocks.7.attn2.to_q,up_blocks.0.attentions.0.transformer_blocks.7.attn2.to_v,up_blocks.0.attentions.0.transformer_blocks.7.ff.net.0.proj,up_blocks.0.attentions.0.transformer_blocks.7.ff.net.2,up_blocks.0.attentions.0.transformer_blocks.8.attn1.to_k,up_blocks.0.attentions.0.transformer_blocks.8.attn1.to_out.0,up_blocks.0.attentions.0.transformer_blocks.8.attn1.to_q,up_blocks.0.attentions.0.transformer_blocks.8.attn1.to_v,up_blocks.0.attentions.0.transformer_blocks.8.attn2.to_k,up_blocks.0.attentions.0.transformer_blocks.8.attn2.to_out.0,up_blocks.0.attentions.0.transformer_blocks.8.attn2.to_q,up_blocks.0.attentions.0.transformer_blocks.8.attn2.to_v,up_blocks.0.attentions.0.transformer_blocks.8.ff.net.0.proj,up_blocks.0.attentions.0.transformer_blocks.8.ff.net.2,up_blocks.0.attentions.0.transformer_blocks.9.attn1.to_k,up_blocks.0.attentions.0.transformer_blocks.9.attn1.to_out.0,up_blocks.0.attentions.0.transformer_blocks.9.attn1.to_q,up_blocks.0.attentions.0.transformer_blocks.9.attn1.to_v,up_blocks.0.attentions.0.transformer_blocks.9.attn2.to_k,up_blocks.0.attentions.0.transformer_blocks.9.attn2.to_out.0,up_blocks.0.attentions.0.transformer_blocks.9.attn2.to_q,up_blocks.0.attentions.0.transformer_blocks.9.attn2.to_v,up_blocks.0.attentions.0.transformer_blocks.9.ff.net.0.proj,up_blocks.0.attentions.0.transformer_blocks.9.ff.net.2,up_blocks.0.attentions.1.proj_in,up_blocks.0.attentions.1.proj_out,up_blocks.0.attentions.1.transformer_blocks.0.attn1.to_k,up_blocks.0.attentions.1.transformer_blocks.0.attn1.to_out.0,up_blocks.0.attentions.1.transformer_blocks.0.attn1.to_q,up_blocks.0.attentions.1.transformer_blocks.0.attn1.to_v,up_blocks.0.attentions.1.transformer_blocks.0.attn2.to_k,up_blocks.0.attentions.1.transformer_blocks.0.attn2.to_out.0,up_blocks.0.attentions.1.transformer_blocks.0.attn2.to_q,up_blocks.0.attentions.1.transformer_blocks.0.attn2.to_v,up_blocks.0.attentions.1.transformer_blocks.0.ff.net.0.proj,up_blocks.0.attentions.1.transformer_blocks.0.ff.net.2,up_blocks.0.attentions.1.transformer_blocks.1.attn1.to_k,up_blocks.0.attentions.1.transformer_blocks.1.attn1.to_out.0,up_blocks.0.attentions.1.transformer_blocks.1.attn1.to_q,up_blocks.0.attentions.1.transformer_blocks.1.attn1.to_v,up_blocks.0.attentions.1.transformer_blocks.1.attn2.to_k,up_blocks.0.attentions.1.transformer_blocks.1.attn2.to_out.0,up_blocks.0.attentions.1.transformer_blocks.1.attn2.to_q,up_blocks.0.attentions.1.transformer_blocks.1.attn2.to_v,up_blocks.0.attentions.1.transformer_blocks.1.ff.net.0.proj,up_blocks.0.attentions.1.transformer_blocks.1.ff.net.2,up_blocks.0.attentions.1.transformer_blocks.2.attn1.to_k,up_blocks.0.attentions.1.transformer_blocks.2.attn1.to_out.0,up_blocks.0.attentions.1.transformer_blocks.2.attn1.to_q,up_blocks.0.attentions.1.transformer_blocks.2.attn1.to_v,up_blocks.0.attentions.1.transformer_blocks.2.attn2.to_k,up_blocks.0.attentions.1.transformer_blocks.2.attn2.to_out.0,up_blocks.0.attentions.1.transformer_blocks.2.attn2.to_q,up_blocks.0.attentions.1.transformer_blocks.2.attn2.to_v,up_blocks.0.attentions.1.transformer_blocks.2.ff.net.0.proj,up_blocks.0.attentions.1.transformer_blocks.2.ff.net.2,up_blocks.0.attentions.1.transformer_blocks.3.attn1.to_k,up_blocks.0.attentions.1.transformer_blocks.3.attn1.to_out.0,up_blocks.0.attentions.1.transformer_blocks.3.attn1.to_q,up_blocks.0.attentions.1.transformer_blocks.3.attn1.to_v,up_blocks.0.attentions.1.transformer_blocks.3.attn2.to_k,up_blocks.0.attentions.1.transformer_blocks.3.attn2.to_out.0,up_blocks.0.attentions.1.transformer_blocks.3.attn2.to_q,up_blocks.0.attentions.1.transformer_blocks.3.attn2.to_v,up_blocks.0.attentions.1.transformer_blocks.3.ff.net.0.proj,up_blocks.0.attentions.1.transformer_blocks.3.ff.net.2,up_blocks.0.attentions.1.transformer_blocks.4.attn1.to_k,up_blocks.0.attentions.1.transformer_blocks.4.attn1.to_out.0,up_blocks.0.attentions.1.transformer_blocks.4.attn1.to_q,up_blocks.0.attentions.1.transformer_blocks.4.attn1.to_v,up_blocks.0.attentions.1.transformer_blocks.4.attn2.to_k,up_blocks.0.attentions.1.transformer_blocks.4.attn2.to_out.0,up_blocks.0.attentions.1.transformer_blocks.4.attn2.to_q,up_blocks.0.attentions.1.transformer_blocks.4.attn2.to_v,up_blocks.0.attentions.1.transformer_blocks.4.ff.net.0.proj,up_blocks.0.attentions.1.transformer_blocks.4.ff.net.2,up_blocks.0.attentions.1.transformer_blocks.5.attn1.to_k,up_blocks.0.attentions.1.transformer_blocks.5.attn1.to_out.0,up_blocks.0.attentions.1.transformer_blocks.5.attn1.to_q,up_blocks.0.attentions.1.transformer_blocks.5.attn1.to_v,up_blocks.0.attentions.1.transformer_blocks.5.attn2.to_k,up_blocks.0.attentions.1.transformer_blocks.5.attn2.to_out.0,up_blocks.0.attentions.1.transformer_blocks.5.attn2.to_q,up_blocks.0.attentions.1.transformer_blocks.5.attn2.to_v,up_blocks.0.attentions.1.transformer_blocks.5.ff.net.0.proj,up_blocks.0.attentions.1.transformer_blocks.5.ff.net.2,up_blocks.0.attentions.1.transformer_blocks.6.attn1.to_k,up_blocks.0.attentions.1.transformer_blocks.6.attn1.to_out.0,up_blocks.0.attentions.1.transformer_blocks.6.attn1.to_q,up_blocks.0.attentions.1.transformer_blocks.6.attn1.to_v,up_blocks.0.attentions.1.transformer_blocks.6.attn2.to_k,up_blocks.0.attentions.1.transformer_blocks.6.attn2.to_out.0,up_blocks.0.attentions.1.transformer_blocks.6.attn2.to_q,up_blocks.0.attentions.1.transformer_blocks.6.attn2.to_v,up_blocks.0.attentions.1.transformer_blocks.6.ff.net.0.proj,up_blocks.0.attentions.1.transformer_blocks.6.ff.net.2,up_blocks.0.attentions.1.transformer_blocks.7.attn1.to_k,up_blocks.0.attentions.1.transformer_blocks.7.attn1.to_out.0,up_blocks.0.attentions.1.transformer_blocks.7.attn1.to_q,up_blocks.0.attentions.1.transformer_blocks.7.attn1.to_v,up_blocks.0.attentions.1.transformer_blocks.7.attn2.to_k,up_blocks.0.attentions.1.transformer_blocks.7.attn2.to_out.0,up_blocks.0.attentions.1.transformer_blocks.7.attn2.to_q,up_blocks.0.attentions.1.transformer_blocks.7.attn2.to_v,up_blocks.0.attentions.1.transformer_blocks.7.ff.net.0.proj,up_blocks.0.attentions.1.transformer_blocks.7.ff.net.2,up_blocks.0.attentions.1.transformer_blocks.8.attn1.to_k,up_blocks.0.attentions.1.transformer_blocks.8.attn1.to_out.0,up_blocks.0.attentions.1.transformer_blocks.8.attn1.to_q,up_blocks.0.attentions.1.transformer_blocks.8.attn1.to_v,up_blocks.0.attentions.1.transformer_blocks.8.attn2.to_k,up_blocks.0.attentions.1.transformer_blocks.8.attn2.to_out.0,up_blocks.0.attentions.1.transformer_blocks.8.attn2.to_q,up_blocks.0.attentions.1.transformer_blocks.8.attn2.to_v,up_blocks.0.attentions.1.transformer_blocks.8.ff.net.0.proj,up_blocks.0.attentions.1.transformer_blocks.8.ff.net.2,up_blocks.0.attentions.1.transformer_blocks.9.attn1.to_k,up_blocks.0.attentions.1.transformer_blocks.9.attn1.to_out.0,up_blocks.0.attentions.1.transformer_blocks.9.attn1.to_q,up_blocks.0.attentions.1.transformer_blocks.9.attn1.to_v,up_blocks.0.attentions.1.transformer_blocks.9.attn2.to_k,up_blocks.0.attentions.1.transformer_blocks.9.attn2.to_out.0,up_blocks.0.attentions.1.transformer_blocks.9.attn2.to_q,up_blocks.0.attentions.1.transformer_blocks.9.attn2.to_v,up_blocks.0.attentions.1.transformer_blocks.9.ff.net.0.proj,up_blocks.0.attentions.1.transformer_blocks.9.ff.net.2,up_blocks.0.attentions.2.proj_in,up_blocks.0.attentions.2.proj_out,up_blocks.0.attentions.2.transformer_blocks.0.attn1.to_k,up_blocks.0.attentions.2.transformer_blocks.0.attn1.to_out.0,up_blocks.0.attentions.2.transformer_blocks.0.attn1.to_q,up_blocks.0.attentions.2.transformer_blocks.0.attn1.to_v,up_blocks.0.attentions.2.transformer_blocks.0.attn2.to_k,up_blocks.0.attentions.2.transformer_blocks.0.attn2.to_out.0,up_blocks.0.attentions.2.transformer_blocks.0.attn2.to_q,up_blocks.0.attentions.2.transformer_blocks.0.attn2.to_v,up_blocks.0.attentions.2.transformer_blocks.0.ff.net.0.proj,up_blocks.0.attentions.2.transformer_blocks.0.ff.net.2,up_blocks.0.attentions.2.transformer_blocks.1.attn1.to_k,up_blocks.0.attentions.2.transformer_blocks.1.attn1.to_out.0,up_blocks.0.attentions.2.transformer_blocks.1.attn1.to_q,up_blocks.0.attentions.2.transformer_blocks.1.attn1.to_v,up_blocks.0.attentions.2.transformer_blocks.1.attn2.to_k,up_blocks.0.attentions.2.transformer_blocks.1.attn2.to_out.0,up_blocks.0.attentions.2.transformer_blocks.1.attn2.to_q,up_blocks.0.attentions.2.transformer_blocks.1.attn2.to_v,up_blocks.0.attentions.2.transformer_blocks.1.ff.net.0.proj,up_blocks.0.attentions.2.transformer_blocks.1.ff.net.2,up_blocks.0.attentions.2.transformer_blocks.2.attn1.to_k,up_blocks.0.attentions.2.transformer_blocks.2.attn1.to_out.0,up_blocks.0.attentions.2.transformer_blocks.2.attn1.to_q,up_blocks.0.attentions.2.transformer_blocks.2.attn1.to_v,up_blocks.0.attentions.2.transformer_blocks.2.attn2.to_k,up_blocks.0.attentions.2.transformer_blocks.2.attn2.to_out.0,up_blocks.0.attentions.2.transformer_blocks.2.attn2.to_q,up_blocks.0.attentions.2.transformer_blocks.2.attn2.to_v,up_blocks.0.attentions.2.transformer_blocks.2.ff.net.0.proj,up_blocks.0.attentions.2.transformer_blocks.2.ff.net.2,up_blocks.0.attentions.2.transformer_blocks.3.attn1.to_k,up_blocks.0.attentions.2.transformer_blocks.3.attn1.to_out.0,up_blocks.0.attentions.2.transformer_blocks.3.attn1.to_q,up_blocks.0.attentions.2.transformer_blocks.3.attn1.to_v,up_blocks.0.attentions.2.transformer_blocks.3.attn2.to_k,up_blocks.0.attentions.2.transformer_blocks.3.attn2.to_out.0,up_blocks.0.attentions.2.transformer_blocks.3.attn2.to_q,up_blocks.0.attentions.2.transformer_blocks.3.attn2.to_v,up_blocks.0.attentions.2.transformer_blocks.3.ff.net.0.proj,up_blocks.0.attentions.2.transformer_blocks.3.ff.net.2,up_blocks.0.attentions.2.transformer_blocks.4.attn1.to_k,up_blocks.0.attentions.2.transformer_blocks.4.attn1.to_out.0,up_blocks.0.attentions.2.transformer_blocks.4.attn1.to_q,up_blocks.0.attentions.2.transformer_blocks.4.attn1.to_v,up_blocks.0.attentions.2.transformer_blocks.4.attn2.to_k,up_blocks.0.attentions.2.transformer_blocks.4.attn2.to_out.0,up_blocks.0.attentions.2.transformer_blocks.4.attn2.to_q,up_blocks.0.attentions.2.transformer_blocks.4.attn2.to_v,up_blocks.0.attentions.2.transformer_blocks.4.ff.net.0.proj,up_blocks.0.attentions.2.transformer_blocks.4.ff.net.2,up_blocks.0.attentions.2.transformer_blocks.5.attn1.to_k,up_blocks.0.attentions.2.transformer_blocks.5.attn1.to_out.0,up_blocks.0.attentions.2.transformer_blocks.5.attn1.to_q,up_blocks.0.attentions.2.transformer_blocks.5.attn1.to_v,up_blocks.0.attentions.2.transformer_blocks.5.attn2.to_k,up_blocks.0.attentions.2.transformer_blocks.5.attn2.to_out.0,up_blocks.0.attentions.2.transformer_blocks.5.attn2.to_q,up_blocks.0.attentions.2.transformer_blocks.5.attn2.to_v,up_blocks.0.attentions.2.transformer_blocks.5.ff.net.0.proj,up_blocks.0.attentions.2.transformer_blocks.5.ff.net.2,up_blocks.0.attentions.2.transformer_blocks.6.attn1.to_k,up_blocks.0.attentions.2.transformer_blocks.6.attn1.to_out.0,up_blocks.0.attentions.2.transformer_blocks.6.attn1.to_q,up_blocks.0.attentions.2.transformer_blocks.6.attn1.to_v,up_blocks.0.attentions.2.transformer_blocks.6.attn2.to_k,up_blocks.0.attentions.2.transformer_blocks.6.attn2.to_out.0,up_blocks.0.attentions.2.transformer_blocks.6.attn2.to_q,up_blocks.0.attentions.2.transformer_blocks.6.attn2.to_v,up_blocks.0.attentions.2.transformer_blocks.6.ff.net.0.proj,up_blocks.0.attentions.2.transformer_blocks.6.ff.net.2,up_blocks.0.attentions.2.transformer_blocks.7.attn1.to_k,up_blocks.0.attentions.2.transformer_blocks.7.attn1.to_out.0,up_blocks.0.attentions.2.transformer_blocks.7.attn1.to_q,up_blocks.0.attentions.2.transformer_blocks.7.attn1.to_v,up_blocks.0.attentions.2.transformer_blocks.7.attn2.to_k,up_blocks.0.attentions.2.transformer_blocks.7.attn2.to_out.0,up_blocks.0.attentions.2.transformer_blocks.7.attn2.to_q,up_blocks.0.attentions.2.transformer_blocks.7.attn2.to_v,up_blocks.0.attentions.2.transformer_blocks.7.ff.net.0.proj,up_blocks.0.attentions.2.transformer_blocks.7.ff.net.2,up_blocks.0.attentions.2.transformer_blocks.8.attn1.to_k,up_blocks.0.attentions.2.transformer_blocks.8.attn1.to_out.0,up_blocks.0.attentions.2.transformer_blocks.8.attn1.to_q,up_blocks.0.attentions.2.transformer_blocks.8.attn1.to_v,up_blocks.0.attentions.2.transformer_blocks.8.attn2.to_k,up_blocks.0.attentions.2.transformer_blocks.8.attn2.to_out.0,up_blocks.0.attentions.2.transformer_blocks.8.attn2.to_q,up_blocks.0.attentions.2.transformer_blocks.8.attn2.to_v,up_blocks.0.attentions.2.transformer_blocks.8.ff.net.0.proj,up_blocks.0.attentions.2.transformer_blocks.8.ff.net.2,up_blocks.0.attentions.2.transformer_blocks.9.attn1.to_k,up_blocks.0.attentions.2.transformer_blocks.9.attn1.to_out.0,up_blocks.0.attentions.2.transformer_blocks.9.attn1.to_q,up_blocks.0.attentions.2.transformer_blocks.9.attn1.to_v,up_blocks.0.attentions.2.transformer_blocks.9.attn2.to_k,up_blocks.0.attentions.2.transformer_blocks.9.attn2.to_out.0,up_blocks.0.attentions.2.transformer_blocks.9.attn2.to_q,up_blocks.0.attentions.2.transformer_blocks.9.attn2.to_v,up_blocks.0.attentions.2.transformer_blocks.9.ff.net.0.proj,up_blocks.0.attentions.2.transformer_blocks.9.ff.net.2,up_blocks.1.attentions.0.proj_in,up_blocks.1.attentions.0.proj_out,up_blocks.1.attentions.0.transformer_blocks.0.attn1.to_k,up_blocks.1.attentions.0.transformer_blocks.0.attn1.to_out.0,up_blocks.1.attentions.0.transformer_blocks.0.attn1.to_q,up_blocks.1.attentions.0.transformer_blocks.0.attn1.to_v,up_blocks.1.attentions.0.transformer_blocks.0.attn2.to_k,up_blocks.1.attentions.0.transformer_blocks.0.attn2.to_out.0,up_blocks.1.attentions.0.transformer_blocks.0.attn2.to_q,up_blocks.1.attentions.0.transformer_blocks.0.attn2.to_v,up_blocks.1.attentions.0.transformer_blocks.0.ff.net.0.proj,up_blocks.1.attentions.0.transformer_blocks.0.ff.net.2,up_blocks.1.attentions.0.transformer_blocks.1.attn1.to_k,up_blocks.1.attentions.0.transformer_blocks.1.attn1.to_out.0,up_blocks.1.attentions.0.transformer_blocks.1.attn1.to_q,up_blocks.1.attentions.0.transformer_blocks.1.attn1.to_v,up_blocks.1.attentions.0.transformer_blocks.1.attn2.to_k,up_blocks.1.attentions.0.transformer_blocks.1.attn2.to_out.0,up_blocks.1.attentions.0.transformer_blocks.1.attn2.to_q,up_blocks.1.attentions.0.transformer_blocks.1.attn2.to_v,up_blocks.1.attentions.0.transformer_blocks.1.ff.net.0.proj,up_blocks.1.attentions.0.transformer_blocks.1.ff.net.2,up_blocks.1.attentions.1.proj_in,up_blocks.1.attentions.1.proj_out,up_blocks.1.attentions.1.transformer_blocks.0.attn1.to_k,up_blocks.1.attentions.1.transformer_blocks.0.attn1.to_out.0,up_blocks.1.attentions.1.transformer_blocks.0.attn1.to_q,up_blocks.1.attentions.1.transformer_blocks.0.attn1.to_v,up_blocks.1.attentions.1.transformer_blocks.0.attn2.to_k,up_blocks.1.attentions.1.transformer_blocks.0.attn2.to_out.0,up_blocks.1.attentions.1.transformer_blocks.0.attn2.to_q,up_blocks.1.attentions.1.transformer_blocks.0.attn2.to_v,up_blocks.1.attentions.1.transformer_blocks.0.ff.net.0.proj,up_blocks.1.attentions.1.transformer_blocks.0.ff.net.2,up_blocks.1.attentions.1.transformer_blocks.1.attn1.to_k,up_blocks.1.attentions.1.transformer_blocks.1.attn1.to_out.0,up_blocks.1.attentions.1.transformer_blocks.1.attn1.to_q,up_blocks.1.attentions.1.transformer_blocks.1.attn1.to_v,up_blocks.1.attentions.1.transformer_blocks.1.attn2.to_k,up_blocks.1.attentions.1.transformer_blocks.1.attn2.to_out.0,up_blocks.1.attentions.1.transformer_blocks.1.attn2.to_q,up_blocks.1.attentions.1.transformer_blocks.1.attn2.to_v,up_blocks.1.attentions.1.transformer_blocks.1.ff.net.0.proj,up_blocks.1.attentions.1.transformer_blocks.1.ff.net.2,up_blocks.1.attentions.2.proj_in,up_blocks.1.attentions.2.proj_out,up_blocks.1.attentions.2.transformer_blocks.0.attn1.to_k,up_blocks.1.attentions.2.transformer_blocks.0.attn1.to_out.0,up_blocks.1.attentions.2.transformer_blocks.0.attn1.to_q,up_blocks.1.attentions.2.transformer_blocks.0.attn1.to_v,up_blocks.1.attentions.2.transformer_blocks.0.attn2.to_k,up_blocks.1.attentions.2.transformer_blocks.0.attn2.to_out.0,up_blocks.1.attentions.2.transformer_blocks.0.attn2.to_q,up_blocks.1.attentions.2.transformer_blocks.0.attn2.to_v,up_blocks.1.attentions.2.transformer_blocks.0.ff.net.0.proj,up_blocks.1.attentions.2.transformer_blocks.0.ff.net.2,up_blocks.1.attentions.2.transformer_blocks.1.attn1.to_k,up_blocks.1.attentions.2.transformer_blocks.1.attn1.to_out.0,up_blocks.1.attentions.2.transformer_blocks.1.attn1.to_q,up_blocks.1.attentions.2.transformer_blocks.1.attn1.to_v,up_blocks.1.attentions.2.transformer_blocks.1.attn2.to_k,up_blocks.1.attentions.2.transformer_blocks.1.attn2.to_out.0,up_blocks.1.attentions.2.transformer_blocks.1.attn2.to_q,up_blocks.1.attentions.2.transformer_blocks.1.attn2.to_v,up_blocks.1.attentions.2.transformer_blocks.1.ff.net.0.proj,up_blocks.1.attentions.2.transformer_blocks.1.ff.net.2 \ + --lora_rank 32 \ + --use_gradient_checkpointing \ + --align_to_opensource_format \ + --offload_models stabilityai/stable-diffusion-xl-base-1.0:vae/diffusion_pytorch_model.safetensors \ + --task sft:train diff --git a/examples/stable_diffusion_xl/model_training/special/split_training/validate.py b/examples/stable_diffusion_xl/model_training/special/split_training/validate.py new file mode 100644 index 000000000..a25953aaf --- /dev/null +++ b/examples/stable_diffusion_xl/model_training/special/split_training/validate.py @@ -0,0 +1,37 @@ +from pathlib import Path + + +def latest_checkpoint(directory): + checkpoints = list(Path(directory).glob("epoch-*.safetensors")) + if not checkpoints: + raise FileNotFoundError(f"No checkpoint found in {directory}") + return max(checkpoints, key=lambda path: int(path.stem.rsplit("-", 1)[-1])) + + +import torch +from diffsynth.core import ModelConfig +from diffsynth.pipelines.stable_diffusion_xl import StableDiffusionXLPipeline + +pipe = StableDiffusionXLPipeline.from_pretrained( + torch_dtype=torch.float32, + model_configs=[ + ModelConfig(model_id="stabilityai/stable-diffusion-xl-base-1.0", origin_file_pattern="text_encoder/model.safetensors"), + ModelConfig(model_id="stabilityai/stable-diffusion-xl-base-1.0", origin_file_pattern="text_encoder_2/model.safetensors"), + ModelConfig(model_id="stabilityai/stable-diffusion-xl-base-1.0", origin_file_pattern="unet/diffusion_pytorch_model.safetensors"), + ModelConfig(model_id="stabilityai/stable-diffusion-xl-base-1.0", origin_file_pattern="vae/diffusion_pytorch_model.safetensors"), + ], + tokenizer_config=ModelConfig(model_id="stabilityai/stable-diffusion-xl-base-1.0", origin_file_pattern="tokenizer/"), + tokenizer_2_config=ModelConfig(model_id="stabilityai/stable-diffusion-xl-base-1.0", origin_file_pattern="tokenizer_2/"), +) +pipe.load_lora(pipe.unet, str(latest_checkpoint('./models/train/stable-diffusion-xl-base-1.0_split'))) + +image = pipe( + prompt="a dog", + negative_prompt="", + cfg_scale=7.0, + height=1024, + width=1024, + seed=42, + num_inference_steps=50, +) +image.save('split_training_stable-diffusion-xl.jpg') diff --git a/examples/wanvideo/model_training/special/split_training/Wan2.1-T2V-1.3B.sh b/examples/wanvideo/model_training/special/split_training/Wan2.1-T2V-1.3B.sh new file mode 100644 index 000000000..5c552a04a --- /dev/null +++ b/examples/wanvideo/model_training/special/split_training/Wan2.1-T2V-1.3B.sh @@ -0,0 +1,38 @@ +set -e + +modelscope download --dataset DiffSynth-Studio/diffsynth_example_dataset --include "wanvideo/Wan2.1-T2V-1.3B/*" --local_dir ./data/diffsynth_example_dataset + +# Stage 1: cache deterministic preprocessing outputs. +accelerate launch examples/wanvideo/model_training/train.py \ + --dataset_base_path data/diffsynth_example_dataset/wanvideo/Wan2.1-T2V-1.3B \ + --dataset_metadata_path data/diffsynth_example_dataset/wanvideo/Wan2.1-T2V-1.3B/metadata.csv \ + --height 480 \ + --width 832 \ + --dataset_repeat 1 \ + --model_id_with_origin_paths 'Wan-AI/Wan2.1-T2V-1.3B:diffusion_pytorch_model*.safetensors,Wan-AI/Wan2.1-T2V-1.3B:models_t5_umt5-xxl-enc-bf16.pth,Wan-AI/Wan2.1-T2V-1.3B:Wan2.1_VAE.pth' \ + --learning_rate 1e-4 \ + --num_epochs 5 \ + --remove_prefix_in_ckpt pipe.dit. \ + --output_path ./models/train/Wan2.1-T2V-1.3B_split_cache \ + --lora_base_model dit \ + --lora_target_modules q,k,v,o,ffn.0,ffn.2 \ + --lora_rank 32 \ + --offload_models 'Wan-AI/Wan2.1-T2V-1.3B:diffusion_pytorch_model*.safetensors' \ + --task sft:data_process + +# Stage 2: train LoRA from the cached dataset. +accelerate launch examples/wanvideo/model_training/train.py \ + --dataset_base_path ./models/train/Wan2.1-T2V-1.3B_split_cache \ + --height 480 \ + --width 832 \ + --dataset_repeat 100 \ + --model_id_with_origin_paths 'Wan-AI/Wan2.1-T2V-1.3B:diffusion_pytorch_model*.safetensors,Wan-AI/Wan2.1-T2V-1.3B:models_t5_umt5-xxl-enc-bf16.pth,Wan-AI/Wan2.1-T2V-1.3B:Wan2.1_VAE.pth' \ + --learning_rate 1e-4 \ + --num_epochs 5 \ + --remove_prefix_in_ckpt pipe.dit. \ + --output_path ./models/train/Wan2.1-T2V-1.3B_split \ + --lora_base_model dit \ + --lora_target_modules q,k,v,o,ffn.0,ffn.2 \ + --lora_rank 32 \ + --offload_models Wan-AI/Wan2.1-T2V-1.3B:models_t5_umt5-xxl-enc-bf16.pth,Wan-AI/Wan2.1-T2V-1.3B:Wan2.1_VAE.pth \ + --task sft:train diff --git a/examples/wanvideo/model_training/special/split_training/validate.py b/examples/wanvideo/model_training/special/split_training/validate.py index 737727751..b4533efd4 100644 --- a/examples/wanvideo/model_training/special/split_training/validate.py +++ b/examples/wanvideo/model_training/special/split_training/validate.py @@ -1,28 +1,33 @@ +from pathlib import Path + + +def latest_checkpoint(directory): + checkpoints = list(Path(directory).glob("epoch-*.safetensors")) + if not checkpoints: + raise FileNotFoundError(f"No checkpoint found in {directory}") + return max(checkpoints, key=lambda path: int(path.stem.rsplit("-", 1)[-1])) + + import torch from PIL import Image from diffsynth.utils.data import save_video, VideoData from diffsynth.pipelines.wan_video import WanVideoPipeline, ModelConfig -from modelscope import dataset_snapshot_download pipe = WanVideoPipeline.from_pretrained( torch_dtype=torch.bfloat16, device="cuda", model_configs=[ - ModelConfig(model_id="Wan-AI/Wan2.1-I2V-14B-480P", origin_file_pattern="diffusion_pytorch_model*.safetensors"), - ModelConfig(model_id="Wan-AI/Wan2.1-I2V-14B-480P", origin_file_pattern="models_t5_umt5-xxl-enc-bf16.pth"), - ModelConfig(model_id="Wan-AI/Wan2.1-I2V-14B-480P", origin_file_pattern="Wan2.1_VAE.pth"), - ModelConfig(model_id="Wan-AI/Wan2.1-I2V-14B-480P", origin_file_pattern="models_clip_open-clip-xlm-roberta-large-vit-huge-14.pth"), + ModelConfig(model_id="Wan-AI/Wan2.1-T2V-1.3B", origin_file_pattern="diffusion_pytorch_model*.safetensors"), + ModelConfig(model_id="Wan-AI/Wan2.1-T2V-1.3B", origin_file_pattern="models_t5_umt5-xxl-enc-bf16.pth"), + ModelConfig(model_id="Wan-AI/Wan2.1-T2V-1.3B", origin_file_pattern="Wan2.1_VAE.pth"), ], ) -pipe.load_lora(pipe.dit, "models/train/Wan2.1-I2V-14B-480P_lora_split/epoch-4.safetensors", alpha=1) - -input_image = VideoData("data/example_video_dataset/video1.mp4", height=480, width=832)[0] +pipe.load_lora(pipe.dit, str(latest_checkpoint('./models/train/Wan2.1-T2V-1.3B_split')), alpha=1) video = pipe( prompt="from sunset to night, a small town, light, house, river", negative_prompt="色调艳丽,过曝,静态,细节模糊不清,字幕,风格,作品,画作,画面,静止,整体发灰,最差质量,低质量,JPEG压缩残留,丑陋的,残缺的,多余的手指,画得不好的手部,画得不好的脸部,畸形的,毁容的,形态畸形的肢体,手指融合,静止不动的画面,杂乱的背景,三条腿,背景人很多,倒着走", - input_image=input_image, seed=1, tiled=True ) -save_video(video, "video_Wan2.1-I2V-14B-480P.mp4", fps=15, quality=5) +save_video(video, 'split_training_Wan2.1-T2V-1.3B.mp4', fps=15, quality=5) diff --git a/examples/z_image/model_training/special/split_training/Z-Image.sh b/examples/z_image/model_training/special/split_training/Z-Image.sh new file mode 100644 index 000000000..e38ce5605 --- /dev/null +++ b/examples/z_image/model_training/special/split_training/Z-Image.sh @@ -0,0 +1,40 @@ +set -e + +modelscope download --dataset DiffSynth-Studio/diffsynth_example_dataset --include "z_image/Z-Image/*" --local_dir ./data/diffsynth_example_dataset + +# Stage 1: cache deterministic preprocessing outputs. +accelerate launch examples/z_image/model_training/train.py \ + --dataset_base_path data/diffsynth_example_dataset/z_image/Z-Image \ + --dataset_metadata_path data/diffsynth_example_dataset/z_image/Z-Image/metadata.csv \ + --max_pixels 1048576 \ + --dataset_repeat 1 \ + --model_id_with_origin_paths 'Tongyi-MAI/Z-Image:transformer/*.safetensors,Tongyi-MAI/Z-Image-Turbo:text_encoder/*.safetensors,Tongyi-MAI/Z-Image-Turbo:vae/diffusion_pytorch_model.safetensors' \ + --learning_rate 1e-4 \ + --num_epochs 5 \ + --remove_prefix_in_ckpt pipe.dit. \ + --output_path ./models/train/Z-Image_split_cache \ + --lora_base_model dit \ + --lora_target_modules to_q,to_k,to_v,to_out.0,w1,w2,w3 \ + --lora_rank 32 \ + --use_gradient_checkpointing \ + --dataset_num_workers 8 \ + --offload_models 'Tongyi-MAI/Z-Image:transformer/*.safetensors' \ + --task sft:data_process + +# Stage 2: train LoRA from the cached dataset. +accelerate launch examples/z_image/model_training/train.py \ + --dataset_base_path ./models/train/Z-Image_split_cache \ + --max_pixels 1048576 \ + --dataset_repeat 50 \ + --model_id_with_origin_paths 'Tongyi-MAI/Z-Image:transformer/*.safetensors,Tongyi-MAI/Z-Image-Turbo:text_encoder/*.safetensors,Tongyi-MAI/Z-Image-Turbo:vae/diffusion_pytorch_model.safetensors' \ + --learning_rate 1e-4 \ + --num_epochs 5 \ + --remove_prefix_in_ckpt pipe.dit. \ + --output_path ./models/train/Z-Image_split \ + --lora_base_model dit \ + --lora_target_modules to_q,to_k,to_v,to_out.0,w1,w2,w3 \ + --lora_rank 32 \ + --use_gradient_checkpointing \ + --dataset_num_workers 8 \ + --offload_models 'Tongyi-MAI/Z-Image-Turbo:text_encoder/*.safetensors,Tongyi-MAI/Z-Image-Turbo:vae/diffusion_pytorch_model.safetensors' \ + --task sft:train diff --git a/examples/z_image/model_training/special/split_training/validate.py b/examples/z_image/model_training/special/split_training/validate.py new file mode 100644 index 000000000..2b30ad692 --- /dev/null +++ b/examples/z_image/model_training/special/split_training/validate.py @@ -0,0 +1,28 @@ +from pathlib import Path + + +def latest_checkpoint(directory): + checkpoints = list(Path(directory).glob("epoch-*.safetensors")) + if not checkpoints: + raise FileNotFoundError(f"No checkpoint found in {directory}") + return max(checkpoints, key=lambda path: int(path.stem.rsplit("-", 1)[-1])) + + +from diffsynth.pipelines.z_image import ZImagePipeline, ModelConfig +import torch + + +pipe = ZImagePipeline.from_pretrained( + torch_dtype=torch.bfloat16, + device="cuda", + model_configs=[ + ModelConfig(model_id="Tongyi-MAI/Z-Image", origin_file_pattern="transformer/*.safetensors"), + ModelConfig(model_id="Tongyi-MAI/Z-Image-Turbo", origin_file_pattern="text_encoder/*.safetensors"), + ModelConfig(model_id="Tongyi-MAI/Z-Image-Turbo", origin_file_pattern="vae/diffusion_pytorch_model.safetensors"), + ], + tokenizer_config=ModelConfig(model_id="Tongyi-MAI/Z-Image-Turbo", origin_file_pattern="tokenizer/"), +) +pipe.load_lora(pipe.dit, str(latest_checkpoint('./models/train/Z-Image_split'))) +prompt = "a dog" +image = pipe(prompt=prompt, seed=42, rand_device="cuda", num_inference_steps=50, cfg_scale=4) +image.save('split_training_Z-Image.jpg')