"""Step 3: Build search training data with train/valid/test split. Creates (query_text, target_semantic_id) pairs from ESCI qrels, with optional PQ augmentation. Usage: uv run python experiments/exp_022_tiger_semantic_id/scripts/step3_build_search_data.py \ --sid_type esci_original uv run python experiments/exp_022_tiger_semantic_id/scripts/step3_build_search_data.py \ --sid_type esci_pq_mean --augment_pq """ import argparse import json import random import sys from collections import defaultdict from pathlib import Path sys.path.insert(0, str(Path(__file__).parent.parent.parent.parent)) DATA_DIR = Path("experiments/exp_022_tiger_semantic_id/data") ENRICHED_DIR = Path("data/enriched") QRELS_PATH = Path("data/processed/esci/qrels.jsonl") QUERIES_PATH = Path("data/processed/esci/queries.jsonl") def main(): parser = argparse.ArgumentParser() parser.add_argument("--sid_type", required=True, choices=["esci_original", "esci_pq_mean"], help="Which Semantic ID to use") parser.add_argument("--augment_pq", action="store_true", help="Add PQs as additional training queries") parser.add_argument("--min_relevance", type=int, default=2, help="Minimum relevance for positive pairs (2=Substitute, 3=Exact)") parser.add_argument("--train_ratio", type=float, default=0.8) parser.add_argument("--valid_ratio", type=float, default=0.1) parser.add_argument("--seed", type=int, default=42) args = parser.parse_args() random.seed(args.seed) # Load Semantic IDs sid_path = DATA_DIR / args.sid_type / "index_rqvae.json" print(f"Loading Semantic IDs from {sid_path}...") with open(sid_path) as f: sid_data = json.load(f) # {int_item_id: [code1, code2, code3, ...]} print(f" {len(sid_data)} items with SIDs") # Load item mapping (ASIN → int_id) mapping_path = DATA_DIR / args.sid_type / "mappings.json" with open(mapping_path) as f: mappings = json.load(f) item_to_idx = mappings["item_to_idx"] # {ASIN: int_id} idx_to_item = {v: k for k, v in item_to_idx.items()} # Load queries print("Loading queries...") queries = {} with open(QUERIES_PATH) as f: for line in f: d = json.loads(line) queries[d["id"]] = d["text"] print(f" {len(queries)} queries") # Load qrels (positive only) print(f"Loading qrels (min_relevance={args.min_relevance})...") qrels = defaultdict(list) with open(QRELS_PATH) as f: for line in f: d = json.loads(line) if d["relevance"] >= args.min_relevance: qrels[d["query_id"]].append(d["item_id"]) # Filter to queries that have at least one item with SID valid_qrels = {} for qid, item_ids in qrels.items(): valid_items = [iid for iid in item_ids if iid in item_to_idx and str(item_to_idx[iid]) in sid_data] if valid_items and qid in queries: valid_qrels[qid] = valid_items print(f" Queries with valid SIDs: {len(valid_qrels)}") total_pairs = sum(len(v) for v in valid_qrels.values()) print(f" Total (query, item) pairs: {total_pairs}") # Split queries into train/valid/test all_query_ids = sorted(valid_qrels.keys()) random.shuffle(all_query_ids) n_train = int(len(all_query_ids) * args.train_ratio) n_valid = int(len(all_query_ids) * args.valid_ratio) train_qids = set(all_query_ids[:n_train]) valid_qids = set(all_query_ids[n_train:n_train + n_valid]) test_qids = set(all_query_ids[n_train + n_valid:]) print(f" Split: train={len(train_qids)}, valid={len(valid_qids)}, test={len(test_qids)}") # Load PQs for augmentation pq_data = {} if args.augment_pq: pq_path = ENRICHED_DIR / "esci_pseudo_queries_gemma.jsonl" print(f"Loading PQs for augmentation from {pq_path}...") with open(pq_path) as f: for line in f: d = json.loads(line) pqs = d.get("pseudo_queries", []) if pqs: pq_data[d["id"]] = [pq.strip() for pq in pqs if pq.strip()] print(f" {len(pq_data)} items with PQs") # Build datasets def build_split(query_ids, split_name, include_pq_aug=False): records = [] for qid in sorted(query_ids): query_text = queries[qid] for item_id in valid_qrels[qid]: int_id = item_to_idx[item_id] semantic_id = sid_data[str(int_id)] records.append({ "query": query_text, "target_sid": semantic_id, "item_id": item_id, "source": "real_query", }) # PQ augmentation (train only) n_real = len(records) if include_pq_aug and pq_data: augmented_items = set() for qid in query_ids: for item_id in valid_qrels[qid]: augmented_items.add(item_id) for item_id in augmented_items: if item_id not in pq_data or item_id not in item_to_idx: continue int_id = item_to_idx[item_id] if str(int_id) not in sid_data: continue semantic_id = sid_data[str(int_id)] for pq in pq_data[item_id]: records.append({ "query": pq, "target_sid": semantic_id, "item_id": item_id, "source": "pq_augment", }) n_aug = len(records) - n_real print(f" {split_name}: {n_real} real + {n_aug} PQ augmented = {len(records)} total") else: print(f" {split_name}: {len(records)} pairs") return records # Build train (with optional PQ augmentation), valid, test (no augmentation) suffix = f"_{args.sid_type}" if args.augment_pq: suffix += "_aug" out_dir = DATA_DIR / "search_data" / f"{args.sid_type}{'_aug' if args.augment_pq else ''}" out_dir.mkdir(parents=True, exist_ok=True) train_data = build_split(train_qids, "train", include_pq_aug=args.augment_pq) valid_data = build_split(valid_qids, "valid", include_pq_aug=False) test_data = build_split(test_qids, "test", include_pq_aug=False) for split_name, data in [("train", train_data), ("valid", valid_data), ("test", test_data)]: out_path = out_dir / f"{split_name}.jsonl" with open(out_path, "w") as f: for r in data: f.write(json.dumps(r, ensure_ascii=False) + "\n") print(f" → {out_path}") # Save split info split_info = { "sid_type": args.sid_type, "augment_pq": args.augment_pq, "min_relevance": args.min_relevance, "train_queries": len(train_qids), "valid_queries": len(valid_qids), "test_queries": len(test_qids), "train_pairs": len(train_data), "valid_pairs": len(valid_data), "test_pairs": len(test_data), "train_real": sum(1 for r in train_data if r["source"] == "real_query"), "train_aug": sum(1 for r in train_data if r["source"] == "pq_augment"), "seed": args.seed, } info_path = out_dir / "split_info.json" with open(info_path, "w") as f: json.dump(split_info, f, indent=2) print(f" → {info_path}") # Save query split for reproducibility split_path = out_dir / "query_split.json" with open(split_path, "w") as f: json.dump({ "train": sorted(train_qids), "valid": sorted(valid_qids), "test": sorted(test_qids), }, f) print(f" → {split_path}") print(f"\n=== Summary ===") print(f" SID type: {args.sid_type}") print(f" PQ augmentation: {args.augment_pq}") print(f" Train: {len(train_data)} ({split_info['train_real']} real + {split_info['train_aug']} aug)") print(f" Valid: {len(valid_data)}") print(f" Test: {len(test_data)}") if __name__ == "__main__": main()