Skip to content

fix(kg_emb): migrate SampleKGDataset off the removed 1.x dataset API - #1202

Open
AxelNoun wants to merge 27 commits into
sunlabuiuc:masterfrom
AxelNoun:fix/kg-emb-2.0-migration
Open

AxelNoun wants to merge 27 commits into
sunlabuiuc:masterfrom
AxelNoun:fix/kg-emb-2.0-migration

Conversation

@AxelNoun

@AxelNoun AxelNoun commented Aug 25, 2026 •

Copy link
Copy Markdown
Contributor

Closes #952. Part of #1201.

This PR makes pyhealth.medcode.pretrained_embeddings.kg_emb work again on the 2.0 dataset API. On master the package fails at import time (ImportError: cannot import name 'SampleBaseDataset'): two dataset modules and the five model modules (KGEBaseModel and the four models) import SampleBaseDataset, which 2.0 removed.

Design

This follows the review of #953 (@jhnwu3, 8 Apr): keep the knowledge-graph dataset in the PyHealth SampleDataset pipeline rather than as a separate class.

  • SampleKGDataset inherits from InMemorySampleDataset. It fits its processors and transforms the samples once, at construction, and keeps them in memory; get_dataloader accepts it.
  • InMemorySampleDataset is an interim base. Its docstring describes it as intended for testing and debugging on data that fits in memory, and a streaming (litdata) variant is not attempted here. The move towards pyhealth/graph suggested in the same review is not part of this PR.
  • A registered processor, kg_entity_list (KGProcessor), pads ground_truth_head and ground_truth_tail, each to its own maximum length seen in fit(), and returns {"value": ..., "mask": ...}.
  • collate_fn_dict_with_padding gains one branch: a dict whose values are all tensors is collated key by key. Dicts with other values, such as the hyperparameters field added by split(), are still collated as lists.
  • KGEBaseModel recovers the unpadded lists through the mask before filtering negatives, so a real entity id 0 is not mistaken for padding.
  • KGEBaseModel and the four models annotate dataset with SampleKGDataset.

Other fixes in kg_emb

  • SampleKGDataset no longer raises AttributeError on its default entity2id=None, and stat() returns its report instead of None. It now checks entity_num / relation_num against the vocabularies and raises ValueError on a mismatch.
  • split() draws from a local numpy.random.default_rng(seed) instead of reseeding NumPy's global state, and raises ValueError on malformed ratios instead of using assert. For a given seed, the partition differs from the previous implementation.
  • The models' __main__ demos imported SampleKGDataset from pyhealth.datasets, where it does not exist. They now run, and load data with get_dataloader.
  • base_kg_dataset.py and umls.py no longer import pandarallel or call pandarallel.initialize(): the package is not a declared dependency, and nothing in kg_emb calls parallel_apply. Both modules now import their siblings relatively.
  • kg_emb/examples/train_kge_model.py wraps the train fold returned by split() in torch.utils.data.DataLoader (the commented-out validation and test loaders likewise), because get_dataloader calls set_shuffle(), which a list does not have.

Documentation: a "Knowledge graph embeddings" section is added to docs/api/medcode.rst, with an example in examples/kg_emb_sample_dataset.py.

Tests

tests/core/test_kg_emb.py, 24 tests: imports; construction, vocabularies and cardinality validation; split() partition, reproducibility and global NumPy state; collation into padded tensors with masks; one TransE training step; BaseKGDataset.set_task() on a synthetic graph; DistMult symmetry and the TransE margin on an exact triple; unpadding with a real entity 0.

Local run on Windows, Python 3.13.5, torch 2.14.1 (CPU), litdata 0.2.76, with the branch merged with master at e94ca1d:

  • python -m unittest discover -t tests -s tests/core -p 'test_kg_emb.py': 24 passed.
  • python tools/check_pr_rules.py --base e94ca1d --head HEAD: passed.
  • The four model __main__ demos and examples/kg_emb_sample_dataset.py run.
  • test_caching.py, test_litdata_merge.py, test_partial_processors.py and test_split_set_task.py, which Fit processors on the training split in set_task (split=PatientSplit) #1273 added or changed: same results as on unmodified master in the same environment, where the only errors are Windows file locks (WinError 32) on temporary directories.
  • Full tests/core suite, run on the earlier merge with master at 5568155: 1421 tests, 17 errors. The same 17 errors occur on unmodified master in the same environment; they come from Windows file locking on temporary directories and from a URL joined with backslashes.

