visual-search / generate_catalog.py
jashp2323's picture
Initial commit: local neural visual search engine with high-res fashion products
3d21552
Raw History Blame Contribute Delete
5.83 kB
import os
import shutil
import random
from PIL import Image
from datasets import load_dataset
from models import CLIPEmbedder
from vector_store import VectorStore
def main():
print("Clearing old catalog database and image files...")
if os.path.exists("catalog.json"):
os.remove("catalog.json")
if os.path.exists("embeddings.npy"):
os.remove("embeddings.npy")
if os.path.exists("static/images"):
shutil.rmtree("static/images")
os.makedirs("static/images", exist_ok=True)
print("Loading CLIP embedder (this may take a few seconds)...")
embedder = CLIPEmbedder()
store = VectorStore()
print("Streaming metadata from ashraq/fashion-product-images-small...")
try:
ashraq_ds = load_dataset("ashraq/fashion-product-images-small", split="train", streaming=True)
metadata_lookup = {}
# Pre-load metadata for matching
print("Gathering metadata fields...")
for item in ashraq_ds:
if len(metadata_lookup) >= 2500:
break
item_id = str(item["id"])
metadata_lookup[item_id] = {
"title": item.get("productDisplayName", "Fashion Item"),
"articleType": item.get("articleType", ""),
"baseColour": item.get("baseColour", ""),
"gender": item.get("gender", "Unisex"),
"usage": item.get("usage", "Casual"),
"season": item.get("season", "All-Season"),
"subCategory": item.get("subCategory", "")
}
except Exception as e:
print(f"Error loading metadata dataset: {e}")
return
print("Streaming high-resolution images from ceyda/fashion-products-small...")
try:
ceyda_ds = load_dataset("ceyda/fashion-products-small", split="train", streaming=True)
except Exception as e:
print(f"Error loading high-res images dataset: {e}")
return
# Select a balanced set of 100 products
categories_count = {
"Apparel": 0,
"Footwear": 0,
"Accessories": 0,
"Bags": 0,
"Other": 0
}
target_per_category = 25
total_target = 100
selected_products = []
print("Selecting diverse high-resolution products...")
for item in ceyda_ds:
if len(selected_products) >= total_target:
break
item_id = str(item["id"])
if item_id not in metadata_lookup:
continue
master_cat = item.get("masterCategory")
meta = metadata_lookup[item_id]
sub_cat = meta.get("subCategory")
# Determine target category
category = master_cat
if sub_cat and "bag" in sub_cat.lower():
category = "Bags"
elif category not in ["Apparel", "Footwear", "Accessories"]:
category = "Other"
# Limit count per category to ensure diversity
if categories_count.get(category, 0) < target_per_category:
# Validate that there is a valid PIL image
if item.get("image") is not None:
categories_count[category] += 1
selected_products.append((item, meta, category))
print(f"Selected {len(selected_products)} high-resolution products:")
for cat, count in categories_count.items():
print(f" - {cat}: {count}")
print("\nProcessing images and computing embeddings...")
for idx, (item, meta, category) in enumerate(selected_products, 1):
prod_id = f"prod-{item['id']}"
title = meta["title"]
image = item.get("image")
# Save image locally
image_path = f"static/images/{prod_id}.png"
image.save(image_path, "PNG")
# Generate a realistic price based on item ID seed
random.seed(int(item['id']))
sub_cat = meta.get("subCategory")
if category == "Apparel":
price = round(random.uniform(25.0, 95.0), 2)
elif category == "Footwear":
price = round(random.uniform(50.0, 150.0), 2)
elif category == "Accessories":
if sub_cat == "Watches":
price = round(random.uniform(90.0, 300.0), 2)
else:
price = round(random.uniform(30.0, 120.0), 2)
elif category == "Bags":
price = round(random.uniform(35.0, 110.0), 2)
else:
price = round(random.uniform(15.0, 50.0), 2)
# Build clean description
colour = meta.get("baseColour", "")
article_type = meta.get("articleType", "")
gender = meta.get("gender", "Unisex")
usage = meta.get("usage", "Casual")
season = meta.get("season", "All-Season")
desc_parts = []
if colour:
desc_parts.append(colour)
if article_type:
desc_parts.append(article_type)
if gender:
desc_parts.append(f"for {gender}")
desc = " ".join(desc_parts)
if usage and usage != "NaN":
desc += f". Designed for {usage.lower()} wear"
if season and season != "NaN":
desc += f" in {season.lower()}."
else:
desc += "."
# Compute CLIP embedding
print(f"[{idx}/{total_target}] Embedding: {title} ({category})")
try:
emb = embedder.get_image_embedding(image)
prod_data = {
"id": prod_id,
"title": title,
"desc": desc,
"price": price,
"category": category,
"image_url": f"/static/images/{prod_id}.png"
}
store.add_product(prod_data, emb)
except Exception as e:
print(f" Error processing {prod_id}: {e}")
print("\nCatalog precomputation complete!")
if __name__ == "__main__":
main()