diff --git a/library/flux_models.py b/library/flux_models.py index 3c7766b..ef90172 100644 --- a/library/flux_models.py +++ b/library/flux_models.py @@ -863,7 +863,8 @@ class Flux(nn.Module): self.time_in.enable_gradient_checkpointing() self.vector_in.enable_gradient_checkpointing() - self.guidance_in.enable_gradient_checkpointing() + if self.guidance_in.__class__ != nn.Identity: + self.guidance_in.enable_gradient_checkpointing() for block in self.double_blocks + self.single_blocks: block.enable_gradient_checkpointing() @@ -875,7 +876,8 @@ class Flux(nn.Module): self.time_in.disable_gradient_checkpointing() self.vector_in.disable_gradient_checkpointing() - self.guidance_in.disable_gradient_checkpointing() + if self.guidance_in.__class__ != nn.Identity: + self.guidance_in.enable_gradient_checkpointing() for block in self.double_blocks + self.single_blocks: block.disable_gradient_checkpointing() @@ -972,7 +974,8 @@ class FluxUpper(nn.Module): self.time_in.enable_gradient_checkpointing() self.vector_in.enable_gradient_checkpointing() - self.guidance_in.enable_gradient_checkpointing() + if self.guidance_in.__class__ != nn.Identity: + self.guidance_in.enable_gradient_checkpointing() for block in self.double_blocks: block.enable_gradient_checkpointing() @@ -984,7 +987,8 @@ class FluxUpper(nn.Module): self.time_in.disable_gradient_checkpointing() self.vector_in.disable_gradient_checkpointing() - self.guidance_in.disable_gradient_checkpointing() + if self.guidance_in.__class__ != nn.Identity: + self.guidance_in.enable_gradient_checkpointing() for block in self.double_blocks: block.disable_gradient_checkpointing() diff --git a/nodes.py b/nodes.py index dfc37e2..b3250a4 100644 --- a/nodes.py +++ b/nodes.py @@ -28,6 +28,7 @@ logger = logging.getLogger(__name__) class FluxNetworkTrainer(NetworkTrainer): def __init__(self): super().__init__() + self.sample_prompts_te_outputs = None def assert_extra_args(self, args, train_dataset_group): super().assert_extra_args(args, train_dataset_group) @@ -41,11 +42,17 @@ class FluxNetworkTrainer(NetworkTrainer): args.network_train_unet_only or not args.cache_text_encoder_outputs ), "network for Text Encoder cannot be trained with caching Text Encoder outputs" + if args.max_token_length is not None: + logger.warning("max_token_length is not used in Flux training") + train_dataset_group.verify_bucket_reso_steps(32) # TODO check this + def get_flux_model_name(self, args): + return "schnell" if "schnell" in args.pretrained_model_name_or_path else "dev" + def load_target_model(self, args, weight_dtype, accelerator): # currently offload to cpu for some models - name = "schnell" if "schnell" in args.pretrained_model_name_or_path else "dev" # TODO change this to a more robust way + name = self.get_flux_model_name(args) # if we load to cpu, flux.to(fp8) takes a long time model = flux_utils.load_flow_model(name, args.pretrained_model_name_or_path, weight_dtype, "cpu") @@ -101,7 +108,18 @@ class FluxNetworkTrainer(NetworkTrainer): return flux_lower def get_tokenize_strategy(self, args): - return strategy_flux.FluxTokenizeStrategy(args.max_token_length, args.tokenizer_cache_dir) + name = self.get_flux_model_name(args) + + if args.t5xxl_max_token_length is None: + if name == "schnell": + t5xxl_max_token_length = 256 + else: + t5xxl_max_token_length = 512 + else: + t5xxl_max_token_length = args.t5xxl_max_token_length + + logger.info(f"t5xxl_max_token_length: {t5xxl_max_token_length}") + return strategy_flux.FluxTokenizeStrategy(t5xxl_max_token_length, args.tokenizer_cache_dir) def get_tokenizers(self, tokenize_strategy: strategy_flux.FluxTokenizeStrategy): return [tokenize_strategy.clip_l, tokenize_strategy.t5xxl] @@ -145,7 +163,7 @@ class FluxNetworkTrainer(NetworkTrainer): dataset.new_cache_text_encoder_outputs(text_encoders, accelerator.is_main_process) # cache sample prompts - self.sample_prompts_te_outputs = None + if args.sample_prompts is not None: logger.info(f"cache Text Encoder outputs for sample prompt: {args.sample_prompts}") @@ -549,6 +567,7 @@ class InitFluxTraining: "network_train_unet_only": ("BOOLEAN", {"default": True, "tooltip": "wheter to train the text encoder"}), "text_encoder_lr": ("FLOAT", {"default": 1e-4, "min": 0.0, "max": 10.0, "step": 0.00001, "tooltip": "text encoder learning rate"}), "apply_t5_attn_mask": ("BOOLEAN", {"default": True, "tooltip": "apply t5 attention mask"}), + "t5xxl_max_token_length": ("INT", {"default": 512, "min": 64, "max": 4096, "step": 8, "tooltip": "dev uses 512, schnell 256"}), "cache_latents": (["disk", "memory", "disabled"], {"tooltip": "caches text encoder outputs"}), "cache_text_encoder_outputs": (["disk", "memory", "disabled"], {"tooltip": "caches text encoder outputs"}), "split_mode": ("BOOLEAN", {"default": False, "tooltip": "[EXPERIMENTAL] use split mode for Flux model, network arg `train_blocks=single` is required"}), @@ -561,7 +580,7 @@ class InitFluxTraining: "model_prediction_type": (["raw", "additive", "sigma_scaled"], {"tooltip": "How to interpret and process the model prediction: raw (use as is), additive (add to noisy input), sigma_scaled (apply sigma scaling)."}), "discrete_flow_shift": ("FLOAT", {"default": 1.0, "min": 0.0, "max": 10.0, "step": 0.01, "tooltip": "for the Euler Discrete Scheduler, default is 3.0"}), "highvram": ("BOOLEAN", {"default": False, "tooltip": "memory mode"}), - "attention_mode": (["sdpa", "xformers", "disabled"], {"default": "default", "tooltip": "memory efficient attention mode"}), + "attention_mode": (["sdpa", "xformers", "disabled"], {"default": "sdpa", "tooltip": "memory efficient attention mode"}), "sample_prompts": ("STRING", {"multiline": True, "default": "illustration of a kitten | photograph of a turtle", "tooltip": "validation sample prompts, for multiple prompts, separate by `|`"}), }, } @@ -685,27 +704,11 @@ class FluxTrainLoop: target_global_step = network_trainer.global_step + steps pbar = comfy.utils.ProgressBar(steps) while network_trainer.global_step < target_global_step: - epoch = training_loop( - break_at_steps=target_global_step, - epoch=network_trainer.current_epoch.value, - num_train_epochs=network_trainer.num_train_epochs, - accelerator=network_trainer.accelerator, - network=network_trainer.network, - text_encoder=network_trainer.text_encoder, - unet=network_trainer.unet, - vae=network_trainer.vae, - tokenizers=network_trainer.tokenizers, - args=network_trainer.args, - train_dataloader=network_trainer.train_dataloader, - initial_step=network_trainer.initial_step, - global_step=network_trainer.global_step, - current_epoch=network_trainer.current_epoch, - metadata=network_trainer.metadata, - optimizer=network_trainer.optimizer, - lr_scheduler=network_trainer.lr_scheduler, - loss_recorder=network_trainer.loss_recorder + steps_done = training_loop( + break_at_steps = target_global_step, + epoch = network_trainer.current_epoch.value, ) - pbar.update(network_trainer.global_step - initial_global_step) + pbar.update(steps_done) # Also break if the global steps have reached the max train steps if network_trainer.global_step >= network_trainer.args.max_train_steps: diff --git a/train_network.py b/train_network.py index 65e7759..23f9213 100644 --- a/train_network.py +++ b/train_network.py @@ -1025,13 +1025,30 @@ class NetworkTrainer: clean_memory_on_device(accelerator.device) + self.epoch_to_start = epoch_to_start + self.num_train_epochs = num_train_epochs + self.accelerator = accelerator + self.network = network + self.text_encoder = text_encoder + self.unet = unet + self.vae = vae + self.tokenizers = tokenizers + self.args = args + self.train_dataloader = train_dataloader + self.initial_step = initial_step + self.current_epoch = current_epoch + self.metadata = metadata + self.optimizer = optimizer + self.lr_scheduler = lr_scheduler + self.loss_recorder = loss_recorder + self.save_model = save_model + self.remove_model = remove_model + progress_bar = tqdm( range(args.max_train_steps - initial_step), smoothing=0, disable=False, desc="steps" ) - def training_loop(break_at_steps, epoch, num_train_epochs, accelerator, network, text_encoder, - unet, vae, tokenizers, args, train_dataloader, initial_step, global_step, - current_epoch, metadata, optimizer, lr_scheduler, loss_recorder): - + def training_loop(break_at_steps, epoch): + steps_done = 0 accelerator.print(f"\nepoch {epoch+1}/{num_train_epochs}") current_epoch.value = epoch + 1 @@ -1040,14 +1057,14 @@ class NetworkTrainer: accelerator.unwrap_model(network).on_epoch_start(text_encoder, unet) skipped_dataloader = None - if initial_step > 0: - skipped_dataloader = accelerator.skip_first_batches(train_dataloader, initial_step - 1) - initial_step = 1 + if self.initial_step > 0: + skipped_dataloader = accelerator.skip_first_batches(train_dataloader, self.initial_step - 1) + self.initial_step = 1 for step, batch in enumerate(skipped_dataloader or train_dataloader): - current_step.value = global_step - if initial_step > 0: - initial_step -= 1 + current_step.value = self.global_step + if self.initial_step > 0: + self.initial_step -= 1 continue with accelerator.accumulate(training_model): @@ -1158,7 +1175,7 @@ class NetworkTrainer: # Checks if the accelerator has performed an optimization step behind the scenes if accelerator.sync_gradients: progress_bar.update(1) - global_step += 1 + self.global_step += 1 current_loss = loss.detach().item() loss_recorder.add(epoch=epoch, step=step, loss=current_loss) @@ -1173,35 +1190,19 @@ class NetworkTrainer: logs = self.generate_step_logs( args, current_loss, avr_loss, lr_scheduler, lr_descriptions, keys_scaled, mean_norm, maximum_norm ) - accelerator.log(logs, step=global_step) + accelerator.log(logs, step=self.global_step) - if global_step >= break_at_steps: + if self.global_step >= break_at_steps: break + steps_done += 1 if args.logging_dir is not None: logs = {"loss/epoch": loss_recorder.moving_average} accelerator.log(logs, step=epoch + 1) - self.global_step = global_step - return current_epoch.value + + return steps_done + - self.epoch_to_start = epoch_to_start - self.num_train_epochs = num_train_epochs - self.accelerator = accelerator - self.network = network - self.text_encoder = text_encoder - self.unet = unet - self.vae = vae - self.tokenizers = tokenizers - self.args = args - self.train_dataloader = train_dataloader - self.initial_step = initial_step - self.current_epoch = current_epoch - self.metadata = metadata - self.optimizer = optimizer - self.lr_scheduler = lr_scheduler - self.loss_recorder = loss_recorder - self.save_model = save_model - self.remove_model = remove_model return training_loop