Datasets#

This module contains dataset classes and utilities for loading and processing molecular data using a modern, composable architecture.

Core Dataset Classes#

class atomworks.ml.datasets.ArrowMetadataIndex(*, data: DataFrame | PathLike, name: str, id_column: str = 'example_id', filters: list[str] | None = None, columns_to_load: list[str] | None = None, load_kwargs: dict | None = None, local_drive_mount: str = '/tmp', job_id_env_var: str = 'SLURM_JOB_ID')[source]#

Bases: object

Memory-mapped Arrow metadata with the same interface as MetadataIndex.

Converts the DataFrame to an uncompressed feather file on local storage and reads it back as a memory-mapped pyarrow.Table, keeping Python heap usage minimal. ID lookups are O(N) — acceptable because they are not on the DataLoader hot path.

get_column_values(column: str) ndarray[source]#

Get values of a column as numpy array.

get_example_id(idx: int) str[source]#

Get example ID by positional index.

get_row(idx: int) Series[source]#

Get metadata row by positional index.

id_to_idx(example_id: str | list[str]) int | list[int][source]#

Convert example ID(s) to positional index(es). O(N) per call.

idx_to_id(idx: int | list[int]) str | ndarray[source]#

Convert positional index(es) to example ID(s).

class atomworks.ml.datasets.ConcatDatasetWithID(datasets: list[ExampleIDProtocol])[source]#

Bases: ConcatDataset

Equivalent to torch.utils.data.ConcatDataset but allows accessing examples by ID.

Provides ID-based access across multiple datasets that implement ExampleIDProtocol.

datasets: list[ExampleIDProtocol]#
get_dataset_by_id(example_id: str) Dataset[source]#

Retrieves the dataset containing the example ID.

Parameters:

example_id – The ID to find.

Returns:

The sub-dataset containing the ID.

Warning

Assumes that the example ID is unique within the dataset. If not, the first occurrence of the example ID is returned.

get_dataset_by_idx(idx: int) Dataset[source]#

Retrieves the dataset containing the index.

Parameters:

idx – The index to find.

Returns:

The sub-dataset containing the index.

Raises:

ValueError – If the index is out of bounds.

id_to_idx(example_id: str) int[source]#

Retrieves the index corresponding to the example ID.

Parameters:

example_id – The ID to convert.

Returns:

The corresponding index.

Raises:

ValueError – If the example ID is not found.

Warning

Assumes that the example ID is unique within the dataset. If not, the first occurrence of the example ID is returned.

idx_to_id(idx: int) str[source]#

Retrieves the example ID corresponding to the index.

Parameters:

idx – The index to convert.

Returns:

The corresponding example ID.

Raises:

ValueError – If the index is out of bounds.

class atomworks.ml.datasets.ExampleIDProtocol(*args, **kwargs)[source]#

Bases: Protocol

Structural interface for datasets that support ID-based access.

id_to_idx(example_id: str | list[str]) int | list[int][source]#

Convert example ID(s) to index(es).

idx_to_id(idx: int | list[int]) str | ndarray[source]#

Convert index(es) to example ID(s).

class atomworks.ml.datasets.FallbackDatasetWrapper(dataset: Dataset, fallback_dataset: Dataset)[source]#

Bases: Dataset

A wrapper around a dataset that allows for a fallback dataset to be used when an error occurs.

Meant to be used with a FallbackSamplerWrapper.

class atomworks.ml.datasets.FileDataset(*, file_paths: list[str | PathLike], name: str, filter_fn: Callable[[PathLike], bool] | None = None, **kwargs: Any)[source]#

Bases: MolecularDataset

Dataset that loads molecular data from individual files.

Each file represents one example in the dataset. If creating a dataset from a directory, use the from_directory() class method instead of the default constructor.

classmethod from_directory(*, directory: PathLike, name: str, max_depth: int = 3, **kwargs: Any) FileDataset[source]#

Create a FileDataset by scanning a directory for files.

Parameters:
  • directory – Path to directory to scan for files.

  • name – Descriptive name for this dataset.

  • max_depth – Maximum depth to scan for files in subdirectories.

  • **kwargs – Additional arguments passed to FileDataset.

Returns:

FileDataset instance with files discovered from the directory.

Example

Create from directory:
>>> dataset = FileDataset.from_directory(directory="/path/to/files", name="my_dataset", max_depth=2)
classmethod from_file_list(*, file_paths: list[str | PathLike], name: str, **kwargs: Any) FileDataset[source]#

Create a FileDataset from an explicit list of file paths.

This is an alias for the main constructor for clarity and consistency with from_directory().

Parameters:
  • file_paths – List of file paths for the dataset. Each file represents one example.

  • name – Descriptive name for this dataset.

  • **kwargs – Additional arguments passed to FileDataset.

Returns:

FileDataset instance with the provided file paths.

id_to_idx(example_id: str | list[str]) int | list[int][source]#

Convert example ID(s) to index(es).

idx_to_id(idx: int | list[int]) str | list[str][source]#

Convert index(es) to example ID(s).

class atomworks.ml.datasets.MetadataIndex(*, data: DataFrame | PathLike, name: str, id_column: str = 'example_id', filters: list[str] | None = None, columns_to_load: list[str] | None = None, load_kwargs: dict | None = None)[source]#

