diff --git a/flux_train_comfy.py b/flux_train_comfy.py index 18b18b2..22527fc 100644 --- a/flux_train_comfy.py +++ b/flux_train_comfy.py @@ -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) diff --git a/flux_train_network_comfy.py b/flux_train_network_comfy.py index 04366d9..c276fae 100644 --- a/flux_train_network_comfy.py +++ b/flux_train_network_comfy.py @@ -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 diff --git a/train_network.py b/train_network.py index 61ae8d2..1950b62 100644 --- a/train_network.py +++ b/train_network.py @@ -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")