From a57883a7bbb8df6ff9a257d97b9383ae15af1256 Mon Sep 17 00:00:00 2001 From: Felipe Bonchristiano Date: Wed, 19 Aug 2026 22:23:59 -0400 Subject: [PATCH 1/5] medlink-fix --- pyhealth/models/medlink/utils.py | 13 +++++----- tests/core/test_medlink.py | 41 ++++++++++++++++++++++++++++++++ 2 files changed, 47 insertions(+), 7 deletions(-) diff --git a/pyhealth/models/medlink/utils.py b/pyhealth/models/medlink/utils.py index 6b3263655..d5bf5f779 100644 --- a/pyhealth/models/medlink/utils.py +++ b/pyhealth/models/medlink/utils.py @@ -131,13 +131,12 @@ def get_bm25_hard_negatives(bm25_model, corpus, queries, qrels): qrels_w_neg = {} for q_id, q in tqdm.tqdm(queries.items()): d_ids = [d_id for d_id in qrels[q_id] if qrels[q_id][d_id] > 0] - ds = [corpus[d_id] for d_id in d_ids] - for d_id, d in zip(d_ids, ds): - scores = bm25_model.get_scores(d) - for (ned_d_id, neg_s) in sorted(scores.items(), key=lambda x: x[1], - reverse=True): - if ned_d_id != d_id: - qrels_w_neg[q_id] = {d_id: 1, ned_d_id: -1} + scores = bm25_model.get_scores(q) + ranked = sorted(scores.items(), key=lambda x: x[1], reverse=True) + for d_id in d_ids: + for (neg_d_id, neg_s) in ranked: + if neg_d_id not in d_ids: # exclude every positive, not just d_id + qrels_w_neg[q_id] = {d_id: 1, neg_d_id: -1} break return qrels_w_neg diff --git a/tests/core/test_medlink.py b/tests/core/test_medlink.py index 9a530a11c..30b98f000 100644 --- a/tests/core/test_medlink.py +++ b/tests/core/test_medlink.py @@ -129,5 +129,46 @@ def test_feature_key_inference(self): self.assertEqual(model.feature_key, "conditions") + def test_hard_negatives_score_query_not_positive_doc(self): + """Regression: hard negatives must be mined by scoring the QUERY, not + the positive document. The old code called get_scores(d) on the + positive doc, so "hard negatives" were docs similar to the answer. + """ + from pyhealth.models.medlink.utils import get_bm25_hard_negatives + + class FakeBM25: + def get_scores(self, text): + if text == "QUERY": # query ranks pos top, then neg_q + return {"pos": 10.0, "neg_q": 9.0, "neg_d": 1.0} + return {"pos": 10.0, "neg_d": 9.0, "neg_q": 1.0} # doc-scoring picks neg_d + + corpus = {"pos": "POS", "neg_q": "NQ", "neg_d": "ND"} + queries = {"q1": "QUERY"} + qrels = {"q1": {"pos": 1}} + out = get_bm25_hard_negatives(FakeBM25(), corpus, queries, qrels) + # scoring the query picks neg_q; scoring the positive doc would pick neg_d + self.assertEqual(out["q1"], {"pos": 1, "neg_q": -1}) + + def test_hard_negatives_exclude_all_positives(self): + """Regression: with multiple positives, no positive may be chosen as a + negative. Scoring the query ranks the positives on top, so excluding + only the current positive would pick another positive as a false negative. + """ + from pyhealth.models.medlink.utils import get_bm25_hard_negatives + + class FakeBM25: + def get_scores(self, text): + return {"pos1": 10.0, "pos2": 9.0, "neg": 8.0} + + corpus = {"pos1": "P1", "pos2": "P2", "neg": "N"} + queries = {"q1": "QUERY"} + qrels = {"q1": {"pos1": 1, "pos2": 1}} + out = get_bm25_hard_negatives(FakeBM25(), corpus, queries, qrels) + neg_ids = [d for d, lbl in out["q1"].items() if lbl == -1] + self.assertNotIn("pos1", neg_ids) + self.assertNotIn("pos2", neg_ids) + self.assertEqual(neg_ids, ["neg"]) + + if __name__ == "__main__": unittest.main() From 0794b21cdc2b627da65744a42881c36ea07834d0 Mon Sep 17 00:00:00 2001 From: Felipe Bonchristiano Date: Thu, 20 Aug 2026 19:16:23 -0500 Subject: [PATCH 2/5] added example docstring to get_bm25_hard_negatives --- pyhealth/models/medlink/utils.py | 9 +++++++++ 1 file changed, 9 insertions(+) diff --git a/pyhealth/models/medlink/utils.py b/pyhealth/models/medlink/utils.py index d5bf5f779..da8071690 100644 --- a/pyhealth/models/medlink/utils.py +++ b/pyhealth/models/medlink/utils.py @@ -127,6 +127,15 @@ def get_bm25_hard_negatives(bm25_model, corpus, queries, qrels): Returns: qrels_w_neg: Updated qrels dictionary containing both positives (1) and negatives (-1). + + Examples: + >>> # bm25_model.get_scores(query) -> {doc_id: score} + >>> corpus = {"d0": ["fever", "cough"], "d1": ["fever", "rash"]} + >>> queries = {"q0": ["fever", "cough"]} + >>> qrels = {"q0": {"d0": 1}} # d0 is q0's positive match + >>> qrels_w_neg = get_bm25_hard_negatives(bm25_model, corpus, queries, qrels) + >>> qrels_w_neg["q0"] # positive kept; top query-ranked non-positive labeled -1 + {'d0': 1, 'd1': -1} """ qrels_w_neg = {} for q_id, q in tqdm.tqdm(queries.items()): From 8d482b4f229846e1e082abe4992c19ef0aaca920 Mon Sep 17 00:00:00 2001 From: Felipe Bonchristiano Date: Thu, 27 Aug 2026 17:05:32 -0500 Subject: [PATCH 3/5] fixed get_bm25_hard_negatives not preserving all positives --- pyhealth/models/medlink/utils.py | 13 ++++++++----- tests/core/test_medlink.py | 22 +++++++++++++++++++++- 2 files changed, 29 insertions(+), 6 deletions(-) diff --git a/pyhealth/models/medlink/utils.py b/pyhealth/models/medlink/utils.py index da8071690..f3b20a7f7 100644 --- a/pyhealth/models/medlink/utils.py +++ b/pyhealth/models/medlink/utils.py @@ -140,13 +140,16 @@ def get_bm25_hard_negatives(bm25_model, corpus, queries, qrels): qrels_w_neg = {} for q_id, q in tqdm.tqdm(queries.items()): d_ids = [d_id for d_id in qrels[q_id] if qrels[q_id][d_id] > 0] + + qrels_w_neg[q_id] = {d_id: 1 for d_id in d_ids} + scores = bm25_model.get_scores(q) ranked = sorted(scores.items(), key=lambda x: x[1], reverse=True) - for d_id in d_ids: - for (neg_d_id, neg_s) in ranked: - if neg_d_id not in d_ids: # exclude every positive, not just d_id - qrels_w_neg[q_id] = {d_id: 1, neg_d_id: -1} - break + for neg_d_id, neg_s in ranked: + if neg_d_id not in d_ids: # exclude every positive, not just d_id + qrels_w_neg[q_id][neg_d_id] = -1 + break + return qrels_w_neg diff --git a/tests/core/test_medlink.py b/tests/core/test_medlink.py index 30b98f000..4d23727cb 100644 --- a/tests/core/test_medlink.py +++ b/tests/core/test_medlink.py @@ -128,7 +128,6 @@ def test_feature_key_inference(self): ) self.assertEqual(model.feature_key, "conditions") - def test_hard_negatives_score_query_not_positive_doc(self): """Regression: hard negatives must be mined by scoring the QUERY, not the positive document. The old code called get_scores(d) on the @@ -169,6 +168,27 @@ def get_scores(self, text): self.assertNotIn("pos2", neg_ids) self.assertEqual(neg_ids, ["neg"]) + def test_hard_negatives_preserves_all_positives(self): + """Regression: with multiple positives, all positives must remain in the + output when a hard negative is added. + """ + from pyhealth.models.medlink.utils import get_bm25_hard_negatives + + class FakeBM25: + def get_scores(self, text): + return {"pos1": 10.0, "pos2": 9.0, "neg": 8.0} + + corpus = {"pos1": "P1", "pos2": "P2", "neg": "N"} + queries = {"q1": "QUERY"} + qrels = {"q1": {"pos1": 1, "pos2": 1}} + + out = get_bm25_hard_negatives(FakeBM25(), corpus, queries, qrels) + + self.assertEqual( + out["q1"], + {"pos1": 1, "pos2": 1, "neg": -1}, + ) + if __name__ == "__main__": unittest.main() From 16e0bb4ed0ac71790653455158c000d03bc70b68 Mon Sep 17 00:00:00 2001 From: Felipe Bonchristiano Date: Fri, 28 Aug 2026 17:18:26 -0500 Subject: [PATCH 4/5] Handle multi-positive MedLink training qrels --- pyhealth/models/medlink/utils.py | 35 +++++++++++--------------------- tests/core/test_medlink.py | 12 ++++++++++- 2 files changed, 23 insertions(+), 24 deletions(-) diff --git a/pyhealth/models/medlink/utils.py b/pyhealth/models/medlink/utils.py index f3b20a7f7..737b23292 100644 --- a/pyhealth/models/medlink/utils.py +++ b/pyhealth/models/medlink/utils.py @@ -184,33 +184,22 @@ def get_train_dataloader( train_samples = [] for query_id in query_ids: s_q = queries[query_id] - id_p, s_p, s_n = None, None, None - assert len(qrels[query_id]) <= 2 + positive_ids, s_n = [], None for corpus_id, score in qrels[query_id].items(): if score == 1: - id_p = corpus_id - s_p = corpus[corpus_id] + positive_ids.append(corpus_id) if score == -1: s_n = corpus[corpus_id] - if s_n is not None: - train_samples.append( - { - "query_id": query_id, - "id_p": id_p, - "s_q": s_q, - "s_p": s_p, - "s_n": s_n, - } - ) - else: - train_samples.append( - { - "query_id": query_id, - "id_p": id_p, - "s_q": s_q, - "s_p": s_p, - } - ) + for id_p in positive_ids: + sample = { + "query_id": query_id, + "id_p": id_p, + "s_q": s_q, + "s_p": corpus[id_p], + } + if s_n is not None: + sample["s_n"] = s_n + train_samples.append(sample) print("Loaded {} training pairs.".format(len(train_samples))) train_dataloader = DataLoader( train_samples, shuffle=shuffle, batch_size=batch_size, collate_fn=collate_fn diff --git a/tests/core/test_medlink.py b/tests/core/test_medlink.py index 4d23727cb..6acbfaf52 100644 --- a/tests/core/test_medlink.py +++ b/tests/core/test_medlink.py @@ -172,7 +172,10 @@ def test_hard_negatives_preserves_all_positives(self): """Regression: with multiple positives, all positives must remain in the output when a hard negative is added. """ - from pyhealth.models.medlink.utils import get_bm25_hard_negatives + from pyhealth.models.medlink.utils import ( + get_bm25_hard_negatives, + get_train_dataloader, + ) class FakeBM25: def get_scores(self, text): @@ -189,6 +192,13 @@ def get_scores(self, text): {"pos1": 1, "pos2": 1, "neg": -1}, ) + dataloader = get_train_dataloader( + corpus, queries, out, batch_size=2, shuffle=False + ) + batch = next(iter(dataloader)) + self.assertEqual(batch["id_p"], ["pos1", "pos2"]) + self.assertEqual(batch["s_n"], ["N", "N"]) + if __name__ == "__main__": unittest.main() From 758a5f0f0cb71a879f9f632942b083164cb8e8f1 Mon Sep 17 00:00:00 2001 From: Felipe Bonchristiano Date: Fri, 28 Aug 2026 17:24:01 -0500 Subject: [PATCH 5/5] Document MedLink train dataloader usage --- pyhealth/models/medlink/utils.py | 9 +++++++++ 1 file changed, 9 insertions(+) diff --git a/pyhealth/models/medlink/utils.py b/pyhealth/models/medlink/utils.py index 737b23292..63fb155ed 100644 --- a/pyhealth/models/medlink/utils.py +++ b/pyhealth/models/medlink/utils.py @@ -178,6 +178,15 @@ def get_train_dataloader( Returns: DataLoader returning batches of dicts. + + Examples: + >>> corpus = {"p1": "positive one", "p2": "positive two", "n": "negative"} + >>> queries = {"q1": "query"} + >>> qrels = {"q1": {"p1": 1, "p2": 1, "n": -1}} + >>> loader = get_train_dataloader(corpus, queries, qrels, batch_size=2, shuffle=False) + Loaded 2 training pairs. + >>> next(iter(loader))["id_p"] + ['p1', 'p2'] """ query_ids = list(queries.keys())