Update Flux2 Control Cfg Distill && Fix Bug in Lora Training Register Hook (#445)
This commit is contained in:
@@ -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()
|
||||
|
||||
|
||||
@@ -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()
|
||||
|
||||
|
||||
@@ -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"
|
||||
```
|
||||
@@ -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"
|
||||
@@ -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()
|
||||
|
||||
|
||||
@@ -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()
|
||||
|
||||
|
||||
@@ -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()
|
||||
|
||||
|
||||
@@ -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()
|
||||
|
||||
|
||||
@@ -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
|
||||
|
||||
|
||||
@@ -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)
|
||||
|
||||
|
||||
@@ -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()
|
||||
|
||||
|
||||
@@ -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()
|
||||
|
||||
|
||||
@@ -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()
|
||||
|
||||
|
||||
@@ -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()
|
||||
|
||||
|
||||
@@ -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)
|
||||
|
||||
|
||||
@@ -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()
|
||||
|
||||
|
||||
@@ -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()
|
||||
|
||||
|
||||
@@ -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()
|
||||
|
||||
|
||||
@@ -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()
|
||||
|
||||
|
||||
@@ -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()
|
||||
|
||||
|
||||
@@ -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"
|
||||
Reference in New Issue
Block a user