Spaces:
Sleeping
Sleeping
File size: 2,143 Bytes
143710c | 1 2 3 4 5 6 7 8 9 10 11 12 13 14 15 16 17 18 19 20 21 22 23 24 25 26 27 28 29 30 31 32 33 34 35 36 37 38 39 40 41 42 43 44 45 46 47 48 49 50 51 52 53 54 55 56 57 58 59 60 61 62 63 | """Shared preprocessing functions.
These are pure functions used identically at training, validation, test,
and inference time. Never duplicate this logic elsewhere (e.g. inside the
API layer) -- always import from here, so train/inference never drift.
"""
from __future__ import annotations
import re
import string
from PIL import Image
from torchvision import transforms
IMAGENET_MEAN = [0.485, 0.456, 0.406]
IMAGENET_STD = [0.229, 0.224, 0.225]
_PUNCT_TABLE = str.maketrans("", "", string.punctuation)
_WHITESPACE_RE = re.compile(r"\s+")
def get_image_transform(train: bool = False) -> transforms.Compose:
"""Return the torchvision transform pipeline for ResNet-style encoders.
Same resize/crop/normalize is used for train, val, test, and inference.
`train=True` is accepted for future augmentation experiments, but the
baseline intentionally does not augment (see README: augmenting would
invalidate the cached, frozen-CNN feature strategy).
"""
return transforms.Compose(
[
transforms.Resize(256),
transforms.CenterCrop(224),
transforms.ToTensor(),
transforms.Normalize(mean=IMAGENET_MEAN, std=IMAGENET_STD),
]
)
def load_image(path: str) -> Image.Image:
"""Load an image from disk and force RGB (some Flickr8k jpgs are grayscale/CMYK)."""
return Image.open(path).convert("RGB")
def clean_caption(caption: str) -> list[str]:
"""Normalize + tokenize a raw caption string.
Steps: lowercase -> strip punctuation -> collapse whitespace -> split.
This is intentionally simple (whitespace tokenization) -- sufficient for
Flickr8k's short descriptive sentences. Swap in a real tokenizer later
if experimenting with subword/transformer decoders.
"""
caption = caption.lower().strip()
caption = caption.translate(_PUNCT_TABLE)
caption = _WHITESPACE_RE.sub(" ", caption).strip()
if not caption:
return []
return caption.split(" ")
def add_special_tokens(tokens: list[str]) -> list[str]:
"""Wrap a token list with <start> / <end>."""
return ["<start>", *tokens, "<end>"] |