API Reference

API reference for launching distributed GPU workloads and initializing Ray from Databricks notebooks.

distributed

databricks.air.distributed(num_accelerators=None, accelerator_type=None, timeout=10800)

Decorator for running a function across GPUs connected to the notebook.

Parameters:
  • num_accelerators (Optional[int], optional) – Number of accelerators to use. Defaults to None (all GPUs connected to the notebook).

  • accelerator_type (Optional[Union[AcceleratorType, str]], optional) – The accelerator type to use: "GPU_1xA10", "GPU_1xH100", "GPU_8xH100" or "GPU_8xB300". If provided, must match the accelerator the notebook is connected to. Defaults to None (automatically detected).

  • timeout (Optional[float], optional) – Wall-clock seconds to allow local execution to run before subprocesses are terminated and LocalExecutionTimeoutError is raised. Defaults to 10800 seconds (3 hours). Pass None to disable.

The decorator returns a DistributedFunction object. Call .distributed() on it with the same arguments you would pass to your original function to launch execution on GPU.

Accelerator Types

Note

accelerator_type is auto-detected from the notebook’s attached accelerator. If you do specify it, you can pass the value as a string, such as "GPU_8xH100".

class databricks.air.compute.AcceleratorType(*values)

Bases: Enum

Accelerator types supported for distributed computing on Databricks AI Runtime.

GPU_1xA10

Single NVIDIA A10 GPU nodes.

GPU_1xH100

Single NVIDIA H100 GPU node.

GPU_8xH100

8x NVIDIA H100 GPU nodes.

GPU_8xB300

8x NVIDIA B300 GPU nodes.

ray_init

databricks.air.ray_init(*args, **kwargs)

Initialize Ray with dashboard access and display the dashboard URL.

Parameters:
  • *args – Positional arguments passed to ray.init().

  • **kwargs – Keyword arguments passed to ray.init().

Returns:

The Ray context returned by ray.init(), with the dashboard URL configured for notebook access.

Raises:

ImportError – If Ray is not installed.

Example

from databricks.air import ray_init

ray_init()

# Use standard Ray APIs after initialization.
@ray.remote
def square(value):
    return value * value

ray.get(square.remote(4))