Load data on AI Runtime

Important

This feature is in Public Preview.

Data and model assets are critical to deep learning and post-training workloads for large language models (LLMs) and vision-language models (VLMs). With AI Runtime, all data and model assets are accessed through Unity Catalog:

  • Unity Catalog volumes: used primarily for large datasets and unstructured files, including images, audio, and text.
  • Unity Catalog tables: used for structured and tabular data, accessed through Spark Connect.

Your volumes and tables must be registered in Unity Catalog and accessible to your user or service principal.

Unity Catalog volume for unstructured data

Unity Catalog volumes provide governed access to non-tabular data in any format, including structured, semi-structured, and unstructured data. In AI Runtime, volumes are the primary mechanism for accessing large datasets, text, model assets, and model checkpoints.

Users can list, read, and write files in Unity Catalog volumes using familiar file-system operations, similar to working with files on a local disk:

import os

dir_path = "/Volumes/<catalog-name>/<schema-name>/<volume-name>/sub-dir"
file_path = os.path.join(dir_path, "test_file")

os.makedirs(dir_path, exist_ok=True)

# Write to the file
with open(file_path, "w") as file:
    file.write("Hello, World!")

Similarly, shell operations work the same way:

%sh ls -l /Volumes/<catalog-name>/<schema-name>/<volume-name>
%sh mkdir -p /Volumes/<catalog-name>/<schema-name>/<volume-name>/sub-dir
%sh touch /Volumes/<catalog-name>/<schema-name>/<volume-name>/sub-dir/test_file

A few characteristics of Unity Catalog volumes make them well suited for machine learning workloads:

  • Distributed storage: Unity Catalog is backed by distributed storage, allowing AI Runtime workloads to read and write data and model assets across the platform, from both notebooks and CLI-based workloads.
  • Optimized for ML access patterns: The underlying storage and access paths are optimized for common ML workloads, particularly large files with sequential reads and writes. This makes Unity Catalog well suited for training data loading, model asset loading, and model checkpoint writing.
  • File-system-like access: Users can list, read, and write files in Unity Catalog volumes using familiar file-system operations, similar to working with files on a local disk.

Due to automatic background commits, users can expect consistent access to Unity Catalog volume data:

  • Writes: AI Runtime automatically commits writes, making changes visible to other applications and workloads accessing the same Unity Catalog volume.
  • Reads: AI Runtime automatically picks up changes to the volume without requiring any explicit refresh or synchronization operation from the user.

Tune volume performance

As mentioned, Unity Catalog volumes are backed by distributed storage and optimized for large files with sequential reads and writes.

A few tips can help you get the best performance from AI Runtime:

  • Concatenate data into larger files: When possible, consolidate data into fewer, larger files, roughly 1 GiB to 10 GiB per file. This allows AI Runtime to aggressively prefetch data and automatically achieve near-optimal sequential read performance.

  • For small-file workloads, use local disk: If your workload involves many small files, consider copying the files to the local disk using parallel copies before processing them. This can reduce the overhead of repeatedly accessing many small files through the volume.

    # Recommended using parallel copy (256 concurrency in this example, you can tune)
    #
    # This takes only 22 seconds to copy 15,375 150KiB small image files.
    %sh cd /Volumes/<catalog-name>/<schema-name>/<volume-name>/sub-dir/ && find . -type f -print0 | xargs -0 -P 256 -I {} cp --parents "{}" /tmp/
    
    
    # !!! Avoid doing this !!!
    #
    # Because the files are copied in serial, this copies the same 15,375 150KiB small image files much more slowly.
    # %sh cp -r /Volumes/<catalog-name>/<schema-name>/<volume-name>/sub-dir/* /tmp
    
  • You can use UCVolumeDataset for your machine learning workloads. It incorporates the optimizations described above to provide efficient data access and loading from Unity Catalog volumes. See the following sections for more details.

Load unstructured data with UCVolumeDataset

For unstructured data such as images, audio, and text files stored in Unity Catalog volumes, use UCVolumeDataset from the databricks.air.data module. UCVolumeDataset is a PyTorch IterableDataset that copies each file from the volume to a fast local cache on first access and yields the cached local file path. It handles the performance and distribution concerns you would otherwise implement by hand:

  • Local caching. Files are copied from the FUSE mount to a local cache directory on first access and served from the cache afterward, so multi-epoch training does not re-read the volume.
  • Automatic partitioning. When torch.distributed is initialized, files are partitioned across ranks and then further divided across DataLoader workers, so each (rank, worker) pair receives a non-overlapping slice with no extra setup.

Note

UCVolumeDataset and databricks.air.data.DataLoader come from the databricks-sdk-air package. Install it with the data extra, which also pulls in a compatible torch:

%pip install "databricks-sdk-air[data]"

UCVolumeDataset yields raw local file paths. To decode those files into tensors, wrap it in a second IterableDataset that consumes the path stream and applies your parsing logic. This keeps I/O and parsing concerns separate.

from databricks.air.data import UCVolumeDataset
from torch.utils.data import IterableDataset
from PIL import Image
import torchvision.transforms.functional as TF

