Skip to content

Data API

The public data contracts are also exported from torch_batteries.

torch_batteries.data

Event-driven dataset and DataLoader construction.

Public API

  • DataPack — base contract for charged data lifecycle methods.
  • DataPackHandler — discovers and dispatches charged DataPack methods.
  • DatasetBundle and DataLoaderBundle — resolved data containers.
  • DataLoaderConfig — validated DataLoader construction options.
  • DataContext and ResolvedData — workflow context and resolution result.

DataPack Source

Base class for charged dataset and DataLoader configuration.

Subclasses define data lifecycle methods with :func:torch_batteries.charge. The default checkpoint contract is stateless; subclasses may override :meth:state_dict and :meth:load_state_dict when dataset construction relies on persistent values such as split indices or streaming positions.

resolve(stage, *, device='cpu') Source

Resolve datasets and DataLoaders without constructing a Battery.

Parameters:

Name Type Description Default
stage DataStage

Complete workflow stage to resolve: "fit", "test", or "predict".

required
device str | device

PyTorch device used for device-aware loader defaults. Standalone resolution defaults to CPU and also accepts "auto".

'cpu'

Returns:

Type Description
AbstractContextManager[ResolvedData]

A context manager yielding ResolvedData and guaranteeing teardown.

Note

Preparation runs once per standalone call. Keep preparation idempotent.

state_dict() Source

Return state that should be stored in a full training checkpoint.

load_state_dict(state_dict) Source

Restore state previously returned by :meth:state_dict.

Parameters:

Name Type Description Default
state_dict dict[str, Any]

DataPack-specific checkpoint state.

required

DataPackHandler Source

Bases: _ChargedHandlerBase

Discover and dispatch lifecycle methods charged on one DataPack.

Parameters:

Name Type Description Default
data_pack DataPack

DataPack whose charged lifecycle methods are discovered.

required

has_handler(event) Source

Return whether the DataPack handles an event.

Parameters:

Name Type Description Default
event Event

Data lifecycle event to inspect.

required

Returns:

Type Description
bool

True when at least one charged handler is registered.

call(event, context) Source

Call all ordered handlers for a side-effect data event.

Parameters:

Name Type Description Default
event Event

Side-effect event to dispatch.

required
context DataContext

Data lifecycle context passed to handlers.

required

provide(event, context, *, default) Source

Return a provider result or the supplied default.

Parameters:

Name Type Description Default
event Event

Provider event to dispatch.

required
context DataContext

Data lifecycle context passed to the provider.

required
default Any

Value returned when no provider is registered.

required

setup(context) Source

Construct and validate datasets for one workflow invocation.

Parameters:

Name Type Description Default
context DataContext

Setup context passed to the dataset provider.

required

build_loader(context, dataset) Source

Resolve a custom loader or materialize a DataLoaderConfig.

Parameters:

Name Type Description Default
context DataContext

Loader configuration context.

required
dataset DatasetType

Dataset for which a loader is required.

required

resolve(stage, *, device='cpu', battery=None, dataset_name=None) Source

Resolve one DataPack stage and guarantee workflow teardown.

Parameters:

Name Type Description Default
stage DataStage

"fit", "test", or "predict".

required
device str | device

Device used for loader configuration.

'cpu'
battery Battery | None

Optional owning Battery included in event contexts.

None
dataset_name str | None

Optional named test or prediction dataset selection.

None

Yields:

Type Description
Generator[ResolvedData]

Resolved datasets and loaders for the requested stage.

DataContext Source

Bases: TypedDict

Context passed to methods charged for DataPack lifecycle events.

Every event receives data_pack, stage, and device. Battery-managed workflows additionally receive battery. stage identifies the workflow as "fit", "test", or "predict". A configured DataPack seed adds seed and a fresh generator initialized with that seed. Loader configuration additionally receives phase, datasets, dataset, and dataset_name. Teardown receives datasets only when setup succeeded.

DataLoaderBundle Source dataclass

DataLoaders resolved for one DataPack stage.

Training and validation contain at most one loader. Test and prediction retain whether their datasets were configured as a bare value or a named mapping.

__post_init__() Source

Validate every configured loader against its phase contract.

for_phase(phase) Source

Return the loader or named loaders configured for a workflow phase.

Parameters:

Name Type Description Default
phase DataPhase

"train", "validation", "test", or "predict".

required

loaders_for_phase(phase) Source

Return phase loaders normalized to a mapping.

Parameters:

Name Type Description Default
phase DataPhase

Workflow phase to normalize.

required

DataLoaderConfig Source dataclass

Validated high-level configuration used to construct a DataLoader.

shuffle=None selects the phase default and pin_memory="auto" lets the runtime select pinning from the Battery device. Setting batch_sampler requires batch_size=None and conflicts with shuffle, sampler, and drop-last options.

__post_init__() Source

Reject types and combinations PyTorch cannot materialize safely.

DatasetBundle Source dataclass

Datasets made available by a charged SETUP_DATA provider.

Training and validation accept one PyTorch dataset. Test and prediction also accept a non-empty mapping of non-blank names to PyTorch datasets.

__post_init__() Source

Validate every configured dataset against its phase contract.

for_phase(phase) Source

Return the dataset or named datasets configured for a workflow phase.

Parameters:

Name Type Description Default
phase DataPhase

"train", "validation", "test", or "predict".

required

datasets_for_phase(phase) Source

Return phase datasets normalized to a mapping.

Parameters:

Name Type Description Default
phase DataPhase

Workflow phase to normalize.

required

ResolvedData Source dataclass

Datasets and DataLoaders materialized for one DataPack stage.