Limits

  • Functional verification only: tested on synthetic graphs, with no MRR or Hits@k compared against the original papers (UMLS requires a licence).
  • KGProcessor truncates a list longer than its fitted width without the warning that Truncate nested inputs longer than the fitted width #1263 added to the nested sequence processors.
  • kg_base.py still calls np.in1d, which warns under NumPy 2.2 and no longer exists in NumPy 2.4. It works while numpy~=2.2.0 is pinned.

Relation to #1192

#1192 (@userjuma) also closes #952 and touches 9 of the same files. Its diagnosis of the import chain was the starting point for this work. The order between the two PRs is discussed in the comments.

AxelNoun and others added 8 commits August 25, 2026 10:46
SampleDataset is now a litdata.StreamingDataset that expects schema.pkl. A knowledge-graph task is an in-memory list of triples, so SampleKGDataset subclasses torch.utils.data.Dataset and exposes KGDatasetProtocol for the models.

Co-authored-by: Cursor <cursoragent@cursor.com>
KGE models only need entity_num, relation_num and task_spec_param. Annotate that structural contract with KGDatasetProtocol so the model layer no longer imports the removed SampleBaseDataset.

Co-authored-by: Cursor <cursoragent@cursor.com>
SampleKGDataset was never exported from pyhealth.datasets. The examples now import it from kg_emb.datasets and build a torch DataLoader with collate_fn_dict_with_padding, because get_dataloader requires litdata.StreamingDataset.set_shuffle().

Co-authored-by: Cursor <cursoragent@cursor.com>
Validate ratios with ValueError so the check survives python -O, and shuffle with a local Generator so the function no longer mutates global NumPy state.

Co-authored-by: Cursor <cursoragent@cursor.com>
pandarallel was never declared in pyproject.toml or pixi.lock. The undeclared import in umls.py and base_kg_dataset.py is what produced the ModuleNotFoundError on the kg_emb import path in issue sunlabuiuc#952. initialize() ran in umls.py with no parallel_apply in kg_emb; mimicextract's parallel_apply calls are unreachable on the empty BaseEHRDataset stub. There is no lockfile entry to regenerate.

Co-authored-by: Cursor <cursoragent@cursor.com>
Cover construction, split reproducibility, generic collation of variable-length ground truths, set_task on a synthetic graph, and scoring invariants. Tests instantiate SampleKGDataset so a rename-only fix cannot go green.

Co-authored-by: Cursor <cursoragent@cursor.com>
Document the map-style SampleKGDataset path in the MedCode API page and add a synthetic TransE example that uses DataLoader instead of get_dataloader.

Co-authored-by: Cursor <cursoragent@cursor.com>
Replace double-hyphen asides in five docstrings with periods or commas so they remain readable in a terminal and under Sphinx.

Co-authored-by: Cursor <cursoragent@cursor.com>
SampleKGDataset failed to satisfy its own KGDatasetProtocol under mypy:
the Protocol declared task_spec_param as a plain attribute
(Mapping[str, Any] | None), which Protocol treats as read-write and
therefore invariant, while SampleKGDataset declares it as
dict[str, Any] | None. Models only ever read task_spec_param, so
declare it as a read-only property instead: read-only Protocol members
are covariant, and a concrete dict satisfies it.
Each model's __main__ block builds an untyped list of dict literals
and then adds a "train" key with a bool value, which mypy rejects
because it infers the dict's value type from the first literal.
Annotate samples as list[dict[str, Any]] in all four demo blocks.
@AxelNoun
AxelNoun marked this pull request as draft August 29, 2026 16:46
@AxelNoun

Copy link
Copy Markdown
Contributor Author

Update: after discussing direction with @jhnwu3, the target is to keep
SampleKGDataset inside the SampleDataset hierarchy rather than decoupling from it,
with in-memory loading as an interim step for negative sampling. Moving this to draft
while I rework it, and suggesting #1192 merge first as the unblocking fix for #952.

Carrying over regardless of the base class: the split() fixes, the constructor
raising AttributeError on its own default entity2id=None, stat() returning
None, the phantom SampleKGDataset import in the models' __main__ blocks, and the
test suite. Not carrying over: KGDatasetProtocol and the standalone Dataset.

@AxelNoun

AxelNoun commented Aug 31, 2026 •

Copy link
Copy Markdown
Contributor Author

Architecture & Implementation Plan Update

