From 66a5fa96c9520be36c1b6e4304307f17f2eb563c Mon Sep 17 00:00:00 2001 From: kijai <40791699+kijai@users.noreply.github.com> Date: Sun, 8 Sep 2024 17:48:02 +0300 Subject: [PATCH] Update nodes.py --- nodes.py | 8 +++++--- 1 file changed, 5 insertions(+), 3 deletions(-) diff --git a/nodes.py b/nodes.py index e57ff56..bddc7c1 100644 --- a/nodes.py +++ b/nodes.py @@ -277,7 +277,7 @@ class OptimizerConfigProdigy: def create_config(self, weight_decay, decouple, min_snr_gamma, use_bias_correction, extra_optimizer_args, **kwargs): kwargs["optimizer_type"] = "prodigy" - extra_args = [arg.strip() for arg in extra_optimizer_args.strip().split(',') if arg.strip()] + extra_args = [arg.strip() for arg in extra_optimizer_args.strip().split('|') if arg.strip()] node_args = [ f"weight_decay={weight_decay}", f"decouple={decouple}", @@ -1068,6 +1068,7 @@ class FluxKohyaInferenceSampler: "guidance_scale": ("FLOAT", {"default": 3.5, "min": 1.0, "max": 32.0, "step": 0.05, "tooltip": "guidance scale"}), "seed": ("INT", {"default": 42,"min": 0, "max": 0xffffffffffffffff, "step": 1}), "use_fp8": ("BOOLEAN", {"default": True, "tooltip": "use fp8 weights"}), + "apply_t5_attn_mask": ("BOOLEAN", {"default": True, "tooltip": "use t5 attention mask"}), "prompt": ("STRING", {"multiline": True, "default": "illustration of a kitten", "tooltip": "prompt"}), }, @@ -1078,7 +1079,7 @@ class FluxKohyaInferenceSampler: FUNCTION = "sample" CATEGORY = "FluxTrainer" - def sample(self, flux_models, lora_name, steps, width, height, guidance_scale, seed, prompt, use_fp8, lora_method): + def sample(self, flux_models, lora_name, steps, width, height, guidance_scale, seed, prompt, use_fp8, lora_method, apply_t5_attn_mask): from .library import flux_utils as flux_utils from .library import strategy_flux as strategy_flux @@ -1091,7 +1092,7 @@ class FluxKohyaInferenceSampler: import gc device = "cuda" - apply_t5_attn_mask = True + if use_fp8: accelerator = accelerate.Accelerator(mixed_precision="bf16") @@ -1284,6 +1285,7 @@ class FluxKohyaInferenceSampler: flux_dtype: torch.dtype, ): timesteps = get_schedule(num_steps, img.shape[1], shift=not is_schnell) + print(timesteps) # denoise initial noise if accelerator: