Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
2 changes: 1 addition & 1 deletion chebai_graph/preprocessing/datasets/base.py
Original file line number Diff line number Diff line change
Expand Up @@ -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 = []
Expand Down
10 changes: 7 additions & 3 deletions chebai_graph/preprocessing/datasets/utils.py
Original file line number Diff line number Diff line change
Expand Up @@ -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.
Expand All @@ -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.
Expand All @@ -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)
40 changes: 26 additions & 14 deletions chebai_graph/preprocessing/properties/augmented_properties.py
Original file line number Diff line number Diff line change
Expand Up @@ -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:
"""
Expand All @@ -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:
"""
Expand All @@ -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:
"""
Expand Down Expand Up @@ -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] = {
Expand All @@ -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] = {
Expand All @@ -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:
"""
Expand Down Expand Up @@ -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:
"""
Expand Down
14 changes: 10 additions & 4 deletions chebai_graph/preprocessing/properties/base.py
Original file line number Diff line number Diff line change
Expand Up @@ -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:
Expand Down Expand Up @@ -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)
Expand Down
72 changes: 48 additions & 24 deletions chebai_graph/preprocessing/properties/properties.py
Original file line number Diff line number Diff line change
Expand Up @@ -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:
"""
Expand All @@ -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:
"""
Expand All @@ -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:
"""
Expand All @@ -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:
"""
Expand All @@ -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:
"""
Expand All @@ -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:
"""
Expand All @@ -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:
"""
Expand All @@ -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:
"""
Expand All @@ -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:
"""
Expand All @@ -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:
"""
Expand All @@ -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]:
"""
Expand All @@ -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")
Expand Down
16 changes: 11 additions & 5 deletions chebai_graph/preprocessing/property_encoder.py
Original file line number Diff line number Diff line change
Expand Up @@ -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)
Expand All @@ -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
Expand Down
Loading