Just dropping a quick update here to keep a clear record of the architectural decisions we aligned on via Discord, which will guide the next commits for #1202:

1. Base Architecture (In-Memory)
As discussed, SampleKGDataset will inherit directly from InMemorySampleDataset. This keeps us correctly under the SampleDataset umbrella while cleanly sidestepping the immediate complexities of handling negative sampling in a pure streaming setup.

2. Data Shape & Serialization (The "Tensor Trick")
To handle the KG triples and the variable-length entity lists (ground_truth_head / ground_truth_tail), we are going to avoid slow JSON/Pickle serialization. I dug into the tuple_time_text_processor.py and will adapt their pattern to keep the litdata backend happy:

  • I will implement a custom KGProcessor.
  • fit() phase: Calculate the global max_length of the entity lists.
  • process() phase: Pad the variable-length lists and convert everything directly into pure PyTorch tensors before serialization.

Next Steps:
I will be working on drafting the KGProcessor and the base class integration tonight. I'll push the new commits once the core pipeline is running so we have concrete code to review!

AxelNoun and others added 2 commits August 31, 2026 22:46
…the Tensor Trick

Per the architecture pivot agreed in the PR discussion: SampleKGDataset moves
back under InMemorySampleDataset instead of standalone torch.utils.data.Dataset,
while keeping its full public surface (entity2id/relation2id, cardinality
validation, split(), stat(), dev/task_spec_param) unchanged.

- Add KGProcessor ("kg_entity_list"), pre-padding ground_truth_head/tail to
  each field's own max length and emitting {"value", "mask"} pure tensors
  ahead of litdata's pickle-based caching, instead of raw variable-length
  Python lists. "triple" goes through the existing "tensor" processor.

- Fix a correctness issue the padding introduces: pad_token_id (0) is not a
  reserved sentinel and can collide with a real entity id. kg_base.py's
  train_neg_sample_gen and test_neg_sample_filter_bias_gen now reconstruct
  the exact unpadded entity list via the mask (_unpad_ground_truth) before
  doing set-membership filtering, so negative sampling and filtered ranking
  stay correct whenever entity 0 is legitimate.

- Add a nested-dict collation branch to collate_fn_dict_with_padding,
  restricted to all-tensor dicts, so {"value","mask"} pairs batch via a
  plain stack (shape is already uniform per field) without disturbing the
  existing list-of-dicts collation used by heterogeneous per-sample dicts
  such as "hyperparameters".

- Update tests/core/test_kg_emb.py for the new shapes: triple is now a
  Tensor (not a tuple), ground_truth_* collate to {"value","mask"}, and
  set_shuffle is now expected (SampleKGDataset is intentionally back under
  the SampleDataset umbrella). Adds TestGroundTruthUnpadding, a regression
  test for the padding/entity-0 collision fix above.

24/24 tests in tests/core/test_kg_emb.py pass, plus the sample_kg_dataset.py
and splitter.py doctests.

Co-Authored-By: Claude Sonnet 5 <noreply@anthropic.com>
…mple)

The contribution-rules check (tools/check_pr_rules.py) failed on 5f66ff4,
something I hadn't run locally before pushing — only pytest, not ruff or
the repo's own gate. It enforces ruff-clean added/modified lines plus a
'>>>' doctest on every new/modified top-level public class or function.

- kg_processor.py: replace typing.Dict/List/Iterable with PEP 585 builtins
  and collections.abc.Iterable (ruff UP035/UP006, target-version py313).
  Entirely new file, so every line was in scope.
- datasets/utils.py: add a runnable '>>>' example to
  collate_fn_dict_with_padding's docstring, since the function body was
  modified (the nested all-tensor-dict branch) and had none.

Verified against the same tooling and base/head SHAs CI used
(0a75f99..<new commit>): `python tools/check_pr_rules.py --base --head`
now reports "All PR contribution rules passed." tests/core/test_kg_emb.py
still 24/24, plus the collate_fn_dict_with_padding doctest.

Co-Authored-By: Claude Sonnet 5 <noreply@anthropic.com>
@AxelNoun

Copy link
Copy Markdown
Contributor Author

CI/CD Green & Ready for Review

