""" Trains LightFM in HYBRID mode (collaborative + content-based). Uses both the user-item interactions (reviews) and the item features (product categories + price buckets) to generate richer embeddings. Input: artifacts/reviews_cleaned.json artifacts/metadata_cleaned.json Output: artifacts/recommender_model.pkl Estimated time: 2-7 minutes. """ import json import os import sys import pickle from collections import Counter from lightfm import LightFM from lightfm.data import Dataset PROJECT_ROOT = os.path.dirname(os.path.dirname(os.path.abspath(__file__))) REVIEWS_FILE = os.path.join(PROJECT_ROOT, "artifacts", "reviews_cleaned.json") METADATA_FILE = os.path.join(PROJECT_ROOT, "artifacts", "metadata_cleaned.json") MODEL_FILE = os.path.join(PROJECT_ROOT, "artifacts", "recommender_model.pkl") # --- Feature configuration --- # Category present on 99.98% of products, zero discriminating value EXCLUDED_EXACT = {"Clothing, Shoes & Jewelry"} # Keywords to filter out promotional categories / noise NOISE_KEYWORDS = { "sale", "off", "save", "deal", "prime", "gift", "clearance", "black friday", "holiday", "test", "exclusion", "shopbop", "cyber", "top 50", "top rated", "featured", "our brands" } # Minimum number of products to consider a category MIN_CATEGORY_COUNT = 10 # Price buckets PRICE_BUCKETS = [ ("price:budget", 0, 15), ("price:affordable", 15, 30), ("price:mid", 30, 50), ("price:premium", 50, 100), ("price:luxury", 100, 10000), ] def get_price_feature(price): if price is None or price == "N/A": return None try: p = float(price) for name, lo, hi in PRICE_BUCKETS: if lo <= p < hi: return name return "price:luxury" except (ValueError, TypeError): return None def flatten_categories(categories): """Flattens the categories (they can be strings or nested lists).""" result = [] for c in categories: if isinstance(c, str): result.append(c) elif isinstance(c, list): for sub in c: if isinstance(sub, str): result.append(sub) return result def main(): print("--- TRAINING HYBRID RECOMMENDER (LightFM + Item Features) ---") # --- 1. Load reviews --- if not os.path.exists(REVIEWS_FILE): print(f"Error: {REVIEWS_FILE} not found!") sys.exit(1) raw_data = [] users_set = set() items_set = set() print("-> Reading interactions...") with open(REVIEWS_FILE, "r", encoding="utf-8") as f: for line in f: try: d = json.loads(line) u, p = d["reviewerID"], d["asin"] raw_data.append((u, p)) users_set.add(u) items_set.add(p) except Exception: continue print(f" Interactions: {len(raw_data):,}") print(f" Unique users: {len(users_set):,}") print(f" Unique products: {len(items_set):,}") if len(raw_data) == 0: print("Error: no interactions found!") sys.exit(1) # --- 2. Load metadata for the features --- if not os.path.exists(METADATA_FILE): print(f"Error: {METADATA_FILE} not found!") sys.exit(1) print("-> Loading product metadata...") with open(METADATA_FILE, "r", encoding="utf-8") as f: products = json.load(f) meta_by_asin = {p["asin"]: p for p in products} # --- 3. Category analysis and filtering --- print("-> Analyzing categories...") cat_counter = Counter() for asin in items_set: p = meta_by_asin.get(asin) if p: for c in flatten_categories(p.get("categories", [])): if c not in EXCLUDED_EXACT: cat_counter[c] += 1 allowed_categories = set() # Empty set where we will put the "good" categories excluded_noise = 0 excluded_rare = 0 for cat, count in cat_counter.items(): if count < MIN_CATEGORY_COUNT: # MIN_CATEGORY_COUNT = 10 excluded_rare += 1 continue # Skip, do not add it lower = cat.lower() if any(kw in lower for kw in NOISE_KEYWORDS): excluded_noise += 1 continue allowed_categories.add(cat) print(f" Total categories found: {len(cat_counter)}") print(f" Excluded (noise/promo): {excluded_noise}") print(f" Excluded (too rare <{MIN_CATEGORY_COUNT}): {excluded_rare}") print(f" Useful categories: {len(allowed_categories)}") # --- 4. Prepare feature names --- price_feature_names = [name for name, _, _ in PRICE_BUCKETS] all_feature_names = sorted(allowed_categories) + price_feature_names print(f" Price features: {len(price_feature_names)}") print(f" Total features: {len(all_feature_names)}") # --- 5. Build LightFM Dataset with features --- print("-> Building LightFM Dataset...") dataset = Dataset() dataset.fit( users=list(users_set), items=list(items_set), item_features=all_feature_names ) (interactions, _) = dataset.build_interactions(raw_data) # --- 6. Build item features matrix --- print("-> Building item features matrix...") item_feature_tuples = [] # List of pairs (ASIN, [feature list]) items_with_features = 0 items_with_price = 0 for asin in items_set: p = meta_by_asin.get(asin) # Look up this product's metadata inside "metadata_cleaned.json" features = [] if p: for c in flatten_categories(p.get("categories", [])): if c in allowed_categories: features.append(c) price_feat = get_price_feature(p.get("price")) if price_feat: features.append(price_feat) items_with_price += 1 if features: item_feature_tuples.append((asin, features)) items_with_features += 1 # normalize=True makes the feature weights normalized (they sum to 1 per product), # so a product with 8 categories does not carry more "weight" than one with 2. item_features_matrix = dataset.build_item_features(item_feature_tuples, normalize=True) print(f" Products with features: {items_with_features:,} / {len(items_set):,}") print(f" Products with price bucket: {items_with_price:,}") print(f" Matrix shape: {item_features_matrix.shape}") # --- 7. Hybrid training --- print("-> Training Hybrid Model (WARP, 30 epochs)...") model = LightFM(loss="warp", no_components=10) model.fit( interactions, # sparse matrix [users x products] (485k interactions) item_features=item_features_matrix, # sparse matrix [products x (products + 524 features)] epochs=30, # passes 30 times over all the data num_threads=2 # parallelizes over 2 CPU threads ) print(f" Shape item_embeddings: {model.item_embeddings.shape}") # --- 8. Saving --- print("-> Saving model...") user_id_map, _, item_id_map, _ = dataset.mapping() idx_to_item = {v: k for k, v in item_id_map.items()} package = { "model": model, "item_id_map": item_id_map, "idx_to_item": idx_to_item, "item_features_matrix": item_features_matrix, } with open(MODEL_FILE, "wb") as f: pickle.dump(package, f) print(f"\n--- DONE ---") print(f"Model saved to: {MODEL_FILE}") print(f"Type: Hybrid (collaborative + {len(allowed_categories)} categories + {len(price_feature_names)} price buckets)") if __name__ == "__main__": main()