From 07762fb4f05f6ccd15e5651edf67fab3a97ba6d4 Mon Sep 17 00:00:00 2001 From: kijai <40791699+kijai@users.noreply.github.com> Date: Tue, 27 Aug 2024 12:21:11 +0300 Subject: [PATCH] expose mask, resume resume doesn't seem functional in kohya for Flux yet --- nodes.py | 59 +++++++++++++++++++++++++++++++++++++++--------- train_network.py | 12 ++++------ 2 files changed, 52 insertions(+), 19 deletions(-) diff --git a/nodes.py b/nodes.py index ed54cb6..efb86a5 100644 --- a/nodes.py +++ b/nodes.py @@ -111,6 +111,7 @@ class TrainDatasetAdd: "num_repeats": ("INT", {"default": 1, "min": 1, "tooltip": "number of times to repeat dataset for an epoch"}), "min_bucket_reso": ("INT", {"default": 256, "min": 64, "max": 4096, "step": 8, "tooltip": "min bucket resolution"}), "max_bucket_reso": ("INT", {"default": 1024, "min": 64, "max": 4096, "step": 8, "tooltip": "max bucket resolution"}), + "alpha_mask": ("BOOLEAN",{"default": False, "tooltip": "use alpha channel as mask for training"}), }, } @@ -120,7 +121,7 @@ class TrainDatasetAdd: CATEGORY = "FluxTrainer" def create_config(self, dataset_config, dataset_path, class_tokens, width, height, batch_size, num_repeats, enable_bucket, - bucket_no_upscale, min_bucket_reso, max_bucket_reso): + bucket_no_upscale, min_bucket_reso, max_bucket_reso, alpha_mask): dataset = { "datasets": [ @@ -131,6 +132,7 @@ class TrainDatasetAdd: "bucket_no_upscale": bucket_no_upscale, "min_bucket_reso": min_bucket_reso, "max_bucket_reso": max_bucket_reso, + "alpha_mask": alpha_mask, "subsets": [ { @@ -287,6 +289,7 @@ class InitFluxLoRATraining: }, "optional": { "additional_args": ("STRING", {"multiline": True, "default": "", "tooltip": "additional args to pass to the training command"}), + "resume_args": ("ARGS", {"default": "", "tooltip": "resume args to pass to the training command"}), }, } @@ -296,7 +299,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,**kwargs,): + gradient_dtype, save_dtype, split_mode, additional_args=None, resume_args=None, **kwargs,): mm.soft_empty_cache() output_dir = os.path.abspath(kwargs.get("output_dir")) @@ -388,6 +391,9 @@ class InitFluxLoRATraining: config_dict.update(kwargs) config_dict.update(optimizer_settings) + if resume_args: + config_dict.update(resume_args) + for key, value in config_dict.items(): setattr(args, key, value) @@ -397,7 +403,6 @@ class InitFluxLoRATraining: epochs_count = network_trainer.num_train_epochs - saved_args_file_path = os.path.join(output_dir, f"{output_name}_args.json") with open(saved_args_file_path, 'w') as f: json.dump(vars(args), f, indent=4) @@ -446,6 +451,7 @@ class InitFluxTraining: }, "optional": { "additional_args": ("STRING", {"multiline": True, "default": "", "tooltip": "additional args to pass to the training command"}), + "resume_args": ("ARGS", {"default": "", "tooltip": "resume args to pass to the training command"}), }, } @@ -455,7 +461,7 @@ class InitFluxTraining: CATEGORY = "FluxTrainer" def init_training(self, flux_models, optimizer_settings, dataset, sample_prompts, output_name, - attention_mode, gradient_dtype, save_dtype, optimizer_fusing, additional_args=None, **kwargs,): + attention_mode, gradient_dtype, save_dtype, optimizer_fusing, additional_args=None, resume_args=None, **kwargs,): mm.soft_empty_cache() output_dir = os.path.abspath(kwargs.get("output_dir")) @@ -543,6 +549,9 @@ class InitFluxTraining: config_dict.update(kwargs) config_dict.update(optimizer_settings) + if resume_args: + config_dict.update(resume_args) + for key, value in config_dict.items(): setattr(args, key, value) @@ -768,8 +777,8 @@ class FluxTrainEnd: }, } - RETURN_TYPES = ("STRING", "STRING",) - RETURN_NAMES = ("lora_path", "metadata",) + RETURN_TYPES = ("STRING", "STRING", "STRING",) + RETURN_NAMES = ("lora_name", "metadata", "lora_path",) FUNCTION = "endtrain" CATEGORY = "FluxTrainer" @@ -792,7 +801,8 @@ class FluxTrainEnd: network_trainer.save_model(ckpt_name, network, network_trainer.global_step, network_trainer.num_train_epochs, force_sync_upload=True) logger.info("model saved.") - final_output_lora_path = os.path.join(network_trainer.args.output_dir, network_trainer.args.output_name) + final_lora_name = str(network_trainer.args.output_name) + final_lora_path = os.path.join(network_trainer.args.output_dir, ckpt_name) # metadata metadata = json.dumps(network_trainer.metadata, indent=2) @@ -801,7 +811,29 @@ class FluxTrainEnd: network_trainer = None mm.soft_empty_cache() - return (final_output_lora_path, metadata) + return (final_lora_name, metadata, final_lora_path) + +class FluxTrainResume: + @classmethod + def INPUT_TYPES(s): + return {"required": { + "load_state_path": ("STRING", {"default": "", "multiline": True, "tooltip": "path to load state from"}), + "skip_until_initial_step" : ("BOOLEAN", {"default": False}), + }, + } + + RETURN_TYPES = ("ARGS", ) + RETURN_NAMES = ("resume_args", ) + FUNCTION = "resume" + CATEGORY = "FluxTrainer" + + def resume(self, load_state_path, skip_until_initial_step): + resume_args ={ + "resume": load_state_path, + "skip_until_initial_step": skip_until_initial_step + } + + return (resume_args, ) class FluxTrainValidationSettings: @classmethod @@ -884,6 +916,7 @@ class VisualizeLoss: "normalize_y": ("BOOLEAN", {"default": True, "tooltip": "normalize the y-axis to 0"}), "width": ("INT", {"default": 768, "min": 256, "max": 4096, "step": 2, "tooltip": "width of the plot in pixels"}), "height": ("INT", {"default": 512, "min": 256, "max": 4096, "step": 2, "tooltip": "height of the plot in pixels"}), + "log_scale": ("BOOLEAN", {"default": False, "tooltip": "use log scale on the y-axis"}), }, } @@ -892,7 +925,7 @@ class VisualizeLoss: FUNCTION = "draw" CATEGORY = "FluxTrainer" - def draw(self, network_trainer, window_size, plot_style, normalize_y, width, height): + def draw(self, network_trainer, window_size, plot_style, normalize_y, width, height, log_scale): import numpy as np loss_values = network_trainer["network_trainer"].loss_recorder.global_loss_list @@ -915,6 +948,8 @@ class VisualizeLoss: ax.set_ylabel('Loss') if normalize_y: plt.ylim(bottom=0) + if log_scale: + ax.set_yscale('log') ax.set_title('Training Loss Over Time') ax.legend() ax.grid(True) @@ -1361,7 +1396,8 @@ NODE_CLASS_MAPPINGS = { "OptimizerConfigAdafactor": OptimizerConfigAdafactor, "FluxTrainSaveModel": FluxTrainSaveModel, "ExtractFluxLoRA": ExtractFluxLoRA, - "OptimizerConfigProdigy": OptimizerConfigProdigy + "OptimizerConfigProdigy": OptimizerConfigProdigy, + "FluxTrainResume": FluxTrainResume } NODE_DISPLAY_NAME_MAPPINGS = { "InitFluxLoRATraining": "Init Flux LoRA Training", @@ -1381,5 +1417,6 @@ NODE_DISPLAY_NAME_MAPPINGS = { "OptimizerConfigAdafactor": "Optimizer Config Adafactor", "FluxTrainSaveModel": "Flux Train Save Model", "ExtractFluxLoRA": "Extract Flux LoRA", - "OptimizerConfigProdigy": "Optimizer Config Prodigy" + "OptimizerConfigProdigy": "Optimizer Config Prodigy", + "FluxTrainResume": "Flux Train Resume" } diff --git a/train_network.py b/train_network.py index 5ae6060..33af595 100644 --- a/train_network.py +++ b/train_network.py @@ -1014,10 +1014,9 @@ class NetworkTrainer: self.global_step = 0 # training loop if initial_step > 0: # only if skip_until_initial_step is specified - for skip_epoch in range(epoch_to_start): # skip epochs - logger.info(f"skipping epoch {skip_epoch+1} because initial_step (multiplied) is {initial_step}") - initial_step -= len(train_dataloader) self.global_step = initial_step + logger.info(f"skipping epoch {epoch_to_start} because initial_step (multiplied) is {initial_step}") + initial_step -= epoch_to_start * len(train_dataloader) # log device and dtype for each model logger.info(f"unet dtype: {unet_weight_dtype}, device: {unet.device}") @@ -1058,14 +1057,11 @@ class NetworkTrainer: skipped_dataloader = None if self.initial_step > 0: - skipped_dataloader = accelerator.skip_first_batches(train_dataloader, self.initial_step - 1) - self.initial_step = 1 + skipped_dataloader = accelerator.skip_first_batches(train_dataloader, initial_step) + initial_step = 0 for step, batch in enumerate(skipped_dataloader or train_dataloader): current_step.value = self.global_step - if self.initial_step > 0: - self.initial_step -= 1 - continue with accelerator.accumulate(training_model): on_step_start(text_encoder, unet)