Everything is fully green! I've completed the implementation of the architectural plan we discussed:

  • Reverted SampleKGDataset to correctly inherit from InMemorySampleDataset.
  • Integrated the KGProcessor (using the "Tensor Trick" with independent max-length padding) to bypass the litdata serialization bottlenecks.
  • Ensured negative sampling correctness in kg_base.py by safely reconstructing the unpadded entity lists via attention masks prior to np.in1d filtering.
  • Passed all local repo linting rules (Ruff Python 3.13 typing updates and newly required doctests).

The PR is out of draft and ready for your final review whenever you have time!

@AxelNoun
AxelNoun marked this pull request as ready for review August 31, 2026 21:10
…KGDataset

SampleKGDataset is back in the SampleDataset hierarchy, so the structural
protocol introduced for the standalone design no longer has a purpose. As
announced on 29 Aug, remove protocols.py and its export, and annotate the
`dataset` parameter of KGEBaseModel, TransE, RotatE, DistMult and ComplEx
with SampleKGDataset, imported at runtime. There is no import cycle: the
dataset modules do not import the models.

The postponed-evaluation import is dropped from the five model files so
that the annotation is the class itself rather than a string.
The API page, the example and the test name still described the
standalone map-style Dataset of the first version, and the module
docstring justified KGProcessor by litdata's pickle-based caching, which
InMemorySampleDataset does not go through: it transforms every sample
once at construction and keeps the result in memory.

- docs/api/medcode.rst: SampleKGDataset is an InMemorySampleDataset and
  get_dataloader accepts it; split() returns plain lists, which is why a
  fold is wrapped in torch DataLoader directly.
- examples/kg_emb_sample_dataset.py: same correction in the docstring.
- sample_kg_dataset.py, kg_base.py: KGProcessor pads the ground-truth
  lists so that they collate into fixed-shape tensors.
- tests: rename test_is_a_map_style_dataset, since SampleKGDataset is an
  IterableDataset subclass.
Since SampleKGDataset inherits from InMemorySampleDataset, it is an
IterableDataset, and torch's DataLoader rejects shuffle=True for it:
running any of the four model modules raised "DataLoader with
IterableDataset: expected unspecified shuffle option". Use get_dataloader,
which calls InMemorySampleDataset.set_shuffle() instead, as the demos did
before the 2.0 migration. The local SampleKGDataset import is dropped
because the module now imports it at the top.
Brings the branch up to sunlabuiuc/PyHealth master at 5568155. No textual
conflicts: docs/api/medcode.rst and pyhealth/processors/__init__.py were
merged automatically.
torch pads integer tensor elements to a common width, so the example
prints tensor([16,  0,  0]); the docstring showed tensor([16, 0, 0]) and
failed under doctest.
@AxelNoun

AxelNoun commented Oct 7, 2026

Copy link
Copy Markdown
Contributor Author

@jhnwu3 @joshuasteier A status update on this PR, and two questions.

Current state: SampleKGDataset inherits from InMemorySampleDataset, with a registered KGProcessor (kg_entity_list) that pads ground_truth_head and ground_truth_tail and returns {value, mask}; kg_base.py removes the padding through the mask before filtering, so entity id 0 is never treated as padding. KGDatasetProtocol has been removed, as announced on 29 Aug, and KGEBaseModel and the four models are annotated with SampleKGDataset. I also fixed the models' __main__ demos, which had failed since the move to InMemorySampleDataset (torch's DataLoader rejects shuffle=True on an IterableDataset), and rewrote the description, which still described the standalone Dataset design. Current master (e94ca1d) is merged into the branch; tests/core/test_kg_emb.py (24 tests) and tools/check_pr_rules.py pass locally, and CI is green (build and contribution-rules).

A correction to my comments of 31 Aug: InMemorySampleDataset does not go through litdata, so the padding serves to collate these fields into tensors, not serialization. The docstrings now say so.

Relation to your review of #953 (8 Apr: stay in the PyHealth pipeline and move towards graph/): this PR does the first part only. It stays in the SampleDataset hierarchy and does not touch pyhealth/graph; InMemorySampleDataset is an interim base. Is that acceptable as a first step, or should the move to graph/ be part of this PR?

#1192 and this PR both close #952 and overlap on 9 files. There are two options:

A. Merge this PR alone (it also removes the SampleBaseDataset imports). #1192 would then be closed, with credit to @userjuma, if they agree.
B. Merge #1192 first, then this PR. This requires CI to run on the head of #1192 (it has only run on its first commit) and a reconciliation afterwards: #1192 adds pandarallel to pyproject.toml, whereas this PR removes the only pandarallel imports under pyhealth/, and both add tests/core/test_kg_emb.py. The six tests in the #1192 version of that file pass against this branch.

