cpu_offload_checkpointing

This commit is contained in:
kijai
2024-09-05 16:17:11 +03:00
parent ba58d55a6c
commit 77587e0193
3 changed files with 16 additions and 2 deletions
+1 -1
View File
@@ -276,7 +276,7 @@ class FluxTrainer:
flux = flux_utils.load_flow_model(name, args.pretrained_model_name_or_path, weight_dtype, "cpu")
if args.gradient_checkpointing:
flux.enable_gradient_checkpointing(args.cpu_offload_checkpointing)
flux.enable_gradient_checkpointing(cpu_offload=args.cpu_offload_checkpointing)
flux.requires_grad_(True)
+4
View File
@@ -36,6 +36,10 @@ class FluxNetworkTrainer(NetworkTrainer):
if args.max_token_length is not None:
logger.warning("max_token_length is not used in Flux training")
assert not args.split_mode or not args.cpu_offload_checkpointing, (
"split_mode and cpu_offload_checkpointing cannot be used together"
)
train_dataset_group.verify_bucket_reso_steps(32) # TODO check this
+11 -1
View File
@@ -444,7 +444,11 @@ class NetworkTrainer:
accelerator.print(f"load network weights from {args.network_weights}: {info}")
if args.gradient_checkpointing:
unet.enable_gradient_checkpointing()
if args.cpu_offload_checkpointing:
unet.enable_gradient_checkpointing(cpu_offload=True)
else:
unet.enable_gradient_checkpointing()
for t_enc, flag in zip(text_encoders, self.get_text_encoders_train_flags(args, text_encoders)):
if flag:
if t_enc.supports_gradient_checkpointing:
@@ -1367,6 +1371,12 @@ def setup_parser() -> argparse.ArgumentParser:
help="use fp8 for U-Net (or DiT), Text Encoder is fp16 or bf16"
" / U-Net(またはDiT)にfp8を使用する。Text Encoderはfp16またはbf16",
)
parser.add_argument(
"--cpu_offload_checkpointing",
action="store_true",
help="[EXPERIMENTAL] enable offloading of tensors to CPU during checkpointing for U-Net or DiT, if supported"
" / 勾配チェックポイント時にテンソルをCPUにオフロードする(U-NetまたはDiTのみ、サポートされている場合)",
)
# parser.add_argument("--loraplus_lr_ratio", default=None, type=float, help="LoRA+ learning rate ratio")
# parser.add_argument("--loraplus_unet_lr_ratio", default=None, type=float, help="LoRA+ UNet learning rate ratio")
# parser.add_argument("--loraplus_text_encoder_lr_ratio", default=None, type=float, help="LoRA+ text encoder learning rate ratio")