ProCreations's picture
download
raw
26 kB
# #!/usr/bin/env python3
# # -*- coding: utf-8 -*-
# import argparse
# import json
# import math
# import os
# import random
# from datetime import datetime
# import numpy as np
# import torch
# from tqdm.auto import tqdm
# from agent_rec.cli_common import add_shared_training_args
# from agent_rec.config import EVAL_TOPK, POS_TOPK, POS_TOPK_BY_PART, TFIDF_MAX_FEATURES
# from agent_rec.data import build_training_pairs, stratified_train_valid_split
# from agent_rec.features import (
# build_feature_cache,
# build_unified_corpora,
# feature_cache_exists,
# build_agent_content_view,
# save_feature_cache,
# load_feature_cache,
# save_vectorizers,
# load_vectorizers,
# UNK_TOOL_TOKEN,
# UNK_LLM_TOKEN,
# build_agent_tool_id_buffers,
# )
# from agent_rec.eval import evaluate_sampled_direct_top10, split_eval_qids_by_part
# from agent_rec.models.dnn import SimpleBPRDNN, bpr_loss
# from agent_rec.run_common import bootstrap_run, cache_key_from_meta, load_or_build_training_cache, shared_cache_dir
# from utils import print_metrics_table
# def main():
# parser = argparse.ArgumentParser()
# add_shared_training_args(
# parser,
# exp_name_default="bpr_dnn",
# device_default="cpu",
# epochs_default=5,
# batch_size_default=1024,
# lr_default=1e-3,
# )
# parser.add_argument("--text_hidden", type=int, default=256)
# parser.add_argument("--id_dim", type=int, default=32)
# parser.add_argument("--max_features", type=int, default=TFIDF_MAX_FEATURES)
# parser.add_argument("--rebuild_feature_cache", type=int, default=0)
# parser.add_argument("--use_query_id_emb", type=int, default=0, help="1 to add optional query-ID embedding")
# parser.add_argument("--use_agent_id_emb", type=int, default=0, help="1 to add learnable per-agent ID embedding")
# parser.add_argument("--use_llm_id_emb", type=int, default=1)
# parser.add_argument("--use_tool_id_emb", type=int, default=1)
# parser.add_argument("--use_model_content_vector", type=int, default=1)
# parser.add_argument("--use_tool_content_vector", type=int, default=1)
# args = parser.parse_args()
# boot = bootstrap_run(
# data_root=args.data_root,
# exp_name=args.exp_name,
# topk=args.topk,
# with_tools=True,
# )
# bundle = boot.bundle
# tools = boot.tools
# all_agents = bundle.all_agents
# all_questions = bundle.all_questions
# all_rankings = bundle.all_rankings
# qid_to_part = bundle.qid_to_part
# tool_names = list(tools.keys())
# q_ids = boot.q_ids
# a_ids = boot.a_ids
# qid2idx = boot.qid2idx
# aid2idx = boot.aid2idx
# qids_in_rank = boot.qids_in_rank
# data_sig = boot.data_sig
# exp_cache_dir = boot.exp_cache_dir
# feature_cache_dir = shared_cache_dir(
# args.data_root,
# "features",
# f"tfidf_{args.max_features}_{data_sig}",
# )
# if feature_cache_exists(feature_cache_dir) and args.rebuild_feature_cache == 0:
# # IMPORTANT: load both the dense feature cache and the vectorizers.
# # The previous version only loaded vectorizers, so `feature_cache` was
# # undefined when the cache already existed.
# feature_cache = load_feature_cache(feature_cache_dir)
# vecs = load_vectorizers(feature_cache_dir)
# if vecs is None:
# raise RuntimeError(
# f"[cache] feature cache exists but vectorizers are missing in {feature_cache_dir}. "
# f"Please rebuild with --rebuild_feature_cache 1."
# )
# q_vectorizer_runtime = vecs.q_vec
# print(f"[cache] loaded features from {feature_cache_dir}")
# else:
# feature_cache, vecs = build_feature_cache(
# all_agents, all_questions, tools, max_features=args.max_features
# )
# # Save BOTH feature arrays and vectorizers. Without save_feature_cache(),
# # future runs may see vectorizers but have no feature_cache object to load.
# os.makedirs(feature_cache_dir, exist_ok=True)
# save_feature_cache(feature_cache_dir, feature_cache)
# save_vectorizers(feature_cache_dir, vecs) # writes q/model/tool pkl files
# q_vectorizer_runtime = vecs.q_vec
# print(f"[cache] rebuilt & saved features to {feature_cache_dir}")
# # Sanity check: rows in Q/A feature matrices must match the id maps used by
# # training pairs and evaluation candidate indices.
# if list(feature_cache.q_ids) != list(q_ids):
# raise RuntimeError(
# "[cache] q_ids in feature_cache do not match bootstrap q_ids. "
# "Please rebuild with --rebuild_feature_cache 1."
# )
# if list(feature_cache.a_ids) != list(a_ids):
# raise RuntimeError(
# "[cache] a_ids in feature_cache do not match bootstrap a_ids. "
# "Please rebuild with --rebuild_feature_cache 1."
# )
# Q_np = feature_cache.Q.astype(np.float32)
# A_text_full_np = build_agent_content_view(
# cache=feature_cache,
# use_model_content_vector=bool(args.use_model_content_vector),
# use_tool_content_vector=bool(args.use_tool_content_vector),
# )
# tool_ids_np = feature_cache.agent_tool_idx_padded
# tool_mask_np = feature_cache.agent_tool_mask
# llm_idx_np = feature_cache.agent_llm_idx
# want_meta = {
# "data_sig": data_sig,
# "pos_topk_by_part": POS_TOPK_BY_PART,
# "neg_per_pos": int(args.neg_per_pos),
# "rng_seed_pairs": int(args.rng_seed_pairs),
# "split_seed": int(args.split_seed),
# "valid_ratio": float(args.valid_ratio),
# }
# training_cache_dir = shared_cache_dir(args.data_root, "training", f"{data_sig}_{cache_key_from_meta(want_meta)}")
# def build_cache():
# train_qids, valid_qids = stratified_train_valid_split(
# qids_in_rank, qid_to_part=qid_to_part, valid_ratio=args.valid_ratio, seed=args.split_seed
# )
# print(f"[split] train={len(train_qids)} valid={len(valid_qids)}")
# rankings_train = {qid: all_rankings[qid] for qid in train_qids}
# pairs = build_training_pairs(
# rankings_train,
# a_ids,
# qid_to_part=qid_to_part,
# pos_topk_by_part=POS_TOPK_BY_PART,
# pos_topk_default=POS_TOPK,
# neg_per_pos=args.neg_per_pos,
# rng_seed=args.rng_seed_pairs,
# )
# pairs_idx = [(qid2idx[q], aid2idx[p], aid2idx[n]) for (q, p, n) in pairs]
# pairs_idx_np = np.array(pairs_idx, dtype=np.int64)
# return train_qids, valid_qids, pairs_idx_np
# train_qids, valid_qids, pairs_idx_np = load_or_build_training_cache(
# training_cache_dir,
# args.rebuild_training_cache,
# want_meta,
# build_cache,
# )
# device = torch.device(args.device)
# model = SimpleBPRDNN(
# d_q=int(Q_np.shape[1]),
# d_a=int(A_text_full_np.shape[1]),
# num_tools=int(len(feature_cache.tool_id_vocab)),
# num_llm_ids=int(len(feature_cache.llm_vocab)),
# agent_tool_indices_padded=torch.tensor(tool_ids_np, dtype=torch.long, device=device),
# agent_tool_mask=torch.tensor(tool_mask_np, dtype=torch.float32, device=device),
# agent_llm_idx=torch.tensor(feature_cache.agent_llm_idx, dtype=torch.long, device=device),
# text_hidden=args.text_hidden,
# id_dim=args.id_dim,
# num_queries=len(q_ids),
# num_agents=len(a_ids),
# use_query_id_emb=bool(args.use_query_id_emb),
# use_agent_id_emb=bool(args.use_agent_id_emb),
# use_tool_id_emb=bool(args.use_tool_id_emb),
# use_llm_id_emb=bool(args.use_llm_id_emb),
# ).to(device)
# optimizer = torch.optim.Adam(model.parameters(), lr=args.lr)
# pairs = pairs_idx_np.tolist()
# num_pairs = len(pairs)
# num_batches = math.ceil(num_pairs / args.batch_size)
# print(f"Training pairs: {num_pairs}, batches/epoch: {num_batches}")
# Q_t = torch.tensor(Q_np, dtype=torch.float32, device=device)
# A_t = torch.tensor(A_text_full_np, dtype=torch.float32, device=device)
# for epoch in range(1, args.epochs + 1):
# random.shuffle(pairs)
# total_loss = 0.0
# pbar = tqdm(range(num_batches), desc=f"Epoch {epoch}/{args.epochs}", leave=True, dynamic_ncols=True)
# model.train()
# for b in pbar:
# batch = pairs[b * args.batch_size:(b + 1) * args.batch_size]
# if not batch:
# continue
# q_idx = torch.tensor([t[0] for t in batch], dtype=torch.long, device=device)
# pos_idx = torch.tensor([t[1] for t in batch], dtype=torch.long, device=device)
# neg_idx = torch.tensor([t[2] for t in batch], dtype=torch.long, device=device)
# q_vec = Q_t[q_idx]
# pos_vec = A_t[pos_idx]
# neg_vec = A_t[neg_idx]
# pos, neg = model(q_vec, pos_vec, neg_vec, pos_idx, neg_idx, q_idx=q_idx)
# loss = bpr_loss(pos, neg)
# optimizer.zero_grad()
# loss.backward()
# optimizer.step()
# total_loss += float(loss.item())
# pbar.set_postfix({"batch_loss": f"{loss.item():.4f}", "avg_loss": f"{(total_loss / (b + 1)):.4f}"})
# print(f"Epoch {epoch}/{args.epochs} - BPR loss: {(total_loss / num_batches if num_batches else 0.0):.4f}")
# model_dir = os.path.join(exp_cache_dir, "models")
# os.makedirs(model_dir, exist_ok=True)
# data_sig = want_meta["data_sig"]
# ckpt_path = os.path.join(model_dir, f"{args.exp_name}_{data_sig}.pt")
# meta_path = os.path.join(model_dir, f"meta_{args.exp_name}_{data_sig}.json")
# ckpt = {
# "state_dict": model.state_dict(),
# "data_sig": data_sig,
# "saved_at": datetime.now().isoformat(timespec="seconds"),
# "dims": {
# "d_q": int(Q_np.shape[1]),
# "d_a": int(A_text_full_np.shape[1]),
# "num_agents": len(a_ids),
# "num_tools": int(len(feature_cache.tool_id_vocab)),
# "text_hidden": args.text_hidden,
# "id_dim": args.id_dim,
# },
# "flags": {
# "use_llm_id_emb": bool(args.use_llm_id_emb),
# "use_tool_id_emb": bool(args.use_tool_id_emb),
# "use_model_content_vector": bool(args.use_model_content_vector),
# "use_tool_content_vector": bool(args.use_tool_content_vector),
# "use_query_id_emb": bool(args.use_query_id_emb),
# "use_agent_id_emb": bool(args.use_agent_id_emb),
# },
# "mappings": {"q_ids": q_ids, "a_ids": a_ids, "tool_names": tool_names},
# "args": vars(args),
# }
# torch.save(ckpt, ckpt_path)
# with open(meta_path, "w", encoding="utf-8") as f:
# json.dump({"data_sig": data_sig, "q_ids": q_ids, "a_ids": a_ids}, f, ensure_ascii=False, indent=2)
# print(f"[save] model -> {ckpt_path}")
# print(f"[save] meta -> {meta_path}")
# model.eval()
# topk = int(args.topk)
# part_splits = split_eval_qids_by_part(valid_qids, qid_to_part=qid_to_part)
# for part in ["PartI", "PartII", "PartIII"]:
# qids_part = part_splits.get(part, [])
# if not qids_part:
# continue
# m_part = evaluate_sampled_direct_top10(
# model=model,
# aid2idx=aid2idx,
# qid2idx=qid2idx,
# all_rankings=all_rankings,
# all_questions=all_questions,
# eval_qids=qids_part,
# q_vectorizer=q_vectorizer_runtime,
# A_text_full=A_text_full_np,
# cand_size=args.eval_cand_size,
# qid_to_part=qid_to_part,
# pos_topk_by_part=POS_TOPK_BY_PART,
# pos_topk_default=POS_TOPK,
# topk=topk,
# desc=f"Valid {part} (direct q-vector, top{topk})",
# )
# print_metrics_table(f"Validation {part} (direct q-vector)", m_part, ks=(topk,), filename=args.exp_name)
# if __name__ == "__main__":
# main()
#!/usr/bin/env python3
# -*- coding: utf-8 -*-
import argparse
import json
import math
import os
import random
from datetime import datetime
import numpy as np
import torch
from tqdm.auto import tqdm
from agent_rec.cli_common import add_shared_training_args
from agent_rec.config import EVAL_TOPK, POS_TOPK, POS_TOPK_BY_PART, TFIDF_MAX_FEATURES
from agent_rec.data import build_training_pairs, stratified_train_valid_split
from agent_rec.features import (
build_feature_cache,
build_unified_corpora,
feature_cache_exists,
build_agent_content_view,
save_feature_cache,
load_feature_cache,
save_vectorizers,
load_vectorizers,
UNK_TOOL_TOKEN,
UNK_LLM_TOKEN,
build_agent_tool_id_buffers,
)
from agent_rec.eval import evaluate_sampled_direct_top10, split_eval_qids_by_part
from agent_rec.models.dnn import SimpleBPRDNN, bpr_loss
from agent_rec.run_common import bootstrap_run, cache_key_from_meta, load_or_build_training_cache, shared_cache_dir
from utils import print_metrics_table
def main():
parser = argparse.ArgumentParser()
add_shared_training_args(
parser,
exp_name_default="bpr_dnn",
device_default="cpu",
epochs_default=5,
batch_size_default=1024,
lr_default=1e-3,
)
parser.add_argument("--text_hidden", type=int, default=256)
parser.add_argument("--id_dim", type=int, default=32)
parser.add_argument("--max_features", type=int, default=TFIDF_MAX_FEATURES)
parser.add_argument("--rebuild_feature_cache", type=int, default=0)
parser.add_argument("--use_query_id_emb", type=int, default=0, help="1 to add optional query-ID embedding")
parser.add_argument("--use_agent_id_emb", type=int, default=0, help="1 to add learnable per-agent ID embedding")
parser.add_argument("--use_llm_id_emb", type=int, default=1)
parser.add_argument("--use_tool_id_emb", type=int, default=1)
parser.add_argument("--use_model_content_vector", type=int, default=1)
parser.add_argument("--use_tool_content_vector", type=int, default=1)
args = parser.parse_args()
boot = bootstrap_run(
data_root=args.data_root,
exp_name=args.exp_name,
topk=args.topk,
with_tools=True,
)
bundle = boot.bundle
tools = boot.tools
all_agents = bundle.all_agents
all_questions = bundle.all_questions
all_rankings = bundle.all_rankings
qid_to_part = bundle.qid_to_part
tool_names = list(tools.keys())
q_ids = boot.q_ids
a_ids = boot.a_ids
qid2idx = boot.qid2idx
aid2idx = boot.aid2idx
qids_in_rank = boot.qids_in_rank
data_sig = boot.data_sig
exp_cache_dir = boot.exp_cache_dir
feature_cache_dir = shared_cache_dir(
args.data_root,
"features",
f"tfidf_{args.max_features}_{data_sig}",
)
if feature_cache_exists(feature_cache_dir) and args.rebuild_feature_cache == 0:
# IMPORTANT: load both the dense feature cache and the vectorizers.
# The previous version only loaded vectorizers, so `feature_cache` was
# undefined when the cache already existed.
feature_cache = load_feature_cache(feature_cache_dir)
vecs = load_vectorizers(feature_cache_dir)
if vecs is None:
raise RuntimeError(
f"[cache] feature cache exists but vectorizers are missing in {feature_cache_dir}. "
f"Please rebuild with --rebuild_feature_cache 1."
)
q_vectorizer_runtime = vecs.q_vec
print(f"[cache] loaded features from {feature_cache_dir}")
else:
feature_cache, vecs = build_feature_cache(
all_agents, all_questions, tools, max_features=args.max_features
)
# Save BOTH feature arrays and vectorizers. Without save_feature_cache(),
# future runs may see vectorizers but have no feature_cache object to load.
os.makedirs(feature_cache_dir, exist_ok=True)
save_feature_cache(feature_cache_dir, feature_cache)
save_vectorizers(feature_cache_dir, vecs) # writes q/model/tool pkl files
q_vectorizer_runtime = vecs.q_vec
print(f"[cache] rebuilt & saved features to {feature_cache_dir}")
# Sanity check: rows in Q/A feature matrices must match the id maps used by
# training pairs and evaluation candidate indices.
if list(feature_cache.q_ids) != list(q_ids):
raise RuntimeError(
"[cache] q_ids in feature_cache do not match bootstrap q_ids. "
"Please rebuild with --rebuild_feature_cache 1."
)
if list(feature_cache.a_ids) != list(a_ids):
raise RuntimeError(
"[cache] a_ids in feature_cache do not match bootstrap a_ids. "
"Please rebuild with --rebuild_feature_cache 1."
)
Q_np = feature_cache.Q.astype(np.float32)
A_text_full_np = build_agent_content_view(
cache=feature_cache,
use_model_content_vector=bool(args.use_model_content_vector),
use_tool_content_vector=bool(args.use_tool_content_vector),
)
tool_ids_np = feature_cache.agent_tool_idx_padded
tool_mask_np = feature_cache.agent_tool_mask
# Align with the TF-IDF/Table-5 style flow:
# 1) --train_parts controls which question parts can enter the train/valid split.
# 2) --eval_parts controls final reporting independently of the train valid split.
# 3) train_parts is part of the training-cache key, so stale PartI+II+III
# pair caches cannot be silently reused for PartIII-only runs.
train_parts = list(args.train_parts)
eval_parts = list(args.eval_parts)
train_part_set = set(train_parts)
qids_for_training = [
qid for qid in qids_in_rank
if qid_to_part.get(qid, "") in train_part_set
]
if not qids_for_training:
raise RuntimeError(
f"No qids found for train_parts={train_parts}. "
f"Available parts include: {sorted(set(qid_to_part.values()))}"
)
print(
f"[parts] train_parts={train_parts} -> qids={len(qids_for_training)}; "
f"eval_parts={eval_parts}"
)
want_meta = {
"data_sig": data_sig,
"pos_topk_by_part": POS_TOPK_BY_PART,
"neg_per_pos": int(args.neg_per_pos),
"rng_seed_pairs": int(args.rng_seed_pairs),
"split_seed": int(args.split_seed),
"valid_ratio": float(args.valid_ratio),
"pair_type": "q_pos_neg_posTopK",
"train_parts": train_parts,
}
training_cache_dir = shared_cache_dir(args.data_root, "training", f"{data_sig}_{cache_key_from_meta(want_meta)}")
def build_cache():
train_qids, valid_qids = stratified_train_valid_split(
qids_for_training,
qid_to_part=qid_to_part,
valid_ratio=args.valid_ratio,
seed=args.split_seed,
)
print(
f"[split] train_parts={train_parts} "
f"train={len(train_qids)} valid={len(valid_qids)}"
)
rankings_train = {qid: all_rankings[qid] for qid in train_qids}
pairs = build_training_pairs(
rankings_train,
a_ids,
qid_to_part=qid_to_part,
pos_topk_by_part=POS_TOPK_BY_PART,
pos_topk_default=POS_TOPK,
neg_per_pos=args.neg_per_pos,
rng_seed=args.rng_seed_pairs,
)
pairs_idx = [(qid2idx[q], aid2idx[p], aid2idx[n]) for (q, p, n) in pairs]
pairs_idx_np = np.array(pairs_idx, dtype=np.int64)
return train_qids, valid_qids, pairs_idx_np
train_qids, valid_qids, pairs_idx_np = load_or_build_training_cache(
training_cache_dir,
args.rebuild_training_cache,
want_meta,
build_cache,
)
device = torch.device(args.device)
model = SimpleBPRDNN(
d_q=int(Q_np.shape[1]),
d_a=int(A_text_full_np.shape[1]),
num_tools=int(len(feature_cache.tool_id_vocab)),
num_llm_ids=int(len(feature_cache.llm_vocab)),
agent_tool_indices_padded=torch.tensor(tool_ids_np, dtype=torch.long, device=device),
agent_tool_mask=torch.tensor(tool_mask_np, dtype=torch.float32, device=device),
agent_llm_idx=torch.tensor(feature_cache.agent_llm_idx, dtype=torch.long, device=device),
text_hidden=args.text_hidden,
id_dim=args.id_dim,
num_queries=len(q_ids),
num_agents=len(a_ids),
use_query_id_emb=bool(args.use_query_id_emb),
use_agent_id_emb=bool(args.use_agent_id_emb),
use_tool_id_emb=bool(args.use_tool_id_emb),
use_llm_id_emb=bool(args.use_llm_id_emb),
).to(device)
optimizer = torch.optim.Adam(model.parameters(), lr=args.lr)
pairs = pairs_idx_np.tolist()
num_pairs = len(pairs)
num_batches = math.ceil(num_pairs / args.batch_size)
print(f"Training pairs: {num_pairs}, batches/epoch: {num_batches}")
Q_t = torch.tensor(Q_np, dtype=torch.float32, device=device)
A_t = torch.tensor(A_text_full_np, dtype=torch.float32, device=device)
for epoch in range(1, args.epochs + 1):
random.shuffle(pairs)
total_loss = 0.0
pbar = tqdm(range(num_batches), desc=f"Epoch {epoch}/{args.epochs}", leave=True, dynamic_ncols=True)
model.train()
for b in pbar:
batch = pairs[b * args.batch_size:(b + 1) * args.batch_size]
if not batch:
continue
q_idx = torch.tensor([t[0] for t in batch], dtype=torch.long, device=device)
pos_idx = torch.tensor([t[1] for t in batch], dtype=torch.long, device=device)
neg_idx = torch.tensor([t[2] for t in batch], dtype=torch.long, device=device)
q_vec = Q_t[q_idx]
pos_vec = A_t[pos_idx]
neg_vec = A_t[neg_idx]
pos, neg = model(q_vec, pos_vec, neg_vec, pos_idx, neg_idx, q_idx=q_idx)
loss = bpr_loss(pos, neg)
optimizer.zero_grad()
loss.backward()
optimizer.step()
total_loss += float(loss.item())
pbar.set_postfix({"batch_loss": f"{loss.item():.4f}", "avg_loss": f"{(total_loss / (b + 1)):.4f}"})
print(f"Epoch {epoch}/{args.epochs} - BPR loss: {(total_loss / num_batches if num_batches else 0.0):.4f}")
model_dir = os.path.join(exp_cache_dir, "models")
os.makedirs(model_dir, exist_ok=True)
data_sig = want_meta["data_sig"]
ckpt_path = os.path.join(model_dir, f"{args.exp_name}_{data_sig}.pt")
meta_path = os.path.join(model_dir, f"meta_{args.exp_name}_{data_sig}.json")
ckpt = {
"state_dict": model.state_dict(),
"data_sig": data_sig,
"saved_at": datetime.now().isoformat(timespec="seconds"),
"dims": {
"d_q": int(Q_np.shape[1]),
"d_a": int(A_text_full_np.shape[1]),
"num_agents": len(a_ids),
"num_tools": int(len(feature_cache.tool_id_vocab)),
"text_hidden": args.text_hidden,
"id_dim": args.id_dim,
},
"flags": {
"use_llm_id_emb": bool(args.use_llm_id_emb),
"use_tool_id_emb": bool(args.use_tool_id_emb),
"use_model_content_vector": bool(args.use_model_content_vector),
"use_tool_content_vector": bool(args.use_tool_content_vector),
"use_query_id_emb": bool(args.use_query_id_emb),
"use_agent_id_emb": bool(args.use_agent_id_emb),
},
"parts": {
"train_parts": train_parts,
"eval_parts": eval_parts,
},
"mappings": {"q_ids": q_ids, "a_ids": a_ids, "tool_names": tool_names},
"args": vars(args),
}
torch.save(ckpt, ckpt_path)
with open(meta_path, "w", encoding="utf-8") as f:
json.dump(
{
"data_sig": data_sig,
"q_ids": q_ids,
"a_ids": a_ids,
"train_parts": train_parts,
"eval_parts": eval_parts,
},
f,
ensure_ascii=False,
indent=2,
)
print(f"[save] model -> {ckpt_path}")
print(f"[save] meta -> {meta_path}")
model.eval()
topk = int(args.topk)
# Final evaluation follows --eval_parts, not valid_qids.
# valid_qids only belongs to the split inside --train_parts; using it here
# would incorrectly prevent cross-part evaluation such as:
# --train_parts PartIII --eval_parts PartI PartII PartIII
eval_part_set = set(eval_parts)
eval_qids = [
qid for qid in qids_in_rank
if qid_to_part.get(qid, "") in eval_part_set
]
part_splits = split_eval_qids_by_part(eval_qids, qid_to_part=qid_to_part)
print(
"[eval] "
+ " | ".join(
f"{part}={len(part_splits.get(part, []))}"
for part in eval_parts
)
)
for part in eval_parts:
qids_part = part_splits.get(part, [])
if not qids_part:
print(f"[eval] skip {part}: no qids")
continue
m_part = evaluate_sampled_direct_top10(
model=model,
aid2idx=aid2idx,
qid2idx=qid2idx,
all_rankings=all_rankings,
all_questions=all_questions,
eval_qids=qids_part,
q_vectorizer=q_vectorizer_runtime,
A_text_full=A_text_full_np,
cand_size=args.eval_cand_size,
qid_to_part=qid_to_part,
pos_topk_by_part=POS_TOPK_BY_PART,
pos_topk_default=POS_TOPK,
topk=topk,
desc=f"Valid {part} (direct q-vector, top{topk})",
)
print_metrics_table(f"Validation {part} (direct q-vector)", m_part, ks=(topk,), filename=args.exp_name)
if __name__ == "__main__":
main()

Xet Storage Details

Size:
26 kB
·
Xet hash:
e04176d0ef340869af200eb0ce588838b284762b6634d567feff9adbb276add1

Xet efficiently stores files, intelligently splitting them into unique chunks and accelerating uploads and downloads. More info.