From 2eccba4e33b21d1d080cb2f415f76a93488120f0 Mon Sep 17 00:00:00 2001 From: melMass Date: Thu, 10 Aug 2023 16:34:36 +0200 Subject: [PATCH] =?UTF-8?q?fix:=20=E2=9A=A1=EF=B8=8F=20move=20getbatchfrom?= =?UTF-8?q?history=20to=20graphutils?= MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit Fixes #59 --- nodes/graph_utils.py | 107 ++++++++++++++++++++++++++++++++- nodes/image_interpolation.py | 112 +---------------------------------- 2 files changed, 107 insertions(+), 112 deletions(-) diff --git a/nodes/graph_utils.py b/nodes/graph_utils.py index 0288a99..3704a2d 100644 --- a/nodes/graph_utils.py +++ b/nodes/graph_utils.py @@ -1,4 +1,109 @@ from ..log import log +from PIL import Image +import urllib.request +import urllib.parse +import torch +import json +from comfy.cli_args import args +from ..utils import pil2tensor +import io + + +def get_image(filename, subfolder, folder_type): + data = {"filename": filename, "subfolder": subfolder, "type": folder_type} + url_values = urllib.parse.urlencode(data) + with urllib.request.urlopen( + f"http://{args.listen}:{args.port}/view?{url_values}" + ) as response: + return io.BytesIO(response.read()) + + +class GetBatchFromHistory: + """Very experimental node to load images from the history of the server. + + Queue items without output are ignored in the count.""" + + @classmethod + def INPUT_TYPES(cls): + return { + "required": { + "enable": ("BOOLEAN", {"default": True}), + "count": ("INT", {"default": 1, "min": 0}), + "offset": ("INT", {"default": 0, "min": -1e9, "max": 1e9}), + "internal_count": ("INT", {"default": 0}), + }, + "optional": { + "passthrough_image": ("IMAGE",), + }, + } + + RETURN_TYPES = ("IMAGE",) + RETURN_NAMES = "images" + CATEGORY = "mtb/animation" + FUNCTION = "load_from_history" + + def load_from_history( + self, + enable=True, + count=0, + offset=0, + internal_count=0, # hacky way to invalidate the node + passthrough_image=None, + ): + if not enable or count == 0: + if passthrough_image is not None: + log.debug("Using passthrough image") + return (passthrough_image,) + log.debug("Load from history is disabled for this iteration") + return (torch.zeros(0),) + frames = [] + + with urllib.request.urlopen( + f"http://{args.listen}:{args.port}/history" + ) as response: + return self.load_batch_frames(response, offset, count, frames) + + def load_batch_frames(self, response, offset, count, frames): + history = json.loads(response.read()) + + output_images = [] + for k, run in history.items(): + for o in run["outputs"]: + for node_id in run["outputs"]: + node_output = run["outputs"][node_id] + if "images" in node_output: + images_output = [] + for image in node_output["images"]: + image_data = get_image( + image["filename"], image["subfolder"], image["type"] + ) + images_output.append(image_data) + output_images.extend(images_output) + if not output_images: + return (torch.zeros(0),) + for i, image in enumerate(list(reversed(output_images))): + if i < offset: + continue + if i >= offset + count: + break + # Decode image as tensor + img = Image.open(image) + log.debug(f"Image from history {i} of shape {img.size}") + frames.append(img) + + # Display the shape of the tensor + # print("Tensor shape:", image_tensor.shape) + + # return (output_images,) + if not frames: + return (torch.zeros(0),) + elif len(frames) != count: + log.warning(f"Expected {count} images, got {len(frames)} instead") + output = pil2tensor( + list(reversed(frames)), + ) + + return (output,) class StringReplace: @@ -72,4 +177,4 @@ class FitNumber: return (res,) -__nodes__ = [StringReplace, FitNumber] +__nodes__ = [StringReplace, FitNumber, GetBatchFromHistory] diff --git a/nodes/image_interpolation.py b/nodes/image_interpolation.py index 8b1e013..6bd39ae 100644 --- a/nodes/image_interpolation.py +++ b/nodes/image_interpolation.py @@ -9,113 +9,8 @@ from frame_interpolation.eval import util, interpolator import numpy as np import comfy import comfy.utils -from PIL import Image -import urllib.request -import urllib.parse -import json import tensorflow as tf import comfy.model_management as model_management -import io - -from comfy.cli_args import args -from ..utils import pil2tensor - - -def get_image(filename, subfolder, folder_type): - data = {"filename": filename, "subfolder": subfolder, "type": folder_type} - url_values = urllib.parse.urlencode(data) - with urllib.request.urlopen( - f"http://{args.listen}:{args.port}/view?{url_values}" - ) as response: - return io.BytesIO(response.read()) - - -class GetBatchFromHistory: - """Very experimental node to load images from the history of the server. - - Queue items without output are ignored in the count.""" - - @classmethod - def INPUT_TYPES(cls): - return { - "required": { - "enable": ("BOOLEAN", {"default": True}), - "count": ("INT", {"default": 1, "min": 0}), - "offset": ("INT", {"default": 0, "min": -1e9, "max": 1e9}), - "internal_count": ("INT", {"default": 0}), - }, - "optional": { - "passthrough_image": ("IMAGE",), - }, - } - - RETURN_TYPES = ("IMAGE",) - RETURN_NAMES = "images" - CATEGORY = "mtb/animation" - FUNCTION = "load_from_history" - - def load_from_history( - self, - enable=True, - count=0, - offset=0, - internal_count=0, # hacky way to invalidate the node - passthrough_image=None, - ): - if not enable or count == 0: - if passthrough_image is not None: - log.debug("Using passthrough image") - return (passthrough_image,) - log.debug("Load from history is disabled for this iteration") - return (torch.zeros(0),) - frames = [] - - with urllib.request.urlopen( - f"http://{args.listen}:{args.port}/history" - ) as response: - return self.load_batch_frames(response, offset, count, frames) - - def load_batch_frames(self, response, offset, count, frames): - history = json.loads(response.read()) - - output_images = [] - for k, run in history.items(): - for o in run["outputs"]: - for node_id in run["outputs"]: - node_output = run["outputs"][node_id] - if "images" in node_output: - images_output = [] - for image in node_output["images"]: - image_data = get_image( - image["filename"], image["subfolder"], image["type"] - ) - images_output.append(image_data) - output_images.extend(images_output) - if not output_images: - return (torch.zeros(0),) - for i, image in enumerate(list(reversed(output_images))): - if i < offset: - continue - if i >= offset + count: - break - # Decode image as tensor - img = Image.open(image) - log.debug(f"Image from history {i} of shape {img.size}") - frames.append(img) - - # Display the shape of the tensor - # print("Tensor shape:", image_tensor.shape) - - # return (output_images,) - if not frames: - return (torch.zeros(0),) - elif len(frames) != count: - log.warning(f"Expected {count} images, got {len(frames)} instead") - output = pil2tensor( - list(reversed(frames)), - ) - - return (output,) class LoadFilmModel: @@ -256,9 +151,4 @@ class ConcatImages: return (self.concatenate_tensors(imageA, imageB),) -__nodes__ = [ - LoadFilmModel, - FilmInterpolation, - ConcatImages, - GetBatchFromHistory, -] +__nodes__ = [LoadFilmModel, FilmInterpolation, ConcatImages]