I have no preference between the two.

@jhnwu3 jhnwu3 left a comment •

Copy link
Copy Markdown
Collaborator

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

Quick thoughts, please don't hesitate to share your thoughts too @AxelNoun

Do you think we can have everything inherit BaseDataset and SampleDataset?

Basically, it'd be nice if we could unify all of our systems into one streaming backend here for time/memory complexity reasons.

Similarly, that way we can assume SampleDataset for each downstream model still that just leverages your kg_processor representations.

@AxelNoun

AxelNoun commented Oct 7, 2026

Copy link
Copy Markdown
Contributor Author

Quick thoughts, please don't hesitate to share your thoughts too @AxelNoun

Do you think we can have everything inherit BaseDataset and SampleDataset?

Basically, it'd be nice if we could unify all of our systems into one streaming backend here for time/memory complexity reasons.

Similarly, that way we can assume SampleDataset for each downstream model still that just leverages your kg_processor representations.

Thanks @jhnwu3, I agree with the direction. Below is where the branch stands today, a design I have checked on the current master, and two points where I would like your input before implementing it.

Where the branch stands

  • SampleKGDataset already subclasses SampleDataset, through InMemorySampleDataset. The variable-length ground_truth_head / ground_truth_tail fields go through the registered kg_entity_list processor (KGProcessor), so batches are built by get_dataloader and collate_fn_dict_with_padding like in any other task. The second half of your suggestion is therefore largely in place.
  • The first half is not. BaseKGDataset is still a standalone ABC: it loads the graph with pandas, runs link_prediction_fn in a single process and pickles the result. InMemorySampleDataset also keeps every transformed sample in memory, and its docstring says it is meant for testing and debugging. So neither the dataset nor the samples use the streaming backend yet.

