Data Loading and Checkpointing
Utilities for streaming training data from Unity Catalog volumes and checkpointing to them.
They are available in GPU environment version 6 and above. Elsewhere, install the data extra to pull
in their dependencies: pip install "databricks-sdk-air[data]".
UCVolumeDataset
- class databricks.air.data.UCVolumeDataset(path)
Bases:
IterableDatasetIterableDataset that streams files from a UC FUSE directory as local paths.
Files are cached locally on first access via a CachingUCFuseFSLayer. The file list is automatically partitioned so that each (rank, worker) pair receives a non-overlapping slice:
If
torch.distributedis initialised, files are first split across ranks.Within each rank, if multiple DataLoader workers are active the rank’s slice is further divided across those workers.
The combined effect is that with
Rranks andWDataLoader workers per rank there areR * Windependent streams, each covering1 / (R * W)of the files.Each dataset instance gets its own private cache directory. When the dataset is handed to a
DataLoaderthe worker processes are forked from this one instance and therefore share that single cache directory, which is what we want, since the workers each take a disjoint slice of the same shard and so never download the same file twice. Independently constructed datasets do not share a cache.The yielded values are local paths into that cache, and the cache evicts old files to bound disk usage (see
CachingUCFuseFSLayer). A yielded path is only guaranteed valid until the next item is pulled from the stream, so open/consume it immediately; do not stash paths to reopen later.Example:
dataset = UCVolumeDataset("/Volumes/my-catalog/my-schema/my-volume/data") loader = DataLoader(dataset, num_workers=4) for local_path in loader: ... # use local_path now; don't keep it for later
Checkpointing
The dataset implements
databricks.air.data.checkpoint.Checkpointable. Progress is tracked per (rank, worker) stream as the number of files that stream has delivered, because with multiple DataLoader workers each worker is a separate forked process advancing its own disjoint slice at its own rate; there is no single global cursor that describes them all.state_dict()reports the calling stream’s count (stamped with its(rank, worker_id));load_state_dict()seeds a per-worker resume map that the next__iter__uses to skip the already-delivered prefix of each stream’s slice, and rejects a blob whose stamped identity does not match the restoring stream. Because the counts are keyed to a specificworld_size/num_workersgrid, resuming is only defined on the same grid the checkpoint was taken on, anddatabricks.air.data.DataLoaderenforces that. Typically you checkpoint through theDataLoader(which harvests every worker’s state and restores it as a set); the dataset methods are the primitive that wrapper datasets delegate to.- type path:
str- param path:
UC FUSE directory to read from (must start with
/Volumes).
- load_state_dict(state_dict)
Restore this one stream’s progress previously returned by
state_dict().Restores a single stream’s offset (
{version, delivered, rank, worker_id}); the next__iter__skipsdeliveredfiles at the head of this stream’s slice. The loader routes each worker’s own state back to it individually, so this method never sees other workers’ state.The stored
rank/worker_idare checked against this stream’s own identity so a blob can only be restored into the stream that produced it; restoring elsewhere would skip or duplicate files on that stream’s slice. RaisesValueErroron an unrecognised version, a negative count, or a stream-identity mismatch.- Return type:
None
- state_dict()
Return a checkpoint of the calling stream’s iteration progress.
Called inside a DataLoader worker (or the main process for
num_workers == 0), so it describes just that stream. The loader carries this dict opaque (it tracks worker identity and the grid itself), so thedeliveredcount is just this stream’s own progress. The returned dict is JSON-serialisable and safe to fold into a larger training checkpoint.The dict also stamps the
(rank, worker_id)of the stream that produced it. These are not needed to resume (the loader already routes each blob back to the right worker and pins the grid), but they letload_state_dict()fail loud if a stream’s state is ever restored into a different stream (a loader-routing bug, or a user hand-wiring one rank’s/worker’s slice into another), which would otherwise silently skip or duplicate files.- Return type:
Dict[str,Any]
DataLoader
- class databricks.air.data.DataLoader(*args, num_workers=-1, prefetch_factor=4, **kwargs)
Bases:
DataLoadertorch.utils.data.DataLoaderthat logs iterator timing to MLflow.Differs from the upstream DataLoader in three ways:
num_workersdefault is chosen from the distributed world size andprefetch_factordefaults to 4 (PyTorch’s defaults are 0 and 2, which leave the GPU idle while batches are prepared). Whentorch.distributedis initialised the default is 6 workers per rank (8 ranks/node × 6 = 48 fetches/node); when it is not (single-process / notebook) the default is 48, so the lone loader still drives ~50-way node parallelism. An explicitly passednum_workersalways wins.persistent_workersis forced on whenevernum_workers > 0.UCVolumeDatasetcaches files through a per-process eviction tracker that is rebuilt from scratch each time a worker is (re)forked. With the upstream default (persistent_workers=False) the loader re-forks workers every epoch, so files cached in earlier epochs become untracked and can never be evicted, so the local cache disk leaks and eventually fills. Keeping workers alive preserves their eviction state across epochs. (PyTorch requiresnum_workers > 0for persistent workers, so it stays off when iterating in the main process, which is already safe because that single process keeps its tracker across epochs.)__iter__returns a wrapper that times eachnext()call and logs the timings to the active MLflow run, if one exists. Metrics are namespaced bytorch.distributedrank so multiple ranks can share a single MLflow run.If the dataset is
Checkpointable(e.g.UCVolumeDataset), the loader supportsstate_dict()/load_state_dict(). Each worker’s dataset state rides back with its batches (via acollate_fnshim) and is harvested as batches are delivered, so a resumed run continues from the next not-yet-delivered file per stream, not the next prefetched one. Works with or without batching (batch_size=Nonedelivers single items; the shim wraps the collate on both paths). The checkpoint pinsworld_sizeandnum_workers; resuming on a different grid raises.
Requires the
forkmultiprocessing start method whennum_workers > 0(aValueErroris raised otherwise).UCVolumeDatasetshares a single local cache directory across workers via fork’s copy-on-write inheritance, and relies on forked workers exiting without running the cache’s cleanup finalizer.spawn/forkserverpickle the dataset into each worker, duplicating the cache’sTemporaryDirectory(and its delete-on-exit finalizer) per worker, which can wipe the shared cache out from under the other workers.- load_state_dict(state_dict)
Restore progress previously returned by
state_dict().Validates the loader-level version and asserts the checkpoint’s grid matches this loader’s (
world_sizeandnum_workers), since per-worker counts are meaningless on a different grid. It then arranges for each worker’s dataset replica to skip its already-delivered files. Forked workers (num_workers > 0) are seeded by the installedworker_init_fnat fork time; the single-process case seeds the dataset directly here.With
persistent_workersthe live workers hold a dataset snapshot from their original fork, so any existing worker pool is torn down here; the next__iter__re-forks workers that inherit the restored state.Raises
ValueErroron an unrecognised version or a grid mismatch, andTypeErrorif the dataset is notCheckpointable.- Return type:
None
- state_dict()
Return a checkpoint of per-worker iteration progress.
Aggregates the dataset state harvested from each worker as batches were delivered, so the checkpoint reflects consumed (not prefetched) progress. The grid (
world_size,num_workers) is recorded soload_state_dict()can reject a resume on a different grid.Raises
TypeErrorif the dataset is notCheckpointable.- Return type:
Dict[str,Any]
UCVolumeWriter
- class databricks.air.data.UCVolumeWriter(path, **kwargs)
Bases:
FileSystemWriterStorageWriter that uploads checkpoint files to a UC volume.
Inherits all of torch’s
FileSystemWriterbehavior, including theBlockingAsyncStagermixin used bydcp.async_save.- Parameters:
path (
Union[str,PathLike]) – Remote UC volume directory the checkpoint is written to.**kwargs (
Any) – Forwarded totorch.distributed.checkpoint.FileSystemWriter(e.g.single_file_per_rank,thread_count,cache_staged_state_dict,serialization_format).
- classmethod validate_checkpoint_id(checkpoint_id)
- Return type:
bool
UCVolumeReader
- class databricks.air.data.UCVolumeReader(path, **kwargs)
Bases:
FileSystemReaderStorageReader that downloads checkpoint files from a UC volume.
- Parameters:
path (
Union[str,PathLike]) – Remote UC volume directory the checkpoint was written to.**kwargs (
Any) – Forwarded totorch.distributed.checkpoint.FileSystemReader.
- classmethod validate_checkpoint_id(checkpoint_id)
- Return type:
bool
Checkpointable
- class databricks.air.data.Checkpointable(*args, **kwargs)
Bases:
ProtocolA dataset whose iteration progress can be checkpointed and resumed.
isinstance(obj, Checkpointable)is true for any object exposing both methods, which is howdatabricks.air.data.DataLoaderdecides whether to harvest and restore dataset state.- load_state_dict(state_dict)
Restore progress from a dict previously returned by
state_dict.- Return type:
None
- state_dict()
Return a JSON-serialisable snapshot of iteration progress.
- Return type:
Dict[str,Any]