From 282ecc4cafc900700de260388f1ce6d3907260f9 Mon Sep 17 00:00:00 2001 From: kijai <40791699+kijai@users.noreply.github.com> Date: Mon, 24 Feb 2025 16:06:02 +0200 Subject: [PATCH] init --- .gitignore | 12 ++++ __init__.py | 3 + nodes.py | 167 +++++++++++++++++++++++++++++++++++++++++++++++ requirements.txt | 1 + 4 files changed, 183 insertions(+) create mode 100644 .gitignore create mode 100644 __init__.py create mode 100644 nodes.py create mode 100644 requirements.txt diff --git a/.gitignore b/.gitignore new file mode 100644 index 0000000..49a4514 --- /dev/null +++ b/.gitignore @@ -0,0 +1,12 @@ +output/ +*__pycache__/ +samples*/ +runs/ +checkpoints/ +master_ip +logs/ +*.DS_Store +.idea +tools/ +.vscode/ +convert_* \ No newline at end of file diff --git a/__init__.py b/__init__.py new file mode 100644 index 0000000..2e96bd6 --- /dev/null +++ b/__init__.py @@ -0,0 +1,3 @@ +from .nodes import NODE_CLASS_MAPPINGS, NODE_DISPLAY_NAME_MAPPINGS + +__all__ = ["NODE_CLASS_MAPPINGS", "NODE_DISPLAY_NAME_MAPPINGS"] \ No newline at end of file diff --git a/nodes.py b/nodes.py new file mode 100644 index 0000000..b72c8ef --- /dev/null +++ b/nodes.py @@ -0,0 +1,167 @@ +import torch +import json +import base64 +import requests +from PIL import Image +from typing import cast, List, Literal, Optional, Union + +from diffusers.image_processor import VaeImageProcessor +from diffusers.video_processor import VideoProcessor +from safetensors.torch import _tobytes + +DTYPE_MAP = { + "float16": torch.float16, + "float32": torch.float32, + "bfloat16": torch.bfloat16, + "uint8": torch.uint8, +} + +def remote_decode( + endpoint: str, + tensor: torch.Tensor, + processor: Optional[Union[VaeImageProcessor, VideoProcessor]] = None, + do_scaling: bool = True, + output_type: Literal["mp4", "pil", "pt"] = "pil", + image_format: Literal["png", "jpg"] = "jpg", + partial_postprocess: bool = False, + input_tensor_type: Literal["base64", "binary"] = "base64", + output_tensor_type: Literal["base64", "binary"] = "base64", + height: Optional[int] = None, + width: Optional[int] = None, +) -> Union[Image.Image, List[Image.Image], bytes, torch.Tensor]: + if tensor.ndim == 3 and height is None and width is None: + raise ValueError("`height` and `width` required for packed latents.") + if output_type == "pt" and partial_postprocess is True and processor is None: + raise ValueError( + "`processor` is required with `output_type='pt' and `partial_postprocess=False`." + ) + headers = {} + parameters = { + "do_scaling": do_scaling, + "output_type": output_type, + "partial_postprocess": partial_postprocess, + "shape": list(tensor.shape), + "dtype": str(tensor.dtype).split(".")[-1], + } + if height is not None and width is not None: + parameters["height"] = height + parameters["width"] = width + tensor_data = _tobytes(tensor, "tensor") + if input_tensor_type == "base64": + headers["Content-Type"] = "tensor/base64" + elif input_tensor_type == "binary": + headers["Content-Type"] = "tensor/binary" + if output_type == "pil" and image_format == "jpg" and processor is None: + headers["Accept"] = "image/jpeg" + elif output_type == "pil" and image_format == "png" and processor is None: + headers["Accept"] = "image/png" + elif (output_tensor_type == "base64" and output_type == "pt") or ( + output_tensor_type == "base64" + and output_type == "pil" + and processor is not None + ): + headers["Accept"] = "tensor/base64" + elif (output_tensor_type == "binary" and output_type == "pt") or ( + output_tensor_type == "binary" + and output_type == "pil" + and processor is not None + ): + headers["Accept"] = "tensor/binary" + elif output_type == "mp4": + headers["Accept"] = "text/plain" + if input_tensor_type == "base64": + kwargs = {"json": {"inputs": base64.b64encode(tensor_data).decode("utf-8")}} + elif input_tensor_type == "binary": + kwargs = {"data": tensor_data} + response = requests.post(endpoint, params=parameters, **kwargs, headers=headers) + if not response.ok: + raise RuntimeError(response.json()) + if output_type == "pt" or (output_type == "pil" and processor is not None): + if output_tensor_type == "base64": + content = response.json() + output_tensor = base64.b64decode(content["inputs"]) + parameters = content["parameters"] + shape = parameters["shape"] + dtype = parameters["dtype"] + elif output_tensor_type == "binary": + output_tensor = response.content + parameters = response.headers + shape = json.loads(parameters["shape"]) + dtype = parameters["dtype"] + torch_dtype = DTYPE_MAP[dtype] + output_tensor = torch.frombuffer( + bytearray(output_tensor), dtype=torch_dtype + ).reshape(shape) + if output_type == "pt": + if partial_postprocess: + output = [Image.fromarray(image.numpy()) for image in output_tensor] + if len(output) == 1: + output = output[0] + else: + if processor is None: + output = output_tensor + else: + if isinstance(processor, VideoProcessor): + output = cast( + List[Image.Image], + processor.postprocess_video(output_tensor, output_type="pil")[0], + ) + else: + output = cast( + Image.Image, + processor.postprocess(output_tensor, output_type="pil")[0], + ) + elif output_type == "mp4": + output = response.content + return output + +#region VideoDecode +class HFRemoteVAEDecode: + @classmethod + def INPUT_TYPES(s): + return {"required": { + "samples": ("LATENT",), + }, + } + + RETURN_TYPES = ("IMAGE",) + RETURN_NAMES = ("images",) + FUNCTION = "decode" + CATEGORY = "HFRemoteVae" + + def decode(self, samples): + latents = samples["samples"]#.squeeze(0).permute(1, 0, 2, 3).contiguous() + if len(latents.shape) == 5: + endpoint ="https://o7ywnmrahorts457.us-east-1.aws.endpoints.huggingface.cloud/" + else: + endpoint="https://whhx50ex1aryqvw6.us-east-1.aws.endpoints.huggingface.cloud/" + result = remote_decode( + endpoint=endpoint, + tensor=latents, + #height=latents.shape[2] * 8, + width=latents.shape[3] * 8, + processor=None, + output_type="pt", + partial_postprocess=False, + input_tensor_type="binary", + output_tensor_type="binary", + do_scaling=False + ) + if len(latents.shape) == 5: + video_processor = VideoProcessor(vae_scale_factor=8) + video_processor.config.do_resize = False + + video = video_processor.postprocess_video(video=result, output_type="pt") + out = video[0].permute(0, 2, 3, 1).cpu().float() + else: + out = result.permute(0, 2, 3, 1).cpu().float() + + return (out,) + + +NODE_CLASS_MAPPINGS = { + "HFRemoteVAEDecode": HFRemoteVAEDecode, + } +NODE_DISPLAY_NAME_MAPPINGS = { + "HFRemoteVAEDecode": "HFRemoteVAEDecode", + } diff --git a/requirements.txt b/requirements.txt new file mode 100644 index 0000000..81caff4 --- /dev/null +++ b/requirements.txt @@ -0,0 +1 @@ +diffusers \ No newline at end of file