Proposed design

  1. BaseKGDataset(BaseDataset) with a single triples table, declared in YAML with patient_id: null and timestamp: null, so that each triple is one record. ClinVarDataset and MedicalTranscriptionsDataset already do this. UMLSDataset then becomes a config file plus a thin subclass.
  2. A KGLinkPrediction(BaseTask). link_prediction_fn needs statistics over the whole graph: the filter sets gt_head[(r, t)] and gt_tail[(h, r)], and the frequencies behind subsampling_weight. _task_transform, however, splits the work by patient_id across workers. Since BaseTask.pre_filter runs on the global LazyFrame before that split, these statistics can be computed there with Polars group_by / join. __call__ then reads a single row.
  3. set_task(..., split=PatientSplit(...)) replaces the custom splitter. With one record per triple, a patient split is a triple-level split, which is the usual transductive protocol. Processors are then fitted on the training part only (Fit processors on the training split in set_task (split=PatientSplit) #1273).
  4. The models stop depending on attributes specific to the KG dataset and on the train / hyperparameters keys that the current splitter adds to each sample. e_num / r_num come from the fitted processors, train vs. eval from self.training, and negative_sampling becomes a model argument.

I checked points 1 to 3 on a toy graph against the current master. set_task with two workers returns a streaming SampleDataset, the filter sets match those produced by link_prediction_fn, and PatientSplit returns three disjoint parts. One detail: triple needs ("tensor", {"dtype": torch.long}), because the default tensor processor casts to float.

Two points for your input

  • Filtered negatives during training. gt_head / gt_tail are currently computed on all triples before the split, and train_neg_sample_gen uses them to remove false negatives during training. Validation and test triples therefore affect training. The standard protocol (Bordes et al., 2013; Sun et al., 2019) uses all known triples only for filtered evaluation, and only training triples for training. I suggest keeping the full sets for evaluation and building the training filter from the training part. Do you agree with this change of behavior, or should it go in a separate PR?
  • Memory used by the filter sets. KGProcessor pads each list to the largest observed length, so storage for these fields grows as N·d_max, where d_max is the largest set size. On UMLS, hub entities make d_max large. Streaming moves this cost to disk but does not remove it. An alternative is to keep a single CSR index of the graph, O(N), alongside the dataset and let the model look up the filter sets there. This is a bigger change, so I would only do it if you think the cost matters.

Scope. This touches the dataset, the task, the four models and the tests. I suggest doing it in this PR, in separate commits, and keeping link_prediction_fn and the old import paths for one release with a deprecation warning. Does that work for you, or would you prefer to merge the current state and do the migration in a follow-up PR?

AxelNoun and others added 2 commits October 7, 2026 21:17
A knowledge graph is now one `triples` table declared in YAML with
`patient_id: null` and `timestamp: null`, so the loader numbers the rows
and each triple is one record (the ClinVarDataset pattern). A PatientSplit
passed to set_task therefore splits triples.

- Entity and relation ids are global: assigned on the full graph, names
  sorted, ids 0..n-1. The Polars helpers (triples_frame,
  entity_vocabulary, relation_vocabulary, index_triples) will be reused by
  the link-prediction task.
- Triples with a missing head, relation or tail raise instead of being
  dropped.
- UMLSDataset is the bundled configs/umls.yaml plus a thin subclass.
  prepare_metadata copies the headerless graph.txt to umls-pyhealth.tsv
  under a header line; graph.txt is never modified, and the copy is made
  again when graph.txt is newer or differs in size. URL roots are rejected
  with a hint: the 1.x bucket serves graph.txt only.
- Ids no longer follow the order of first appearance used in 1.x, so 1.x
  embeddings must be mapped through the id2entity saved with them.
- The 1.x API keeps working for one release with DeprecationWarning:
  set_task(task_fn, ...) returns a SampleKGDataset built from the
  triples in file order, and triples, stat() and info() are kept.
  refresh_cache is accepted and ignored. The legacy path rejects split=
  and the other 2.0 arguments instead of storing them as hyper-parameters.

Co-Authored-By: Claude Opus 5.5 <noreply@anthropic.com>
KGTripleProcessor ("kg_triple") returns a triple of ids as a LongTensor
and, from the triples it is fitted on, keeps plain dicts: true_head,
true_tail and the word2vec frequencies behind subsampling weights, built
as in the reference RotatE implementation of Sun et al. (2019), which
computes them from training triples only. With
set_task(task, split=PatientSplit(...)), SampleBuilder fits processors on
the training part, so these dicts never see validation or test triples.
learns_statistics is set, so set_task without split= warns.

Entity and relation counts are given, not fitted, because ids are global;
size() returns the number of entities. The fitted dicts sit in a holder
whose str() is a digest of their content: set_task puts vars() of
pre-fitted processors into a JSON cache key, which cannot hold
tuple-keyed dicts.

KGProcessor no longer truncates lists longer than its fitted length.
Fitted on the training split, that length depends on the seed, and a
truncated filter set leaves known positives among the ranked candidates.
collate_fn_dict_with_padding already pads such a batch with mask 0.

Co-Authored-By: Claude Opus 5.5 <noreply@anthropic.com>
AxelNoun and others added 7 commits October 9, 2026 22:24
KGLinkPrediction(BaseTask) gives one sample per triple: the triple of ids
and its ground_truth_head / ground_truth_tail lists. Those lists cover the
whole graph, as the filter sets of filtered ranking evaluation (Bordes et
al., 2013; Sun et al., 2019); training will draw its negatives from the
kg_triple processor fitted on the training part instead.

set_task shards records by patient_id across workers, so the graph-wide
lists are computed in pre_filter with Polars group_by and join, and
__call__ reads one row. pre_filter materialises its frame once: the
workers re-run the lazy plan for every batch of records, which would
otherwise repeat the aggregations each time. It removes duplicate triples,
keeping the first in file order, logs how many, raises on a triple with a
missing field, and checks the graph has the entity and relation counts
the task was built for.

BaseKGDataset gains default_task, and its set_task:
- warns at every call when the kg_triple processor would be fitted
  without split=, i.e. on every triple (not when the processor is
  supplied, since supplied processors are never refitted);
- with split=, logs per held-out part how many triples involve an entity
  or a relation absent from the training triples. Ids are global, so
  these triples are kept and scored; their embeddings simply got no
  training signal. The parts are read sequentially, because litdata's
  indexed access on a subset is linear in the index.

Co-Authored-By: Claude Opus 5.5 <noreply@anthropic.com>
- The numbers of entities and relations come from the dataset's fitted
  kg_triple processor, which the model keeps as triple_processor.
- Training and evaluation follow model.train() / model.eval(), as the
  Trainer sets them, instead of the per-sample train flag. A 1.x batch
  whose flag disagrees with the mode triggers a warning.
- negative_sampling is a model argument. It defaults to the value a 1.x
  SampleKGDataset carries, else to 128, the default of the reference
  implementation of Sun et al. (2019).
- With use_subsampling_weight, the weights come from the processor
  (training triples) when the samples do not carry them.

A 1.x SampleKGDataset is still accepted. Training negatives are still
filtered with the samples' ground-truth lists; the next commit changes
that.

Co-Authored-By: Claude Opus 5.5 <noreply@anthropic.com>
Training negatives were filtered with the samples' ground-truth lists,
which cover the whole graph. Every validation and test positive was thus
kept out of the training negatives, so training depended on the held-out
triples.

With a kg_triple processor, the filter now comes from its true_head /
true_tail dicts, fitted on the training part only, as in the reference
implementation of Sun et al. (2019). The full-graph lists remain what
they are for: filtered ranking at evaluation (Bordes et al., 2013; Sun et
al., 2019). A 1.x SampleKGDataset has no such processor and keeps its
former behaviour until it is removed.

The regression tests fail when the full-graph filter is restored: a tail
true only in a test triple, and on a real split a head true only via a
test triple, must remain drawable as training negatives.

Co-Authored-By: Claude Opus 5.5 <noreply@anthropic.com>
link_prediction_fn, kg_emb.datasets.split, SampleKGDataset and building
a KGE model from a dataset without a kg_triple processor now raise a
DeprecationWarning that names the replacement: BaseKGDataset.set_task
with the KGLinkPrediction task and split=PatientSplit(...). Behaviour and
import paths are unchanged; they will be removed in the next release.

The new import in link_prediction.py goes at the end of the existing,
unsorted block so that the PR lint gate, which checks added lines only,
does not flag the block's pre-existing violations.

Co-Authored-By: Claude Opus 5.5 <noreply@anthropic.com>
- docs/api/medcode.rst describes BaseKGDataset / UMLSDataset, the
  KGLinkPrediction task with split=PatientSplit(...), the kg_triple
  processor that keeps training negatives free of held-out triples, and
  the deprecated 1.x entry points, including the change of id order.
- examples/kg_emb_sample_dataset.py, added earlier in this PR for the
  in-memory SampleKGDataset, becomes examples/kg_emb_link_prediction.py:
  a self-contained run of the 2.0 pipeline with the Trainer on a small
  synthetic graph whose structure a translation model can learn, so the
  filtered test metrics end well above those of a random ranking.
- kg_emb/examples/train_kge_model.py trains on UMLS with the new API. It
  no longer loads a 1.x checkpoint, whose rows follow the old id order.

Co-Authored-By: Claude Opus 5.5 <noreply@anthropic.com>
Two regressions of the model refactor, for 1.x callers:
- 1.x samples may carry hyperparameters["negative_sampling"], which the
  models read before. The model's value now only defaults to it when
  neither the model argument nor the dataset chose one (as in the
  models' __main__ demos); when both disagree, the model's value is used
  and a warning says so. Before this fix the batch value was ignored and
  128 negatives were drawn.
- negative_sampling accepts any integral number, e.g. np.int64 from a
  hyperparameter grid, instead of int only.

Tests now check that the chosen number of negatives reaches the sampler.

Co-Authored-By: Claude Opus 5.5 <noreply@anthropic.com>
…tests

- A KGE model built from a BaseKGDataset itself, instead of a set_task
  part, took the 1.x fallback (the dataset has the 1.x count attributes)
  and filtered training negatives with the batches' full-graph ground
  truths, i.e. with held-out triples. Base datasets are now refused with
  a message pointing to the training part.
- The leak regression tests go through forward on both sides and check
  that every training positive of a pair stays filtered while held-out
  positives are drawable; restoring the full-graph filter on either side,
  or keeping only the triple's own entity, now fails them.
- A legacy set_task(link_prediction_fn) call emits one DeprecationWarning
  instead of three, two of which pointed inside pyhealth.
- Deprecation notes use a warning admonition: Sphinx's deprecated
  directive needs a version number, which is not known yet.
- The docs say that the deprecated 1.x path keeps filtering training
  negatives with every triple, and the example no longer claims that all
  top tails of (e0, r1, ?) come from the second cluster.

Co-Authored-By: Claude Opus 5.5 <noreply@anthropic.com>

This branch has not been deployed

No deployments
Sign up for free to join this conversation on GitHub. Already have an account? Sign in to comment

Labels

None yet

Projects

None yet

Development

Successfully merging this pull request may close these issues.

fix: kg_emb broken import in PyHealth 2.0

3 participants