From c3aa4ea889153519f7be40636b44d5a03a060816 Mon Sep 17 00:00:00 2001 From: kijai <40791699+kijai@users.noreply.github.com> Date: Mon, 9 Sep 2024 20:27:31 +0300 Subject: [PATCH] fix cpu_offload_checkpointing --- nodes.py | 50 ++++++++++++++++++++++++++++++++++++++++---------- 1 file changed, 40 insertions(+), 10 deletions(-) diff --git a/nodes.py b/nodes.py index 53e7809..e70fbf3 100644 --- a/nodes.py +++ b/nodes.py @@ -88,7 +88,6 @@ class TrainDatasetGeneralConfig: }, "optional": { "reset_on_queue": ("BOOLEAN",{"default": False, "tooltip": "Force refresh of everything for cleaner queueing"}), - "reg_data_dir": ("STRING",{"multiline": True, "forceInput": True, "default": "", "tooltip": "reg data dir"}), } } @@ -97,7 +96,7 @@ class TrainDatasetGeneralConfig: FUNCTION = "create_config" CATEGORY = "FluxTrainer" - def create_config(self, shuffle_caption, caption_dropout_rate, color_aug, flip_aug, alpha_mask, reset_on_queue=False, reg_data_dir=""): + def create_config(self, shuffle_caption, caption_dropout_rate, color_aug, flip_aug, alpha_mask, reset_on_queue=False): dataset = { "general": { @@ -114,11 +113,37 @@ class TrainDatasetGeneralConfig: #print(dataset_json) dataset_config = { "datasets": dataset_json, - "alpha_mask": alpha_mask, - "reg_data_dir": reg_data_dir + "alpha_mask": alpha_mask } return (dataset_config,) +class TrainDatasetRegularization: + + @classmethod + def INPUT_TYPES(s): + return {"required": { + "dataset_path": ("STRING",{"multiline": True, "default": "", "tooltip": "path to dataset, root is the 'ComfyUI' folder, with windows portable 'ComfyUI_windows_portable'"}), + "class_tokens": ("STRING",{"multiline": True, "default": "", "tooltip": "aka trigger word, if specified, will be added to the start of each caption, if no captions exist, will be used on it's own"}), + "num_repeats": ("INT", {"default": 1, "min": 1, "tooltip": "number of times to repeat dataset for an epoch"}), + }, + } + + RETURN_TYPES = ("JSON",) + RETURN_NAMES = ("subset",) + FUNCTION = "create_config" + CATEGORY = "FluxTrainer" + + def create_config(self, dataset_path, class_tokens, num_repeats): + + reg_subset = { + "image_dir": dataset_path, + "class_tokens": class_tokens, + "num_repeats": num_repeats, + "is_reg": True + } + + return reg_subset, + class TrainDatasetAdd: def __init__(self): self.previous_dataset_signature = None @@ -138,6 +163,9 @@ class TrainDatasetAdd: "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"}), }, + "optional": { + "regularization": ("JSON", {"tooltip": "reg data dir"}), + } } RETURN_TYPES = ("JSON",) @@ -146,7 +174,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, regularization=None): new_dataset = { "resolution": (width, height), @@ -163,6 +191,8 @@ class TrainDatasetAdd: } ] } + if regularization is not None: + new_dataset["subsets"].append(regularization) # Generate a signature for the new dataset new_dataset_signature = self.generate_signature(new_dataset) @@ -448,6 +478,7 @@ class InitFluxLoRATraining: if gradient_checkpointing == "disabled": config_dict["gradient_checkpointing"] = False elif gradient_checkpointing == "enabled_with_cpu_offloading": + config_dict["gradient_checkpointing"] = True config_dict["cpu_offload_checkpointing"] = True else: config_dict["gradient_checkpointing"] = True @@ -455,9 +486,6 @@ class InitFluxLoRATraining: if flux_models["lora_path"]: config_dict["network_weights"] = flux_models["lora_path"] - if dataset["reg_data_dir"]: - config_dict["reg_data_dir"] = dataset["reg_data_dir"] - config_dict.update(kwargs) config_dict.update(optimizer_settings) @@ -1491,7 +1519,8 @@ NODE_CLASS_MAPPINGS = { "ExtractFluxLoRA": ExtractFluxLoRA, "OptimizerConfigProdigy": OptimizerConfigProdigy, "FluxTrainResume": FluxTrainResume, - "FluxTrainBlockSelect": FluxTrainBlockSelect + "FluxTrainBlockSelect": FluxTrainBlockSelect, + "TrainDatasetRegularization": TrainDatasetRegularization } NODE_DISPLAY_NAME_MAPPINGS = { "InitFluxLoRATraining": "Init Flux LoRA Training", @@ -1513,5 +1542,6 @@ NODE_DISPLAY_NAME_MAPPINGS = { "ExtractFluxLoRA": "Extract Flux LoRA", "OptimizerConfigProdigy": "Optimizer Config Prodigy", "FluxTrainResume": "Flux Train Resume", - "FluxTrainBlockSelect": "Flux Train Block Select" + "FluxTrainBlockSelect": "Flux Train Block Select", + "TrainDatasetRegularization": "Train Dataset Regularization" }