crosspaper / scripts /recommender.py
Kattine
fix: force CPU device so ZeroGPU cannot strand the encoder
098a435
Raw
History Blame Contribute Delete
12.2 kB
"""
CrossPaper recommendation engine.
Retrieves candidate papers via FAISS nearest-neighbor search on fine-tuned
embeddings, then reranks using Maximal Marginal Relevance (MMR) to balance
relevance with disciplinary diversity.
Supports both the fine-tuned model (main recommendations) and the base model
(for before/after comparison in the demo).
AI Attribution: MMR implementation and diversity scoring logic
assisted by Claude (Anthropic).
"""
import pickle
from pathlib import Path
import faiss
import numpy as np
import pandas as pd
from sentence_transformers import SentenceTransformer
DATA_DIR = Path("data/processed")
BASE_MODEL_DIR = Path("models/base")
FINETUNED_MODEL_DIR = Path("models/fine_tuned")
# Subtracted from a candidate's MMR score when its field already appears in the
# selected set. Sized against the score's own range: relevance and max_sim are
# cosine similarities, so scores sit roughly in [-1, 1] and 0.15 is a real but
# not overwhelming push toward an unrepresented field.
FIELD_REPEAT_PENALTY = 0.15
class CrossPaperRecommender:
"""Recommends cross-disciplinary papers with diversity-aware reranking.
Uses FAISS for fast retrieval and MMR for ensuring recommendations
span multiple fields rather than clustering in one field.
"""
def __init__(self, data_dir=DATA_DIR, model_dir=FINETUNED_MODEL_DIR):
"""Initialize the recommender.
Args:
data_dir: Directory containing FAISS indexes and paper metadata.
model_dir: Directory containing the sentence-transformer model.
"""
self.data_dir = Path(data_dir)
self.model_dir = Path(model_dir)
self.model = None
self.index = None
self.metadata = None
self.embeddings = None
def load(self, index_name="finetuned"):
"""Load model, index, and metadata into memory.
Args:
index_name: Which index to load ('base' or 'finetuned').
"""
model_path = (
BASE_MODEL_DIR if index_name == "base" else FINETUNED_MODEL_DIR
)
print(f"Loading {index_name} model from {model_path}...")
# device="cpu" is required, not a preference.
#
# sentence-transformers picks its device from torch.cuda.is_available()
# when none is given. The ZeroGPU runtime patches that to return True so
# apps believe a GPU exists, but the real device is only attached inside
# a @spaces.GPU call. A model loaded outside one lands on a device that
# is never materialised, and encode() then returns zero vectors — which
# fails silently, because FAISS still returns k results ranked by index
# order rather than similarity.
#
# This model is 22M parameters and retrieval is a CPU FAISS index, so
# there is nothing to gain from a GPU here anyway.
self.model = SentenceTransformer(str(model_path), device="cpu")
index_path = self.data_dir / f"{index_name}.index"
print(f"Loading FAISS index from {index_path}...")
self.index = faiss.read_index(str(index_path))
embeddings_path = self.data_dir / f"{index_name}_embeddings.npy"
self.embeddings = np.load(str(embeddings_path))
metadata_path = self.data_dir / "paper_metadata.pkl"
self.metadata = pd.read_pickle(str(metadata_path))
print(f" Ready: {self.index.ntotal} papers indexed")
def retrieve(self, query, top_k=50):
"""Retrieve top-k candidate papers by embedding similarity.
Args:
query: Natural language query string.
top_k: Number of candidates to retrieve (before reranking).
Returns:
Tuple of (similarity scores array, candidate indices array).
"""
query_embedding = self.model.encode(
[query], normalize_embeddings=True
).astype(np.float32)
scores, indices = self.index.search(query_embedding, top_k)
# A degenerate query vector fails silently: FAISS still returns k
# results, every inner product is zero, and index order decides the
# ranking. The corpus is ordered by field, so a broken encoder returns
# the head of the index and looks exactly like a model that only knows
# one field. This check makes that failure loud.
norm = float(np.linalg.norm(query_embedding))
healthy = (
np.isfinite(query_embedding).all()
and abs(norm - 1.0) < 0.01
and indices[0][0] >= 5
)
if healthy:
print(
f"[encoder-check] ok norm={norm:.4f} "
f"top_idx={indices[0][:3].tolist()} "
f"top_scores={[round(float(s), 3) for s in scores[0][:3]]}",
flush=True,
)
else:
print(
f"[encoder-check] FAILED norm={norm:.4f} "
f"finite={bool(np.isfinite(query_embedding).all())} "
f"emb_head={[round(float(v), 4) for v in query_embedding[0][:5]]} "
f"idx_head={indices[0][:5].tolist()} "
f"score_head={[round(float(s), 4) for s in scores[0][:5]]}",
flush=True,
)
return scores[0], indices[0]
def mmr_rerank(self, query, candidates_idx, candidates_scores, top_n=10, lambda_param=0.6):
"""Rerank candidates using Maximal Marginal Relevance.
Balances relevance (similarity to query) with diversity (dissimilarity
to already-selected papers), with an additional field diversity
bonus for papers from underrepresented fields.
MMR(d) = lambda * sim(q, d) - (1 - lambda) * max(sim(d, d_j) for d_j in selected)
An additional field penalty is applied: if a paper's field
already appears in the selected set, its MMR score is reduced. This
encourages the final list to span multiple fields.
Args:
query: Original query string (unused, scores pre-computed).
candidates_idx: Array of candidate paper indices.
candidates_scores: Array of similarity scores for candidates.
top_n: Number of papers to return after reranking.
lambda_param: Relevance vs. diversity tradeoff (0=pure diversity, 1=pure relevance).
Returns:
List of dictionaries with paper info and scores.
"""
selected = []
selected_indices = []
remaining = list(range(len(candidates_idx)))
for _ in range(min(top_n, len(candidates_idx))):
best_score = -float("inf")
best_idx = -1
for i in remaining:
paper_idx = candidates_idx[i]
relevance = candidates_scores[i]
# Diversity: max similarity to any already-selected paper
if selected_indices:
candidate_emb = self.embeddings[paper_idx].reshape(1, -1)
selected_embs = self.embeddings[selected_indices]
similarities = np.dot(selected_embs, candidate_emb.T).flatten()
max_sim = np.max(similarities)
else:
max_sim = 0.0
mmr_score = lambda_param * relevance - (1 - lambda_param) * max_sim
# Field diversity penalty, applied additively.
#
# A multiplicative penalty inverts: mmr_score is negative
# whenever the diversity term dominates (low lambda), and
# scaling a negative number down by a factor raises it, which
# turns the penalty into a reward for repeating a field.
# Subtracting a constant keeps the direction stable at every
# lambda and on both sides of zero.
paper_field = self.metadata.iloc[paper_idx]["field"]
selected_fields = [
self.metadata.iloc[idx]["field"] for idx in selected_indices
]
if paper_field in selected_fields:
mmr_score -= FIELD_REPEAT_PENALTY
if mmr_score > best_score:
best_score = mmr_score
best_idx = i
if best_idx == -1:
break
paper_idx = candidates_idx[best_idx]
selected_indices.append(paper_idx)
paper_row = self.metadata.iloc[paper_idx]
selected.append({
"title": paper_row["title"],
"abstract": paper_row.get("abstract", "")[:300],
"field": paper_row["field"],
"year": int(paper_row.get("year", 0)),
"cited_by_count": int(paper_row.get("cited_by_count", 0)),
"relevance_score": float(candidates_scores[best_idx]),
"mmr_score": float(best_score),
})
remaining.remove(best_idx)
return selected
def recommend(self, query, top_n=10, lambda_param=0.6):
"""Generate recommendations for a query with diversity reranking.
This is the main entry point. Retrieves candidates via FAISS,
then applies MMR reranking to balance relevance with field
diversity.
Args:
query: Natural language description of research interest.
top_n: Number of recommendations to return.
lambda_param: Relevance vs. diversity tradeoff.
Returns:
Dictionary with recommendations list and diversity metrics.
"""
scores, indices = self.retrieve(query, top_k=top_n * 5)
recommendations = self.mmr_rerank(
query, indices, scores, top_n=top_n, lambda_param=lambda_param
)
diversity_metrics = self._compute_diversity(recommendations)
return {
"recommendations": recommendations,
"diversity": diversity_metrics,
}
def _compute_diversity(self, recommendations):
"""Compute field diversity metrics for a recommendation set.
Args:
recommendations: List of recommendation dictionaries.
Returns:
Dictionary with diversity metrics including Shannon entropy,
field distribution, and cross-field hit rate.
"""
if not recommendations:
return {"entropy": 0.0, "distribution": {}, "cross_field_rate": 0.0}
fields = [r["field"] for r in recommendations]
unique, counts = np.unique(fields, return_counts=True)
probs = counts / counts.sum()
# Shannon entropy (higher = more diverse)
entropy = -np.sum(probs * np.log2(probs + 1e-10))
# Distribution as percentages
distribution = {
disc: float(count / len(fields))
for disc, count in zip(unique, counts)
}
# Cross-field rate (fraction of results NOT from the dominant field)
dominant_fraction = max(probs)
cross_rate = 1.0 - dominant_fraction
return {
"entropy": float(entropy),
"distribution": distribution,
"cross_field_rate": float(cross_rate),
"num_fields": int(len(unique)),
}
def main():
"""Quick smoke test for the recommender."""
recommender = CrossPaperRecommender()
recommender.load(index_name="finetuned")
test_queries = [
"attention mechanism in visual processing",
"reinforcement learning for decision making",
"gene expression regulation in neural development",
]
for query in test_queries:
print(f"\nQuery: {query}")
print("-" * 60)
result = recommender.recommend(query, top_n=5)
for i, rec in enumerate(result["recommendations"], 1):
print(f" {i}. [{rec['field']}] {rec['title'][:80]}")
print(f" relevance={rec['relevance_score']:.3f} mmr={rec['mmr_score']:.3f}")
div = result["diversity"]
print(f" Diversity: entropy={div['entropy']:.2f}, "
f"fields={div['num_fields']}, "
f"cross_rate={div['cross_field_rate']:.0%}")
if __name__ == "__main__":
main()