Files
melMass-comfy_mtb/nodes/graph_utils.py
T
2023-07-20 00:00:22 +02:00

98 lines
2.7 KiB
Python

import torch
import folder_paths
import os
from ..log import log
class SaveTensors:
"""Debug node that will probably be removed in the future"""
def __init__(self):
self.output_dir = folder_paths.get_output_directory()
self.type = "output"
@classmethod
def INPUT_TYPES(cls):
return {
"required": {
"filename_prefix": ("STRING", {"default": "ComfyPickle"}),
},
"optional": {
"image": ("IMAGE",),
"mask": ("MASK",),
"latent": ("LATENT",),
},
}
FUNCTION = "save"
OUTPUT_NODE = True
RETURN_TYPES = ()
CATEGORY = "utils"
def save(
self,
filename_prefix,
image: torch.Tensor = None,
mask: torch.Tensor = None,
latent: torch.Tensor = None,
):
(
full_output_folder,
filename,
counter,
subfolder,
filename_prefix,
) = folder_paths.get_save_image_path(filename_prefix, self.output_dir)
if image is not None:
image_file = f"{filename}_image_{counter:05}.pt"
torch.save(image, os.path.join(full_output_folder, image_file))
# np.save(os.path.join(full_output_folder, image_file), image.cpu().numpy())
if mask is not None:
mask_file = f"{filename}_mask_{counter:05}.pt"
torch.save(mask, os.path.join(full_output_folder, mask_file))
# np.save(os.path.join(full_output_folder, mask_file), mask.cpu().numpy())
if latent is not None:
# for latent we must use pickle
latent_file = f"{filename}_latent_{counter:05}.pt"
torch.save(latent, os.path.join(full_output_folder, latent_file))
# pickle.dump(latent, open(os.path.join(full_output_folder, latent_file), "wb"))
# np.save(os.path.join(full_output_folder, latent_file), latent[""].cpu().numpy())
return f"{filename_prefix}_{counter:05}"
class StringReplace:
"""Basic string replacement"""
@classmethod
def INPUT_TYPES(cls):
return {
"required": {
"string": ("STRING", {"forceInput": True}),
"old": ("STRING", {"default": ""}),
"new": ("STRING", {"default": ""}),
}
}
FUNCTION = "replace_str"
RETURN_TYPES = ("STRING",)
CATEGORY = "string"
def replace_str(self, string: str, old: str, new: str):
log.debug(f"Current string: {string}")
log.debug(f"Find string: {old}")
log.debug(f"Replace string: {new}")
string = string.replace(old, new)
log.debug(f"New string: {string}")
return (string,)
__nodes__ = [SaveTensors, StringReplace]