still doesn't work

This commit is contained in:
kijai
2024-09-02 22:30:18 +03:00
parent b0fe9c2d14
commit d1b1306ceb
3 changed files with 81 additions and 31 deletions
+20 -18
View File
@@ -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
View File
@@ -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
+6 -3
View File
@@ -331,6 +331,7 @@ class InitFluxLoRATraining:
"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"}),
},
}
@@ -340,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', train_blocks=None, **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"))
@@ -404,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}",
@@ -415,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}
@@ -434,7 +437,7 @@ class InitFluxLoRATraining:
}
config_dict.update(split_mode_settings.get(split_mode, {}))
else:
config_dict["split_mode"] = False
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}")