class ImageDataset(IterableDataset):
    """Decodes each cached file path from UCVolumeDataset into a tensor."""

    def __init__(self, path_dataset: UCVolumeDataset):
        self._path_dataset = path_dataset

    def __iter__(self):
        for local_path in self._path_dataset:
            image = Image.open(local_path).convert("RGB")
            yield TF.to_tensor(image)

path_dataset = UCVolumeDataset("/Volumes/catalog/schema/my_volume/images")
dataset = ImageDataset(path_dataset)

The wrapper receives already-cached local paths, so the parsing step never touches the volume. You can chain additional wrappers for augmentation, tokenization, or filtering.

For optimal performance, pair UCVolumeDataset with databricks.air.data.DataLoader rather than the stock PyTorch DataLoader. It is tuned for AI Runtime I/O and fetches and caches files concurrently while the GPU computes.

Checkpoint models on volumes

To checkpoint your model so you can resume training from the latest snapshot or recover from a crash, you can use Unity Catalog volumes just like a local file system.

Databricks recommends using a distributed checkpoint (DCP) for better performance on both single-GPU and multi-GPU workloads. See Fast, fault-tolerant PyTorch training on AI Runtime from the Databricks engineering blog.

import torch.distributed.checkpoint as dcp
from torch.distributed.checkpoint.state_dict import get_state_dict, set_state_dict

import databricks.air.data

checkpoint_path = "/Volumes/my-catalog/my-schema/my-volume/checkpoints/step_1000"

# Save
model_sd, optim_sd = get_state_dict(model, optimizer)
state_dict = {"model": model_sd, "optim": optim_sd, "step": 1000}
dcp.async_save(
    state_dict,
    storage_writer=databricks.air.data.UCVolumeWriter(checkpoint_path))

# Load
model_sd, optim_sd = get_state_dict(model, optimizer)
state_dict = {"model": model_sd, "optim": optim_sd}
dcp.load(
    state_dict,
    storage_reader=databricks.air.data.UCVolumeReader(checkpoint_path))

set_state_dict(
    model,
    optimizer,
    model_state_dict=state_dict["model"],
    optim_state_dict=state_dict["optim"],
)

The monolithic torch.save approach also works.

  • For single-GPU model checkpointing,

    # The monolithic torch.save approach for single GPU chip
    
    # Save
    torch.save({"model": model.state_dict(), "opt": optimizer.state_dict()},
               "/Volumes/<catalog-name>/<schema-name>/<volume-name>/sub-dir/ckpt.pt")
    
    
    # Load
    ckpt = torch.load(
        "/Volumes/<catalog-name>/<schema-name>/<volume-name>/sub-dir/ckpt.pt",
        weights_only=True)
    model.load_state_dict(ckpt["model"])
    optimizer.load_state_dict(ckpt["opt"])
    
  • For distributed training launched via torchrun,

    # The monolithic torch.save approach for multi-GPU distributed training.
    # This snippet assumes your launcher has already called
    # dist.init_process_group(...).
    
    import os
    import torch.distributed as dist
    
    # Save only on rank 0.
    if dist.get_rank() == 0:
        torch.save({"model": model.state_dict(), "opt": optimizer.state_dict()},
                   "/Volumes/<catalog-name>/<schema-name>/<volume-name>/sub-dir/ckpt.pt")
    
    # Wait for rank 0 to finish writing before any rank reads.
    dist.barrier()
    
    # Load on ALL ranks (map to current rank's local GPU).
    local_rank = int(os.environ["LOCAL_RANK"])
    ckpt = torch.load(
        "/Volumes/<catalog-name>/<schema-name>/<volume-name>/sub-dir/ckpt.pt",
        map_location=f"cuda:{local_rank}",
        weights_only=True)
    model.load_state_dict(ckpt["model"])
    optimizer.load_state_dict(ckpt["opt"])
    

Load tabular data

Use Spark Connect to load tabular machine learning data from Delta tables.

For single-node training, you can convert Apache Spark DataFrames into pandas DataFrames using the PySpark method toPandas(), and then optionally convert to NumPy format using the PySpark method to_numpy().

Note

Spark Connect defers analysis and name resolution to execution time, which may change the behavior of your code. See Compare Spark Connect to Spark Classic.

Spark Connect supports most PySpark APIs, including Spark SQL, Pandas API on Spark, Structured Streaming, and MLlib (DataFrame-based). See the PySpark API reference documentation for the latest supported APIs.

For other limitations, see Serverless compute limitations.

Load large Delta tables using Unity Catalog volumes

For large Delta tables that are too big to convert with toPandas(), export the data to a Unity Catalog volume and load it directly using PyTorch or Hugging Face:

# Step 1: Export the Delta table to Parquet files in a UC volume
output_path = "/Volumes/catalog/schema/my_volume/training_data"
spark.table("catalog.schema.my_table").write.mode("overwrite").parquet(output_path)
# Step 2: Load the exported data directly using Hugging Face datasets
from datasets import load_dataset

dataset = load_dataset("parquet", data_files="/Volumes/catalog/schema/my_volume/training_data/*.parquet")

This approach avoids Spark overhead during training and works well for both single-GPU and distributed training workflows.