From 7095626d6aa60f15e24fa551b2d53f0b05dd47c4 Mon Sep 17 00:00:00 2001 From: kijai <40791699+kijai@users.noreply.github.com> Date: Sat, 31 Aug 2024 16:50:18 +0300 Subject: [PATCH] fix the ability to disable fp8_base --- nodes.py | 2 +- train_network.py | 2 +- 2 files changed, 2 insertions(+), 2 deletions(-) diff --git a/nodes.py b/nodes.py index e645ef5..d114290 100644 --- a/nodes.py +++ b/nodes.py @@ -396,7 +396,7 @@ class InitFluxLoRATraining: "t5xxl_max_token_length": 512, "alpha_mask": dataset["alpha_mask"], "network_train_unet_only": True if train_clip_l == 'disabled' else False, - "fp8_base_unet": True if train_clip_l!='use_fp8' else False, + "fp8_base_unet": True if train_clip_l=='use_gradient_dtype' else False, } attention_settings = { "sdpa": {"mem_eff_attn": True, "xformers": False, "spda": True}, diff --git a/train_network.py b/train_network.py index 2292e17..754b19b 100644 --- a/train_network.py +++ b/train_network.py @@ -539,7 +539,7 @@ class NetworkTrainer: accelerator.print("enable fp8 training for U-Net.") unet_weight_dtype = torch.float8_e4m3fn - if not args.fp8_base_unet: + if not args.fp8_base_unet and not args.network_train_unet_only: accelerator.print("enable fp8 training for Text Encoder.") te_weight_dtype = weight_dtype if args.fp8_base_unet else torch.float8_e4m3fn