2025-03-03 20:52:21 +00:00

32 lines
1.1 KiB
Python

import base64
import PIL
import numpy as np
from PIL import Image
from torch import Tensor
import torch
def tensor2pil(image: Tensor) -> PIL.Image.Image:
return Image.fromarray(np.clip(255. * image.cpu().numpy().squeeze(), 0, 255).astype(np.uint8))
def pil2base64(image: PIL.Image.Image) -> str:
from io import BytesIO
buffered = BytesIO()
image.save(buffered, format="JPEG")
return base64.b64encode(buffered.getvalue()).decode("utf-8")
def pil2tensor(images: Image.Image | list[Image.Image]) -> torch.Tensor:
"""Converts a PIL Image or a list of PIL Images to a tensor."""
def single_pil2tensor(image: Image.Image) -> torch.Tensor:
np_image = np.array(image).astype(np.float32) / 255.0
if np_image.ndim == 2: # Grayscale
return torch.from_numpy(np_image).unsqueeze(0) # (1, H, W)
else: # RGB or RGBA
return torch.from_numpy(np_image).unsqueeze(0) # (1, H, W, C)
if isinstance(images, Image.Image):
return single_pil2tensor(images)
else:
return torch.cat([single_pil2tensor(img) for img in images], dim=0)