Fix network args, update example folder path to support template loader

This commit is contained in:
kijai
2025-02-02 16:29:04 +02:00
parent f6af45a169
commit 6a91611a2b
4 changed files with 21 additions and 15 deletions
+3
View File
@@ -504,6 +504,9 @@ class LoRANetwork(torch.nn.Module):
logger.info(f"train T5XXL as well") logger.info(f"train T5XXL as well")
#self.only_if_contains = ["lora_unet_single_blocks_20_linear2"] #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 # create module instances
def create_modules( def create_modules(
+18 -15
View File
@@ -499,6 +499,7 @@ class InitFluxLoRATraining:
parser = train_network_setup_parser() parser = train_network_setup_parser()
flux_train_utils.add_flux_train_arguments(parser) flux_train_utils.add_flux_train_arguments(parser)
if additional_args is not None: if additional_args is not None:
print(f"additional_args: {additional_args}") print(f"additional_args: {additional_args}")
args, _ = parser.parse_known_args(args=shlex.split(additional_args)) args, _ = parser.parse_known_args(args=shlex.split(additional_args))
@@ -576,21 +577,6 @@ class InitFluxLoRATraining:
if T5_lr != "NaN": if T5_lr != "NaN":
config_dict["text_encoder_lr"] = [clip_l_lr, T5_lr] 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": if gradient_checkpointing == "disabled":
config_dict["gradient_checkpointing"] = False config_dict["gradient_checkpointing"] = False
elif gradient_checkpointing == "enabled_with_cpu_offloading": elif gradient_checkpointing == "enabled_with_cpu_offloading":
@@ -613,6 +599,21 @@ class InitFluxLoRATraining:
for key, value in config_dict.items(): for key, value in config_dict.items():
setattr(args, key, value) 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") saved_args_file_path = os.path.join(output_dir, f"{output_name}_args.json")
with open(saved_args_file_path, 'w') as f: with open(saved_args_file_path, 'w') as f:
@@ -702,6 +703,8 @@ class InitFluxTraining:
dataset_toml = toml.dumps(json.loads(dataset_config)) dataset_toml = toml.dumps(json.loads(dataset_config))
parser = train_setup_parser() parser = train_setup_parser()
flux_train_utils.add_flux_train_arguments(parser)
if additional_args is not None: if additional_args is not None:
print(f"additional_args: {additional_args}") print(f"additional_args: {additional_args}")
args, _ = parser.parse_known_args(args=shlex.split(additional_args)) args, _ = parser.parse_known_args(args=shlex.split(additional_args))