diff --git a/chebai_graph/preprocessing/bin/AtomCharge/indices_one_hot.txt b/chebai_graph/preprocessing/bin/chebi/AtomCharge/indices_one_hot.txt similarity index 100% rename from chebai_graph/preprocessing/bin/AtomCharge/indices_one_hot.txt rename to chebai_graph/preprocessing/bin/chebi/AtomCharge/indices_one_hot.txt diff --git a/chebai_graph/preprocessing/bin/AtomFunctionalGroup/indices_one_hot.txt b/chebai_graph/preprocessing/bin/chebi/AtomFunctionalGroup/indices_one_hot.txt similarity index 100% rename from chebai_graph/preprocessing/bin/AtomFunctionalGroup/indices_one_hot.txt rename to chebai_graph/preprocessing/bin/chebi/AtomFunctionalGroup/indices_one_hot.txt diff --git a/chebai_graph/preprocessing/bin/AtomHybridization/indices_one_hot.txt b/chebai_graph/preprocessing/bin/chebi/AtomHybridization/indices_one_hot.txt similarity index 100% rename from chebai_graph/preprocessing/bin/AtomHybridization/indices_one_hot.txt rename to chebai_graph/preprocessing/bin/chebi/AtomHybridization/indices_one_hot.txt diff --git a/chebai_graph/preprocessing/bin/AtomNodeLevel/indices_one_hot.txt b/chebai_graph/preprocessing/bin/chebi/AtomNodeLevel/indices_one_hot.txt similarity index 100% rename from chebai_graph/preprocessing/bin/AtomNodeLevel/indices_one_hot.txt rename to chebai_graph/preprocessing/bin/chebi/AtomNodeLevel/indices_one_hot.txt diff --git a/chebai_graph/preprocessing/bin/AtomNumHs/indices_one_hot.txt b/chebai_graph/preprocessing/bin/chebi/AtomNumHs/indices_one_hot.txt similarity index 100% rename from chebai_graph/preprocessing/bin/AtomNumHs/indices_one_hot.txt rename to chebai_graph/preprocessing/bin/chebi/AtomNumHs/indices_one_hot.txt diff --git a/chebai_graph/preprocessing/bin/AtomType/indices_one_hot.txt b/chebai_graph/preprocessing/bin/chebi/AtomType/indices_one_hot.txt similarity index 100% rename from chebai_graph/preprocessing/bin/AtomType/indices_one_hot.txt rename to chebai_graph/preprocessing/bin/chebi/AtomType/indices_one_hot.txt diff --git a/chebai_graph/preprocessing/bin/BondLevel/indices_one_hot.txt b/chebai_graph/preprocessing/bin/chebi/BondLevel/indices_one_hot.txt similarity index 100% rename from chebai_graph/preprocessing/bin/BondLevel/indices_one_hot.txt rename to chebai_graph/preprocessing/bin/chebi/BondLevel/indices_one_hot.txt diff --git a/chebai_graph/preprocessing/bin/BondType/indices_one_hot.txt b/chebai_graph/preprocessing/bin/chebi/BondType/indices_one_hot.txt similarity index 100% rename from chebai_graph/preprocessing/bin/BondType/indices_one_hot.txt rename to chebai_graph/preprocessing/bin/chebi/BondType/indices_one_hot.txt diff --git a/chebai_graph/preprocessing/bin/NumAtomBonds/indices_one_hot.txt b/chebai_graph/preprocessing/bin/chebi/NumAtomBonds/indices_one_hot.txt similarity index 100% rename from chebai_graph/preprocessing/bin/NumAtomBonds/indices_one_hot.txt rename to chebai_graph/preprocessing/bin/chebi/NumAtomBonds/indices_one_hot.txt diff --git a/chebai_graph/preprocessing/datasets/base.py b/chebai_graph/preprocessing/datasets/base.py index 26bc507..4e8e6c3 100644 --- a/chebai_graph/preprocessing/datasets/base.py +++ b/chebai_graph/preprocessing/datasets/base.py @@ -50,7 +50,7 @@ def __init__( super().__init__(**kwargs) # atom_properties and bond_properties are given as lists containing class_paths if properties is not None: - properties = [resolve_property(prop) for prop in properties] + properties = [resolve_property(prop, self.data_type) for prop in properties] properties = self._sort_properties(properties) else: properties = [] diff --git a/chebai_graph/preprocessing/datasets/utils.py b/chebai_graph/preprocessing/datasets/utils.py index 3ec6515..79f2fb3 100644 --- a/chebai_graph/preprocessing/datasets/utils.py +++ b/chebai_graph/preprocessing/datasets/utils.py @@ -4,7 +4,9 @@ from chebai_graph.preprocessing.properties import MolecularProperty -def resolve_property(property: str | MolecularProperty) -> MolecularProperty: +def resolve_property( + property: str | MolecularProperty, data_type: str +) -> MolecularProperty: """ Resolves a molecular property specification (either as a class instance or class path string) into a MolecularProperty instance. @@ -19,6 +21,8 @@ def resolve_property(property: str | MolecularProperty) -> MolecularProperty: property (str | MolecularProperty): The property to resolve. Can be a class instance, a fully qualified class name (e.g. "module.ClassName"), or a class name assumed to be in `chebai_graph.preprocessing.properties`. + data_type (str): The data type associated with the property. This used to determine or set + tokens file path for the property if applicable. Returns: MolecularProperty: An instance of the resolved MolecularProperty. @@ -37,7 +41,7 @@ def resolve_property(property: str | MolecularProperty) -> MolecularProperty: module_name = property[:last_dot] class_name = property[last_dot + 1 :] module = importlib.import_module(module_name) - return getattr(module, class_name)() + return getattr(module, class_name)(data_type=data_type) except ValueError: # if only a class name is given, assume the module is chebai_graph.processing.properties - return getattr(graph_properties, property)() + return getattr(graph_properties, property)(data_type=data_type) diff --git a/chebai_graph/preprocessing/properties/augmented_properties.py b/chebai_graph/preprocessing/properties/augmented_properties.py index f5f7b1d..6f112ea 100644 --- a/chebai_graph/preprocessing/properties/augmented_properties.py +++ b/chebai_graph/preprocessing/properties/augmented_properties.py @@ -23,14 +23,17 @@ class AtomNodeLevel(AllNodeTypeProperty): - def __init__(self, encoder: PropertyEncoder | None = None): + def __init__(self, encoder: PropertyEncoder | None = None, **kwargs) -> None: """ Initialize AtomNodeLevel with an optional encoder. Args: encoder (PropertyEncoder | None): Property encoder to use. Defaults to OneHotEncoder. """ - super().__init__(encoder or OneHotEncoder(self)) + data_type = kwargs.get("data_type") + super().__init__( + encoder=encoder or OneHotEncoder(self, data_type=data_type), **kwargs + ) def get_atom_value(self, atom: Chem.rdchem.Atom | dict) -> str | int | bool: """ @@ -46,14 +49,17 @@ def get_atom_value(self, atom: Chem.rdchem.Atom | dict) -> str | int | bool: class AtomFunctionalGroup(FGNodeTypeProperty): - def __init__(self, encoder: PropertyEncoder | None = None): + def __init__(self, encoder: PropertyEncoder | None = None, **kwargs) -> None: """ Initialize AtomFunctionalGroup with an optional encoder. Args: encoder (PropertyEncoder | None): Property encoder to use. Defaults to OneHotEncoder. """ - super().__init__(encoder or OneHotEncoder(self)) + data_type = kwargs.get("data_type") + super().__init__( + encoder=encoder or OneHotEncoder(self, data_type=data_type), **kwargs + ) def get_atom_value(self, atom: Chem.rdchem.Atom | dict) -> str | int | bool: """ @@ -69,14 +75,17 @@ def get_atom_value(self, atom: Chem.rdchem.Atom | dict) -> str | int | bool: class AtomRingSize(FGNodeTypeProperty): - def __init__(self, encoder: PropertyEncoder | None = None): + def __init__(self, encoder: PropertyEncoder | None = None, **kwargs) -> None: """ Initialize AtomRingSize with an optional encoder. Args: encoder (PropertyEncoder | None): Property encoder to use. Defaults to OneHotEncoder. """ - super().__init__(encoder or OneHotEncoder(self)) + data_type = kwargs.get("data_type") + super().__init__( + encoder=encoder or OneHotEncoder(self, data_type=data_type), **kwargs + ) def get_atom_value(self, atom: Chem.rdchem.Atom | dict) -> int: """ @@ -113,14 +122,14 @@ def _check_modify_atom_prop_value( class IsHydrogenBondDonorFG(FGNodeTypeProperty): - def __init__(self, encoder: PropertyEncoder | None = None): + def __init__(self, encoder: PropertyEncoder | None = None, **kwargs) -> None: """ Initialize IsHydrogenBondDonorFG with an optional encoder. Args: encoder (PropertyEncoder | None): Property encoder to use. Defaults to BoolEncoder. """ - super().__init__(encoder or BoolEncoder(self)) + super().__init__(encoder=encoder or BoolEncoder(self), **kwargs) # fmt: off # https://github.com/thaonguyen217/farm_molecular_representation/blob/main/src/(6)gen_FG_KG.py#L26-L31 self._hydrogen_bond_donor: set[str] = { @@ -146,14 +155,14 @@ def get_atom_value(self, atom: Chem.rdchem.Atom | dict) -> bool: class IsHydrogenBondAcceptorFG(FGNodeTypeProperty): - def __init__(self, encoder: PropertyEncoder | None = None): + def __init__(self, encoder: PropertyEncoder | None = None, **kwargs) -> None: """ Initialize IsHydrogenBondAcceptorFG with an optional encoder. Args: encoder (PropertyEncoder | None): Property encoder to use. Defaults to BoolEncoder. """ - super().__init__(encoder or BoolEncoder(self)) + super().__init__(encoder=encoder or BoolEncoder(self), **kwargs) # fmt: off # https://github.com/thaonguyen217/farm_molecular_representation/blob/main/src/(6)gen_FG_KG.py#L33-L39 self._hydrogen_bond_acceptor: set[str] = { @@ -180,13 +189,13 @@ def get_atom_value(self, atom: Chem.rdchem.Atom | dict) -> bool: class IsFGAlkyl(FGNodeTypeProperty): - def __init__(self, encoder: PropertyEncoder | None = None): + def __init__(self, encoder: PropertyEncoder | None = None, **kwargs) -> None: """ Args: encoder (PropertyEncoder | None): Optional encoder to use for this property. Defaults to BoolEncoder if not provided. """ - super().__init__(encoder or BoolEncoder(self)) + super().__init__(encoder=encoder or BoolEncoder(self), **kwargs) def get_atom_value(self, atom: Chem.rdchem.Atom | dict) -> int: """ @@ -323,12 +332,15 @@ class AugAtomAromaticity(AugNodeValueDefaulter, pr.AtomAromaticity): class BondLevel(AugmentedBondProperty): - def __init__(self, encoder: PropertyEncoder | None = None): + def __init__(self, encoder: PropertyEncoder | None = None, **kwargs) -> None: """ Args: encoder (PropertyEncoder | None): Optional encoder to use. Defaults to OneHotEncoder. """ - super().__init__(encoder or OneHotEncoder(self)) + data_type = kwargs.get("data_type") + super().__init__( + encoder=encoder or OneHotEncoder(self, data_type=data_type), **kwargs + ) def get_bond_value(self, bond: Chem.rdchem.Bond | dict) -> str: """ diff --git a/chebai_graph/preprocessing/properties/base.py b/chebai_graph/preprocessing/properties/base.py index b28c415..f148df8 100644 --- a/chebai_graph/preprocessing/properties/base.py +++ b/chebai_graph/preprocessing/properties/base.py @@ -31,10 +31,16 @@ class MolecularProperty(ABC): Defaults to IndexEncoder if not provided. """ - def __init__(self, encoder: PropertyEncoder | None = None) -> None: + def __init__( + self, + data_type: str, + encoder: PropertyEncoder | None = None, + ) -> None: + assert data_type is not None, "data_type must be provided for MolecularProperty" if encoder is None: - encoder = IndexEncoder(self) + encoder = IndexEncoder(self, data_type=data_type) self.encoder: PropertyEncoder = encoder + self._data_type = data_type @property def name(self) -> str: @@ -174,8 +180,8 @@ class AugAtomType(FrozenPropertyAlias, AtomType): ... ValueError: If new tokens are added to the frozen encoder during processing. """ - def __init__(self, encoder: PropertyEncoder | None = None) -> None: - super().__init__(encoder) + def __init__(self, encoder: PropertyEncoder, **kwargs) -> None: + super().__init__(encoder=encoder, **kwargs) # Lock the encoder's cache to prevent adding new tokens if hasattr(self.encoder, "cache") and isinstance(self.encoder.cache, dict): self.encoder.cache = MappingProxyType(self.encoder.cache) diff --git a/chebai_graph/preprocessing/properties/properties.py b/chebai_graph/preprocessing/properties/properties.py index 2154f9c..87112da 100644 --- a/chebai_graph/preprocessing/properties/properties.py +++ b/chebai_graph/preprocessing/properties/properties.py @@ -19,8 +19,11 @@ class AtomType(AtomProperty): Uses a one-hot encoder by default. """ - def __init__(self, encoder: PropertyEncoder | None = None) -> None: - super().__init__(encoder or OneHotEncoder(self)) + def __init__(self, encoder: PropertyEncoder | None = None, **kwargs) -> None: + data_type = kwargs.get("data_type") + super().__init__( + encoder=encoder or OneHotEncoder(self, data_type=data_type), **kwargs + ) def get_atom_value(self, atom: Chem.rdchem.Atom) -> int: """ @@ -42,8 +45,11 @@ class NumAtomBonds(AtomProperty): Uses a one-hot encoder by default. """ - def __init__(self, encoder: PropertyEncoder | None = None) -> None: - super().__init__(encoder or OneHotEncoder(self)) + def __init__(self, encoder: PropertyEncoder | None = None, **kwargs) -> None: + data_type = kwargs.get("data_type") + super().__init__( + encoder=encoder or OneHotEncoder(self, data_type=data_type), **kwargs + ) def get_atom_value(self, atom: Chem.rdchem.Atom) -> int: """ @@ -65,8 +71,11 @@ class AtomCharge(AtomProperty): Uses a one-hot encoder by default. """ - def __init__(self, encoder: PropertyEncoder | None = None) -> None: - super().__init__(encoder or OneHotEncoder(self)) + def __init__(self, encoder: PropertyEncoder | None = None, **kwargs) -> None: + data_type = kwargs.get("data_type") + super().__init__( + encoder=encoder or OneHotEncoder(self, data_type=data_type), **kwargs + ) def get_atom_value(self, atom: Chem.rdchem.Atom) -> int: """ @@ -88,8 +97,11 @@ class AtomChirality(AtomProperty): Uses a one-hot encoder by default. """ - def __init__(self, encoder: PropertyEncoder | None = None) -> None: - super().__init__(encoder or OneHotEncoder(self)) + def __init__(self, encoder: PropertyEncoder | None = None, **kwargs) -> None: + data_type = kwargs.get("data_type") + super().__init__( + encoder=encoder or OneHotEncoder(self, data_type=data_type), **kwargs + ) def get_atom_value(self, atom: Chem.rdchem.Atom) -> Chem.rdchem.ChiralType: """ @@ -111,8 +123,11 @@ class AtomHybridization(AtomProperty): Uses a one-hot encoder by default. """ - def __init__(self, encoder: PropertyEncoder | None = None) -> None: - super().__init__(encoder or OneHotEncoder(self)) + def __init__(self, encoder: PropertyEncoder | None = None, **kwargs) -> None: + data_type = kwargs.get("data_type") + super().__init__( + encoder=encoder or OneHotEncoder(self, data_type=data_type), **kwargs + ) def get_atom_value(self, atom: Chem.rdchem.Atom) -> Chem.rdchem.HybridizationType: """ @@ -134,8 +149,11 @@ class AtomNumHs(AtomProperty): Uses a one-hot encoder by default. """ - def __init__(self, encoder: PropertyEncoder | None = None) -> None: - super().__init__(encoder or OneHotEncoder(self)) + def __init__(self, encoder: PropertyEncoder | None = None, **kwargs) -> None: + data_type = kwargs.get("data_type") + super().__init__( + encoder=encoder or OneHotEncoder(self, data_type=data_type), **kwargs + ) def get_atom_value(self, atom: Chem.rdchem.Atom) -> int: """ @@ -157,8 +175,8 @@ class AtomAromaticity(AtomProperty): Uses a boolean encoder by default. """ - def __init__(self, encoder: PropertyEncoder | None = None) -> None: - super().__init__(encoder or BoolEncoder(self)) + def __init__(self, encoder: PropertyEncoder | None = None, **kwargs) -> None: + super().__init__(encoder=encoder or BoolEncoder(self), **kwargs) def get_atom_value(self, atom: Chem.rdchem.Atom) -> bool: """ @@ -180,8 +198,8 @@ class BondAromaticity(BondProperty): Uses a boolean encoder by default. """ - def __init__(self, encoder: PropertyEncoder | None = None) -> None: - super().__init__(encoder or BoolEncoder(self)) + def __init__(self, encoder: PropertyEncoder | None = None, **kwargs) -> None: + super().__init__(encoder=encoder or BoolEncoder(self), **kwargs) def get_bond_value(self, bond: Chem.rdchem.Bond) -> bool: """ @@ -203,8 +221,11 @@ class BondType(BondProperty): Uses a one-hot encoder by default. """ - def __init__(self, encoder: PropertyEncoder | None = None) -> None: - super().__init__(encoder or OneHotEncoder(self)) + def __init__(self, encoder: PropertyEncoder | None = None, **kwargs) -> None: + data_type = kwargs.get("data_type") + super().__init__( + encoder=encoder or OneHotEncoder(self, data_type=data_type), **kwargs + ) def get_bond_value(self, bond: Chem.rdchem.Bond) -> Chem.rdchem.BondType: """ @@ -226,8 +247,8 @@ class BondInRing(BondProperty): Uses a boolean encoder by default. """ - def __init__(self, encoder: PropertyEncoder | None = None) -> None: - super().__init__(encoder or BoolEncoder(self)) + def __init__(self, encoder: PropertyEncoder | None = None, **kwargs) -> None: + super().__init__(encoder=encoder or BoolEncoder(self), **kwargs) def get_bond_value(self, bond: Chem.rdchem.Bond) -> bool: """ @@ -249,8 +270,11 @@ class MoleculeNumRings(MoleculeProperty): Uses a one-hot encoder by default. """ - def __init__(self, encoder: PropertyEncoder | None = None) -> None: - super().__init__(encoder or OneHotEncoder(self)) + def __init__(self, encoder: PropertyEncoder | None = None, **kwargs) -> None: + data_type = kwargs.get("data_type") + super().__init__( + encoder=encoder or OneHotEncoder(self, data_type=data_type), **kwargs + ) def get_property_value(self, mol: Chem.rdchem.Mol) -> list[int]: """ @@ -272,8 +296,8 @@ class RDKit2DNormalized(MoleculeProperty): Uses an identity encoder by default. """ - def __init__(self, encoder: PropertyEncoder | None = None) -> None: - super().__init__(encoder or AsIsEncoder(self)) + def __init__(self, encoder: PropertyEncoder | None = None, **kwargs) -> None: + super().__init__(encoder=encoder or AsIsEncoder(self), **kwargs) self.generator_normalized = rdNormalizedDescriptors.RDKit2DNormalized() # Create a dummy molecule (e.g., methane) to extract the length of descriptor vector dummy_mol = Chem.MolFromSmiles("C") diff --git a/chebai_graph/preprocessing/property_encoder.py b/chebai_graph/preprocessing/property_encoder.py index 38cb279..5daeb4c 100644 --- a/chebai_graph/preprocessing/property_encoder.py +++ b/chebai_graph/preprocessing/property_encoder.py @@ -99,12 +99,15 @@ class IndexEncoder(PropertyEncoder): **kwargs: Additional keyword arguments. """ - def __init__(self, property, indices_dir: str | None = None, **kwargs) -> None: + def __init__( + self, property, data_type: str, indices_dir: str | None = None, **kwargs + ) -> None: super().__init__(property, **kwargs) if indices_dir is None: indices_dir = os.path.dirname(inspect.getfile(self.__class__)) self.dirname = indices_dir # load already existing cache + self._data_type = data_type with open(self.index_path, "r") as pk: self.cache: dict[str, int] = { token.strip(): idx for idx, token in enumerate(pk) @@ -126,12 +129,15 @@ def index_path(self) -> str: Returns: Path to index file. """ + assert self._data_type is not None, "data_type must be set for IndexEncoder" index_path = os.path.join( - self.dirname, "bin", self.property.name, f"indices_{self.name}.txt" - ) - os.makedirs( - os.path.join(self.dirname, "bin", self.property.name), exist_ok=True + self.dirname, + "bin", + self._data_type, + self.property.name, + f"indices_{self.name}.txt", ) + os.makedirs(os.path.dirname(index_path), exist_ok=True) if not os.path.exists(index_path): with open(index_path, "x"): pass