add SaveStateDict
This commit is contained in:
@@ -79,6 +79,7 @@ visualization of 4-channel latent tensor
|
||||
|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|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|CLIPlIter|`CLIP`, `CLIP`|`CLIP`|iterate CLIPs|
|
||||
|model|VAElIter|`VAE`, `VAE`|`VAE`|iterate VAEs|
|
||||
|
||||
@@ -5,6 +5,7 @@ from .model.loader import StateDictLoader, Dict2Model
|
||||
from .model.iter import ModelIter, CLIPIter, VAEIter
|
||||
from .model.merge import StateDictMerger, StateDictMergerBlockWeighted
|
||||
from .model.merge2 import StateDictMergerBlockWeightedMulti
|
||||
from .model.save import SaveStateDict
|
||||
from .image.image import GridImage
|
||||
from .image.latenttoimage import LatentToImage, LatentToHist
|
||||
from .image.blend_extra import Blend2
|
||||
@@ -64,6 +65,9 @@ NODE_CLASS_MAPPINGS = {
|
||||
## weights should be specified by Text
|
||||
'StateDictMergerBlockWeightedMulti': StateDictMergerBlockWeightedMulti,
|
||||
|
||||
## save state_dict
|
||||
'SaveStateDict': SaveStateDict,
|
||||
|
||||
# image
|
||||
|
||||
## extra blend mode
|
||||
@@ -73,6 +77,7 @@ NODE_CLASS_MAPPINGS = {
|
||||
'GridImage': GridImage,
|
||||
|
||||
# others
|
||||
|
||||
'SaveText': SaveText,
|
||||
|
||||
}
|
||||
|
||||
@@ -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)
|
||||
Reference in New Issue
Block a user