From f4586a3e128555dc1de332ecded88a46aa36e9c7 Mon Sep 17 00:00:00 2001 From: hnmr293 Date: Sun, 30 Apr 2023 20:54:45 +0900 Subject: [PATCH] add csv output --- __init__.py | 5 ++++ image/latenttoimage.py | 35 +++++++++++++++++++---- outputs.py | 65 ++++++++++++++++++++++++++++++++++++++++++ 3 files changed, 99 insertions(+), 6 deletions(-) create mode 100644 outputs.py diff --git a/__init__.py b/__init__.py index 7703711..107ec75 100644 --- a/__init__.py +++ b/__init__.py @@ -8,6 +8,7 @@ from .model.merge2 import StateDictMergerBlockWeightedMulti from .image.image import GridImage from .image.latenttoimage import LatentToImage, LatentToHist from .image.blend_extra import Blend2 +from .outputs import SaveText NODE_CLASS_MAPPINGS = { # latent @@ -70,4 +71,8 @@ NODE_CLASS_MAPPINGS = { ## rearrange images to single image with specified columns and gap 'GridImage': GridImage, + + # others + 'SaveText': SaveText, + } diff --git a/image/latenttoimage.py b/image/latenttoimage.py index a91eb6e..00ac067 100644 --- a/image/latenttoimage.py +++ b/image/latenttoimage.py @@ -1,4 +1,7 @@ import colorsys +from io import StringIO +import csv +from typing import Optional import torch import torchvision.transforms.functional import PIL.Image @@ -76,7 +79,7 @@ class LatentToHist: }, } - RETURN_TYPES = ('IMAGE',) + RETURN_TYPES = ('IMAGE', 'STRING') FUNCTION = 'execute' CATEGORY = 'latent' @@ -117,10 +120,12 @@ class LatentToHist: else: ymax = float(ymax) - image_tensors = list(self.plot(hists, ymax)) + sio = StringIO() + image_tensors = list(self.plot(hists, ymax, sio)) images = torch.cat(image_tensors) - return (images,) + + return (images, sio.getvalue()) def hist(self, batch_data: torch.Tensor, min_: float, max_: float, bins: int): assert batch_data.dim() == 2 @@ -138,12 +143,18 @@ class LatentToHist: return hists - def plot(self, hists: list, ymax: float): + def plot(self, hists: list, ymax: float, sio: Optional[StringIO] = None): try: from matplotlib import pyplot as plt except Exception as e: raise RuntimeError('LatentToHist requires matplotlib. Please install it.') + if sio is not None: + writer = csv.writer(sio) + writer.writerow(['ch', 'value', 'degree']) + else: + writer = None + def color(c: int, C: int): h = c / C return colorsys.hsv_to_rgb(h, 1.0, 1.0) @@ -158,14 +169,22 @@ class LatentToHist: for c, (hist, bin_edges) in enumerate(batch_hists): bin_edges = bin_edges.view((1,1,-1)) # batch,in_ch,iW W = W.to(bin_edges.device) - x = torch.nn.functional.conv1d(bin_edges, W) + x = torch.nn.functional.conv1d(bin_edges, W).squeeze() + + assert x.shape == hist.shape plot = ax.plot( - x.squeeze(), hist, color=color(c,C), + x, hist, color=color(c,C), linewidth=2, marker='o', markersize=4, ) plots.append(plot[0]) + if writer is not None: + writer.writerows([ + # ch, v, d + [c, x[i].item(), hist[i].item()] + for i in range(x.size(-1)) + ]) ax.legend(plots, [f'ch={c}' for c in range(C)], loc=2) ax.set_ylim(0, ymax) @@ -182,4 +201,8 @@ class LatentToHist: # (C, H, W) image_tensor = einops.rearrange(image_tensor, 'c h w -> h w c').unsqueeze(0) + + if sio is not None: + sio.flush() + yield image_tensor diff --git a/outputs.py b/outputs.py new file mode 100644 index 0000000..854f80c --- /dev/null +++ b/outputs.py @@ -0,0 +1,65 @@ +import os +import folder_paths + +class SaveText: + + 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': 'ComfyUI' }), + 'ext': ('STRING', { 'default': 'txt' }), + 'text': ('STRING', { 'multiline': True, 'default': '' }), + } + } + + OUTPUT_NODE = True + + RETURN_TYPES = () + + FUNCTION = 'execute' + + CATEGORY = 'utils' + + def execute(self, filename_prefix: str, ext: str, text: str): + def map_filename(filename): + prefix_len = len(os.path.basename(filename_prefix)) + prefix = filename[:prefix_len + 1] + try: + digits = int(filename[prefix_len + 1:].split('_')[0]) + except: + digits = 0 + return (digits, prefix) + + subfolder = os.path.dirname(os.path.normpath(filename_prefix)) + filename = os.path.basename(os.path.normpath(filename_prefix)) + + full_output_folder = os.path.join(self.output_dir, subfolder) + + if os.path.commonpath((self.output_dir, os.path.abspath(full_output_folder))) != self.output_dir: + print("Saving image outside the output folder is not allowed.") + return {} + + try: + counter = max(filter(lambda a: a[1][:-1] == filename and a[1][-1] == "_", map(map_filename, os.listdir(full_output_folder))))[0] + 1 + except ValueError: + counter = 1 + except FileNotFoundError: + os.makedirs(full_output_folder, exist_ok=True) + counter = 1 + + if ext is None or len(ext) == 0 or ext == '.': + ext = '.txt' + if not ext.startswith('.'): + ext = '.' + ext + + file = f"{filename}_{counter:05}_{ext}" + with open(os.path.join(full_output_folder, file), 'w') as io: + io.write(text) + counter += 1 + + return {}