mirror of
https://git.datalinker.icu/comfyanonymous/ComfyUI
synced 2026-09-04 08:07:05 +08:00
Refactor util functions (#20)
This commit is contained in:
parent
6ebad2d81b
commit
740caecda5
@ -1,7 +1,7 @@
|
|||||||
import io
|
import io
|
||||||
from inspect import cleandoc
|
from inspect import cleandoc
|
||||||
from comfy.comfy_types.node_typing import FileLocator
|
from comfy.comfy_types.node_typing import FileLocator
|
||||||
from typing import Literal
|
from typing import Literal, Optional
|
||||||
from comfy.utils import common_upscale
|
from comfy.utils import common_upscale
|
||||||
from comfy.comfy_types.node_typing import IO, ComfyNodeABC, InputTypeDict
|
from comfy.comfy_types.node_typing import IO, ComfyNodeABC, InputTypeDict
|
||||||
from comfy_api_nodes.apis import (
|
from comfy_api_nodes.apis import (
|
||||||
@ -49,6 +49,7 @@ import os
|
|||||||
import time
|
import time
|
||||||
import uuid
|
import uuid
|
||||||
import folder_paths
|
import folder_paths
|
||||||
|
from io import BytesIO
|
||||||
|
|
||||||
def downscale_input(image, total_pixels=1536*1024):
|
def downscale_input(image, total_pixels=1536*1024):
|
||||||
samples = image.movedim(-1,1)
|
samples = image.movedim(-1,1)
|
||||||
@ -121,15 +122,54 @@ def validate_aspect_ratio(aspect_ratio: str, minimum_ratio: float, maximum_ratio
|
|||||||
raise Exception(f"Aspect ratio cannot reduce to any greater than {maximum_ratio_str} ({maximum_ratio}), but was {aspect_ratio} ({calculated_ratio}).")
|
raise Exception(f"Aspect ratio cannot reduce to any greater than {maximum_ratio_str} ({maximum_ratio}), but was {aspect_ratio} ({calculated_ratio}).")
|
||||||
return aspect_ratio
|
return aspect_ratio
|
||||||
|
|
||||||
def process_image_response(response: requests.Response):
|
|
||||||
'''Uses content from a Response object and converts it to a torch.Tensor'''
|
def mimetype_to_extension(mime_type: str) -> str:
|
||||||
image = Image.open(io.BytesIO(response.content)).convert("RGBA")
|
"""Converts a MIME type to a file extension."""
|
||||||
|
return mime_type.split('/')[-1].lower()
|
||||||
|
|
||||||
|
|
||||||
|
def download_url_to_bytesio(url: str, timeout: int = None) -> BytesIO:
|
||||||
|
"""Downloads content from a URL using requests and returns it as BytesIO.
|
||||||
|
|
||||||
|
Args:
|
||||||
|
url: The URL to download.
|
||||||
|
timeout: Request timeout in seconds. Defaults to None (no timeout).
|
||||||
|
|
||||||
|
Returns:
|
||||||
|
BytesIO object containing the downloaded content.
|
||||||
|
"""
|
||||||
|
response = requests.get(url, stream=True, timeout=timeout)
|
||||||
|
response.raise_for_status() # Raises HTTPError for bad responses (4XX or 5XX)
|
||||||
|
return BytesIO(response.content)
|
||||||
|
|
||||||
|
|
||||||
|
def bytesio_to_image_tensor(image_bytesio: BytesIO, mode: str = "RGBA") -> torch.Tensor:
|
||||||
|
"""Converts image data from BytesIO to a torch.Tensor.
|
||||||
|
|
||||||
|
Args:
|
||||||
|
image_bytesio: BytesIO object containing the image data.
|
||||||
|
mode: The PIL mode to convert the image to (e.g., "RGB", "RGBA").
|
||||||
|
|
||||||
|
Returns:
|
||||||
|
A torch.Tensor representing the image (1, H, W, C).
|
||||||
|
|
||||||
|
Raises:
|
||||||
|
PIL.UnidentifiedImageError: If the image data cannot be identified.
|
||||||
|
ValueError: If the specified mode is invalid.
|
||||||
|
"""
|
||||||
|
image = Image.open(image_bytesio)
|
||||||
|
image = image.convert(mode)
|
||||||
image_array = np.array(image).astype(np.float32) / 255.0
|
image_array = np.array(image).astype(np.float32) / 255.0
|
||||||
return torch.from_numpy(image_array).unsqueeze(0)
|
return torch.from_numpy(image_array).unsqueeze(0)
|
||||||
|
|
||||||
def convert_image_to_bytesio(image: torch.Tensor, name: str=None, allow_alpha=True, total_pixels=2048*2048):
|
|
||||||
img_binary = None
|
def process_image_response(response: requests.Response):
|
||||||
# only care about first image, if it is a batch
|
'''Uses content from a Response object and converts it to a torch.Tensor'''
|
||||||
|
return bytesio_to_image_tensor(BytesIO(response.content))
|
||||||
|
|
||||||
|
|
||||||
|
def _tensor_to_pil(image: torch.Tensor, total_pixels: int = 2048*2048) -> Image.Image:
|
||||||
|
"""Converts a single torch.Tensor image [H, W, C] to a PIL Image, optionally downscaling."""
|
||||||
if len(image.shape) > 3:
|
if len(image.shape) > 3:
|
||||||
image = image[0]
|
image = image[0]
|
||||||
# TODO: remove alpha if not allowed and present
|
# TODO: remove alpha if not allowed and present
|
||||||
@ -137,13 +177,73 @@ def convert_image_to_bytesio(image: torch.Tensor, name: str=None, allow_alpha=Tr
|
|||||||
input_tensor = downscale_input(input_tensor.unsqueeze(0), total_pixels=total_pixels).squeeze()
|
input_tensor = downscale_input(input_tensor.unsqueeze(0), total_pixels=total_pixels).squeeze()
|
||||||
image_np = (input_tensor.numpy() * 255).astype(np.uint8)
|
image_np = (input_tensor.numpy() * 255).astype(np.uint8)
|
||||||
img = Image.fromarray(image_np)
|
img = Image.fromarray(image_np)
|
||||||
|
return img
|
||||||
|
|
||||||
|
|
||||||
|
def _pil_to_bytesio(img: Image.Image, mime_type: str = 'image/png') -> BytesIO:
|
||||||
|
"""Converts a PIL Image to a BytesIO object."""
|
||||||
img_byte_arr = io.BytesIO()
|
img_byte_arr = io.BytesIO()
|
||||||
img.save(img_byte_arr, format='PNG')
|
# Derive PIL format from MIME type (e.g., 'image/png' -> 'PNG')
|
||||||
|
pil_format = mime_type.split('/')[-1].upper()
|
||||||
|
if pil_format == 'JPG':
|
||||||
|
pil_format = 'JPEG'
|
||||||
|
img.save(img_byte_arr, format=pil_format)
|
||||||
img_byte_arr.seek(0)
|
img_byte_arr.seek(0)
|
||||||
img_binary = img_byte_arr
|
return img_byte_arr
|
||||||
img_binary.name = f"{name if name else uuid.uuid4()}.png"
|
|
||||||
|
|
||||||
|
def tensor_to_bytesio(image: torch.Tensor, name: Optional[str] = None, total_pixels: int = 2048*2048, mime_type: str = 'image/png') -> BytesIO:
|
||||||
|
"""Converts a torch.Tensor image to a named BytesIO object.
|
||||||
|
|
||||||
|
Args:
|
||||||
|
image: Input torch.Tensor image.
|
||||||
|
name: Optional filename for the BytesIO object.
|
||||||
|
total_pixels: Maximum total pixels for potential downscaling.
|
||||||
|
mime_type: Target image MIME type (e.g., 'image/png', 'image/jpeg', 'image/webp', 'video/mp4').
|
||||||
|
|
||||||
|
Returns:
|
||||||
|
Named BytesIO object containing the image data.
|
||||||
|
"""
|
||||||
|
pil_image = _tensor_to_pil(image, total_pixels=total_pixels)
|
||||||
|
img_binary = _pil_to_bytesio(pil_image, mime_type=mime_type)
|
||||||
|
img_binary.name = f"{name if name else uuid.uuid4()}.{mimetype_to_extension(mime_type)}"
|
||||||
return img_binary
|
return img_binary
|
||||||
|
|
||||||
|
|
||||||
|
def tensor_to_base64_string(image_tensor: torch.Tensor, total_pixels: int = 2048*2048, mime_type: str = 'image/png') -> str:
|
||||||
|
"""Convert [B, H, W, C] or [H, W, C] tensor to a base64 string.
|
||||||
|
|
||||||
|
Args:
|
||||||
|
image_tensor: Input torch.Tensor image.
|
||||||
|
total_pixels: Maximum total pixels for potential downscaling.
|
||||||
|
mime_type: Target image MIME type (e.g., 'image/png', 'image/jpeg', 'image/webp', 'video/mp4').
|
||||||
|
|
||||||
|
Returns:
|
||||||
|
Base64 encoded string of the image.
|
||||||
|
"""
|
||||||
|
pil_image = _tensor_to_pil(image_tensor, total_pixels=total_pixels)
|
||||||
|
img_byte_arr = _pil_to_bytesio(pil_image, mime_type=mime_type)
|
||||||
|
img_bytes = img_byte_arr.getvalue()
|
||||||
|
# Encode bytes to base64 string
|
||||||
|
base64_encoded_string = base64.b64encode(img_bytes).decode("utf-8")
|
||||||
|
return base64_encoded_string
|
||||||
|
|
||||||
|
|
||||||
|
def tensor_to_data_uri(image_tensor: torch.Tensor, total_pixels: int = 2048 * 2048, mime_type: str = 'image/png') -> str:
|
||||||
|
"""Converts a tensor image to a Data URI string.
|
||||||
|
|
||||||
|
Args:
|
||||||
|
image_tensor: Input torch.Tensor image.
|
||||||
|
total_pixels: Maximum total pixels for potential downscaling.
|
||||||
|
mime_type: Target image MIME type (e.g., 'image/png', 'image/jpeg', 'image/webp').
|
||||||
|
|
||||||
|
Returns:
|
||||||
|
Data URI string (e.g., 'data:image/png;base64,...').
|
||||||
|
"""
|
||||||
|
base64_string = tensor_to_base64_string(image_tensor, total_pixels, mime_type)
|
||||||
|
return f"data:{mime_type};base64,{base64_string}"
|
||||||
|
|
||||||
|
|
||||||
def upload_images_to_comfyapi(image: torch.Tensor, max_images=8, auth_token=None) -> list[str]:
|
def upload_images_to_comfyapi(image: torch.Tensor, max_images=8, auth_token=None) -> list[str]:
|
||||||
# if batch, try to upload each file if max_images is greater than 0
|
# if batch, try to upload each file if max_images is greater than 0
|
||||||
idx_image = 0
|
idx_image = 0
|
||||||
@ -157,7 +257,7 @@ def upload_images_to_comfyapi(image: torch.Tensor, max_images=8, auth_token=None
|
|||||||
if len(image.shape) > 3:
|
if len(image.shape) > 3:
|
||||||
curr_image = image[idx_image]
|
curr_image = image[idx_image]
|
||||||
# get BytesIO version of image
|
# get BytesIO version of image
|
||||||
img_binary = convert_image_to_bytesio(curr_image)
|
img_binary = tensor_to_bytesio(curr_image)
|
||||||
# first, request upload/download urls from comfy API
|
# first, request upload/download urls from comfy API
|
||||||
operation = SynchronousOperation(
|
operation = SynchronousOperation(
|
||||||
endpoint=ApiEndpoint(
|
endpoint=ApiEndpoint(
|
||||||
|
|||||||
Loading…
x
Reference in New Issue
Block a user