US_Cond-UNet / image_processing_cond_unet.py
Morelli001's picture
Upload folder using huggingface_hub
badc3e1 verified
Raw
History Blame Contribute Delete
3.2 kB
from typing import Optional, Union
import numpy as np
import torch
from PIL import Image
from torchvision.transforms import v2
from transformers.image_processing_utils import BaseImageProcessor, BatchFeature
class CondUNetImageProcessor(BaseImageProcessor):
model_input_names = ["pixel_values"]
def __init__(
self,
image_size=512,
keep_aspect_ratio=True,
self_normalize=True,
mean=None,
std=None,
**kwargs,
):
super().__init__(**kwargs)
self.image_size = image_size
self.keep_aspect_ratio = keep_aspect_ratio
self.self_normalize = self_normalize
self.mean = mean or [123.675, 116.28, 103.53]
self.std = std or [58.395, 57.12, 57.375]
def preprocess(
self,
images: Union[Image.Image, np.ndarray, torch.Tensor, list],
return_tensors: Optional[Union[str, torch.Tensor]] = None,
**kwargs,
):
if not isinstance(images, (list, tuple)):
images = [images]
pixel_values = [self._preprocess_image(image) for image in images]
return BatchFeature(
data={"pixel_values": torch.stack(pixel_values)},
tensor_type=return_tensors,
)
def _preprocess_image(self, image):
if isinstance(image, Image.Image):
image = np.array(image.convert("RGB"), copy=True)
if isinstance(image, np.ndarray):
image = torch.from_numpy(image)
if image.ndim != 3:
raise ValueError("Expected an HWC or CHW RGB image.")
if image.shape[-1] in (1, 3):
image = image.permute(2, 0, 1)
if image.shape[0] == 1:
image = image.expand(3, -1, -1)
if image.shape[0] != 3:
raise ValueError("Cond-UNet requires one or three input channels.")
height, width = image.shape[-2:]
if self.keep_aspect_ratio:
resize_factor = max(height, width) / self.image_size
new_height = int(height / resize_factor)
new_width = int(width / resize_factor)
new_height += new_height % 2
new_width += new_width % 2
image = v2.functional.resize(image, [new_height, new_width])
pad_left = (self.image_size - new_width) // 2
pad_top = (self.image_size - new_height) // 2
image = v2.functional.pad(image, fill=0, padding=[pad_left, pad_top])
else:
image = v2.functional.resize(image, [self.image_size, self.image_size])
image = image.to(dtype=torch.float32)
if image.max() <= 1:
image = image * 255.0
if self.self_normalize:
mask = (image > 0).any(dim=0)
if mask.any():
valid_pixels = image[:, mask]
mean = valid_pixels.mean()
std = valid_pixels.std()
if std > 1e-8:
return (image - mean) / std
return image - mean
return image.clone()
mean = torch.tensor(self.mean, dtype=image.dtype).view(-1, 1, 1)
std = torch.tensor(self.std, dtype=image.dtype).view(-1, 1, 1)
return (image - mean) / std