update
This commit is contained in:
@@ -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:
|
||||
|
||||
@@ -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"]
|
||||
|
||||
Reference in New Issue
Block a user