diff --git a/nodes.py b/nodes.py index 195ce76..fa1b635 100644 --- a/nodes.py +++ b/nodes.py @@ -448,7 +448,7 @@ class FluxNetworkTrainer(NetworkTrainer): metadata["ss_model_prediction_type"] = args.model_prediction_type metadata["ss_discrete_flow_shift"] = args.discrete_flow_shift -class SelectModelsTrainFlux: +class FluxTrainModelSelect: @classmethod def INPUT_TYPES(s): return {"required": { @@ -500,10 +500,10 @@ class TrainDatasetConfig: RETURN_TYPES = ("TOML_DATASET",) RETURN_NAMES = ("dataset",) - FUNCTION = "loadmodel" + FUNCTION = "create_config" CATEGORY = "TrainFlux" - def loadmodel(self, dataset_path, class_tokens, width, height, batch_size, enable_bucket, color_aug, flip_aug, + def create_config(self, dataset_path, class_tokens, width, height, batch_size, enable_bucket, color_aug, flip_aug, bucket_no_upscale, min_bucket_reso, max_bucket_reso): import toml @@ -535,7 +535,7 @@ class TrainDatasetConfig: return (toml.dumps(dataset),) -class TrainFlux: +class InitFluxTraining: @classmethod def INPUT_TYPES(s): return {"required": { @@ -548,8 +548,6 @@ class TrainFlux: #"max_train_epochs": ("INT", {"default": 4, "min": 1, "max": 1000, "step": 1, "tooltip": "max number of training epochs"}), "optimizer_type": (["adamw8bit", "adafactor", "prodigy"], {"default": "adamw8bit", "tooltip": "optimizer type"}), "max_train_steps": ("INT", {"default": 1500, "min": 1, "max": 10000, "step": 1, "tooltip": "max number of training steps"}), - "save_every_n_steps": ("INT", {"default": 250, "min": 1, "max": 10000, "step": 1, "tooltip": "save every n epochs"}), - "sample_every_n_steps": ("INT", {"default": 250, "min": 1, "max": 10000, "step": 1, "tooltip": "sample every n steps"}), "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"}), @@ -572,11 +570,10 @@ class TrainFlux: RETURN_TYPES = ("NETWORKTRAINER",) RETURN_NAMES = ("network_trainer",) - FUNCTION = "loadmodel" + FUNCTION = "init_training" CATEGORY = "TrainFlux" - def loadmodel(self, flux_models, dataset, sample_prompts, output_name, optimizer_type, **kwargs,): - device = mm.get_torch_device() + def init_training(self, flux_models, dataset, sample_prompts, output_name, optimizer_type, **kwargs,): mm.soft_empty_cache() parser = setup_parser() @@ -649,7 +646,7 @@ class TrainFlux: with torch.inference_mode(False): network_trainer = FluxNetworkTrainer() - training_loop = network_trainer.train(args) + training_loop = network_trainer.init_train(args) final_output_lora_path = os.path.join(output_dir, "output", output_name) @@ -751,14 +748,14 @@ class TrainLoop: NODE_CLASS_MAPPINGS = { - "TrainFlux": TrainFlux, - "SelectModelsTrainFlux": SelectModelsTrainFlux, + "InitFluxTraining": InitFluxTraining, + "FluxTrainModelSelect": FluxTrainModelSelect, "TrainDatasetConfig": TrainDatasetConfig, "TrainLoop": TrainLoop } NODE_DISPLAY_NAME_MAPPINGS = { - "TrainFlux": "TrainFlux", - "SelectModelsTrainFlux": "SelectModelsTrainFlux", + "InitFluxTraining": "Init Flux Training", + "FluxTrainModelSelect": "FluxTrain ModelSelect", "TrainDatasetConfig": "Train Dataset Config", "TrainLoop": "Train Loop" } diff --git a/train_network.py b/train_network.py index 547a1ea..1194285 100644 --- a/train_network.py +++ b/train_network.py @@ -238,7 +238,7 @@ class NetworkTrainer: # endregion - def train(self, args): + def init_train(self, args): session_id = random.randint(0, 2**32) training_started_at = time.time() train_util.verify_training_args(args) @@ -993,7 +993,7 @@ class NetworkTrainer: text_encoders = [] # For --sample_at_first - self.sample_images(accelerator, args, 0, global_step, accelerator.device, vae, tokenizers, text_encoder, unet) + #self.sample_images(accelerator, args, 0, global_step, accelerator.device, vae, tokenizers, text_encoder, unet) # training loop if initial_step > 0: # only if skip_until_initial_step is specified