Compare commits
| Author | SHA1 | Date | |
|---|---|---|---|
|
|
d1b1306ceb | ||
|
|
b0fe9c2d14 |
+20
-18
@@ -67,8 +67,9 @@ class FluxNetworkTrainer(NetworkTrainer):
|
||||
|
||||
logger.info("prepare split model")
|
||||
with init_empty_weights():
|
||||
flux_upper = flux_models.FluxUpper(model.params)
|
||||
flux_lower = flux_models.FluxLower(model.params)
|
||||
flux_upper = flux_models.FluxUpper(model.params, flux_lower)
|
||||
|
||||
sd = model.state_dict()
|
||||
|
||||
# lower (trainable)
|
||||
@@ -234,12 +235,12 @@ class FluxNetworkTrainer(NetworkTrainer):
|
||||
self.target_device = device
|
||||
|
||||
def forward(self, img, img_ids, txt, txt_ids, timesteps, y, guidance=None, txt_attention_mask=None):
|
||||
self.flux_lower.to("cpu")
|
||||
clean_memory_on_device(self.target_device)
|
||||
#self.flux_lower.to("cpu")
|
||||
#clean_memory_on_device(self.target_device)
|
||||
self.flux_upper.to(self.target_device)
|
||||
img, txt, vec, pe = self.flux_upper(img, img_ids, txt, txt_ids, timesteps, y, guidance, txt_attention_mask)
|
||||
self.flux_upper.to("cpu")
|
||||
clean_memory_on_device(self.target_device)
|
||||
#self.flux_upper.to("cpu")
|
||||
#clean_memory_on_device(self.target_device)
|
||||
self.flux_lower.to(self.target_device)
|
||||
return self.flux_lower(img, txt, vec, pe, txt_attention_mask)
|
||||
|
||||
@@ -374,16 +375,16 @@ class FluxNetworkTrainer(NetworkTrainer):
|
||||
)
|
||||
else:
|
||||
# split forward to reduce memory usage
|
||||
assert network.train_blocks == "single", "train_blocks must be single for split mode"
|
||||
#assert network.train_blocks == "single", "train_blocks must be single for split mode"
|
||||
with accelerator.autocast():
|
||||
# move flux lower to cpu, and then move flux upper to gpu
|
||||
unet.to("cpu")
|
||||
clean_memory_on_device(accelerator.device)
|
||||
#unet.to("cpu")
|
||||
#clean_memory_on_device(accelerator.device)
|
||||
self.flux_upper.to(accelerator.device)
|
||||
|
||||
# upper model does not require grad
|
||||
with torch.no_grad():
|
||||
intermediate_img, intermediate_txt, vec, pe = self.flux_upper(
|
||||
model_pred = self.flux_upper(
|
||||
img=packed_noisy_model_input,
|
||||
img_ids=img_ids,
|
||||
txt=t5_out,
|
||||
@@ -392,19 +393,20 @@ class FluxNetworkTrainer(NetworkTrainer):
|
||||
timesteps=timesteps / 1000,
|
||||
guidance=guidance_vec,
|
||||
txt_attention_mask=t5_attn_mask,
|
||||
train_lower=True,
|
||||
)
|
||||
|
||||
model_pred.requires_grad_(True)
|
||||
# move flux upper back to cpu, and then move flux lower to gpu
|
||||
self.flux_upper.to("cpu")
|
||||
clean_memory_on_device(accelerator.device)
|
||||
unet.to(accelerator.device)
|
||||
#self.flux_upper.to("cpu")
|
||||
#clean_memory_on_device(accelerator.device)
|
||||
#unet.to(accelerator.device)
|
||||
|
||||
# lower model requires grad
|
||||
intermediate_img.requires_grad_(True)
|
||||
intermediate_txt.requires_grad_(True)
|
||||
vec.requires_grad_(True)
|
||||
pe.requires_grad_(True)
|
||||
model_pred = unet(img=intermediate_img, txt=intermediate_txt, vec=vec, pe=pe, txt_attention_mask=t5_attn_mask)
|
||||
# intermediate_img.requires_grad_(True)
|
||||
# intermediate_txt.requires_grad_(True)
|
||||
# vec.requires_grad_(True)
|
||||
# pe.requires_grad_(True)
|
||||
#model_pred = unet(img=intermediate_img, txt=intermediate_txt, vec=vec, pe=pe, txt_attention_mask=t5_attn_mask)
|
||||
|
||||
# unpack latents
|
||||
model_pred = flux_utils.unpack_latents(model_pred, packed_latent_height, packed_latent_width)
|
||||
|
||||
+55
-10
@@ -1095,9 +1095,9 @@ class FluxUpper(nn.Module):
|
||||
Transformer model for flow matching on sequences.
|
||||
"""
|
||||
|
||||
def __init__(self, params: FluxParams):
|
||||
def __init__(self, params: FluxParams, lower_model):
|
||||
super().__init__()
|
||||
|
||||
self.lower_model = lower_model
|
||||
self.params = params
|
||||
self.in_channels = params.in_channels
|
||||
self.out_channels = self.in_channels
|
||||
@@ -1127,6 +1127,19 @@ class FluxUpper(nn.Module):
|
||||
]
|
||||
)
|
||||
|
||||
self.excluded_blocks = [7]
|
||||
if self.excluded_blocks is None:
|
||||
self.excluded_blocks = [] # default to no blocks excluded
|
||||
|
||||
self.single_blocks = nn.ModuleList(
|
||||
[
|
||||
SingleStreamBlock(self.hidden_size, self.num_heads, mlp_ratio=params.mlp_ratio)
|
||||
for i in range(params.depth_single_blocks) if i not in self.excluded_blocks
|
||||
]
|
||||
)
|
||||
print("UPPER: Single blocks: ", self.single_blocks)
|
||||
|
||||
self.final_layer = LastLayer(self.hidden_size, 1, self.out_channels)
|
||||
self.gradient_checkpointing = False
|
||||
|
||||
@property
|
||||
@@ -1173,6 +1186,7 @@ class FluxUpper(nn.Module):
|
||||
y: Tensor,
|
||||
guidance: Tensor | None = None,
|
||||
txt_attention_mask: Tensor | None = None,
|
||||
train_lower=False
|
||||
) -> Tensor:
|
||||
if img.ndim != 3 or txt.ndim != 3:
|
||||
raise ValueError("Input img and txt tensors must have 3 dimensions.")
|
||||
@@ -1193,7 +1207,20 @@ class FluxUpper(nn.Module):
|
||||
for block in self.double_blocks:
|
||||
img, txt = block(img=img, txt=txt, vec=vec, pe=pe, txt_attention_mask=txt_attention_mask)
|
||||
|
||||
return img, txt, vec, pe
|
||||
img = torch.cat((txt, img), 1)
|
||||
|
||||
for i, block in enumerate(self.single_blocks):
|
||||
if i in self.excluded_blocks:
|
||||
img = self.lower_model(img, txt, vec=vec, pe=pe, txt_attention_mask=txt_attention_mask, train=train_lower)
|
||||
else:
|
||||
img = block(img, vec=vec, pe=pe, txt_attention_mask=txt_attention_mask)
|
||||
print(img.shape)
|
||||
|
||||
img = img[:, txt.shape[1]:, ...]
|
||||
|
||||
img = self.final_layer(img, vec) # (N, T, patch_size ** 2 * out_channels)
|
||||
|
||||
return img
|
||||
|
||||
|
||||
class FluxLower(nn.Module):
|
||||
@@ -1207,14 +1234,23 @@ class FluxLower(nn.Module):
|
||||
self.num_heads = params.num_heads
|
||||
self.out_channels = params.in_channels
|
||||
|
||||
selected_blocks = [7]
|
||||
|
||||
if selected_blocks is None:
|
||||
selected_blocks = range(params.depth_single_blocks) # default to all blocks
|
||||
|
||||
self.single_blocks = nn.ModuleList(
|
||||
[
|
||||
SingleStreamBlock(self.hidden_size, self.num_heads, mlp_ratio=params.mlp_ratio)
|
||||
for _ in range(params.depth_single_blocks)
|
||||
for i in selected_blocks
|
||||
]
|
||||
)
|
||||
|
||||
for i, block in enumerate(self.single_blocks):
|
||||
print(f"LOWER: Single block {i}: {block.__class__.__name__}")
|
||||
print("LOWER: Single blocks: ", self.single_blocks)
|
||||
|
||||
self.final_layer = LastLayer(self.hidden_size, 1, self.out_channels)
|
||||
#self.final_layer = LastLayer(self.hidden_size, 1, self.out_channels)
|
||||
|
||||
self.gradient_checkpointing = False
|
||||
|
||||
@@ -1249,11 +1285,20 @@ class FluxLower(nn.Module):
|
||||
vec: Tensor | None = None,
|
||||
pe: Tensor | None = None,
|
||||
txt_attention_mask: Tensor | None = None,
|
||||
train: bool = False,
|
||||
) -> Tensor:
|
||||
img = torch.cat((txt, img), 1)
|
||||
for block in self.single_blocks:
|
||||
img = block(img, vec=vec, pe=pe, txt_attention_mask=txt_attention_mask)
|
||||
img = img[:, txt.shape[1] :, ...]
|
||||
|
||||
if train:
|
||||
img.requires_grad_(True)
|
||||
txt.requires_grad_(True)
|
||||
vec.requires_grad_(True)
|
||||
pe.requires_grad_(True)
|
||||
#img = torch.cat((txt, img), 1)
|
||||
print("img.shape to lower: ", img.shape)
|
||||
with torch.enable_grad():
|
||||
for block in self.single_blocks:
|
||||
img = block(img, vec=vec, pe=pe, txt_attention_mask=txt_attention_mask)
|
||||
#img = img[:, txt.shape[1] :, ...]
|
||||
|
||||
img = self.final_layer(img, vec) # (N, T, patch_size ** 2 * out_channels)
|
||||
#img = self.final_layer(img, vec) # (N, T, patch_size ** 2 * out_channels)
|
||||
return img
|
||||
+75
-48
@@ -253,8 +253,8 @@ def create_network(
|
||||
|
||||
# single or double blocks
|
||||
train_blocks = kwargs.get("train_blocks", None) # None (default), "all" (same as None), "single", "double"
|
||||
if train_blocks is not None:
|
||||
assert train_blocks in ["all", "single", "double"], f"invalid train_blocks: {train_blocks}"
|
||||
#if train_blocks is not None:
|
||||
# assert train_blocks in ["all", "single", "double"], f"invalid train_blocks: {train_blocks}"
|
||||
|
||||
# すごく引数が多いな ( ^ω^)・・・
|
||||
network = LoRANetwork(
|
||||
@@ -383,53 +383,77 @@ class LoRANetwork(torch.nn.Module):
|
||||
else (self.LORA_PREFIX_TEXT_ENCODER_CLIP if text_encoder_idx == 0 else self.LORA_PREFIX_TEXT_ENCODER_T5)
|
||||
)
|
||||
|
||||
def process_module(name, module, prefix, modules_dim, modules_alpha, dropout, rank_dropout, module_dropout, target_replace_modules):
|
||||
loras = []
|
||||
skipped = []
|
||||
|
||||
def process_child(child_name, child_module):
|
||||
is_linear = child_module.__class__.__name__ == "Linear"
|
||||
is_conv2d = child_module.__class__.__name__ == "Conv2d"
|
||||
is_conv2d_1x1 = is_conv2d and child_module.kernel_size == (1, 1)
|
||||
|
||||
if is_linear or is_conv2d:
|
||||
lora_name = prefix + "." + name + "." + child_name
|
||||
lora_name = lora_name.replace(".", "_")
|
||||
|
||||
dim = None
|
||||
alpha = None
|
||||
|
||||
if modules_dim is not None:
|
||||
# Module specified
|
||||
if lora_name in modules_dim:
|
||||
dim = modules_dim[lora_name]
|
||||
alpha = modules_alpha[lora_name]
|
||||
else:
|
||||
# Normally, target all
|
||||
if is_linear or is_conv2d_1x1:
|
||||
dim = self.lora_dim
|
||||
alpha = self.alpha
|
||||
elif self.conv_lora_dim is not None:
|
||||
dim = self.conv_lora_dim
|
||||
alpha = self.conv_alpha
|
||||
|
||||
if dim is None or dim == 0:
|
||||
# Output skipped information
|
||||
if is_linear or is_conv2d_1x1 or (self.conv_lora_dim is not None):
|
||||
skipped.append(lora_name)
|
||||
return
|
||||
|
||||
lora = module_class(
|
||||
lora_name,
|
||||
child_module,
|
||||
self.multiplier,
|
||||
dim,
|
||||
alpha,
|
||||
dropout=dropout,
|
||||
rank_dropout=rank_dropout,
|
||||
module_dropout=module_dropout,
|
||||
)
|
||||
loras.append(lora)
|
||||
|
||||
for child_name, child_module in module.named_modules():
|
||||
process_child(child_name, child_module)
|
||||
|
||||
return loras, skipped
|
||||
|
||||
|
||||
loras = []
|
||||
skipped = []
|
||||
|
||||
target_replace_modules = [module.strip() for module in target_replace_modules]
|
||||
|
||||
for name, module in root_module.named_modules():
|
||||
if module.__class__.__name__ in target_replace_modules:
|
||||
for child_name, child_module in module.named_modules():
|
||||
is_linear = child_module.__class__.__name__ == "Linear"
|
||||
is_conv2d = child_module.__class__.__name__ == "Conv2d"
|
||||
is_conv2d_1x1 = is_conv2d and child_module.kernel_size == (1, 1)
|
||||
|
||||
if is_linear or is_conv2d:
|
||||
lora_name = prefix + "." + name + "." + child_name
|
||||
lora_name = lora_name.replace(".", "_")
|
||||
|
||||
dim = None
|
||||
alpha = None
|
||||
|
||||
if modules_dim is not None:
|
||||
# モジュール指定あり
|
||||
if lora_name in modules_dim:
|
||||
dim = modules_dim[lora_name]
|
||||
alpha = modules_alpha[lora_name]
|
||||
else:
|
||||
# 通常、すべて対象とする
|
||||
if is_linear or is_conv2d_1x1:
|
||||
dim = self.lora_dim
|
||||
alpha = self.alpha
|
||||
elif self.conv_lora_dim is not None:
|
||||
dim = self.conv_lora_dim
|
||||
alpha = self.conv_alpha
|
||||
|
||||
if dim is None or dim == 0:
|
||||
# skipした情報を出力
|
||||
if is_linear or is_conv2d_1x1 or (self.conv_lora_dim is not None):
|
||||
skipped.append(lora_name)
|
||||
continue
|
||||
|
||||
lora = module_class(
|
||||
lora_name,
|
||||
child_module,
|
||||
self.multiplier,
|
||||
dim,
|
||||
alpha,
|
||||
dropout=dropout,
|
||||
rank_dropout=rank_dropout,
|
||||
module_dropout=module_dropout,
|
||||
)
|
||||
loras.append(lora)
|
||||
if any("blocks" in part for part in target_replace_modules) and name in target_replace_modules:
|
||||
print(module)
|
||||
module_loras, module_skipped = process_module(name, module, prefix, modules_dim, modules_alpha, dropout, rank_dropout, module_dropout, target_replace_modules)
|
||||
loras.extend(module_loras)
|
||||
skipped.extend(module_skipped)
|
||||
elif module.__class__.__name__ in target_replace_modules:
|
||||
print(module)
|
||||
module_loras, module_skipped = process_module(name, module, prefix, modules_dim, modules_alpha, dropout, rank_dropout, module_dropout, target_replace_modules)
|
||||
loras.extend(module_loras)
|
||||
skipped.extend(module_skipped)
|
||||
|
||||
return loras, skipped
|
||||
|
||||
# create LoRA for text encoder
|
||||
@@ -444,9 +468,12 @@ class LoRANetwork(torch.nn.Module):
|
||||
self.text_encoder_loras.extend(text_encoder_loras)
|
||||
skipped_te += skipped
|
||||
logger.info(f"create LoRA for Text Encoder: {len(self.text_encoder_loras)} modules.")
|
||||
|
||||
print("TRAIN BLOCKS:", self.train_blocks)
|
||||
# create LoRA for U-Net
|
||||
if self.train_blocks == "all":
|
||||
if any("blocks" in part for part in self.train_blocks.split(',')):
|
||||
target_replace_modules = self.train_blocks.split(',')
|
||||
print("TARGET_REPLACE_MODULES:", target_replace_modules)
|
||||
elif self.train_blocks == "all":
|
||||
target_replace_modules = LoRANetwork.FLUX_TARGET_REPLACE_MODULE_DOUBLE + LoRANetwork.FLUX_TARGET_REPLACE_MODULE_SINGLE
|
||||
elif self.train_blocks == "single":
|
||||
target_replace_modules = LoRANetwork.FLUX_TARGET_REPLACE_MODULE_SINGLE
|
||||
|
||||
@@ -270,7 +270,23 @@ class OptimizerConfigProdigy:
|
||||
kwargs["optimizer_args"] = node_args + extra_args
|
||||
kwargs["min_snr_gamma"] = min_snr_gamma if min_snr_gamma != 0.0 else None
|
||||
|
||||
return (kwargs,)
|
||||
return (kwargs,)
|
||||
|
||||
class FluxLoRATrainBlocks:
|
||||
@classmethod
|
||||
def INPUT_TYPES(s):
|
||||
return {"required": {
|
||||
"train_blocks": ("STRING",{"default": "single_blocks.7", "multiline": True, "tooltip": "specify individual blocks to include in the training, for example 'single_blocks.7'"}),
|
||||
},
|
||||
}
|
||||
|
||||
RETURN_TYPES = ("BLOCKS",)
|
||||
RETURN_NAMES = ("train_blocks",)
|
||||
FUNCTION = "create_config"
|
||||
CATEGORY = "FluxTrainer"
|
||||
|
||||
def create_config(self, train_blocks):
|
||||
return (train_blocks,)
|
||||
|
||||
class InitFluxLoRATraining:
|
||||
@classmethod
|
||||
@@ -314,6 +330,8 @@ class InitFluxLoRATraining:
|
||||
"resume_args": ("ARGS", {"default": "", "tooltip": "resume args to pass to the training command"}),
|
||||
"train_clip_l": (['disabled', 'use_gradient_dtype', 'use_fp8'], {"default": 'disabled', "tooltip": "also train the clip_l text encoder using specified dtype"}),
|
||||
"text_encoder_lr": ("FLOAT", {"default": 0, "min": 0.0, "max": 10.0, "step": 0.00001, "tooltip": "text encoder learning rate"}),
|
||||
"train_blocks": ("BLOCKS", ),
|
||||
"gradient_checkpointing": ("BOOLEAN", {"default": True, "tooltip": "use gradient checkpointing"}),
|
||||
},
|
||||
}
|
||||
|
||||
@@ -323,7 +341,7 @@ class InitFluxLoRATraining:
|
||||
CATEGORY = "FluxTrainer"
|
||||
|
||||
def init_training(self, flux_models, dataset, optimizer_settings, sample_prompts, output_name, attention_mode,
|
||||
gradient_dtype, save_dtype, split_mode, additional_args=None, resume_args=None, train_clip_l='disabled', **kwargs,):
|
||||
gradient_dtype, save_dtype, split_mode, additional_args=None, resume_args=None, train_clip_l='disabled', train_blocks=None, gradient_checkpointing=True, **kwargs,):
|
||||
mm.soft_empty_cache()
|
||||
|
||||
output_dir = os.path.abspath(kwargs.get("output_dir"))
|
||||
@@ -387,7 +405,7 @@ class InitFluxLoRATraining:
|
||||
"persistent_data_loader_workers": False,
|
||||
"max_data_loader_n_workers": 0,
|
||||
"seed": 42,
|
||||
"gradient_checkpointing": True,
|
||||
#"gradient_checkpointing": True,
|
||||
"network_module": ".networks.lora_flux",
|
||||
"dataset_config": dataset_toml,
|
||||
"output_name": f"{output_name}_rank{kwargs.get('network_dim')}_{save_dtype}",
|
||||
@@ -398,6 +416,8 @@ class InitFluxLoRATraining:
|
||||
"network_train_unet_only": True if train_clip_l == 'disabled' else False,
|
||||
"fp8_base_unet": True if train_clip_l=='use_gradient_dtype' else False,
|
||||
}
|
||||
if gradient_checkpointing:
|
||||
config_dict["gradient_checkpointing"] = True
|
||||
attention_settings = {
|
||||
"sdpa": {"mem_eff_attn": True, "xformers": False, "spda": True},
|
||||
"xformers": {"mem_eff_attn": True, "xformers": True, "spda": False}
|
||||
@@ -410,11 +430,20 @@ class InitFluxLoRATraining:
|
||||
}
|
||||
config_dict.update(gradient_dtype_settings.get(gradient_dtype, {}))
|
||||
|
||||
split_mode_settings = {
|
||||
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, {}))
|
||||
if train_blocks is None:
|
||||
split_mode_settings = {
|
||||
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, {}))
|
||||
else:
|
||||
config_dict["split_mode"] = True
|
||||
if "network_args" not in config_dict:
|
||||
config_dict["network_args"] = []
|
||||
config_dict["network_args"].append(f"train_blocks={train_blocks}")
|
||||
|
||||
print("NETWORK ARGS: ", config_dict["network_args"])
|
||||
|
||||
|
||||
config_dict.update(kwargs)
|
||||
config_dict.update(optimizer_settings)
|
||||
@@ -1425,7 +1454,8 @@ NODE_CLASS_MAPPINGS = {
|
||||
"FluxTrainSaveModel": FluxTrainSaveModel,
|
||||
"ExtractFluxLoRA": ExtractFluxLoRA,
|
||||
"OptimizerConfigProdigy": OptimizerConfigProdigy,
|
||||
"FluxTrainResume": FluxTrainResume
|
||||
"FluxTrainResume": FluxTrainResume,
|
||||
"FluxLoRATrainBlocks": FluxLoRATrainBlocks
|
||||
}
|
||||
NODE_DISPLAY_NAME_MAPPINGS = {
|
||||
"InitFluxLoRATraining": "Init Flux LoRA Training",
|
||||
@@ -1446,5 +1476,6 @@ NODE_DISPLAY_NAME_MAPPINGS = {
|
||||
"FluxTrainSaveModel": "Flux Train Save Model",
|
||||
"ExtractFluxLoRA": "Extract Flux LoRA",
|
||||
"OptimizerConfigProdigy": "Optimizer Config Prodigy",
|
||||
"FluxTrainResume": "Flux Train Resume"
|
||||
"FluxTrainResume": "Flux Train Resume",
|
||||
"FluxLoRATrainBlocks": "FluxLoRATrainBlocks"
|
||||
}
|
||||
|
||||
Reference in New Issue
Block a user