From 0d1d93931f415b8d27add300e092c1c3640adb4d Mon Sep 17 00:00:00 2001 From: kijai <40791699+kijai@users.noreply.github.com> Date: Thu, 5 Sep 2024 15:58:17 +0300 Subject: [PATCH] cpu_offload_checkpointing --- flux_train_comfy.py | 4 ++-- flux_train_network_comfy.py | 4 ++++ train_network.py | 11 ++++++++++- 3 files changed, 16 insertions(+), 3 deletions(-) diff --git a/flux_train_comfy.py b/flux_train_comfy.py index 18b18b2..b6fdc61 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) @@ -680,7 +680,7 @@ class FluxTrainer: else: with torch.no_grad(): # encode images to latents. images are [-1, 1] - latents = ae.encode(batch["images"]) + latents = ae.encode(batch["images"].to(ae.dtype)).to(accelerator.device, dtype=weight_dtype) # NaNが含まれていれば警告を表示し0に置き換える if torch.any(torch.isnan(latents)): diff --git a/flux_train_network_comfy.py b/flux_train_network_comfy.py index 204789c..41488c2 100644 --- a/flux_train_network_comfy.py +++ b/flux_train_network_comfy.py @@ -43,6 +43,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 def get_flux_model_name(self, args): diff --git a/train_network.py b/train_network.py index 615d90d..a6c8dc5 100644 --- a/train_network.py +++ b/train_network.py @@ -457,7 +457,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: @@ -1383,6 +1387,11 @@ def setup_parser() -> argparse.ArgumentParser: help="initial step number including all epochs, 0 means first step (same as not specifying). overwrites initial_epoch." + " / 初期ステップ数、全エポックを含むステップ数、0で最初のステップ(未指定時と同じ)。initial_epochを上書きする", ) + 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", + ) # 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")