From 5029998a01150e55a1a00f50641b36a7f0d99c5b Mon Sep 17 00:00:00 2001 From: B1llsmith Date: Thu, 20 Aug 2026 01:56:24 +0300 Subject: [PATCH] Fix kg_emb broken import: SampleBaseDataset -> SampleDataset, add pandarallel dependency --- .../kg_emb/datasets/sample_kg_dataset.py | 6 +-- .../kg_emb/datasets/splitter.py | 6 +-- .../kg_emb/models/complex.py | 4 +- .../kg_emb/models/distmult.py | 4 +- .../kg_emb/models/kg_base.py | 4 +- .../kg_emb/models/rotate.py | 4 +- .../kg_emb/models/transe.py | 4 +- pyproject.toml | 3 +- tests/core/test_kg_emb.py | 48 +++++++++++++++++++ 9 files changed, 66 insertions(+), 17 deletions(-) create mode 100644 tests/core/test_kg_emb.py diff --git a/pyhealth/medcode/pretrained_embeddings/kg_emb/datasets/sample_kg_dataset.py b/pyhealth/medcode/pretrained_embeddings/kg_emb/datasets/sample_kg_dataset.py index 59d72e888..478eba748 100644 --- a/pyhealth/medcode/pretrained_embeddings/kg_emb/datasets/sample_kg_dataset.py +++ b/pyhealth/medcode/pretrained_embeddings/kg_emb/datasets/sample_kg_dataset.py @@ -1,10 +1,10 @@ -from pyhealth.datasets import SampleBaseDataset +from pyhealth.datasets import SampleDataset -class SampleKGDataset(SampleBaseDataset): +class SampleKGDataset(SampleDataset): """Sample KG dataset class. - This class inherits from `SampleBaseDataset` and is specifically designed + This class inherits from `SampleDataset` and is specifically designed for KG datasets. Args: diff --git a/pyhealth/medcode/pretrained_embeddings/kg_emb/datasets/splitter.py b/pyhealth/medcode/pretrained_embeddings/kg_emb/datasets/splitter.py index 559ecb7c4..58ecff1d5 100644 --- a/pyhealth/medcode/pretrained_embeddings/kg_emb/datasets/splitter.py +++ b/pyhealth/medcode/pretrained_embeddings/kg_emb/datasets/splitter.py @@ -4,18 +4,18 @@ import numpy as np import torch -from pyhealth.datasets import SampleBaseDataset +from pyhealth.datasets import SampleDataset def split( - dataset: SampleBaseDataset, + dataset: SampleDataset, ratios: Union[Tuple[float, float, float], List[float]], seed: Optional[int] = None, ): """Splits the dataset by its outermost indexed items Args: - dataset: a `SampleBaseDataset` object + dataset: a `SampleDataset` object ratios: a list/tuple of ratios for train / val / test seed: random seed for shuffling the dataset diff --git a/pyhealth/medcode/pretrained_embeddings/kg_emb/models/complex.py b/pyhealth/medcode/pretrained_embeddings/kg_emb/models/complex.py index 8fa2a443a..7b2f32658 100644 --- a/pyhealth/medcode/pretrained_embeddings/kg_emb/models/complex.py +++ b/pyhealth/medcode/pretrained_embeddings/kg_emb/models/complex.py @@ -1,5 +1,5 @@ from.kg_base import KGEBaseModel -from pyhealth.datasets import SampleBaseDataset +from pyhealth.datasets import SampleDataset import torch @@ -13,7 +13,7 @@ class ComplEx(KGEBaseModel): def __init__( self, - dataset: SampleBaseDataset, + dataset: SampleDataset, e_dim: int = 600, r_dim: int = 600, ns: str = "adv", diff --git a/pyhealth/medcode/pretrained_embeddings/kg_emb/models/distmult.py b/pyhealth/medcode/pretrained_embeddings/kg_emb/models/distmult.py index e7563137c..bce297139 100644 --- a/pyhealth/medcode/pretrained_embeddings/kg_emb/models/distmult.py +++ b/pyhealth/medcode/pretrained_embeddings/kg_emb/models/distmult.py @@ -1,5 +1,5 @@ from.kg_base import KGEBaseModel -from pyhealth.datasets import SampleBaseDataset +from pyhealth.datasets import SampleDataset import torch @@ -12,7 +12,7 @@ class DistMult(KGEBaseModel): """ def __init__( self, - dataset: SampleBaseDataset, + dataset: SampleDataset, e_dim: int = 300, r_dim: int = 300, ns: str = "adv", diff --git a/pyhealth/medcode/pretrained_embeddings/kg_emb/models/kg_base.py b/pyhealth/medcode/pretrained_embeddings/kg_emb/models/kg_base.py index 2de13afe2..1be30c23d 100644 --- a/pyhealth/medcode/pretrained_embeddings/kg_emb/models/kg_base.py +++ b/pyhealth/medcode/pretrained_embeddings/kg_emb/models/kg_base.py @@ -1,5 +1,5 @@ from abc import ABC -from pyhealth.datasets import SampleBaseDataset +from pyhealth.datasets import SampleDataset import torch import time @@ -32,7 +32,7 @@ def device(self): def __init__( self, - dataset: SampleBaseDataset, + dataset: SampleDataset, e_dim: int = 500, r_dim: int = 500, ns: str = "uniform", diff --git a/pyhealth/medcode/pretrained_embeddings/kg_emb/models/rotate.py b/pyhealth/medcode/pretrained_embeddings/kg_emb/models/rotate.py index df7143a6e..20536f2f5 100644 --- a/pyhealth/medcode/pretrained_embeddings/kg_emb/models/rotate.py +++ b/pyhealth/medcode/pretrained_embeddings/kg_emb/models/rotate.py @@ -1,5 +1,5 @@ from.kg_base import KGEBaseModel -from pyhealth.datasets import SampleBaseDataset +from pyhealth.datasets import SampleDataset import torch @@ -13,7 +13,7 @@ class RotatE(KGEBaseModel): def __init__( self, - dataset: SampleBaseDataset, + dataset: SampleDataset, e_dim: int = 600, r_dim: int = 300, ns='adv', diff --git a/pyhealth/medcode/pretrained_embeddings/kg_emb/models/transe.py b/pyhealth/medcode/pretrained_embeddings/kg_emb/models/transe.py index fbb6e68f6..6cf9a4f89 100644 --- a/pyhealth/medcode/pretrained_embeddings/kg_emb/models/transe.py +++ b/pyhealth/medcode/pretrained_embeddings/kg_emb/models/transe.py @@ -1,5 +1,5 @@ from.kg_base import KGEBaseModel -from pyhealth.datasets import SampleBaseDataset +from pyhealth.datasets import SampleDataset import torch @@ -13,7 +13,7 @@ class TransE(KGEBaseModel): def __init__( self, - dataset: SampleBaseDataset, + dataset: SampleDataset, e_dim: int = 300, r_dim: int = 300, ns: str = "adv", diff --git a/pyproject.toml b/pyproject.toml index b4626e649..d69f40284 100644 --- a/pyproject.toml +++ b/pyproject.toml @@ -45,7 +45,8 @@ dependencies = [ "narwhals~=2.13.0", "more-itertools~=10.8.0", "einops>=0.8.0", - "linear-attention-transformer>=0.19.1", + "linear-attention-transformer>=0.19.1", + "pandarallel>=1.6.5", ] license = "BSD-3-Clause" license-files = ["LICENSE.md"] diff --git a/tests/core/test_kg_emb.py b/tests/core/test_kg_emb.py new file mode 100644 index 000000000..891d1ef9e --- /dev/null +++ b/tests/core/test_kg_emb.py @@ -0,0 +1,48 @@ +import inspect +import unittest + +from pyhealth.datasets import SampleDataset + + +class TestKgEmbImports(unittest.TestCase): + def test_import_pretrained_embeddings(self): + try: + import pyhealth.medcode.pretrained_embeddings # noqa: F401 + except ImportError as e: + self.fail( + f"Importing pyhealth.medcode.pretrained_embeddings failed: {e}" + ) + + def test_model_classes_importable(self): + from pyhealth.medcode.pretrained_embeddings.kg_emb.models import ( + ComplEx, + DistMult, + KGEBaseModel, + RotatE, + TransE, + ) + + for cls in (KGEBaseModel, TransE, RotatE, DistMult, ComplEx): + self.assertTrue( + isinstance(cls, type), + msg=f"{cls} was not importable as a class", + ) + + def test_transe_uses_sample_dataset(self): + from pyhealth.medcode.pretrained_embeddings.kg_emb.models import TransE + + sig = inspect.signature(TransE.__init__) + dataset_annotation = sig.parameters["dataset"].annotation + self.assertEqual( + dataset_annotation, + SampleDataset, + msg=( + "TransE.__init__'s `dataset` parameter is not annotated as " + "SampleDataset — regression check for the " + "SampleBaseDataset -> SampleDataset rename" + ), + ) + + +if __name__ == "__main__": + unittest.main() \ No newline at end of file