Dataset Architecture#
AtomWorks provides a modern, composable dataset architecture that separates data loading, processing, and transformation concerns. This approach replaces the legacy parser-based system with functional loaders and transform pipelines.
Warning
The metadata parser system (atomworks.ml.datasets.parsers) is deprecated and will be removed in a future version.
Use the new loader-based approach with FileDataset and PandasDataset instead.
Modern Dataset Architecture#
The current AtomWorks dataset system consists of three main components:
Datasets: Container classes that manage data access and indexing
Loaders: Functions that process raw data into transform-ready format
Transforms: Pipelines that convert loaded data into model inputs
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:
objectMemory-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.
- class atomworks.ml.datasets.ConcatDatasetWithID(datasets: list[ExampleIDProtocol])[source]#
Bases:
ConcatDatasetEquivalent to
torch.utils.data.ConcatDatasetbut allows accessing examples by ID.Provides ID-based access across multiple datasets that implement
ExampleIDProtocol.- cumulative_sizes: list[int]#
- 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.
- class atomworks.ml.datasets.ExampleIDProtocol(*args, **kwargs)[source]#
Bases:
ProtocolStructural interface for datasets that support ID-based access.
- class atomworks.ml.datasets.FallbackDatasetWrapper(dataset: Dataset, fallback_dataset: Dataset)[source]#
Bases:
DatasetA 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:
MolecularDatasetDataset 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.
- 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:
objectParquet/DataFrame metadata with filtering and ID lookup.
Single source of truth for example IDs across all dataset backends. The
id_columnis set as the pandas index for O(1).loclookups.
- class atomworks.ml.datasets.MetadataIndexProtocol(*args, **kwargs)[source]#
Bases:
ProtocolStructural interface shared by all metadata index backends.
- 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:
DatasetBase 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:
MolecularDatasetDataset 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.DataFrameforMetadataIndexor apa.TableforArrowMetadataIndex.- Raises:
AttributeError – If the metadata backend has no
dataattribute (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:
objectIdentity metadata mapping
"0"…"n-1"to sequential indices.Same interface as
MetadataIndex/ArrowMetadataIndexfor datasets that don’t need an external metadata file.
- 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#
Loaders are functions that process raw dataset output (e.g., pandas Series) into a Transform-ready format. They replace the legacy parser classes with a more flexible, functional approach.
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 withSTANDARD_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 withSTANDARD_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" ... )
Basic Usage Examples#
File-based datasets (replacing simple file parsers):
from atomworks.ml.datasets import FileDataset
from atomworks.io import parse
def simple_loading_fn(raw_data) -> dict:
"""Simple loading function that parses structural data."""
parse_output = parse(raw_data)
return {"atom_array": parse_output["assemblies"]["1"][0]}
dataset = FileDataset.from_directory(
directory="/path/to/structures",
name="my_dataset",
loader=simple_loading_fn
)
Tabular datasets (replacing metadata parsers):
from atomworks.ml.datasets import PandasDataset
from atomworks.ml.datasets.loaders import create_loader_with_query_pn_units
dataset = PandasDataset(
data="metadata.parquet",
name="interfaces_dataset",
loader=create_loader_with_query_pn_units(
pn_unit_iid_colnames=["pn_unit_1_iid", "pn_unit_2_iid"]
)
)
Custom loaders for specialized use cases:
def custom_loader(row: pd.Series) -> dict:
"""Custom loader with specific processing logic."""
# Load structure
structure_path = Path(row["path"])
parse_output = parse(structure_path)
# Extract specific metadata
metadata = {
"resolution": row.get("resolution", None),
"method": row.get("method", "unknown"),
"custom_field": row.get("custom_field", "default_value")
}
return {
"atom_array": parse_output["assemblies"]["1"][0],
"extra_info": metadata,
"example_id": row["example_id"]
}
dataset = PandasDataset(
data=my_dataframe,
name="custom_dataset",
loader=custom_loader
)
Common Loader Patterns#
Base loader for standard structure loading:
from atomworks.ml.datasets.loaders import create_base_loader
loader = create_base_loader(
example_id_colname="example_id",
path_colname="path",
assembly_id_colname="assembly_id",
base_path="/data/structures",
extension=".cif"
)
Interface loader for protein-protein interfaces:
from atomworks.ml.datasets.loaders import create_loader_with_query_pn_units
loader = create_loader_with_query_pn_units(
pn_unit_iid_colnames=["pn_unit_1_iid", "pn_unit_2_iid"],
base_path="/data/pdb",
extension=".cif.gz"
)
Validation loader with scoring targets:
from atomworks.ml.datasets.loaders import create_loader_with_interfaces_and_pn_units_to_score
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"
)
Integration with Transform Pipelines#
Loaders work seamlessly with AtomWorks transform pipelines. The loader output becomes the input to the transform pipeline:
from atomworks.ml.transforms.base import Compose
from atomworks.ml.transforms.crop import CropSpatialLikeAF3
from atomworks.ml.transforms.atom_array import AddGlobalAtomIdAnnotation
# Create a transform pipeline
transform_pipeline = Compose([
AddGlobalAtomIdAnnotation(),
CropSpatialLikeAF3(crop_size=256),
])
# Create dataset with both loader and transforms
dataset = PandasDataset(
data="metadata.parquet",
name="my_dataset",
loader=loader_with_query_pn_units(
pn_unit_iid_colnames=["pn_unit_1_iid", "pn_unit_2_iid"]
),
transform=transform_pipeline
)
# Access processed data
example = dataset[0] # Returns transformed data ready for model input
Data Flow#
The complete data flow in the new architecture is:
Raw Data: File paths or DataFrame rows
Loader: Processes raw data into standardized format with
AtomArrayTransform Pipeline: Converts loaded data into model-ready tensors
Model Input: Final processed data ready for training/inference
This separation allows for: - Reusable loaders across different datasets - Composable transforms that can be mixed and matched - Easy testing of individual components - Clear debugging when issues arise