Update Peft Lora && Update Readme (#376)
This commit is contained in:
@@ -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
|
||||
|
||||
@@ -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
|
||||
```
|
||||
|
||||
|
||||
@@ -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
@@ -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)
|
||||
|
||||
@@ -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
|
||||
Reference in New Issue
Block a user