rename stuff

This commit is contained in:
kijai
2024-08-15 11:47:31 +03:00
parent 596149fc8d
commit 00fee231dd
2 changed files with 13 additions and 16 deletions
+11 -14
View File
@@ -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"
}
+2 -2
View File
@@ -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