caip-vitl256 / image_processing_caip.py
yuvansharma's picture
CAIP ViT-L/16 encoder (fp32)
5ffccf8 verified
Raw
History Blame Contribute Delete
2.62 kB
"""Image processor for CAIP — a bit-exact replica of the open_clip *val* transform
used at training time.
IMPORTANT: this model was trained with **OpenAI-CLIP normalization stats**
(mean=(0.481,0.458,0.408), std=(0.269,0.261,0.276)), NOT SigLIP 0.5/0.5. Using
SigLIP stats here would silently mismatch training. The pipeline is:
Resize(shortest_side -> 256, bicubic, antialias) -> CenterCrop(256)
-> convert RGB -> ToTensor([0,1]) -> Normalize(CLIP mean/std)
We use the same torchvision ops as the training transform so `pixel_values` are
identical to what the model saw.
"""
from typing import List, Union
import torch
from PIL import Image
from torchvision.transforms import CenterCrop, Compose, InterpolationMode, Normalize, Resize, ToTensor
from transformers.image_processing_utils import BaseImageProcessor, BatchFeature
_INTERP = {
"bicubic": InterpolationMode.BICUBIC,
"bilinear": InterpolationMode.BILINEAR,
"nearest": InterpolationMode.NEAREST,
}
def _convert_to_rgb(image):
return image.convert("RGB")
class CaipImageProcessor(BaseImageProcessor):
model_input_names = ["pixel_values"]
def __init__(
self,
size: int = 256,
image_mean=(0.48145466, 0.4578275, 0.40821073),
image_std=(0.26862954, 0.26130258, 0.27577711),
interpolation: str = "bicubic",
resize_mode: str = "shortest",
**kwargs,
):
super().__init__(**kwargs)
self.size = size
self.image_mean = list(image_mean)
self.image_std = list(image_std)
self.interpolation = interpolation
self.resize_mode = resize_mode
def _build_transform(self) -> Compose:
# Rebuilt per call (cheap) so the non-serializable Compose is never stored on
# self / written into preprocessor_config.json. resize_mode 'shortest' == int Resize.
interp = _INTERP[self.interpolation]
return Compose([
Resize(self.size, interpolation=interp, antialias=True),
CenterCrop(self.size),
_convert_to_rgb,
ToTensor(),
Normalize(mean=self.image_mean, std=self.image_std),
])
def preprocess(
self,
images: Union[Image.Image, List[Image.Image]],
return_tensors: str = "pt",
**kwargs,
) -> BatchFeature:
if isinstance(images, Image.Image):
images = [images]
tfm = self._build_transform()
pixel_values = torch.stack([tfm(img) for img in images]) # [B, 3, 256, 256]
return BatchFeature(data={"pixel_values": pixel_values}, tensor_type=return_tensors)