From d1b1306cebe9fca00628c66200f0c5dac10aaa9c Mon Sep 17 00:00:00 2001 From: kijai <40791699+kijai@users.noreply.github.com> Date: Mon, 2 Sep 2024 22:30:18 +0300 Subject: [PATCH] still doesn't work --- flux_train_network_comfy.py | 38 ++++++++++++---------- library/flux_models.py | 65 +++++++++++++++++++++++++++++++------ nodes.py | 9 +++-- 3 files changed, 81 insertions(+), 31 deletions(-) diff --git a/flux_train_network_comfy.py b/flux_train_network_comfy.py index 04366d9..f45a041 100644 --- a/flux_train_network_comfy.py +++ b/flux_train_network_comfy.py @@ -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) diff --git a/library/flux_models.py b/library/flux_models.py index 46c1381..127a66e 100644 --- a/library/flux_models.py +++ b/library/flux_models.py @@ -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 \ No newline at end of file diff --git a/nodes.py b/nodes.py index 91679e7..c6b132d 100644 --- a/nodes.py +++ b/nodes.py @@ -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}")