Update Flux2 Control Cfg Distill && Fix Bug in Lora Training Register Hook (#445)

This commit is contained in:
Bubbliiiing
2026-02-03 10:23:10 +08:00
committed by GitHub
parent 0af07603da
commit 0f0e2bd5ab
81 changed files with 9151 additions and 176 deletions
-3
View File
@@ -1025,9 +1025,6 @@ def main():
accelerator.register_save_state_pre_hook(save_model_hook)
accelerator.register_load_state_pre_hook(load_model_hook)
accelerator.register_save_state_pre_hook(save_model_hook)
accelerator.register_load_state_pre_hook(load_model_hook)
if args.gradient_checkpointing:
transformer3d.enable_gradient_checkpointing()
-3
View File
@@ -1086,9 +1086,6 @@ def main():
accelerator.register_save_state_pre_hook(save_model_hook)
accelerator.register_load_state_pre_hook(load_model_hook)
accelerator.register_save_state_pre_hook(save_model_hook)
accelerator.register_load_state_pre_hook(load_model_hook)
if args.gradient_checkpointing:
transformer3d.enable_gradient_checkpointing()
+3 -3
View File
@@ -71,7 +71,7 @@ accelerate launch --mixed_precision="bf16" scripts/flux2_fun/train_control.py \
--enable_bucket \
--low_vram \
--uniform_sampling \
--transformer_path="models/Personalized_Model/FLUX.2-dev-Fun-Controlnet-Union.safetensors" \
--transformer_path="models/Personalized_Model/FLUX.2-dev-Fun-Controlnet-Union-2602.safetensors" \
--trainable_modules "control" \
--resume_from_checkpoint="latest"
```
@@ -112,7 +112,7 @@ accelerate launch --use_deepspeed --deepspeed_config_file config/zero_stage2_con
--enable_bucket \
--low_vram \
--uniform_sampling \
--transformer_path="models/Personalized_Model/FLUX.2-dev-Fun-Controlnet-Union.safetensors" \
--transformer_path="models/Personalized_Model/FLUX.2-dev-Fun-Controlnet-Union-2602.safetensors" \
--trainable_modules "control" \
--resume_from_checkpoint="latest"
```
@@ -153,7 +153,7 @@ accelerate launch --mixed_precision="bf16" --use_fsdp --fsdp_auto_wrap_policy TR
--enable_bucket \
--low_vram \
--uniform_sampling \
--transformer_path="models/Personalized_Model/FLUX.2-dev-Fun-Controlnet-Union.safetensors" \
--transformer_path="models/Personalized_Model/FLUX.2-dev-Fun-Controlnet-Union-2602.safetensors" \
--trainable_modules "control" \
--resume_from_checkpoint="latest"
```
+1 -1
View File
@@ -31,6 +31,6 @@ accelerate launch --mixed_precision="bf16" scripts/flux2_fun/train_control.py \
--enable_bucket \
--low_vram \
--uniform_sampling \
--transformer_path="models/Personalized_Model/FLUX.2-dev-Fun-Controlnet-Union.safetensors" \
--transformer_path="models/Personalized_Model/FLUX.2-dev-Fun-Controlnet-Union-2602.safetensors" \
--trainable_modules "control" \
--resume_from_checkpoint="latest"
File diff suppressed because it is too large Load Diff
@@ -0,0 +1,37 @@
export MODEL_NAME="models/Diffusion_Transformer/FLUX.2-dev"
export DATASET_NAME="datasets/internal_datasets/"
export DATASET_META_NAME="datasets/internal_datasets/metadata.json"
# NCCL_IB_DISABLE=1 and NCCL_P2P_DISABLE=1 are used in multi nodes without RDMA.
# export NCCL_IB_DISABLE=1
# export NCCL_P2P_DISABLE=1
NCCL_DEBUG=INFO
accelerate launch --mixed_precision="bf16" scripts/flux2_fun/train_control_distill.py \
--config_path="config/flux2/flux2_control.yaml" \
--pretrained_model_name_or_path=$MODEL_NAME \
--train_data_dir=$DATASET_NAME \
--train_data_meta=$DATASET_META_NAME \
--train_batch_size=1 \
--image_sample_size=1328 \
--gradient_accumulation_steps=1 \
--dataloader_num_workers=8 \
--num_train_epochs=100 \
--checkpointing_steps=50 \
--learning_rate=2e-06 \
--lr_scheduler="constant_with_warmup" \
--lr_warmup_steps=100 \
--seed=42 \
--output_dir="output_dir_flux2_control_CFG_Distill" \
--gradient_checkpointing \
--mixed_precision="bf16" \
--adam_weight_decay=3e-2 \
--adam_epsilon=1e-10 \
--vae_mini_batch=1 \
--max_grad_norm=0.05 \
--enable_bucket \
--low_vram \
--uniform_sampling \
--transformer_path="models/Personalized_Model/FLUX.2-dev-Fun-Controlnet-Union-2602.safetensors" \
--trainable_modules "control" \
--random_hw_adapt \
--resume_from_checkpoint="latest"
-3
View File
@@ -1130,9 +1130,6 @@ def main():
accelerator.register_save_state_pre_hook(save_model_hook)
accelerator.register_load_state_pre_hook(load_model_hook)
accelerator.register_save_state_pre_hook(save_model_hook)
accelerator.register_load_state_pre_hook(load_model_hook)
if args.gradient_checkpointing:
transformer3d.enable_gradient_checkpointing()
-3
View File
@@ -1039,9 +1039,6 @@ def main():
accelerator.register_save_state_pre_hook(save_model_hook)
accelerator.register_load_state_pre_hook(load_model_hook)
accelerator.register_save_state_pre_hook(save_model_hook)
accelerator.register_load_state_pre_hook(load_model_hook)
if args.gradient_checkpointing:
transformer3d.enable_gradient_checkpointing()
-3
View File
@@ -966,9 +966,6 @@ def main():
accelerator.register_save_state_pre_hook(save_model_hook)
accelerator.register_load_state_pre_hook(load_model_hook)
accelerator.register_save_state_pre_hook(save_model_hook)
accelerator.register_load_state_pre_hook(load_model_hook)
if args.gradient_checkpointing:
transformer3d.enable_gradient_checkpointing()
-3
View File
@@ -922,9 +922,6 @@ def main():
accelerator.register_save_state_pre_hook(save_model_hook)
accelerator.register_load_state_pre_hook(load_model_hook)
accelerator.register_save_state_pre_hook(save_model_hook)
accelerator.register_load_state_pre_hook(load_model_hook)
if args.gradient_checkpointing:
transformer3d.enable_gradient_checkpointing()
+2
View File
@@ -945,6 +945,8 @@ def main():
state_dict = state_dict["state_dict"] if "state_dict" in state_dict else state_dict
m, u = generator_transformer3d.load_state_dict(state_dict, strict=False)
m, u = real_score_transformer3d.load_state_dict(state_dict, strict=False)
m, u = fake_score_transformer3d.load_state_dict(state_dict, strict=False)
print(f"missing keys: {len(m)}, unexpected keys: {len(u)}")
assert len(u) == 0
+2 -3
View File
@@ -1007,6 +1007,8 @@ def main():
state_dict = state_dict["state_dict"] if "state_dict" in state_dict else state_dict
m, u = generator_transformer3d.load_state_dict(state_dict, strict=False)
m, u = real_score_transformer3d.load_state_dict(state_dict, strict=False)
m, u = fake_score_transformer3d.load_state_dict(state_dict, strict=False)
print(f"missing keys: {len(m)}, unexpected keys: {len(u)}")
assert len(u) == 0
@@ -1122,9 +1124,6 @@ def main():
accelerator.register_save_state_pre_hook(save_model_hook)
accelerator.register_load_state_pre_hook(load_model_hook)
accelerator.register_save_state_pre_hook(save_model_hook)
accelerator.register_load_state_pre_hook(load_model_hook)
accelerator_fake_score_transformer3d.register_save_state_pre_hook(save_model_hook)
accelerator_fake_score_transformer3d.register_load_state_pre_hook(load_model_hook)
-3
View File
@@ -1059,9 +1059,6 @@ def main():
accelerator.register_save_state_pre_hook(save_model_hook)
accelerator.register_load_state_pre_hook(load_model_hook)
accelerator.register_save_state_pre_hook(save_model_hook)
accelerator.register_load_state_pre_hook(load_model_hook)
if args.gradient_checkpointing:
transformer3d.enable_gradient_checkpointing()
-3
View File
@@ -970,9 +970,6 @@ def main():
accelerator.register_save_state_pre_hook(save_model_hook)
accelerator.register_load_state_pre_hook(load_model_hook)
accelerator.register_save_state_pre_hook(save_model_hook)
accelerator.register_load_state_pre_hook(load_model_hook)
if args.gradient_checkpointing:
transformer3d.enable_gradient_checkpointing()
-3
View File
@@ -1029,9 +1029,6 @@ def main():
accelerator.register_save_state_pre_hook(save_model_hook)
accelerator.register_load_state_pre_hook(load_model_hook)
accelerator.register_save_state_pre_hook(save_model_hook)
accelerator.register_load_state_pre_hook(load_model_hook)
if args.gradient_checkpointing:
transformer3d.enable_gradient_checkpointing()
-3
View File
@@ -1045,9 +1045,6 @@ def main():
accelerator.register_save_state_pre_hook(save_model_hook)
accelerator.register_load_state_pre_hook(load_model_hook)
accelerator.register_save_state_pre_hook(save_model_hook)
accelerator.register_load_state_pre_hook(load_model_hook)
if args.gradient_checkpointing:
transformer3d.enable_gradient_checkpointing()
-3
View File
@@ -1169,9 +1169,6 @@ def main():
accelerator.register_save_state_pre_hook(save_model_hook)
accelerator.register_load_state_pre_hook(load_model_hook)
accelerator.register_save_state_pre_hook(save_model_hook)
accelerator.register_load_state_pre_hook(load_model_hook)
accelerator_fake_score_transformer3d.register_save_state_pre_hook(save_model_hook)
accelerator_fake_score_transformer3d.register_load_state_pre_hook(load_model_hook)
-3
View File
@@ -1111,9 +1111,6 @@ def main():
accelerator.register_save_state_pre_hook(save_model_hook)
accelerator.register_load_state_pre_hook(load_model_hook)
accelerator.register_save_state_pre_hook(save_model_hook)
accelerator.register_load_state_pre_hook(load_model_hook)
if args.gradient_checkpointing:
transformer3d.enable_gradient_checkpointing()
-3
View File
@@ -1095,9 +1095,6 @@ def main():
accelerator.register_save_state_pre_hook(save_model_hook)
accelerator.register_load_state_pre_hook(load_model_hook)
accelerator.register_save_state_pre_hook(save_model_hook)
accelerator.register_load_state_pre_hook(load_model_hook)
if args.gradient_checkpointing:
transformer3d.enable_gradient_checkpointing()
-3
View File
@@ -1059,9 +1059,6 @@ def main():
accelerator.register_save_state_pre_hook(save_model_hook)
accelerator.register_load_state_pre_hook(load_model_hook)
accelerator.register_save_state_pre_hook(save_model_hook)
accelerator.register_load_state_pre_hook(load_model_hook)
if args.gradient_checkpointing:
transformer3d.enable_gradient_checkpointing()
-3
View File
@@ -1074,9 +1074,6 @@ def main():
accelerator.register_save_state_pre_hook(save_model_hook)
accelerator.register_load_state_pre_hook(load_model_hook)
accelerator.register_save_state_pre_hook(save_model_hook)
accelerator.register_load_state_pre_hook(load_model_hook)
if args.gradient_checkpointing:
transformer3d.enable_gradient_checkpointing()
-3
View File
@@ -946,9 +946,6 @@ def main():
accelerator.register_save_state_pre_hook(save_model_hook)
accelerator.register_load_state_pre_hook(load_model_hook)
accelerator.register_save_state_pre_hook(save_model_hook)
accelerator.register_load_state_pre_hook(load_model_hook)
if args.gradient_checkpointing:
transformer3d.enable_gradient_checkpointing()
+2 -2
View File
@@ -1,4 +1,4 @@
export MODEL_NAME="models/Diffusion_Transformer/Z-Image-Turbo"
export MODEL_NAME="models/Diffusion_Transformer/Z-Image"
export DATASET_NAME="datasets/internal_datasets/"
export DATASET_META_NAME="datasets/internal_datasets/metadata.json"
# NCCL_IB_DISABLE=1 and NCCL_P2P_DISABLE=1 are used in multi nodes without RDMA.
@@ -31,5 +31,5 @@ accelerate launch --mixed_precision="bf16" scripts/z_image_fun/train_control.py
--enable_bucket \
--uniform_sampling \
--add_inpaint_info \
--transformer_path="models/Personalized_Model/Z-Image-Turbo-Fun-Controlnet-Union-2.1.safetensors" \
--transformer_path="models/Personalized_Model/Z-Image-Fun-Controlnet-Union-2.1.safetensors" \
--trainable_modules "control"
@@ -1659,6 +1659,8 @@ def main():
if args.low_vram and not args.enable_text_encoder_in_dataloader:
text_encoder.to('cpu')
torch.cuda.empty_cache()
if args.low_vram:
real_score_transformer3d = real_score_transformer3d.to(accelerator.device)
with accelerator.accumulate(generator_transformer3d):
def get_sigmas(timesteps, n_dim=4, dtype=torch.float32):
@@ -0,0 +1,35 @@
export MODEL_NAME="models/Diffusion_Transformer/Z-Image-Turbo"
export DATASET_NAME="datasets/internal_datasets/"
export DATASET_META_NAME="datasets/internal_datasets/metadata.json"
# NCCL_IB_DISABLE=1 and NCCL_P2P_DISABLE=1 are used in multi nodes without RDMA.
# export NCCL_IB_DISABLE=1
# export NCCL_P2P_DISABLE=1
NCCL_DEBUG=INFO
accelerate launch --mixed_precision="bf16" scripts/z_image_fun/train_control.py \
--config_path="config/z_image/z_image_control_2.1.yaml" \
--pretrained_model_name_or_path=$MODEL_NAME \
--train_data_dir=$DATASET_NAME \
--train_data_meta=$DATASET_META_NAME \
--train_batch_size=1 \
--image_sample_size=1328 \
--gradient_accumulation_steps=1 \
--dataloader_num_workers=8 \
--num_train_epochs=100 \
--checkpointing_steps=50 \
--learning_rate=2e-05 \
--lr_scheduler="constant_with_warmup" \
--lr_warmup_steps=100 \
--seed=42 \
--output_dir="output_dir_z_image_control" \
--gradient_checkpointing \
--mixed_precision="bf16" \
--adam_weight_decay=3e-2 \
--adam_epsilon=1e-10 \
--vae_mini_batch=1 \
--max_grad_norm=0.05 \
--enable_bucket \
--uniform_sampling \
--add_inpaint_info \
--transformer_path="models/Personalized_Model/Z-Image-Turbo-Fun-Controlnet-Union-2.1.safetensors" \
--trainable_modules "control"