diff --git a/main.py b/main.py index 4d54404..ecb9120 100755 --- a/main.py +++ b/main.py @@ -17,7 +17,7 @@ from peft import LoraConfig, get_peft_model from typing import Union, Iterable, List, Dict, Tuple, Optional, cast from trainer.utils.utils import * -from trainer.lora import save_lora +from trainer.lora import save_checkpoint from trainer.embedding_handler import TokenEmbeddingsHandler from trainer.dataset import PreprocessedDataset from trainer.config import TrainingConfig @@ -353,7 +353,7 @@ def train( config.save_as_json( os.path.join(output_save_dir, "training_args.json") ) - save_lora( + save_checkpoint( output_dir=output_save_dir, global_step=global_step, unet=unet, @@ -426,7 +426,7 @@ def train( if not os.path.exists(output_save_dir): os.makedirs(output_save_dir, exist_ok=True) config.save_as_json(os.path.join(output_save_dir, "training_args.json")) - save_lora( + save_checkpoint( output_dir=output_save_dir, global_step=global_step, unet=unet, diff --git a/trainer/lora.py b/trainer/lora.py index f8c3063..6cccaf4 100644 --- a/trainer/lora.py +++ b/trainer/lora.py @@ -4,6 +4,7 @@ from safetensors.torch import load_file, save_file from typing import Dict from peft import PeftModel from trainer.embedding_handler import TokenEmbeddingsHandler +from .utils.json_stuff import save_as_json from diffusers import StableDiffusionPipeline, StableDiffusionXLPipeline def patch_pipe_with_lora(pipe, lora_path, lora_scale = 1.0): @@ -60,27 +61,69 @@ from diffusers.utils import ( convert_unet_state_dict_to_peft, ) -def save_lora( - output_dir, - global_step, - unet, - embedding_handler, - token_dict, - is_lora, - unet_lora_parameters, - name: str = None, - text_encoder_peft_models: list = None - ): + +def remove_delimiter_characters(name: str): + # Make sure all weird delimiter characters are removed from concept_name before using it as a filepath: + return name.replace(" ", "_").replace("/", "_").replace("\\", "_").replace(":", "_").replace("*", "_").replace("?", "_").replace("\"", "_").replace("<", "_").replace(">", "_").replace("|", "_") + +# Convert to WebUI format + +def convert_pytorch_lora_safetensors_to_webui( + pytorch_lora_weights_filename: str, + output_filename: str +): + assert os.path.exists(pytorch_lora_weights_filename), f"Invalid path: {pytorch_lora_weights_filename}" + lora_state_dict = load_file(pytorch_lora_weights_filename) + peft_state_dict = convert_all_state_dict_to_peft(lora_state_dict) + kohya_state_dict = convert_state_dict_to_kohya(peft_state_dict) + save_file(kohya_state_dict, output_filename) + +def save_checkpoint( + output_dir: str, + global_step: int, + unet, + embedding_handler, + token_dict: dict, + is_lora: bool, + unet_lora_parameters, + name: str = None, + text_encoder_peft_models: list = None +): """ Save the model + embeddings to output_dir + + output_dir/ + - name_embeddings.safetensors + - special_params.json + + ## optional + text_encoder_lora_0/ + - adapter_config.json + - adapter_model.safetensors + - README.md + + ## optional + text_encoder_lora_1/ + - adapter_config.json + - adapter_model.safetensors + - README.md """ print(f"Saving checkpoint at step.. {global_step}") - # Make sure all weird delimiter characters are removed from concept_name before using it as a filepath: - name = name.replace(" ", "_").replace("/", "_").replace("\\", "_").replace(":", "_").replace("*", "_").replace("?", "_").replace("\"", "_").replace("<", "_").replace(">", "_").replace("|", "_") - embedding_handler.save_embeddings(f"{output_dir}/{name}_embeddings.safetensors") + name = remove_delimiter_characters(name) - with open(f"{output_dir}/special_params.json", "w") as f: - json.dump(token_dict, f) + embedding_handler.save_embeddings( + os.path.join( + output_dir, + f"{name}_embeddings.safetensors" + ) + ) + + save_as_json( + token_dict, + filename = os.path.join( + output_dir, "special_params.json" + ) + ) if text_encoder_peft_models is not None: for idx, model in enumerate(text_encoder_peft_models): @@ -103,11 +146,10 @@ def save_lora( #text_encoder_2_lora_layers=text_encoder_two_lora_layers_to_save, ) - # Convert to WebUI format - lora_state_dict = load_file(f"{output_dir}/pytorch_lora_weights.safetensors") - peft_state_dict = convert_all_state_dict_to_peft(lora_state_dict) - kohya_state_dict = convert_state_dict_to_kohya(peft_state_dict) - save_file(kohya_state_dict, f"{output_dir}/{name}.safetensors") + convert_pytorch_lora_safetensors_to_webui( + pytorch_lora_weights_filename=os.path.join(output_dir, "pytorch_lora_weights.safetensors"), + output_filename=os.path.join(output_dir, f"{name}.safetensors") + ) else: unet.save_pretrained(save_directory = output_dir)