diff --git a/networks/lora_flux.py b/networks/lora_flux.py index c06a291..0878346 100644 --- a/networks/lora_flux.py +++ b/networks/lora_flux.py @@ -324,6 +324,10 @@ def create_network( if train_blocks is not None: assert train_blocks in ["all", "single", "double"], f"invalid train_blocks: {train_blocks}" + only_if_contains = kwargs.get("only_if_contains", None) + if only_if_contains is not None: + only_if_contains = [word.strip() for word in only_if_contains.split(',')] + # split qkv split_qkv = kwargs.get("split_qkv", False) if split_qkv is not None: @@ -350,6 +354,7 @@ def create_network( split_qkv=split_qkv, train_t5xxl=train_t5xxl, varbose=True, + only_if_contains=only_if_contains ) loraplus_lr_ratio = kwargs.get("loraplus_lr_ratio", None) @@ -458,6 +463,7 @@ class LoRANetwork(torch.nn.Module): split_qkv: bool = False, train_t5xxl: bool = False, varbose: Optional[bool] = False, + only_if_contains: Optional[List[str]] = None, ) -> None: super().__init__() self.multiplier = multiplier @@ -477,6 +483,8 @@ class LoRANetwork(torch.nn.Module): self.loraplus_unet_lr_ratio = None self.loraplus_text_encoder_lr_ratio = None + self.only_if_contains = only_if_contains + if modules_dim is not None: logger.info(f"create LoRA network from weights") else: @@ -495,6 +503,8 @@ class LoRANetwork(torch.nn.Module): if train_t5xxl: logger.info(f"train T5XXL as well") + #self.only_if_contains = ["lora_unet_single_blocks_20_linear2"] + # create module instances def create_modules( is_flux: bool, text_encoder_idx: Optional[int], root_module: torch.nn.Module, target_replace_modules: List[str] @@ -517,6 +527,10 @@ class LoRANetwork(torch.nn.Module): if is_linear or is_conv2d: lora_name = prefix + "." + name + "." + child_name lora_name = lora_name.replace(".", "_") + #lora_unet_single_blocks_20_linear2 + + if "unet" in lora_name and (self.only_if_contains is not None and not any(word in lora_name for word in self.only_if_contains)): + continue dim = None alpha = None @@ -590,6 +604,7 @@ class LoRANetwork(torch.nn.Module): self.unet_loras: List[Union[LoRAModule, LoRAInfModule]] self.unet_loras, skipped_un = create_modules(True, None, unet, target_replace_modules) logger.info(f"create LoRA for FLUX {self.train_blocks} blocks: {len(self.unet_loras)} modules.") + print(self.unet_loras) skipped = skipped_te + skipped_un if varbose and len(skipped) > 0: diff --git a/nodes.py b/nodes.py index 99c78d2..abee24c 100644 --- a/nodes.py +++ b/nodes.py @@ -9,6 +9,8 @@ import toml import json import time import shutil +import shlex + from pathlib import Path script_directory = os.path.dirname(os.path.abspath(__file__)) @@ -340,7 +342,9 @@ class InitFluxLoRATraining: parser = train_network_setup_parser() if additional_args is not None: - args, _ = parser.parse_known_args(args=[additional_args]) + print(f"additional_args: {additional_args}") + args, _ = parser.parse_known_args(args=shlex.split(additional_args)) + print(args) else: args, _ = parser.parse_known_args() #print(args) @@ -415,7 +419,12 @@ class InitFluxLoRATraining: True: {"split_mode": True, "network_args": ["train_blocks=single"]}, False: {"split_mode": False, "network_args": ["train_blocks=all"]} } - config_dict.update(split_mode_settings.get(split_mode, {})) + + selected_split_mode_settings = split_mode_settings.get(split_mode, {}) + if 'network_args' in config_dict and isinstance(config_dict['network_args'], list): + config_dict['network_args'].extend(selected_split_mode_settings.pop('network_args', [])) + else: + config_dict.update(selected_split_mode_settings) if "T5" in train_text_encoder: additional_network_args = ["train_t5xxl=True"]