enrich / step3_build_search_data.py
EmiliaLee's picture
Upload step3_build_search_data.py with huggingface_hub
4ae72cd verified
Raw
History Blame Contribute Delete
8.21 kB
"""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()