| """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) |
|
|
| |
| 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) |
| print(f" {len(sid_data)} items with SIDs") |
|
|
| |
| 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"] |
| idx_to_item = {v: k for k, v in item_to_idx.items()} |
|
|
| |
| 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") |
|
|
| |
| 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"]) |
|
|
| |
| 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}") |
|
|
| |
| 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)}") |
|
|
| |
| 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") |
|
|
| |
| 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", |
| }) |
|
|
| |
| 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 |
|
|
| |
| 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}") |
|
|
| |
| 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}") |
|
|
| |
| 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() |
|
|