Source code for atomworks.ml.datasets.metadata

"""Metadata index for filtering and ID-mapping in molecular datasets."""

import contextlib
import logging
import os
import time
from os import PathLike
from pathlib import Path
from typing import Any, Protocol, runtime_checkable

import numpy as np
import pandas as pd
import pyarrow as pa
import pyarrow.compute as pc
import pyarrow.feather as feather
from filelock import FileLock

from atomworks.common import as_list
from atomworks.constants import NA_VALUES
from atomworks.ml.utils.io import read_parquet_with_metadata

logger = logging.getLogger("datasets")


[docs] @runtime_checkable class MetadataIndexProtocol(Protocol): """Structural interface shared by all metadata index backends.""" def __len__(self) -> int: ... def __contains__(self, example_id: str) -> bool: ...
[docs] def id_to_idx(self, example_id: str | list[str]) -> int | list[int]: ...
[docs] def idx_to_id(self, idx: int | list[int]) -> str | np.ndarray: ...
[docs] def get_row(self, idx: int) -> pd.Series: ...
[docs] def get_column_values(self, column: str) -> np.ndarray: ...
[docs] def get_example_id(self, idx: int) -> str: ...
[docs] class MetadataIndex: """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. """ def __init__( self, *, data: pd.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, ): """Initialize MetadataIndex. Args: data: Either a pandas DataFrame or path to a CSV/Parquet file. name: Descriptive name for logging. id_column: Column to use as the index for example ID lookups. filters: Optional list of pandas query strings applied sequentially. columns_to_load: Optional list of columns to load from file. load_kwargs: Additional keyword arguments for pandas read functions. """ if isinstance(data, PathLike | str): data = _load_from_path(data, columns_to_load, **(load_kwargs or {})) assert id_column in data.columns, f"Column {id_column} not found. Available: {list(data.columns)}" if filters: data = _apply_filters(data, filters, name) self.name = name data.set_index(id_column, inplace=True, drop=False, verify_integrity=True) self.data: pd.DataFrame = data def __len__(self) -> int: return len(self.data) def __contains__(self, example_id: str) -> bool: return example_id in self.data.index
[docs] def id_to_idx(self, example_id: str | list[str]) -> int | list[int]: """Convert example ID(s) to positional index(es).""" if np.isscalar(example_id): return self.data.index.get_loc(example_id) return [self.data.index.get_loc(eid) for eid in example_id]
[docs] def idx_to_id(self, idx: int | list[int]) -> str | np.ndarray: """Convert positional index(es) to example ID(s).""" _return_single = False if np.isscalar(idx) or (isinstance(idx, np.ndarray) and idx.shape == ()): _return_single = True idx = idx.item() if isinstance(idx, np.ndarray) else idx idx = slice(idx, idx + 1) ids = self.data.iloc[idx].index.values return ids[0] if _return_single else ids
[docs] def get_row(self, idx: int) -> pd.Series: """Get metadata row by positional index.""" return self.data.iloc[idx]
[docs] def get_column_values(self, column: str) -> np.ndarray: """Get values of a column as numpy array.""" return np.array(self.data[column])
[docs] def get_example_id(self, idx: int) -> str: """Get example ID by positional index.""" return str(self.data.iloc[idx].name)
[docs] class ArrowMetadataIndex: """Memory-mapped Arrow metadata with the same interface as :class:`MetadataIndex`. Converts the DataFrame to an uncompressed feather file on local storage and reads it back as a memory-mapped :class:`pyarrow.Table`, keeping Python heap usage minimal. ID lookups are O(N) — acceptable because they are not on the DataLoader hot path. """ def __init__( self, *, data: pd.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", ): """Initialize ArrowMetadataIndex. Args: data: Either a pandas DataFrame or path to a CSV/Parquet file. name: Descriptive name for logging. id_column: Column to use as the ID column. filters: Optional list of pandas query strings applied sequentially. columns_to_load: Optional list of columns to load from file. load_kwargs: Additional keyword arguments for pandas read functions. local_drive_mount: Root directory for feather files. job_id_env_var: Environment variable used to namespace feather files per job (e.g. ``"SLURM_JOB_ID"``). """ self._id_column = id_column self.name = name self.feather_path = _build_feather_index( data=data, name=name, id_column=id_column, filters=filters, columns_to_load=columns_to_load, load_kwargs=load_kwargs, local_drive_mount=local_drive_mount, job_id_env_var=job_id_env_var, ) self.data: pa.Table = feather.read_table(self.feather_path, memory_map=True) assert id_column in self.data.column_names, f"Column {id_column} not found. Available: {self.data.column_names}" def __getstate__(self) -> dict: """Pickle-friendly: store path instead of the full table.""" return { "feather_path": self.feather_path, "_id_column": self._id_column, "name": self.name, } def __setstate__(self, state: dict) -> None: """Restore from pickle by re-opening the feather file.""" self.feather_path = state["feather_path"] self._id_column = state["_id_column"] self.name = state["name"] self.data = feather.read_table(self.feather_path, memory_map=True) def __len__(self) -> int: return self.data.num_rows def __contains__(self, example_id: str) -> bool: return pc.index(self.data.column(self._id_column), example_id).as_py() != -1
[docs] def id_to_idx(self, example_id: str | list[str]) -> int | list[int]: """Convert example ID(s) to positional index(es). O(N) per call.""" col = self.data.column(self._id_column) if np.isscalar(example_id): idx = pc.index(col, example_id).as_py() if idx == -1: raise KeyError(example_id) return idx mask = pc.is_in(col, value_set=pa.array(example_id)) matching_rows = np.where(mask.to_numpy())[0] matching_ids = col.take(matching_rows).to_pylist() id_to_row = {str(v): int(r) for v, r in zip(matching_ids, matching_rows, strict=False)} return [id_to_row[eid] for eid in example_id]
[docs] def idx_to_id(self, idx: int | list[int]) -> str | np.ndarray: """Convert positional index(es) to example ID(s).""" col = self.data.column(self._id_column) if np.isscalar(idx) or (isinstance(idx, np.ndarray) and idx.shape == ()): idx = idx.item() if isinstance(idx, np.ndarray) else idx return str(col[idx].as_py()) return np.array([str(col[i].as_py()) for i in idx])
[docs] def get_row(self, idx: int) -> pd.Series: """Get metadata row by positional index.""" return self.data.slice(idx, 1).to_pandas().iloc[0]
[docs] def get_column_values(self, column: str) -> np.ndarray: """Get values of a column as numpy array.""" return self.data.column(column).to_numpy(zero_copy_only=False)
[docs] def get_example_id(self, idx: int) -> str: """Get example ID by positional index.""" return str(self.data.column(self._id_column)[idx].as_py())
[docs] class SequentialMetadataIndex: """Identity metadata mapping ``"0"`` … ``"n-1"`` to sequential indices. Same interface as :class:`MetadataIndex` / :class:`ArrowMetadataIndex` for datasets that don't need an external metadata file. """ def __init__(self, *, n_entries: int, idx_column: str): self._n = n_entries self._idx_column = idx_column def __len__(self) -> int: return self._n def __contains__(self, example_id: str) -> bool: try: return 0 <= int(example_id) < self._n except (ValueError, TypeError): return False
[docs] def id_to_idx(self, example_id: str | list[str]) -> int | list[int]: """Convert string index(es) to int.""" if np.isscalar(example_id): idx = int(example_id) if not 0 <= idx < self._n: raise KeyError(example_id) return idx return [self.id_to_idx(eid) for eid in example_id]
[docs] def idx_to_id(self, idx: int | list[int]) -> str | np.ndarray: """Convert int index(es) to string.""" if np.isscalar(idx) or (isinstance(idx, np.ndarray) and idx.shape == ()): return str(idx.item() if isinstance(idx, np.ndarray) else idx) return np.array([str(i) for i in idx])
[docs] def get_column_values(self, column: str) -> np.ndarray: """Return sequential indices for ``idx_column``.""" if column != self._idx_column: raise KeyError(f"Column {column!r} not available without a metadata file") return np.arange(self._n)
[docs] def get_row(self, idx: int) -> pd.Series: """Not available — ``SequentialMetadataIndex`` has no row data.""" raise NotImplementedError( "SequentialMetadataIndex has no row data. " "Provide a metadata file to use get_row()." )
[docs] def get_example_id(self, idx: int) -> str: return str(idx)
def _build_feather_index( *, data: pd.DataFrame | PathLike | str, name: str, id_column: str, filters: list[str] | None, columns_to_load: list[str] | None, load_kwargs: dict | None, local_drive_mount: str, job_id_env_var: str, ) -> str: """Build once (under a lock) and return the path to a memory-mappable feather index.""" job_id = os.environ.get(job_id_env_var, f"manual_{os.getpid()}_{int(time.time())}") feather_path = os.path.join(local_drive_mount, job_id, f"{name}.feather") os.makedirs(os.path.dirname(feather_path), exist_ok=True) with FileLock(feather_path + ".lock"): if not os.path.exists(feather_path): # Read the source only here (feather missing) — not once per rank. if isinstance(data, PathLike | str): data = _load_from_path(data, columns_to_load, **(load_kwargs or {})) assert id_column in data.columns, f"Column {id_column} not found. Available: {list(data.columns)}" if filters: data = _apply_filters(data, filters, name) try: import torch.distributed as dist rank = dist.get_rank() if dist.is_available() and dist.is_initialized() else "N/A" except ImportError: rank = "N/A" logger.info(f"Rank {rank}: Converting {name} to feather at {feather_path}") tmp_path = f"{feather_path}.tmp.{os.getpid()}" try: feather.write_feather(pa.Table.from_pandas(data), tmp_path, compression="uncompressed") os.replace(tmp_path, feather_path) # atomic publish (same filesystem) except BaseException: with contextlib.suppress(OSError): os.unlink(tmp_path) # don't leave a partial temp behind on failure raise return feather_path def _load_from_path( path: PathLike | str, columns_to_load: list[str] | None = None, **load_kwargs: Any, ) -> pd.DataFrame: """Load DataFrame from CSV or Parquet file.""" path = Path(path) if columns_to_load is not None: columns_to_load = as_list(columns_to_load) if path.suffix == ".csv": return pd.read_csv( path, usecols=columns_to_load, keep_default_na=False, na_values=NA_VALUES, **load_kwargs, ) elif path.suffix == ".parquet": return read_parquet_with_metadata(path, columns=columns_to_load, **load_kwargs) else: raise ValueError(f"Unsupported file type: {path.suffix}") def _apply_filters( data: pd.DataFrame, filters: list[str], name: str, ) -> pd.DataFrame: """Apply pandas query filters sequentially with waterfall logging.""" initial_count = len(data) filter_results: list[tuple[str, int]] = [] for query in filters: original_count = len(data) data = data.query(query) filtered_count = len(data) if filtered_count == 0: raise ValueError(f"Query '{query}' on dataset {name} removed all rows.") rows_removed = original_count - filtered_count filter_results.append((query, rows_removed)) _log_filter_summary(name, initial_count, len(data), filter_results) return data def _log_filter_summary( name: str, initial_count: int, final_count: int, filter_results: list[tuple[str, int]], max_width: int = 120, ) -> None: """Log a waterfall summary of all applied filters.""" for query, removed in filter_results: if removed == 0: logger.warning(f"Query '{query}' on dataset {name} did not remove any rows.") total_removed = initial_count - final_count total_pct = (total_removed / initial_count) * 100 if initial_count > 0 else 0.0 header = f" {name}: {initial_count:,} \u2192 {final_count:,} rows ({total_pct:.1f}% removed) " bar_width = 22 # inner_width = max_width - 2 (for border chars) inner_width = max_width - 2 # Format: " <query> <count> <pct> <bar> " # 2 + query + 1 + 10 + 1 + 7 + 1 + bar + 2(padding) = inner_width max_query_width = inner_width - 2 - 1 - 10 - 1 - 7 - 1 - bar_width - 2 filter_lines: list[str] = [] for query, removed in filter_results: pct = (removed / initial_count) * 100 if initial_count > 0 else 0.0 filled = round(pct / 100 * bar_width) bar = "\u2588" * filled + "\u2591" * (bar_width - filled) count_str = f"-{removed:,}" pct_str = f"({pct:.1f}%)" if len(query) > max_query_width: query = query[: max_query_width - 3] + "..." filter_lines.append(f" {query:<{max_query_width}s} {count_str:>10s} {pct_str:>7s} {bar}") top = f"\u250c{header:\u2500<{inner_width}}\u2510" bottom = f"\u2514{'\u2500' * inner_width}\u2518" empty = f"\u2502{' ' * inner_width}\u2502" lines = [top, empty] for fl in filter_lines: lines.append(f"\u2502{fl:<{inner_width}}\u2502") lines += [empty, bottom] logger.info("\n" + "\n".join(lines))