AestheticSigLIP / naflex.py
somepago's picture
Upload naflex.py with huggingface_hub
4fa3795 verified
Raw
History Blame
3.79 kB
"""
NaFlex preprocessing — aspect-ratio-aware patching for SigLIP 2.
Pipeline:
PIL Image → compute best (h, w) patch grid → resize → normalize → patchify
"""
import math
import torch
import torch.nn.functional as F
from PIL import Image
from torchvision.transforms.functional import to_tensor
def compute_patch_grid(
img_h: int,
img_w: int,
patch_size: int = 16,
max_patches: int = 256,
) -> tuple[int, int]:
"""Find the (h_patches, w_patches) grid that best preserves the image
aspect ratio while keeping total patches ≤ max_patches."""
aspect = img_h / img_w
best_h, best_w, best_waste = 1, 1, float("inf")
max_side = int(math.sqrt(max_patches * max(aspect, 1 / aspect))) + 1
for h in range(1, max_side + 1):
w = min(int(max_patches / h), max_side)
if w < 1:
continue
if h * w > max_patches:
w = max_patches // h
if w < 1:
continue
cur_aspect = h / w
waste = abs(math.log(cur_aspect / aspect))
used = h * w
# Prefer grids that use more patches (higher resolution)
score = waste - 0.001 * used
if score < best_waste:
best_waste = score
best_h, best_w = h, w
return best_h, best_w
def preprocess_image(
image: Image.Image,
patch_size: int = 16,
max_patches: int = 256,
mean: tuple[float, ...] = (0.5, 0.5, 0.5),
std: tuple[float, ...] = (0.5, 0.5, 0.5),
) -> tuple[torch.Tensor, tuple[int, int]]:
"""Process a single PIL image into flattened patches.
Returns
-------
patches : (num_patches, patch_dim) float32
grid : (h_patches, w_patches)
"""
image = image.convert("RGB")
img_w, img_h = image.size
h_patches, w_patches = compute_patch_grid(img_h, img_w, patch_size, max_patches)
target_h = h_patches * patch_size
target_w = w_patches * patch_size
image = image.resize((target_w, target_h), Image.BILINEAR)
# (3, H, W) float32 in [0, 1]
tensor = to_tensor(image)
# Normalize
m = torch.tensor(mean).view(3, 1, 1)
s = torch.tensor(std).view(3, 1, 1)
tensor = (tensor - m) / s
# Patchify: (3, H, W) -> (h_patches, patch_size, w_patches, patch_size, 3)
# -> (h_patches * w_patches, patch_size * patch_size * 3)
C = tensor.shape[0]
tensor = tensor.reshape(C, h_patches, patch_size, w_patches, patch_size)
tensor = tensor.permute(1, 3, 2, 4, 0) # (hp, wp, ps, ps, C)
patches = tensor.reshape(h_patches * w_patches, patch_size * patch_size * C)
return patches, (h_patches, w_patches)
def naflex_collate(batch: list[dict]) -> dict:
"""Collate function for DataLoader.
Each element in *batch* must have:
- "patches" : (num_patches, patch_dim)
- "grid" : (h_patches, w_patches)
- "score" : float (optional for inference)
Returns dict with padded tensors ready for the model.
"""
max_n = max(b["patches"].shape[0] for b in batch)
patch_dim = batch[0]["patches"].shape[1]
B = len(batch)
padded_patches = torch.zeros(B, max_n, patch_dim)
attention_mask = torch.zeros(B, max_n)
spatial_shapes = torch.zeros(B, 2, dtype=torch.long)
scores = []
for i, b in enumerate(batch):
n = b["patches"].shape[0]
padded_patches[i, :n] = b["patches"]
attention_mask[i, :n] = 1.0
spatial_shapes[i, 0] = b["grid"][0]
spatial_shapes[i, 1] = b["grid"][1]
if "score" in b:
scores.append(b["score"])
out = {
"patches": padded_patches,
"attention_mask": attention_mask,
"spatial_shapes": spatial_shapes,
}
if scores:
out["scores"] = torch.tensor(scores, dtype=torch.float32)
return out