Bases: object

Parquet/DataFrame metadata with filtering and ID lookup.

Single source of truth for example IDs across all dataset backends. The id_column is set as the pandas index for O(1) .loc lookups.

get_column_values(column: str) ndarray[source]#

Get values of a column as numpy array.

get_example_id(idx: int) str[source]#

Get example ID by positional index.

get_row(idx: int) Series[source]#

Get metadata row by positional index.

id_to_idx(example_id: str | list[str]) int | list[int][source]#

Convert example ID(s) to positional index(es).

idx_to_id(idx: int | list[int]) str | ndarray[source]#

Convert positional index(es) to example ID(s).

class atomworks.ml.datasets.MetadataIndexProtocol(*args, **kwargs)[source]#

Bases: Protocol

Structural interface shared by all metadata index backends.

get_column_values(column: str) ndarray[source]#
get_example_id(idx: int) str[source]#
get_row(idx: int) Series[source]#
id_to_idx(example_id: str | list[str]) int | list[int][source]#
idx_to_id(idx: int | list[int]) str | ndarray[source]#
class atomworks.ml.datasets.MolecularDataset(*, name: str, transform: Callable | None = None, loader: Callable | None = None, save_failed_examples_to_dir: str | Path | None = None)[source]#

Bases: Dataset

Base class for AtomWorks molecular datasets.

Handles Transform pipelines and loader functionality for molecular data. Subclasses implement __getitem__() with their own data access patterns.

class atomworks.ml.datasets.PandasDataset(*, data: DataFrame | PathLike, name: str, id_column: str = 'example_id', filters: list[str] | None = None, columns_to_load: list[str] | None = None, transform: Callable | None = None, loader: Callable | None = None, save_failed_examples_to_dir: str | Path | None = None, load_kwargs: dict | tuple | None = None, memory_map: bool = False, metadata: MetadataIndexProtocol | None = None)[source]#

Bases: MolecularDataset

Dataset for tabular data stored as pandas DataFrames.

Delegates all metadata, filtering, and ID-mapping logic to MetadataIndex.

property data: DataFrame | Table#

The underlying metadata table.

Returns a pd.DataFrame for MetadataIndex or a pa.Table for ArrowMetadataIndex.

Raises:

AttributeError – If the metadata backend has no data attribute (e.g. SequentialMetadataIndex).

id_to_idx(example_id: str | list[str]) int | list[int][source]#

Convert an example ID to the corresponding local index.

idx_to_id(idx: int | list[int]) str | ndarray[source]#

Convert a local index to the corresponding example ID.

property metadata: MetadataIndexProtocol#

The metadata index.

class atomworks.ml.datasets.SequentialMetadataIndex(*, n_entries: int, idx_column: str)[source]#

Bases: object

Identity metadata mapping "0""n-1" to sequential indices.

Same interface as MetadataIndex / ArrowMetadataIndex for datasets that don’t need an external metadata file.

get_column_values(column: str) ndarray[source]#

Return sequential indices for idx_column.

get_example_id(idx: int) str[source]#
get_row(idx: int) Series[source]#

Not available — SequentialMetadataIndex has no row data.

id_to_idx(example_id: str | list[str]) int | list[int][source]#

Convert string index(es) to int.

idx_to_id(idx: int | list[int]) str | ndarray[source]#

Convert int index(es) to string.

atomworks.ml.datasets.get_row_and_index_by_example_id(dataset: ExampleIDProtocol, example_id: str) dict[source]#

Retrieve a row and its index from a nested dataset structure by its example ID.

Parameters:
  • dataset – The dataset or concatenated dataset to search. Must have the id_to_idx method.

  • example_id – The example ID to search for.

Returns:

Dictionary containing the row and global index corresponding to the example ID.

Functional Loaders#

Functional loader implementations for AtomWorks datasets.

Loaders are functions that process raw dataset output (e.g., pandas Series) into a Transform-ready format. E.g., converts what may be dataset-specific metadata into a standard format for use in AtomWorks Transform pipelines.

atomworks.ml.datasets.loaders.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[[Series], dict[str, Any]][source]#

Factory function that creates a picklable base loader for AtomWorks datasets.

Parameters:
  • 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 ParseConfig (preferred) or a legacy dict of kwargs merged with STANDARD_PARSER_ARGS.

Returns:

A picklable loader function (via functools.partial) for multiprocessing.

atomworks.ml.datasets.loaders.create_cif_bytes_loader(parser_args: ParseConfig | dict | None = None, assembly_id_colname: str | None = 'assembly_id', altloc_seed_colname: str | None = None) Callable[source]#

Factory for loading CIF bytes (e.g. from LMDB) via parse().

Parameters:
  • parser_args – Parser arguments — a ParseConfig (preferred) or a legacy dict of kwargs merged with STANDARD_PARSER_ARGS.

  • assembly_id_colname – Column name in the metadata row containing the assembly ID. If None, assembly ID defaults to "1".

atomworks.ml.datasets.loaders.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[[Series], dict[str, Any]][source]#

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"
... )
atomworks.ml.datasets.loaders.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[[Series], dict[str, Any]][source]#

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"
... )

Dataset Architecture and Migration Guide#