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()