2 Commits
Author SHA1 Message Date
kijai d1b1306ceb still doesn't work 2024-09-02 22:30:18 +03:00
kijai b0fe9c2d14 block select testing 2024-09-02 03:13:36 +03:00
4 changed files with 191 additions and 86 deletions
+20 -18
View File
@@ -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
View File
@@ -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
View File
@@ -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
+41 -10
View File
@@ -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"
} }