Spaces:
Paused
Paused
File size: 12,173 Bytes
a7f5b86 cb1f3d1 a7f5b86 cb1f3d1 7d49a9d a7f5b86 7d49a9d cb1f3d1 a7f5b86 cb1f3d1 a7f5b86 cb1f3d1 a7f5b86 098a435 a7f5b86 cb1f3d1 a7f5b86 cb1f3d1 a7f5b86 cb1f3d1 a7f5b86 7d49a9d a7f5b86 cb1f3d1 a7f5b86 7d49a9d a7f5b86 cb1f3d1 a7f5b86 cb1f3d1 a7f5b86 7d49a9d a7f5b86 7d49a9d a7f5b86 7d49a9d a7f5b86 cb1f3d1 | 1 2 3 4 5 6 7 8 9 10 11 12 13 14 15 16 17 18 19 20 21 22 23 24 25 26 27 28 29 30 31 32 33 34 35 36 37 38 39 40 41 42 43 44 45 46 47 48 49 50 51 52 53 54 55 56 57 58 59 60 61 62 63 64 65 66 67 68 69 70 71 72 73 74 75 76 77 78 79 80 81 82 83 84 85 86 87 88 89 90 91 92 93 94 95 96 97 98 99 100 101 102 103 104 105 106 107 108 109 110 111 112 113 114 115 116 117 118 119 120 121 122 123 124 125 126 127 128 129 130 131 132 133 134 135 136 137 138 139 140 141 142 143 144 145 146 147 148 149 150 151 152 153 154 155 156 157 158 159 160 161 162 163 164 165 166 167 168 169 170 171 172 173 174 175 176 177 178 179 180 181 182 183 184 185 186 187 188 189 190 191 192 193 194 195 196 197 198 199 200 201 202 203 204 205 206 207 208 209 210 211 212 213 214 215 216 217 218 219 220 221 222 223 224 225 226 227 228 229 230 231 232 233 234 235 236 237 238 239 240 241 242 243 244 245 246 247 248 249 250 251 252 253 254 255 256 257 258 259 260 261 262 263 264 265 266 267 268 269 270 271 272 273 274 275 276 277 278 279 280 281 282 283 284 285 286 287 288 289 290 291 292 293 294 295 296 297 298 299 300 301 302 303 304 305 306 307 308 309 310 311 312 313 314 315 316 317 318 | """
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()
|