From b2ca65c0db272e781ecd82934613320d918c9e9b Mon Sep 17 00:00:00 2001 From: qnsh Date: Tue, 14 Oct 2025 15:59:41 +0800 Subject: [PATCH] update utils --- utils/image_utils.py | 19 +++++++++++++++++-- 1 file changed, 17 insertions(+), 2 deletions(-) diff --git a/utils/image_utils.py b/utils/image_utils.py index ab7e8cd..ec81532 100644 --- a/utils/image_utils.py +++ b/utils/image_utils.py @@ -20,13 +20,28 @@ def tensor_to_base64_string( Returns: Base64 encoded string of the image. """ - pil_image = _tensor_to_pil(image_tensor, total_pixels=total_pixels) + pil_image = tensor_to_pil(image_tensor) + # 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 - +# tensor to pil +def tensor_to_pil(image: torch.Tensor) -> Image.Image: + if len(image.shape) > 3: + image = image[0] + image_np = image.cpu().numpy() + if image_np.shape[0] == 3: + image_np = image_np.transpose(1, 2, 0) + image_np = (image_np * 255).clip(0, 255).astype('uint8') + return Image.fromarray(image_np) +# pil to tensor +def pil_to_tensor(pil_image): + image_np = np.array(pil_image).astype(np.float32) / 255.0 + image_tensor = torch.from_numpy(image_np).unsqueeze(0) + return image_tensor + 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: