Repository navigation
Conversation
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.
|
Update: after discussing direction with @jhnwu3, the target is to keep Carrying over regardless of the base class: the |
Architecture & Implementation Plan UpdateJust 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) 2. Data Shape & Serialization (The "Tensor Trick")
Next Steps: |
…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>
CI/CD Green & Ready for ReviewEverything is fully green! I've completed the implementation of the architectural plan we discussed:
The PR is out of draft and ready for your final review whenever you have time! |
…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.
|
@jhnwu3 @joshuasteier A status update on this PR, and two questions. Current state: A correction to my comments of 31 Aug: Relation to your review of #953 (8 Apr: stay in the PyHealth pipeline and move towards #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 I have no preference between the two. |
There was a problem hiding this comment.
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 Where the branch stands
Proposed design
I checked points 1 to 3 on a toy graph against the current Two points for your input
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 |
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>
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>
Closes #952. Part of #1201.
This PR makes
pyhealth.medcode.pretrained_embeddings.kg_embwork again on the 2.0 dataset API. Onmasterthe package fails at import time (ImportError: cannot import name 'SampleBaseDataset'): two dataset modules and the five model modules (KGEBaseModeland the four models) importSampleBaseDataset, which 2.0 removed.Design
This follows the review of #953 (@jhnwu3, 8 Apr): keep the knowledge-graph dataset in the PyHealth
SampleDatasetpipeline rather than as a separate class.SampleKGDatasetinherits fromInMemorySampleDataset. It fits its processors and transforms the samples once, at construction, and keeps them in memory;get_dataloaderaccepts it.InMemorySampleDatasetis 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 towardspyhealth/graphsuggested in the same review is not part of this PR.kg_entity_list(KGProcessor), padsground_truth_headandground_truth_tail, each to its own maximum length seen infit(), and returns{"value": ..., "mask": ...}.collate_fn_dict_with_paddinggains one branch: a dict whose values are all tensors is collated key by key. Dicts with other values, such as thehyperparametersfield added bysplit(), are still collated as lists.KGEBaseModelrecovers the unpadded lists through the mask before filtering negatives, so a real entity id 0 is not mistaken for padding.KGEBaseModeland the four models annotatedatasetwithSampleKGDataset.Other fixes in
kg_embSampleKGDatasetno longer raisesAttributeErroron its defaultentity2id=None, andstat()returns its report instead ofNone. It now checksentity_num/relation_numagainst the vocabularies and raisesValueErroron a mismatch.split()draws from a localnumpy.random.default_rng(seed)instead of reseeding NumPy's global state, and raisesValueErroron malformed ratios instead of usingassert. For a givenseed, the partition differs from the previous implementation.__main__demos importedSampleKGDatasetfrompyhealth.datasets, where it does not exist. They now run, and load data withget_dataloader.base_kg_dataset.pyandumls.pyno longer importpandarallelor callpandarallel.initialize(): the package is not a declared dependency, and nothing inkg_embcallsparallel_apply. Both modules now import their siblings relatively.kg_emb/examples/train_kge_model.pywraps the train fold returned bysplit()intorch.utils.data.DataLoader(the commented-out validation and test loaders likewise), becauseget_dataloadercallsset_shuffle(), which a list does not have.Documentation: a "Knowledge graph embeddings" section is added to
docs/api/medcode.rst, with an example inexamples/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
masterat 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.__main__demos andexamples/kg_emb_sample_dataset.pyrun.test_caching.py,test_litdata_merge.py,test_partial_processors.pyandtest_split_set_task.py, which Fit processors on the training split in set_task (split=PatientSplit) #1273 added or changed: same results as on unmodifiedmasterin the same environment, where the only errors are Windows file locks (WinError 32) on temporary directories.tests/coresuite, run on the earlier merge withmasterat 5568155: 1421 tests, 17 errors. The same 17 errors occur on unmodifiedmasterin the same environment; they come from Windows file locking on temporary directories and from a URL joined with backslashes.Limits
KGProcessortruncates 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.pystill callsnp.in1d, which warns under NumPy 2.2 and no longer exists in NumPy 2.4. It works whilenumpy~=2.2.0is 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.