← back to Handbag Auth Nextjs
python-matcher/models/handbag_embedder.py
83 lines
"""
Handbag Image Embedder using CLIP
Converts handbag images to 768-dimensional vectors for similarity search
"""
import torch
from PIL import Image
from transformers import CLIPProcessor, CLIPModel
# Use GPU if available
device = "cuda" if torch.cuda.is_available() else "cpu"
print(f"Using device: {device}")
# Load CLIP model (768 dimensions)
model = CLIPModel.from_pretrained("openai/clip-vit-large-patch14").to(device)
processor = CLIPProcessor.from_pretrained("openai/clip-vit-large-patch14")
print("✅ CLIP model loaded successfully")
@torch.inference_mode()
def image_to_vec(image: Image.Image):
"""
Convert PIL Image to normalized 768-dim vector
Args:
image: PIL Image in RGB mode
Returns:
list[float]: Normalized 768-dimensional vector
"""
# Ensure RGB mode
if image.mode != "RGB":
image = image.convert("RGB")
# Process image
inputs = processor(images=image, return_tensors="pt").to(device)
outputs = model.get_image_features(**inputs)
# L2 normalize
vec = outputs[0] / outputs[0].norm()
# Return as list for pgvector
return vec.cpu().numpy().tolist()
def embed_batch(images: list[Image.Image], batch_size: int = 8):
"""
Embed multiple images in batches for efficiency
Args:
images: List of PIL Images
batch_size: Number of images to process at once
Returns:
list[list[float]]: List of normalized vectors
"""
vectors = []
for i in range(0, len(images), batch_size):
batch = images[i:i + batch_size]
# Ensure all RGB
batch = [img.convert("RGB") if img.mode != "RGB" else img for img in batch]
# Process batch
inputs = processor(images=batch, return_tensors="pt").to(device)
outputs = model.get_image_features(**inputs)
# Normalize each vector
for vec in outputs:
normalized = vec / vec.norm()
vectors.append(normalized.cpu().numpy().tolist())
return vectors
if __name__ == "__main__":
# Test the embedder
print("Testing embedder with sample image...")
test_img = Image.new("RGB", (224, 224), color="red")
vec = image_to_vec(test_img)
print(f"✅ Generated vector with {len(vec)} dimensions")
print(f"Sample values: {vec[:5]}")