expose mask, resume

resume doesn't seem functional in kohya for Flux yet
This commit is contained in:
kijai
2024-08-27 12:21:11 +03:00
parent 94eede3dc0
commit 07762fb4f0
2 changed files with 52 additions and 19 deletions
+48 -11
View File
@@ -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"
}
+4 -8
View File
@@ -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)