cleaner checkpoint saving WIP
This commit is contained in:
@@ -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
@@ -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)
|
||||
|
||||
|
||||
Reference in New Issue
Block a user