Compare commits
| Author | SHA1 | Date | |
|---|---|---|---|
|
|
d1b1306ceb | ||
|
|
b0fe9c2d14 |
+20
-18
@@ -67,8 +67,9 @@ class FluxNetworkTrainer(NetworkTrainer):
|
|||||||
|
|
||||||
logger.info("prepare split model")
|
logger.info("prepare split model")
|
||||||
with init_empty_weights():
|
with init_empty_weights():
|
||||||
flux_upper = flux_models.FluxUpper(model.params)
|
|
||||||
flux_lower = flux_models.FluxLower(model.params)
|
flux_lower = flux_models.FluxLower(model.params)
|
||||||
|
flux_upper = flux_models.FluxUpper(model.params, flux_lower)
|
||||||
|
|
||||||
sd = model.state_dict()
|
sd = model.state_dict()
|
||||||
|
|
||||||
# lower (trainable)
|
# lower (trainable)
|
||||||
@@ -234,12 +235,12 @@ class FluxNetworkTrainer(NetworkTrainer):
|
|||||||
self.target_device = device
|
self.target_device = device
|
||||||
|
|
||||||
def forward(self, img, img_ids, txt, txt_ids, timesteps, y, guidance=None, txt_attention_mask=None):
|
def forward(self, img, img_ids, txt, txt_ids, timesteps, y, guidance=None, txt_attention_mask=None):
|
||||||
self.flux_lower.to("cpu")
|
#self.flux_lower.to("cpu")
|
||||||
clean_memory_on_device(self.target_device)
|
#clean_memory_on_device(self.target_device)
|
||||||
self.flux_upper.to(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)
|
img, txt, vec, pe = self.flux_upper(img, img_ids, txt, txt_ids, timesteps, y, guidance, txt_attention_mask)
|
||||||
self.flux_upper.to("cpu")
|
#self.flux_upper.to("cpu")
|
||||||
clean_memory_on_device(self.target_device)
|
#clean_memory_on_device(self.target_device)
|
||||||
self.flux_lower.to(self.target_device)
|
self.flux_lower.to(self.target_device)
|
||||||
return self.flux_lower(img, txt, vec, pe, txt_attention_mask)
|
return self.flux_lower(img, txt, vec, pe, txt_attention_mask)
|
||||||
|
|
||||||
@@ -374,16 +375,16 @@ class FluxNetworkTrainer(NetworkTrainer):
|
|||||||
)
|
)
|
||||||
else:
|
else:
|
||||||
# split forward to reduce memory usage
|
# 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():
|
with accelerator.autocast():
|
||||||
# move flux lower to cpu, and then move flux upper to gpu
|
# move flux lower to cpu, and then move flux upper to gpu
|
||||||
unet.to("cpu")
|
#unet.to("cpu")
|
||||||
clean_memory_on_device(accelerator.device)
|
#clean_memory_on_device(accelerator.device)
|
||||||
self.flux_upper.to(accelerator.device)
|
self.flux_upper.to(accelerator.device)
|
||||||
|
|
||||||
# upper model does not require grad
|
# upper model does not require grad
|
||||||
with torch.no_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=packed_noisy_model_input,
|
||||||
img_ids=img_ids,
|
img_ids=img_ids,
|
||||||
txt=t5_out,
|
txt=t5_out,
|
||||||
@@ -392,19 +393,20 @@ class FluxNetworkTrainer(NetworkTrainer):
|
|||||||
timesteps=timesteps / 1000,
|
timesteps=timesteps / 1000,
|
||||||
guidance=guidance_vec,
|
guidance=guidance_vec,
|
||||||
txt_attention_mask=t5_attn_mask,
|
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
|
# move flux upper back to cpu, and then move flux lower to gpu
|
||||||
self.flux_upper.to("cpu")
|
#self.flux_upper.to("cpu")
|
||||||
clean_memory_on_device(accelerator.device)
|
#clean_memory_on_device(accelerator.device)
|
||||||
unet.to(accelerator.device)
|
#unet.to(accelerator.device)
|
||||||
|
|
||||||
# lower model requires grad
|
# lower model requires grad
|
||||||
intermediate_img.requires_grad_(True)
|
# intermediate_img.requires_grad_(True)
|
||||||
intermediate_txt.requires_grad_(True)
|
# intermediate_txt.requires_grad_(True)
|
||||||
vec.requires_grad_(True)
|
# vec.requires_grad_(True)
|
||||||
pe.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)
|
#model_pred = unet(img=intermediate_img, txt=intermediate_txt, vec=vec, pe=pe, txt_attention_mask=t5_attn_mask)
|
||||||
|
|
||||||
# unpack latents
|
# unpack latents
|
||||||
model_pred = flux_utils.unpack_latents(model_pred, packed_latent_height, packed_latent_width)
|
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.
|
Transformer model for flow matching on sequences.
|
||||||
"""
|
"""
|
||||||
|
|
||||||
def __init__(self, params: FluxParams):
|
def __init__(self, params: FluxParams, lower_model):
|
||||||
super().__init__()
|
super().__init__()
|
||||||
|
self.lower_model = lower_model
|
||||||
self.params = params
|
self.params = params
|
||||||
self.in_channels = params.in_channels
|
self.in_channels = params.in_channels
|
||||||
self.out_channels = self.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
|
self.gradient_checkpointing = False
|
||||||
|
|
||||||
@property
|
@property
|
||||||
@@ -1173,6 +1186,7 @@ class FluxUpper(nn.Module):
|
|||||||
y: Tensor,
|
y: Tensor,
|
||||||
guidance: Tensor | None = None,
|
guidance: Tensor | None = None,
|
||||||
txt_attention_mask: Tensor | None = None,
|
txt_attention_mask: Tensor | None = None,
|
||||||
|
train_lower=False
|
||||||
) -> Tensor:
|
) -> Tensor:
|
||||||
if img.ndim != 3 or txt.ndim != 3:
|
if img.ndim != 3 or txt.ndim != 3:
|
||||||
raise ValueError("Input img and txt tensors must have 3 dimensions.")
|
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:
|
for block in self.double_blocks:
|
||||||
img, txt = block(img=img, txt=txt, vec=vec, pe=pe, txt_attention_mask=txt_attention_mask)
|
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):
|
class FluxLower(nn.Module):
|
||||||
@@ -1207,14 +1234,23 @@ class FluxLower(nn.Module):
|
|||||||
self.num_heads = params.num_heads
|
self.num_heads = params.num_heads
|
||||||
self.out_channels = params.in_channels
|
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(
|
self.single_blocks = nn.ModuleList(
|
||||||
[
|
[
|
||||||
SingleStreamBlock(self.hidden_size, self.num_heads, mlp_ratio=params.mlp_ratio)
|
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
|
self.gradient_checkpointing = False
|
||||||
|
|
||||||
@@ -1249,11 +1285,20 @@ class FluxLower(nn.Module):
|
|||||||
vec: Tensor | None = None,
|
vec: Tensor | None = None,
|
||||||
pe: Tensor | None = None,
|
pe: Tensor | None = None,
|
||||||
txt_attention_mask: Tensor | None = None,
|
txt_attention_mask: Tensor | None = None,
|
||||||
|
train: bool = False,
|
||||||
) -> Tensor:
|
) -> Tensor:
|
||||||
img = torch.cat((txt, img), 1)
|
|
||||||
for block in self.single_blocks:
|
if train:
|
||||||
img = block(img, vec=vec, pe=pe, txt_attention_mask=txt_attention_mask)
|
img.requires_grad_(True)
|
||||||
img = img[:, txt.shape[1] :, ...]
|
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
|
return img
|
||||||
+75
-48
@@ -253,8 +253,8 @@ def create_network(
|
|||||||
|
|
||||||
# single or double blocks
|
# single or double blocks
|
||||||
train_blocks = kwargs.get("train_blocks", None) # None (default), "all" (same as None), "single", "double"
|
train_blocks = kwargs.get("train_blocks", None) # None (default), "all" (same as None), "single", "double"
|
||||||
if train_blocks is not None:
|
#if train_blocks is not None:
|
||||||
assert train_blocks in ["all", "single", "double"], f"invalid train_blocks: {train_blocks}"
|
# assert train_blocks in ["all", "single", "double"], f"invalid train_blocks: {train_blocks}"
|
||||||
|
|
||||||
# すごく引数が多いな ( ^ω^)・・・
|
# すごく引数が多いな ( ^ω^)・・・
|
||||||
network = LoRANetwork(
|
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)
|
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 = []
|
loras = []
|
||||||
skipped = []
|
skipped = []
|
||||||
|
|
||||||
|
target_replace_modules = [module.strip() for module in target_replace_modules]
|
||||||
|
|
||||||
for name, module in root_module.named_modules():
|
for name, module in root_module.named_modules():
|
||||||
if module.__class__.__name__ in target_replace_modules:
|
if any("blocks" in part for part in target_replace_modules) and name in target_replace_modules:
|
||||||
for child_name, child_module in module.named_modules():
|
print(module)
|
||||||
is_linear = child_module.__class__.__name__ == "Linear"
|
module_loras, module_skipped = process_module(name, module, prefix, modules_dim, modules_alpha, dropout, rank_dropout, module_dropout, target_replace_modules)
|
||||||
is_conv2d = child_module.__class__.__name__ == "Conv2d"
|
loras.extend(module_loras)
|
||||||
is_conv2d_1x1 = is_conv2d and child_module.kernel_size == (1, 1)
|
skipped.extend(module_skipped)
|
||||||
|
elif module.__class__.__name__ in target_replace_modules:
|
||||||
if is_linear or is_conv2d:
|
print(module)
|
||||||
lora_name = prefix + "." + name + "." + child_name
|
module_loras, module_skipped = process_module(name, module, prefix, modules_dim, modules_alpha, dropout, rank_dropout, module_dropout, target_replace_modules)
|
||||||
lora_name = lora_name.replace(".", "_")
|
loras.extend(module_loras)
|
||||||
|
skipped.extend(module_skipped)
|
||||||
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)
|
|
||||||
return loras, skipped
|
return loras, skipped
|
||||||
|
|
||||||
# create LoRA for text encoder
|
# create LoRA for text encoder
|
||||||
@@ -444,9 +468,12 @@ class LoRANetwork(torch.nn.Module):
|
|||||||
self.text_encoder_loras.extend(text_encoder_loras)
|
self.text_encoder_loras.extend(text_encoder_loras)
|
||||||
skipped_te += skipped
|
skipped_te += skipped
|
||||||
logger.info(f"create LoRA for Text Encoder: {len(self.text_encoder_loras)} modules.")
|
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
|
# 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
|
target_replace_modules = LoRANetwork.FLUX_TARGET_REPLACE_MODULE_DOUBLE + LoRANetwork.FLUX_TARGET_REPLACE_MODULE_SINGLE
|
||||||
elif self.train_blocks == "single":
|
elif self.train_blocks == "single":
|
||||||
target_replace_modules = LoRANetwork.FLUX_TARGET_REPLACE_MODULE_SINGLE
|
target_replace_modules = LoRANetwork.FLUX_TARGET_REPLACE_MODULE_SINGLE
|
||||||
|
|||||||
@@ -270,7 +270,23 @@ class OptimizerConfigProdigy:
|
|||||||
kwargs["optimizer_args"] = node_args + extra_args
|
kwargs["optimizer_args"] = node_args + extra_args
|
||||||
kwargs["min_snr_gamma"] = min_snr_gamma if min_snr_gamma != 0.0 else None
|
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:
|
class InitFluxLoRATraining:
|
||||||
@classmethod
|
@classmethod
|
||||||
@@ -314,6 +330,8 @@ class InitFluxLoRATraining:
|
|||||||
"resume_args": ("ARGS", {"default": "", "tooltip": "resume args to pass to the training command"}),
|
"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"}),
|
"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"}),
|
"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"
|
CATEGORY = "FluxTrainer"
|
||||||
|
|
||||||
def init_training(self, flux_models, dataset, optimizer_settings, sample_prompts, output_name, attention_mode,
|
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()
|
mm.soft_empty_cache()
|
||||||
|
|
||||||
output_dir = os.path.abspath(kwargs.get("output_dir"))
|
output_dir = os.path.abspath(kwargs.get("output_dir"))
|
||||||
@@ -387,7 +405,7 @@ class InitFluxLoRATraining:
|
|||||||
"persistent_data_loader_workers": False,
|
"persistent_data_loader_workers": False,
|
||||||
"max_data_loader_n_workers": 0,
|
"max_data_loader_n_workers": 0,
|
||||||
"seed": 42,
|
"seed": 42,
|
||||||
"gradient_checkpointing": True,
|
#"gradient_checkpointing": True,
|
||||||
"network_module": ".networks.lora_flux",
|
"network_module": ".networks.lora_flux",
|
||||||
"dataset_config": dataset_toml,
|
"dataset_config": dataset_toml,
|
||||||
"output_name": f"{output_name}_rank{kwargs.get('network_dim')}_{save_dtype}",
|
"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,
|
"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,
|
"fp8_base_unet": True if train_clip_l=='use_gradient_dtype' else False,
|
||||||
}
|
}
|
||||||
|
if gradient_checkpointing:
|
||||||
|
config_dict["gradient_checkpointing"] = True
|
||||||
attention_settings = {
|
attention_settings = {
|
||||||
"sdpa": {"mem_eff_attn": True, "xformers": False, "spda": True},
|
"sdpa": {"mem_eff_attn": True, "xformers": False, "spda": True},
|
||||||
"xformers": {"mem_eff_attn": True, "xformers": True, "spda": False}
|
"xformers": {"mem_eff_attn": True, "xformers": True, "spda": False}
|
||||||
@@ -410,11 +430,20 @@ class InitFluxLoRATraining:
|
|||||||
}
|
}
|
||||||
config_dict.update(gradient_dtype_settings.get(gradient_dtype, {}))
|
config_dict.update(gradient_dtype_settings.get(gradient_dtype, {}))
|
||||||
|
|
||||||
split_mode_settings = {
|
if train_blocks is None:
|
||||||
True: {"split_mode": True, "network_args": ["train_blocks=single"]},
|
split_mode_settings = {
|
||||||
False: {"split_mode": False, "network_args": ["train_blocks=all"]}
|
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, {}))
|
}
|
||||||
|
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(kwargs)
|
||||||
config_dict.update(optimizer_settings)
|
config_dict.update(optimizer_settings)
|
||||||
@@ -1425,7 +1454,8 @@ NODE_CLASS_MAPPINGS = {
|
|||||||
"FluxTrainSaveModel": FluxTrainSaveModel,
|
"FluxTrainSaveModel": FluxTrainSaveModel,
|
||||||
"ExtractFluxLoRA": ExtractFluxLoRA,
|
"ExtractFluxLoRA": ExtractFluxLoRA,
|
||||||
"OptimizerConfigProdigy": OptimizerConfigProdigy,
|
"OptimizerConfigProdigy": OptimizerConfigProdigy,
|
||||||
"FluxTrainResume": FluxTrainResume
|
"FluxTrainResume": FluxTrainResume,
|
||||||
|
"FluxLoRATrainBlocks": FluxLoRATrainBlocks
|
||||||
}
|
}
|
||||||
NODE_DISPLAY_NAME_MAPPINGS = {
|
NODE_DISPLAY_NAME_MAPPINGS = {
|
||||||
"InitFluxLoRATraining": "Init Flux LoRA Training",
|
"InitFluxLoRATraining": "Init Flux LoRA Training",
|
||||||
@@ -1446,5 +1476,6 @@ NODE_DISPLAY_NAME_MAPPINGS = {
|
|||||||
"FluxTrainSaveModel": "Flux Train Save Model",
|
"FluxTrainSaveModel": "Flux Train Save Model",
|
||||||
"ExtractFluxLoRA": "Extract Flux LoRA",
|
"ExtractFluxLoRA": "Extract Flux LoRA",
|
||||||
"OptimizerConfigProdigy": "Optimizer Config Prodigy",
|
"OptimizerConfigProdigy": "Optimizer Config Prodigy",
|
||||||
"FluxTrainResume": "Flux Train Resume"
|
"FluxTrainResume": "Flux Train Resume",
|
||||||
|
"FluxLoRATrainBlocks": "FluxLoRATrainBlocks"
|
||||||
}
|
}
|
||||||
|
|||||||
Reference in New Issue
Block a user