Source code for atomworks.ml.datasets.loaders.cif

"""CIF-based dataset loaders."""

import functools
import io
from collections.abc import Callable
from pathlib import Path
from typing import Any

import pandas as pd

from atomworks.io.config import ParseConfig
from atomworks.io.parser import STANDARD_PARSER_ARGS, parse
from atomworks.io.utils.io_utils import infer_pdb_file_type

from .base import _construct_metadata_hierarchy, _construct_structure_path


def _resolve_parse_config(
    parser_args: ParseConfig | dict | None,
    assembly_id: str,
    source: Path | io.BytesIO | io.StringIO,
    file_type: str | None = None,
    altloc_seed: int | None = None,
) -> ParseConfig:
    """Resolve ``parser_args`` to a :py:class:`ParseConfig`, stamping ``build_assembly`` from ``assembly_id`` for CIF/bCIF inputs.

    Accepts a ``ParseConfig`` (preferred) or a legacy dict of kwargs (merged
    with :py:data:`STANDARD_PARSER_ARGS` via :py:meth:`ParseConfig.from_dict`).

    When ``altloc_seed`` is provided, altloc selection is switched to
    ``"random_clash_aware"`` with the given seed for deterministic sampling.
    """
    if not isinstance(parser_args, ParseConfig):
        # Legacy dict support: merge with STANDARD_PARSER_ARGS and convert to ParseConfig
        parser_args = ParseConfig.from_dict({**STANDARD_PARSER_ARGS, **(parser_args or {})})

    if file_type is not None:
        # Overwrite file_type in parser_args if explicitly provided (e.g. for CIF bytes loader)
        parser_args = parser_args.replace(file_type=file_type)

    if (parser_args.file_type or infer_pdb_file_type(source)) in ("cif", "bcif"):
        parser_args = parser_args.replace(build_assembly=(assembly_id,))

        if altloc_seed is not None:
            # For CIF inputs, if altloc_seed is provided, switch to random_clash_aware altloc selection with the given seed
            parser_args = parser_args.replace(altloc="random_clash_aware", altloc_seed=altloc_seed)

    return parser_args


def _base_loader_function(
    row: pd.Series,
    example_id_colname: str,
    path_colname: str,
    assembly_id_colname: str | None,
    altloc_seed_colname: str | None,
    attrs: dict,
    base_path: str,
    extension: str,
    sharding_pattern: str | None,
    parser_args: ParseConfig | dict | None,
) -> dict[str, Any]:
    """Base loader function (picklable when used with functools.partial)."""
    # Prepare loader-specific attributes
    loader_attrs = attrs.copy()
    if base_path and "base_path" not in loader_attrs:
        loader_attrs["base_path"] = base_path
    if extension and "extension" not in loader_attrs:
        loader_attrs["extension"] = extension

    extra_info = _construct_metadata_hierarchy(row, loader_attrs)

    assembly_id = row[assembly_id_colname] if assembly_id_colname is not None and assembly_id_colname in row else "1"
    altloc_seed = row[altloc_seed_colname] if altloc_seed_colname is not None else None

    path = _construct_structure_path(
        row[path_colname], extra_info.get("base_path"), extra_info.get("extension"), sharding_pattern
    )
    result_dict = parse(path, config=_resolve_parse_config(parser_args, assembly_id, path, altloc_seed=altloc_seed))

    # Remove used columns from extra_info
    exclude_cols = (
        [example_id_colname, path_colname]
        + ([assembly_id_colname] if assembly_id_colname else [])
        + ([altloc_seed_colname] if altloc_seed_colname else [])
        + ["base_path", "extension"]
    )
    extra_info = {k: v for k, v in extra_info.items() if k not in exclude_cols}

    return {
        "example_id": row[example_id_colname],
        "path": path,
        "assembly_id": assembly_id,
        "altloc_seed": altloc_seed,
        "extra_info": extra_info,
        "atom_array": result_dict["assemblies"][assembly_id][0],
        "atom_array_stack": result_dict["assemblies"][assembly_id],
        "chain_info": result_dict["chain_info"],
        "ligand_info": result_dict["ligand_info"],
        "metadata": result_dict["metadata"],
    }


