still doesn't work
This commit is contained in:
+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
|
||||
@@ -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}")
|
||||
|
||||
Reference in New Issue
Block a user