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 toNone(automatically detected).timeout (Optional[float], optional) – Wall-clock seconds to allow local execution to run before subprocesses are terminated and
LocalExecutionTimeoutErroris raised. Defaults to 10800 seconds (3 hours). PassNoneto 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:
EnumAccelerator 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))