save workflow as json first thing as well, allow block ranges for block select

This commit is contained in:
kijai
2024-09-10 17:05:18 +03:00
parent 47c62fd8c9
commit 753f617f26
+47 -9
View File
@@ -361,6 +361,9 @@ class InitFluxLoRATraining:
"block_args": ("ARGS", {"default": "", "tooltip": "limit the blocks used in the LoRA"}),
"gradient_checkpointing": (["enabled", "enabled_with_cpu_offloading", "disabled"], {"default": "enabled", "tooltip": "use gradient checkpointing"}),
},
"hidden": {
"prompt": "PROMPT", "extra_pnginfo": "EXTRA_PNGINFO"
},
}
RETURN_TYPES = ("NETWORKTRAINER", "INT", "KOHYA_ARGS",)
@@ -370,7 +373,7 @@ class InitFluxLoRATraining:
def init_training(self, flux_models, dataset, optimizer_settings, sample_prompts, output_name, attention_mode,
gradient_dtype, save_dtype, split_mode, additional_args=None, resume_args=None, train_text_encoder='disabled',
block_args=None, gradient_checkpointing="enabled", **kwargs,):
block_args=None, gradient_checkpointing="enabled", prompt=None, extra_pnginfo=None, **kwargs,):
mm.soft_empty_cache()
output_dir = os.path.abspath(kwargs.get("output_dir"))
@@ -493,17 +496,27 @@ class InitFluxLoRATraining:
for key, value in config_dict.items():
setattr(args, key, value)
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)
#workflow saving
metadata = {}
if extra_pnginfo is not None:
metadata.update(extra_pnginfo["workflow"])
saved_workflow_file_path = os.path.join(output_dir, f"{output_name}_workflow.json")
with open(saved_workflow_file_path, 'w') as f:
json.dump(metadata, f, indent=4)
#pass args to kohya and initialize trainer
with torch.inference_mode(False):
network_trainer = FluxNetworkTrainer()
training_loop = network_trainer.init_train(args)
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)
trainer = {
"network_trainer": network_trainer,
"training_loop": training_loop,
@@ -942,7 +955,7 @@ class FluxTrainBlockSelect:
@classmethod
def INPUT_TYPES(s):
return {"required": {
"include": ("STRING", {"default": "lora_unet_single_blocks_20_linear2", "multiline": True, "tooltip": "blocks to include in the LoRA network"}),
"include": ("STRING", {"default": "lora_unet_single_blocks_20_linear2", "multiline": True, "tooltip": "blocks to include in the LoRA network, to select multiple blocks either input them as "}),
},
}
@@ -952,10 +965,35 @@ class FluxTrainBlockSelect:
CATEGORY = "FluxTrainer"
def block_select(self, include):
block_args ={
"include": f"only_if_contains={include}",
import re
# Split the input string by commas to handle multiple ranges/blocks
elements = include.split(',')
# Initialize a list to collect block names
blocks = []
# Pattern to find ranges like (10-20)
pattern = re.compile(r'\((\d+)-(\d+)\)')
for element in elements:
element = element.strip()
matches = pattern.findall(element)
if matches:
for start, end in matches:
# Generate block names for the range and add them to the list
blocks.extend([f"lora_unet_single_blocks_{i}_linear2" for i in range(int(start), int(end) + 1)])
else:
# If no range is found, add the block name directly
blocks.append(element)
# Construct the `include` string
include_string = ','.join(blocks)
block_args = {
"include": f"only_if_contains={include_string}",
}
return (block_args, )
class FluxTrainValidationSettings: