Update Peft Lora && Update Readme (#376)

This commit is contained in:
Bubbliiiing
2025-11-20 17:44:22 +08:00
committed by GitHub
parent ccc3b1055e
commit 037a2e8360
79 changed files with 2499 additions and 849 deletions
+12 -44
View File
@@ -9,6 +9,17 @@ Some parameters in the sh file can be confusing, and they are explained in this
- For example, when `random_hw_adapt` is enabled, `image_sample_size=1024`, the resolution of image inputs for training is `512x512` to `1024x1024`
- `resume_from_checkpoint` is used to set the training should be resumed from a previous checkpoint. Use a path or `"latest"` to automatically select the last available checkpoint.
When train model with multi machines, please set the params as follows:
```sh
export MASTER_ADDR="your master address"
export MASTER_PORT=10086
export WORLD_SIZE=1 # The number of machines
export NUM_PROCESS=8 # The number of processes, such as WORLD_SIZE * 8
export RANK=0 # The rank of this machine
accelerate launch --mixed_precision="bf16" --main_process_ip=$MASTER_ADDR --main_process_port=$MASTER_PORT --num_machines=$WORLD_SIZE --num_processes=$NUM_PROCESS --machine_rank=$RANK scripts/xxx/xxx.py
```
Without deepspeed:
Training flux without DeepSpeed may result in insufficient GPU memory.
@@ -47,7 +58,7 @@ accelerate launch --mixed_precision="bf16" scripts/flux/train.py \
--trainable_modules "."
```
With deepspeed zero-2:
With Deepspeed Zero-2:
```sh
export MODEL_NAME="models/Diffusion_Transformer/FLUX.1-dev"
@@ -84,49 +95,6 @@ accelerate launch --use_deepspeed --deepspeed_config_file config/zero_stage2_con
--trainable_modules "."
```
Deepspeed zero-3:
After training, you can use the following command to get the final model:
```sh
python scripts/zero_to_bf16.py output_dir/checkpoint-{our-num-steps} output_dir/checkpoint-{your-num-steps}-outputs --max_shard_size 80GB --safe_serialization
```
Training shell command is as follows:
```sh
export MODEL_NAME="models/Diffusion_Transformer/FLUX.1-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 --zero_stage 3 --zero3_save_16bit_model true --zero3_init_flag true --use_deepspeed --deepspeed_config_file config/zero_stage3_config.json --deepspeed_multinode_launcher standard scripts/flux/train.py \
--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=1024 \
--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" \
--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 \
--trainable_modules "."
```
With FSDP:
```sh
+24 -41
View File
@@ -8,6 +8,21 @@ Some parameters in the sh file can be confusing, and they are explained in this
- `random_hw_adapt` is used to enable automatic height and width scaling for images. When `random_hw_adapt` is enabled, the training images will have their height and width set to `image_sample_size` as the maximum and `512` as the minimum.
- For example, when `random_hw_adapt` is enabled, `image_sample_size=1024`, the resolution of image inputs for training is `512x512` to `1024x1024`
- `resume_from_checkpoint` is used to set the training should be resumed from a previous checkpoint. Use a path or `"latest"` to automatically select the last available checkpoint.
- `target_name` represents the components/modules to which LoRA will be applied, separated by commas.
- `use_peft_lora` indicates whether to use the PEFT module for adding LoRA. Using this module will be more memory-efficient.
- `rank` means the dimension of the LoRA update matrices.
- `network_alpha` means the scale of the LoRA update matrices.
When train model with multi machines, please set the params as follows:
```sh
export MASTER_ADDR="your master address"
export MASTER_PORT=10086
export WORLD_SIZE=1 # The number of machines
export NUM_PROCESS=8 # The number of processes, such as WORLD_SIZE * 8
export RANK=0 # The rank of this machine
accelerate launch --mixed_precision="bf16" --main_process_ip=$MASTER_ADDR --main_process_port=$MASTER_PORT --num_machines=$WORLD_SIZE --num_processes=$NUM_PROCESS --machine_rank=$RANK scripts/xxx/xxx.py
```
Without deepspeed:
@@ -41,10 +56,14 @@ accelerate launch --mixed_precision="bf16" scripts/flux/train_lora.py \
--vae_mini_batch=1 \
--max_grad_norm=0.05 \
--enable_bucket \
--rank=64 \
--network_alpha=32 \
--target_name="to_q,to_k,to_v,ff.0,ff.2,ff_context.0,ff_context.2" \
--use_peft_lora \
--uniform_sampling
```
With deepspeed zero-2:
With Deepspeed Zero-2:
```sh
export MODEL_NAME="models/Diffusion_Transformer/FLUX.1-dev"
@@ -75,46 +94,10 @@ accelerate launch --use_deepspeed --deepspeed_config_file config/zero_stage2_con
--vae_mini_batch=1 \
--max_grad_norm=0.05 \
--enable_bucket \
--uniform_sampling
```
Deepspeed zero-3:
After training, you can use the following command to get the final model:
```sh
python scripts/zero_to_bf16.py output_dir/checkpoint-{our-num-steps} output_dir/checkpoint-{your-num-steps}-outputs --max_shard_size 80GB --safe_serialization
```
Training shell command is as follows:
```sh
export MODEL_NAME="models/Diffusion_Transformer/FLUX.1-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 --zero_stage 3 --zero3_save_16bit_model true --zero3_init_flag true --use_deepspeed --deepspeed_config_file config/zero_stage3_config.json --deepspeed_multinode_launcher standard scripts/flux/train_lora.py \
--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=1024 \
--gradient_accumulation_steps=1 \
--dataloader_num_workers=8 \
--num_train_epochs=100 \
--checkpointing_steps=50 \
--learning_rate=1e-04 \
--seed=42 \
--output_dir="output_dir" \
--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 \
--rank=64 \
--network_alpha=32 \
--target_name="to_q,to_k,to_v,ff.0,ff.2,ff_context.0,ff_context.2" \
--use_peft_lora \
--uniform_sampling
```
+5
View File
@@ -989,7 +989,12 @@ def main():
elif zero_stage == 3:
# create custom saving & loading hooks so that `accelerator.save_state(...)` serializes in a nice format
def save_model_hook(models, weights, output_dir):
accelerate_state_dict = accelerator.get_state_dict(models[-1], unwrap=True)
if accelerator.is_main_process:
from safetensors.torch import save_file
safetensor_save_path = os.path.join(output_dir, f"diffusion_pytorch_model.safetensors")
save_file(accelerate_state_dict, safetensor_save_path, metadata={"format": "pt"})
with open(os.path.join(output_dir, "sampler_pos_start.pkl"), 'wb') as file:
pickle.dump([batch_sampler.sampler._pos_start, first_epoch], file)
+87 -29
View File
@@ -577,6 +577,9 @@ def parse_args():
default=64,
help=("The dimension of the LoRA update matrices."),
)
parser.add_argument(
"--use_peft_lora", action="store_true", help="Whether or not to use peft lora."
)
parser.add_argument(
"--train_text_encoder",
action="store_true",
@@ -710,6 +713,12 @@ def parse_args():
default=3.5,
help="the FLUX.1 dev variant is a guidance distilled model",
)
parser.add_argument(
"--target_name",
type=str,
default=None,
help=("The module is trained in loras. "),
)
args = parser.parse_args()
env_local_rank = int(os.environ.get("LOCAL_RANK", -1))
@@ -894,16 +903,25 @@ def main():
transformer3d.requires_grad_(False)
# Lora will work with this...
network = create_network(
1.0,
args.rank,
args.network_alpha,
text_encoder,
transformer3d,
neuron_dropout=None,
skip_name=args.lora_skip_name,
)
network.apply_to(text_encoder, transformer3d, args.train_text_encoder, True)
if args.use_peft_lora:
from peft import LoraConfig, inject_adapter_in_model, get_peft_model_state_dict
lora_config = LoraConfig(r=args.rank, lora_alpha=args.network_alpha, target_modules=args.target_name.split(","))
transformer3d = inject_adapter_in_model(lora_config, transformer3d)
network = None
else:
network = create_network(
1.0,
args.rank,
args.network_alpha,
text_encoder,
transformer3d,
neuron_dropout=None,
target_name=args.target_name,
skip_name=args.lora_skip_name,
)
network = network.to(weight_dtype)
network.apply_to(text_encoder, transformer3d, args.train_text_encoder and not args.training_with_video_token_length, True)
if args.transformer_path is not None:
print(f"From checkpoint: {args.transformer_path}")
@@ -939,13 +957,14 @@ def main():
accelerate_state_dict = accelerator.get_state_dict(models[-1], unwrap=True)
if accelerator.is_main_process:
from safetensors.torch import save_file
safetensor_save_path = os.path.join(output_dir, f"lora_diffusion_pytorch_model.safetensors")
network_state_dict = {}
for key in accelerate_state_dict:
if "network" in key:
network_state_dict[key.replace("network.", "")] = accelerate_state_dict[key].to(weight_dtype)
if args.use_peft_lora:
network_state_dict = get_peft_model_state_dict(accelerator.unwrap_model(models[-1]), accelerate_state_dict)
else:
network_state_dict = {}
for key in accelerate_state_dict:
if "network" in key:
network_state_dict[key.replace("network.", "")] = accelerate_state_dict[key].to(weight_dtype)
save_file(network_state_dict, safetensor_save_path, metadata={"format": "pt"})
with open(os.path.join(output_dir, "sampler_pos_start.pkl"), 'wb') as file:
@@ -960,8 +979,18 @@ def main():
print(f"Load pkl from {pkl_path}. Get loaded_number = {loaded_number}.")
elif zero_stage == 3:
# create custom saving & loading hooks so that `accelerator.save_state(...)` serializes in a nice format
def save_model_hook(models, weights, output_dir):
accelerate_state_dict = accelerator.get_state_dict(models[-1], unwrap=True)
if accelerator.is_main_process:
from safetensors.torch import save_file
safetensor_save_path = os.path.join(output_dir, f"lora_diffusion_pytorch_model.safetensors")
if args.use_peft_lora:
network_state_dict = get_peft_model_state_dict(accelerator.unwrap_model(models[-1]), accelerate_state_dict)
else:
network_state_dict = accelerate_state_dict
save_file(network_state_dict, safetensor_save_path, metadata={"format": "pt"})
with open(os.path.join(output_dir, "sampler_pos_start.pkl"), 'wb') as file:
pickle.dump([batch_sampler.sampler._pos_start, first_epoch], file)
@@ -973,10 +1002,15 @@ def main():
batch_sampler.sampler._pos_start = max(loaded_number - args.dataloader_num_workers * accelerator.num_processes * 2, 0)
print(f"Load pkl from {pkl_path}. Get loaded_number = {loaded_number}.")
else:
# create custom saving & loading hooks so that `accelerator.save_state(...)` serializes in a nice format
def save_model_hook(models, weights, output_dir):
if accelerator.is_main_process:
safetensor_save_path = os.path.join(output_dir, f"lora_diffusion_pytorch_model.safetensors")
save_model(safetensor_save_path, accelerator.unwrap_model(models[-1]))
if args.use_peft_lora:
save_model(safetensor_save_path, get_peft_model_state_dict(accelerator.unwrap_model(models[-1])))
else:
save_model(safetensor_save_path, accelerator.unwrap_model(models[-1]))
if not args.use_deepspeed:
for _ in range(len(weights)):
weights.pop()
@@ -1030,9 +1064,14 @@ def main():
else:
optimizer_cls = torch.optim.AdamW
logging.info("Add network parameters")
trainable_params = list(filter(lambda p: p.requires_grad, network.parameters()))
trainable_params_optim = network.prepare_optimizer_params(args.learning_rate / 2, args.learning_rate, args.learning_rate)
if args.use_peft_lora:
logging.info("Add peft parameters")
trainable_params = list(filter(lambda p: p.requires_grad, transformer3d.parameters()))
trainable_params_optim = list(filter(lambda p: p.requires_grad, transformer3d.parameters()))
else:
logging.info("Add network parameters")
trainable_params = list(filter(lambda p: p.requires_grad, network.parameters()))
trainable_params_optim = network.prepare_optimizer_params(args.learning_rate / 2, args.learning_rate, args.learning_rate)
if args.use_came:
optimizer = optimizer_cls(
@@ -1252,9 +1291,13 @@ def main():
)
# Prepare everything with our `accelerator`.
if fsdp_stage != 0:
if args.use_peft_lora:
transformer3d, optimizer, train_dataloader, lr_scheduler = accelerator.prepare(
transformer3d, optimizer, train_dataloader, lr_scheduler
)
elif fsdp_stage != 0:
transformer3d.network = network
transformer3d = transformer3d.to(weight_dtype)
transformer3d = transformer3d.to(dtype=weight_dtype)
transformer3d, optimizer, train_dataloader, lr_scheduler = accelerator.prepare(
transformer3d, optimizer, train_dataloader, lr_scheduler
)
@@ -1263,14 +1306,14 @@ def main():
network, optimizer, train_dataloader, lr_scheduler
)
if zero_stage == 3:
if zero_stage != 0 and not args.use_peft_lora:
from functools import partial
from videox_fun.dist import set_multi_gpus_devices, shard_model
shard_fn = partial(shard_model, device_id=accelerator.device, param_dtype=weight_dtype, module_to_wrapper=transformer3d.transformer_blocks)
transformer3d = shard_fn(transformer3d)
if fsdp_stage != 0:
if fsdp_stage != 0 or zero_stage != 0:
from functools import partial
from videox_fun.dist import set_multi_gpus_devices, shard_model
@@ -1417,6 +1460,10 @@ def main():
def save_model(ckpt_file, unwrapped_nw):
os.makedirs(args.output_dir, exist_ok=True)
accelerator.print(f"\nsaving checkpoint: {ckpt_file}")
if isinstance(unwrapped_nw, dict):
from safetensors.torch import save_file
save_file(unwrapped_nw, ckpt_file, metadata={"format": "pt"})
return ckpt_file
unwrapped_nw.save_weights(ckpt_file, weight_dtype, None)
progress_bar = tqdm(
@@ -1644,9 +1691,14 @@ def main():
torch.cuda.empty_cache()
torch.cuda.ipc_collect()
if not args.save_state:
safetensor_save_path = os.path.join(args.output_dir, f"checkpoint-{global_step}.safetensors")
save_model(safetensor_save_path, accelerator.unwrap_model(network))
logger.info(f"Saved safetensor to {safetensor_save_path}")
if args.use_peft_lora:
safetensor_save_path = os.path.join(args.output_dir, f"checkpoint-{global_step}.safetensors")
save_model(safetensor_save_path, get_peft_model_state_dict(accelerator.unwrap_model(transformer3d)))
logger.info(f"Saved safetensor to {safetensor_save_path}")
else:
safetensor_save_path = os.path.join(args.output_dir, f"checkpoint-{global_step}.safetensors")
save_model(safetensor_save_path, accelerator.unwrap_model(network))
logger.info(f"Saved safetensor to {safetensor_save_path}")
else:
accelerator_save_path = os.path.join(args.output_dir, f"checkpoint-{global_step}")
accelerator.save_state(accelerator_save_path)
@@ -1697,8 +1749,14 @@ def main():
torch.cuda.empty_cache()
torch.cuda.ipc_collect()
if not args.save_state:
safetensor_save_path = os.path.join(args.output_dir, f"checkpoint-{global_step}.safetensors")
save_model(safetensor_save_path, accelerator.unwrap_model(network))
if args.use_peft_lora:
safetensor_save_path = os.path.join(args.output_dir, f"checkpoint-{global_step}.safetensors")
save_model(safetensor_save_path, get_peft_model_state_dict(accelerator.unwrap_model(transformer3d)))
logger.info(f"Saved safetensor to {safetensor_save_path}")
else:
safetensor_save_path = os.path.join(args.output_dir, f"checkpoint-{global_step}.safetensors")
save_model(safetensor_save_path, accelerator.unwrap_model(network))
logger.info(f"Saved safetensor to {safetensor_save_path}")
else:
accelerator_save_path = os.path.join(args.output_dir, f"checkpoint-{global_step}")
accelerator.save_state(accelerator_save_path)
+4
View File
@@ -26,4 +26,8 @@ accelerate launch --mixed_precision="bf16" scripts/flux/train_lora.py \
--vae_mini_batch=1 \
--max_grad_norm=0.05 \
--enable_bucket \
--rank=64 \
--network_alpha=32 \
--target_name="to_q,to_k,to_v,ff.0,ff.2,ff_context.0,ff_context.2" \
--use_peft_lora \
--uniform_sampling