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()