add SaveStateDict

This commit is contained in:
hnmr293
2023-05-02 03:26:22 +09:00
parent adabf59e29
commit b3478a44f5
3 changed files with 93 additions and 0 deletions
+1
View File
@@ -79,6 +79,7 @@ visualization of 4-channel latent tensor
|model|StateDictMerger|`DICT`, `DICT`, `FLOAT`|`MODEL`, `CLIP`, `VAE`|merge two or three models| |model|StateDictMerger|`DICT`, `DICT`, `FLOAT`|`MODEL`, `CLIP`, `VAE`|merge two or three models|
|model|StateDictMergerBlockWeighted|`DICT`, `DICT`|`DICT`|merge two models with per-block weights| |model|StateDictMergerBlockWeighted|`DICT`, `DICT`|`DICT`|merge two models with per-block weights|
|model|StateDictMergerBlockWeightedMulti|`MODEL`, `MODEL`, `STRING`|`MODEL`, `CLIP`, `VAE`|merge two models with per-block weights| |model|StateDictMergerBlockWeightedMulti|`MODEL`, `MODEL`, `STRING`|`MODEL`, `CLIP`, `VAE`|merge two models with per-block weights|
|model|SaveStateDict|`MODEL`, `STRING`, `STRING`|`-`|save state_dict to the output directory|
|model|ModelIter|`MODEL`, `MODEL`|`MODEL`|iterate models| |model|ModelIter|`MODEL`, `MODEL`|`MODEL`|iterate models|
|model|CLIPlIter|`CLIP`, `CLIP`|`CLIP`|iterate CLIPs| |model|CLIPlIter|`CLIP`, `CLIP`|`CLIP`|iterate CLIPs|
|model|VAElIter|`VAE`, `VAE`|`VAE`|iterate VAEs| |model|VAElIter|`VAE`, `VAE`|`VAE`|iterate VAEs|
+5
View File
@@ -5,6 +5,7 @@ from .model.loader import StateDictLoader, Dict2Model
from .model.iter import ModelIter, CLIPIter, VAEIter from .model.iter import ModelIter, CLIPIter, VAEIter
from .model.merge import StateDictMerger, StateDictMergerBlockWeighted from .model.merge import StateDictMerger, StateDictMergerBlockWeighted
from .model.merge2 import StateDictMergerBlockWeightedMulti from .model.merge2 import StateDictMergerBlockWeightedMulti
from .model.save import SaveStateDict
from .image.image import GridImage from .image.image import GridImage
from .image.latenttoimage import LatentToImage, LatentToHist from .image.latenttoimage import LatentToImage, LatentToHist
from .image.blend_extra import Blend2 from .image.blend_extra import Blend2
@@ -64,6 +65,9 @@ NODE_CLASS_MAPPINGS = {
## weights should be specified by Text ## weights should be specified by Text
'StateDictMergerBlockWeightedMulti': StateDictMergerBlockWeightedMulti, 'StateDictMergerBlockWeightedMulti': StateDictMergerBlockWeightedMulti,
## save state_dict
'SaveStateDict': SaveStateDict,
# image # image
## extra blend mode ## extra blend mode
@@ -73,6 +77,7 @@ NODE_CLASS_MAPPINGS = {
'GridImage': GridImage, 'GridImage': GridImage,
# others # others
'SaveText': SaveText, 'SaveText': SaveText,
} }
+87
View File
@@ -0,0 +1,87 @@
import os
import re
from typing import Union
import torch
import folder_paths
class SaveStateDict:
@classmethod
def INPUT_TYPES(cls):
return {
'required': {
'weights': ('DICT',),
'filename': ('STRING', { 'multiline': False, 'default': 'merged_model.safetensors' }),
'overwrite': (['False', 'True'],),
}
}
OUTPUT_NODE = True
RETURN_TYPES = ()
FUNCTION = 'execute'
CATEGORY = 'model'
def execute(self, weights: dict, filename: str, overwrite: Union[str,bool]):
if isinstance(overwrite, str):
overwrite = overwrite.lower() == 'true'
subdir = os.path.dirname(os.path.normpath(filename))
basename = os.path.basename(os.path.normpath(filename))
output_dir = folder_paths.get_output_directory()
full_output_dir = os.path.join(output_dir, subdir)
if os.path.commonpath((output_dir, os.path.realpath(full_output_dir))) != output_dir:
print('Saving image outside the output folder is not allowed.')
return {}
full_path = os.path.join(full_output_dir, basename)
if os.path.exists(full_path):
print(f'{full_path} already exists.')
if overwrite:
print(f'overwriting: {full_path}')
else:
base, ext = os.path.splitext(basename)
def replace(m: re.Match):
x = m.group(0)
if len(x) == 0:
return '0'
else:
n = int(x)
return str(n + 1)
base_renamed = re.sub(r'\d+$|$', replace, base)
full_path_renamed = os.path.join(full_output_dir, base_renamed + ext)
print(f'rename {full_path} -> {full_path_renamed}')
full_path = full_path_renamed
os.makedirs(full_output_dir, exist_ok=True)
print(f'Saving the state_dict to {full_path}')
ext = os.path.splitext(full_path)[1]
saver = None
if ext == '.pt' or ext == '.ckpt':
saver = self.save_torch
elif ext == '.safetensor' or ext == '.safetensors':
saver = self.save_safetensors
else:
full_path += '.safetensors'
saver = self.save_safetensors
saver(weights, full_path)
return {}
def save_torch(self, model: dict, path: str):
torch.save(model, path)
def save_safetensors(self, model: dict, path: str):
import safetensors.torch
safetensors.torch.save_file(model, path)