cleaner checkpoint saving WIP

This commit is contained in:
mayukhdeb
2024-04-17 04:56:34 -07:00
parent aaf090c2e8
commit cde4432588
2 changed files with 66 additions and 24 deletions
+3 -3
View File
@@ -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,
+63 -21
View File
@@ -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)