For the complete documentation index, see llms.txt. Markdown versions of all pages are available by appending .md to any URL (e.g. /get-started.md).
Mojo struct
DeviceGraph
struct DeviceGraph
Represents an instantiated device graph that can be replayed.
A DeviceGraph captures a sequence of GPU operations (such as kernel
launches) as a reusable graph. Once instantiated from a
DeviceGraphBuilder, the graph can be replayed multiple times at a
lower overhead than re-enqueueing each operation individually.
To obtain a DeviceGraph, use
DeviceGraph.create().
Graph capture is currently implemented for CUDA and HIP devices only. Creating a graph on any other device, such as an Apple GPU or a CPU, raises.
Implemented traits
AnyType,
Copyable,
Deinitable,
ImplicitlyCopyable,
Movable,
Writable
Methods
__init__
def __init__(out self, *, copy: Self)
Creates a copy of an existing device graph by incrementing its reference count.
Args:
- copy (
Self): The device graph to copy.
__deinit__
def __deinit__(deinit self)
Releases resources associated with this device graph.
replay
def replay(self)
Replays the captured sequence of GPU operations.
Submits the pre-captured sequence of operations for execution on the device. This is more efficient than re-enqueueing each operation individually because the graph has already been compiled and instantiated by the driver.
Example:
from max.gpu.host import DeviceContext, DeviceGraph, DeviceGraphBuilder
def kernel():
print("replaying")
with DeviceContext() as ctx:
var compiled_fn = ctx.compile_function[kernel]()
def build(mut builder: DeviceGraphBuilder) raises {imm}:
_ = builder.add_function(
compiled_fn, grid_dim=1, block_dim=1, dependencies=[]
)
var graph = DeviceGraph.create(ctx, build)
graph.replay()
graph.replay() # replay as many times as needed
ctx.synchronize()Raises:
If replay fails.
create
static def create[*Ts: DeviceGraphInput](ctx: DeviceContext, build: T, *inputs: *Ts.values, *, cache: Pointer[DeviceGraphCache]) -> Self
Builds and instantiates a device graph, reusing a cached one if it can.
Behaves like the uncached overload, except that a graph an earlier call
built from equivalent inputs is returned as-is and build is never
called.
Example:
from max.gpu.host import (
DeviceContext, DeviceGraph, DeviceGraphBuilder, DeviceGraphCache
)
def kernel():
print("replaying")
with DeviceContext() as ctx:
var compiled_fn = ctx.compile_function[kernel]()
var cache = DeviceGraphCache()
def build(mut builder: DeviceGraphBuilder) raises {imm}:
_ = builder.add_function(
compiled_fn, grid_dim=1, block_dim=1, dependencies=[]
)
# The second call reuses the graph the first one built.
var graph = DeviceGraph.create(ctx, build, cache=Pointer(to=cache))
var same = DeviceGraph.create(ctx, build, cache=Pointer(to=cache))
ctx.synchronize()Parameters:
- *Ts (
DeviceGraphInput): Types of the device graph inputs.
Args:
- ctx (
DeviceContext): Device context for the target device. - build (
T): Callback that adds nodes to the supplied builder, called only on a cache miss. - *inputs (
*Ts.values): The device graph inputs the cache key is derived from. - cache (
Pointer[DeviceGraphCache]): The cache to consult, and to store a newly built graph in.
Returns:
Self: The instantiated device graph, which may be one a previous call
built.
Raises:
If ctx is on a device without graph support, or if graph builder
creation, build, or instantiation fails.
static def create(ctx: DeviceContext, build: T) -> Self
Builds and instantiates a device graph within a scoped callback.
Calls build with a fresh DeviceGraphBuilder, then instantiates the
result into a replayable DeviceGraph. The builder, and any
DeviceGraphNode handles obtained from it, are valid only for the
duration of build: their origin is scoped to this call and cannot
escape it, so a node handle cannot be stored beyond the callback or
used with a different graph.
Pass a cache to the overload above to reuse a previously built graph
instead of recording one on every call.
Graph capture is currently implemented for CUDA and HIP devices only.
On any other device, such as an Apple GPU or a CPU, this raises before
build runs.
Example:
from max.gpu.host import DeviceContext, DeviceGraph, DeviceGraphBuilder
def kernel(x: Int):
print("Value:", x)
with DeviceContext() as ctx:
var compiled_fn = ctx.compile_function[kernel]()
def build(mut builder: DeviceGraphBuilder) raises {imm}:
_ = builder.add_function(
compiled_fn, 42, grid_dim=1, block_dim=1, dependencies=[]
)
var graph = DeviceGraph.create(ctx, build)
graph.replay()
ctx.synchronize()Args:
- ctx (
DeviceContext): Device context for the target device. - build (
T): Callback that adds nodes to the supplied builder. It receives the builder by mutable reference and therefore cannot instantiate it directly; instantiation happens here once the callback returns.
Returns:
Self: The instantiated device graph.
Raises:
If ctx is on a device without graph support, or if graph builder
creation, build, or instantiation fails.