init
This commit is contained in:
+12
@@ -0,0 +1,12 @@
|
||||
output/
|
||||
*__pycache__/
|
||||
samples*/
|
||||
runs/
|
||||
checkpoints/
|
||||
master_ip
|
||||
logs/
|
||||
*.DS_Store
|
||||
.idea
|
||||
tools/
|
||||
.vscode/
|
||||
convert_*
|
||||
@@ -0,0 +1,3 @@
|
||||
from .nodes import NODE_CLASS_MAPPINGS, NODE_DISPLAY_NAME_MAPPINGS
|
||||
|
||||
__all__ = ["NODE_CLASS_MAPPINGS", "NODE_DISPLAY_NAME_MAPPINGS"]
|
||||
@@ -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",
|
||||
}
|
||||
@@ -0,0 +1 @@
|
||||
diffusers
|
||||
Reference in New Issue
Block a user