From 6a91611a2b216e6fa51a1300f04b8136e1ccda87 Mon Sep 17 00:00:00 2001 From: kijai <40791699+kijai@users.noreply.github.com> Date: Sun, 2 Feb 2025 16:29:04 +0200 Subject: [PATCH] Fix network args, update example folder path to support template loader --- .../flux_lora_train_example01.json | 0 .../sdxl_train_example_01.json | 0 networks/lora_flux.py | 3 ++ nodes.py | 33 ++++++++++--------- 4 files changed, 21 insertions(+), 15 deletions(-) rename {examples => example_workflows}/flux_lora_train_example01.json (100%) rename {examples => example_workflows}/sdxl_train_example_01.json (100%) diff --git a/examples/flux_lora_train_example01.json b/example_workflows/flux_lora_train_example01.json similarity index 100% rename from examples/flux_lora_train_example01.json rename to example_workflows/flux_lora_train_example01.json diff --git a/examples/sdxl_train_example_01.json b/example_workflows/sdxl_train_example_01.json similarity index 100% rename from examples/sdxl_train_example_01.json rename to example_workflows/sdxl_train_example_01.json diff --git a/networks/lora_flux.py b/networks/lora_flux.py index 5fb8633..fe1c02a 100644 --- a/networks/lora_flux.py +++ b/networks/lora_flux.py @@ -504,6 +504,9 @@ class LoRANetwork(torch.nn.Module): logger.info(f"train T5XXL as well") #self.only_if_contains = ["lora_unet_single_blocks_20_linear2"] + print(self.only_if_contains) + print(self.only_if_contains) + print(self.only_if_contains) # create module instances def create_modules( diff --git a/nodes.py b/nodes.py index c53e04a..e098bbc 100644 --- a/nodes.py +++ b/nodes.py @@ -499,6 +499,7 @@ class InitFluxLoRATraining: parser = train_network_setup_parser() flux_train_utils.add_flux_train_arguments(parser) + if additional_args is not None: print(f"additional_args: {additional_args}") args, _ = parser.parse_known_args(args=shlex.split(additional_args)) @@ -576,21 +577,6 @@ class InitFluxLoRATraining: if T5_lr != "NaN": config_dict["text_encoder_lr"] = [clip_l_lr, T5_lr] - #network args - additional_network_args = [] - - if "T5" in train_text_encoder: - additional_network_args.append("train_t5xxl=True") - - if block_args: - additional_network_args.append(block_args["include"]) - - # Handle network_args in args Namespace - if hasattr(args, 'network_args') and isinstance(args.network_args, list): - args.network_args.extend(additional_network_args) - else: - setattr(args, 'network_args', additional_network_args) - if gradient_checkpointing == "disabled": config_dict["gradient_checkpointing"] = False elif gradient_checkpointing == "enabled_with_cpu_offloading": @@ -613,6 +599,21 @@ class InitFluxLoRATraining: for key, value in config_dict.items(): setattr(args, key, value) + + #network args + additional_network_args = [] + + if "T5" in train_text_encoder: + additional_network_args.append("train_t5xxl=True") + + if block_args: + additional_network_args.append(block_args["include"]) + + # Handle network_args in args Namespace + if hasattr(args, 'network_args') and isinstance(args.network_args, list): + args.network_args.extend(additional_network_args) + else: + setattr(args, 'network_args', additional_network_args) saved_args_file_path = os.path.join(output_dir, f"{output_name}_args.json") with open(saved_args_file_path, 'w') as f: @@ -702,6 +703,8 @@ class InitFluxTraining: dataset_toml = toml.dumps(json.loads(dataset_config)) parser = train_setup_parser() + flux_train_utils.add_flux_train_arguments(parser) + if additional_args is not None: print(f"additional_args: {additional_args}") args, _ = parser.parse_known_args(args=shlex.split(additional_args))