Spaces:
Sleeping
Sleeping
Download generate_catalog.py from jashp2323/visual-search: direct link, hf CLI and curl.
- Browser
- Download file 5.83 kB
-
https://huggingface.co/spaces/jashp2323/visual-search/resolve/main/generate_catalog.py
- Command line
-
hf download hf://spaces/jashp2323/visual-search/generate_catalog.py
-
curl -L -o generate_catalog.py https://huggingface.co/spaces/jashp2323/visual-search/resolve/main/generate_catalog.py
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() | |