[docs] def create_base_loader( example_id_colname: str = "example_id", path_colname: str = "path", assembly_id_colname: str | None = "assembly_id", altloc_seed_colname: str | None = None, attrs: dict | None = None, base_path: str = "", extension: str = "", sharding_pattern: str | None = None, parser_args: ParseConfig | dict | None = None, ) -> Callable[[pd.Series], dict[str, Any]]: """Factory function that creates a picklable base loader for AtomWorks datasets. Args: example_id_colname: Name of column containing unique example identifiers path_colname: Name of column containing paths to structure files assembly_id_colname: Optional column name containing assembly IDs. If None, assembly_id defaults to "1" for all examples. attrs: Additional attributes to merge with highest precedence into the metadata hierarchy (and ultimately included in the output dictionary's "extra_info" key). base_path: Base path to prepend to file paths if not included in path column extension: File extension to add/replace if not included in path column sharding_pattern: Pattern for how files are organized in subdirectories, if not specified in the path - "/1:2/": Use characters 1-2 for first directory level - "/1:2/0:2/": Use chars 1-2 for first dir, then chars 0-2 for second dir - None: No sharding (default) parser_args: Parser arguments — a :py:class:`ParseConfig` (preferred) or a legacy dict of kwargs merged with :py:data:`STANDARD_PARSER_ARGS`. Returns: A picklable loader function (via functools.partial) for multiprocessing. """ return functools.partial( _base_loader_function, example_id_colname=example_id_colname, path_colname=path_colname, assembly_id_colname=assembly_id_colname, altloc_seed_colname=altloc_seed_colname, attrs=attrs or {}, base_path=base_path, extension=extension, sharding_pattern=sharding_pattern, parser_args=parser_args, )
def _loader_with_query_pn_units_function( row: pd.Series, base_loader: Callable, pn_unit_iid_colnames: list[str], ) -> dict[str, Any]: """Loader with query pn_units (picklable when used with functools.partial).""" result = base_loader(row) result["extra_info"] = {k: v for k, v in result["extra_info"].items() if k not in pn_unit_iid_colnames} query_pn_unit_iids = [row[colname] for colname in pn_unit_iid_colnames] result["query_pn_unit_iids"] = query_pn_unit_iids return result
[docs] def create_loader_with_query_pn_units( example_id_colname: str = "example_id", path_colname: str = "path", pn_unit_iid_colnames: str | list[str] | None = None, assembly_id_colname: str | None = "assembly_id", altloc_seed_colname: str | None = None, base_path: str = "", extension: str = "", sharding_pattern: str | None = None, attrs: dict | None = None, parser_args: ParseConfig | dict | None = None, ) -> Callable[[pd.Series], dict[str, Any]]: """Factory function that creates a picklable loader for pipelines with query pn_units (chains). For instance, in the interfaces dataset, each sampled row contains two pn_unit instance IDs that should be included in the cropped structure. Examples: Interfaces dataset: >>> loader = create_loader_with_query_pn_units( ... pn_unit_iid_colnames=["pn_unit_1_iid", "pn_unit_2_iid"], assembly_id_colname="assembly_id" ... ) Chains dataset: >>> loader = create_loader_with_query_pn_units( ... pn_unit_iid_colnames="pn_unit_iid", base_path="/data/structures", extension=".cif.gz" ... ) """ # Normalize pn_unit_iid_colnames to list format if isinstance(pn_unit_iid_colnames, str): pn_unit_iid_colnames = [pn_unit_iid_colnames] pn_unit_iid_colnames = pn_unit_iid_colnames or [] # Create base loader with common parameters base_loader = create_base_loader( example_id_colname=example_id_colname, path_colname=path_colname, assembly_id_colname=assembly_id_colname, altloc_seed_colname=altloc_seed_colname, attrs=attrs, base_path=base_path, extension=extension, sharding_pattern=sharding_pattern, parser_args=parser_args, ) return functools.partial( _loader_with_query_pn_units_function, base_loader=base_loader, pn_unit_iid_colnames=pn_unit_iid_colnames, )
def _loader_with_interfaces_and_pn_units_to_score_function( row: pd.Series, base_loader: Callable, interfaces_to_score_colname: str | None, pn_units_to_score_colname: str | None, ) -> dict[str, Any]: """Loader with scoring info (picklable when used with functools.partial).""" result = base_loader(row) exclude_cols = [interfaces_to_score_colname, pn_units_to_score_colname] result["extra_info"] = {k: v for k, v in result["extra_info"].items() if k not in exclude_cols} interfaces_to_score = row[interfaces_to_score_colname] if interfaces_to_score_colname is not None else None pn_units_to_score = row[pn_units_to_score_colname] if pn_units_to_score_colname is not None else None result.update( { "interfaces_to_score": interfaces_to_score, "pn_units_to_score": pn_units_to_score, } ) return result
[docs] def create_loader_with_interfaces_and_pn_units_to_score( example_id_colname: str = "example_id", path_colname: str = "path", assembly_id_colname: str | None = "assembly_id", altloc_seed_colname: str | None = None, interfaces_to_score_colname: str | None = "interfaces_to_score", pn_units_to_score_colname: str | None = "pn_units_to_score", base_path: str = "", extension: str = "", sharding_pattern: str | None = None, attrs: dict | None = None, parser_args: ParseConfig | dict | None = None, ) -> Callable[[pd.Series], dict[str, Any]]: """Factory function that creates a picklable loader for validation datasets with scoring information. Example: >>> loader = create_loader_with_interfaces_and_pn_units_to_score( ... interfaces_to_score_colname="interfaces_to_score", pn_units_to_score_colname="pn_units_to_score" ... ) """ # Create base loader with common parameters base_loader = create_base_loader( example_id_colname=example_id_colname, path_colname=path_colname, assembly_id_colname=assembly_id_colname, altloc_seed_colname=altloc_seed_colname, attrs=attrs, base_path=base_path, extension=extension, sharding_pattern=sharding_pattern, parser_args=parser_args, ) return functools.partial( _loader_with_interfaces_and_pn_units_to_score_function, base_loader=base_loader, interfaces_to_score_colname=interfaces_to_score_colname, pn_units_to_score_colname=pn_units_to_score_colname, )
def _cif_bytes_loader_function( raw_data: tuple, parser_args: ParseConfig | dict | None, assembly_id_colname: str | None, altloc_seed_colname: str | None = None, ) -> dict[str, Any]: """Loader for CIF bytes (picklable when used with functools.partial).""" cif_bytes, global_idx, metadata_row = raw_data # Extract assembly_id from metadata row when available assembly_id = "1" if metadata_row is not None and assembly_id_colname is not None and assembly_id_colname in metadata_row.index: assembly_id = str(metadata_row[assembly_id_colname]) altloc_seed = None if metadata_row is not None and altloc_seed_colname is not None and altloc_seed_colname in metadata_row.index: altloc_seed = int(metadata_row[altloc_seed_colname]) source = io.StringIO(cif_bytes.decode("utf-8")) result = parse( source, config=_resolve_parse_config(parser_args, assembly_id, source, file_type="cif", altloc_seed=altloc_seed), ) # Build extra_info from metadata row extra_info: dict[str, Any] = {} if metadata_row is not None: extra_info = _construct_metadata_hierarchy(metadata_row, {}) # Remove columns already represented in the output exclude_cols = {"example_id"} if assembly_id_colname: exclude_cols.add(assembly_id_colname) extra_info = {k: v for k, v in extra_info.items() if k not in exclude_cols} return { "assembly_id": assembly_id, "extra_info": extra_info, "atom_array": result["assemblies"][assembly_id][0], "atom_array_stack": result["assemblies"][assembly_id], "chain_info": result["chain_info"], "ligand_info": result["ligand_info"], "metadata": result["metadata"], }
[docs] def create_cif_bytes_loader( parser_args: ParseConfig | dict | None = None, assembly_id_colname: str | None = "assembly_id", altloc_seed_colname: str | None = None, ) -> Callable: """Factory for loading CIF bytes (e.g. from LMDB) via :py:func:`~atomworks.io.parser.parse`. Args: parser_args: Parser arguments — a :py:class:`ParseConfig` (preferred) or a legacy dict of kwargs merged with :py:data:`STANDARD_PARSER_ARGS`. assembly_id_colname: Column name in the metadata row containing the assembly ID. If ``None``, assembly ID defaults to ``"1"``. """ return functools.partial( _cif_bytes_loader_function, parser_args=parser_args, assembly_id_colname=assembly_id_colname, altloc_seed_colname=altloc_seed_colname, )