How to Implement Automatic Differentiation Using the $autodiff Block for Custom Kernels
Use the $autodiff DSL block in LuisaCompute to mark differentiable variables with requires_grad(), compute forward expressions, call backward() to trigger reverse-mode automatic differentiation, and read results with grad()—all without external libraries or manual graph construction.
LuisaCompute implements reverse-mode automatic differentiation (AD) directly in its domain-specific language (DSL), enabling custom kernels to compute gradients efficiently on the GPU. The $autodiff macro introduces a mini-programming context where the compiler records a computational graph for operations inside the block, then automatically generates the backward pass. This guide explains how to implement automatic differentiation using the $autodiff block for custom kernels in the luisagroup/luisacompute repository.
How the $autodiff Block Works
The AD engine builds the graph once per kernel launch, making it significantly cheaper than finite-difference approximations. The process follows four distinct phases according to the source code in src/xir/op.cpp and docs/source/dsl.md.
Marking Inputs with requires_grad
First, call requires_grad(var…) to tell the AD engine which variables need gradients. As documented in [docs/source/dsl.md (lines 78–106)](https://github.com/luisagroup/luisacompute/blob/stable/docs/source/dsl.md#L78-L106), this creates placeholder gradient buffers for each marked variable. You can pass a single variable or a variadic list:
requires_grad(x); // single variable
requires_grad(a, b); // multiple variables
Recording the Forward Pass
All DSL operations (+, *, sin, etc.) inside the block are recorded as intrinsic AD ops (e.g., AUTODIFF_REQUIRES_GRADIENT). The forward values are computed normally while the graph is constructed. This is implemented in [src/xir/op.cpp (lines 269–282)](https://github.com/luisagroup/luisacompute/blob/stable/src/xir/op.cpp#L269-L282), where the IR nodes capture the computational dependencies.
Triggering the Backward Pass
Calling backward(expr) walks the recorded graph in reverse order, applying the chain rule and accumulating partial derivatives into the gradient buffers of the marked variables. This triggers the actual reverse-mode differentiation.
Reading Gradients
Finally, grad(var) returns the computed gradient for a marked variable. The value can be written back to a buffer or used in further host-side calculations.
Implementation Pattern for Custom Kernels
When implementing automatic differentiation using the $autodiff block for custom kernels, follow this structured approach:
- Capture input buffers (or pass them as kernel arguments)
- Create a Callable that contains the
$autodiffblock to isolate AD logic and enable reuse - Pack inputs into a temporary structure (
ArrayFloat<N>or a tuple) if you need more than one differentiable argument - Mark the packed variables with
requires_grad - Compute the forward expression, call
backward, then read gradients withgrad - Write the gradients to output buffers
The AD system is purely DSL-based; no external libraries or manual graph construction is required.
Code Examples
Single Differentiable Input
This minimal kernel computes the gradient of f(x) = x * sin(x):
// kernel1.cpp – computes gradient of f(x)=x*sin(x)
Kernel1D grad_kernel = [](BufferFloat x_buf, BufferFloat grad_buf) noexcept {
auto i = dispatch_id().x;
Float x = x_buf.read(i);
$autodiff {
requires_grad(x); // x needs a gradient
Float y = x * sin(x); // forward expression
backward(y); // back‑propagate
grad_buf.write(i, grad(x)); // write ∂y/∂x
};
};
The forward pass computes y = x·sin(x). After backward(y), grad(x) holds sin(x) + x·cos(x).
Multiple Inputs with Callables
For reusable AD logic with multiple inputs, pack variables into an array and use a Callable:
// reusable callable – computes gradient of z = a * sin(b) + b * cos(a)
Callable diff_ab = [](ArrayFloat<2> args) noexcept {
auto a = args[0];
auto b = args[1];
Float a_grad = 0.f, b_grad = 0.f;
$autodiff {
requires_grad(a, b);
Float z = a * sin(b) + b * cos(a);
backward(z);
a_grad = grad(a);
b_grad = grad(b);
};
return make_float2(a_grad, b_grad);
};
// kernel that uses the callable
Kernel1D kernel = [&](BufferFloat a_buf, BufferFloat b_buf,
BufferFloat2 grad_buf) noexcept {
auto i = dispatch_id().x;
Float a = a_buf.read(i);
Float b = b_buf.read(i);
auto g = diff_ab(ArrayFloat<2>{a, b});
grad_buf.write(i, g);
};
This pattern is demonstrated in [src/tests/test_autodiff.cpp (lines 66–87)](https://github.com/luisagroup/luisacompute/blob/stable/src/tests/test_autodiff.cpp#L66-L87).
Control Flow and Custom Functions
Control flow ($if, $for) works inside $autodiff, but dynamic loops must be manually unrolled because the graph needs a static size:
Callable piecewise = [](Float x) noexcept {
Float y;
$if (x > 0.f) {
y = x * x; // x²
} $else {
y = -x; // -x
};
return y;
};
Kernel1D kernel = [&](BufferFloat in, BufferFloat out, BufferFloat grad) noexcept {
auto i = dispatch_id().x;
Float v = in.read(i);
$autodiff {
requires_grad(v);
Float y = piecewise(v);
backward(y);
grad.write(i, grad(v));
};
};
Key Implementation Details
- Gradient types: The gradient of an array-like variable (e.g.,
ArrayFloat<3> a) is a matching array (grad(a)[i]). - Supported operations: Only operations that have AD definitions can be differentiated. Most built-ins (arithmetic, trig, etc.) are supported; custom host functions are not.
- Graph construction: The AD graph is built once per kernel launch, resulting in overhead comparable to a single forward pass plus a single backward sweep.
- Variadic marking:
requires_gradaccepts multiple arguments simultaneously (requires_grad(x, y, z)).
Summary
- LuisaCompute implements reverse-mode AD directly in the DSL through the
$autodiffblock, eliminating the need for external autodiff libraries. - Mark differentiable variables with
requires_grad(), trigger differentiation withbackward(expr), and retrieve results usinggrad(var). - The underlying implementation emits intrinsic AD ops (e.g.,
AUTODIFF_REQUIRES_GRADIENT) as defined insrc/xir/op.cpp. - Use Callables to encapsulate reusable AD logic, and pack multiple inputs into
ArrayFloat<N>structures when necessary. - Control flow is supported, but loops must be statically unrolled to maintain a fixed graph size.
Frequently Asked Questions
What operations support automatic differentiation in LuisaCompute?
Built-in DSL operations—including arithmetic (+, -, *, /), trigonometric functions (sin, cos, tan), and most standard math intrinsics—have AD definitions and can be differentiated automatically. Custom host functions defined outside the DSL do not have gradient definitions and cannot be used inside $autodiff blocks.
Can I use dynamic loops inside $autodiff blocks?
No, dynamic loops must be manually unrolled. The AD engine constructs a static computational graph at compile time, so loop bounds must be known constants. Use $for with constant ranges or manually unroll iterations to ensure the graph has a fixed size.
How do I handle multiple differentiable variables?
Pass multiple variables to requires_grad() as a variadic list (requires_grad(a, b, c)), or pack them into an array (e.g., ArrayFloat<2>) and unpack them inside a Callable. The latter approach is recommended for reusable AD components, as shown in src/tests/test_autodiff.cpp.
Where is the automatic differentiation logic implemented?
The DSL syntax is documented in docs/source/dsl.md, while the underlying IR intrinsics (such as AUTODIFF_REQUIRES_GRADIENT) are defined in src/xir/op.cpp. The public API headers in include/luisa/xir/op.h provide the C++ symbols used by the DSL, connecting the high-level $autodiff syntax to the low-level reverse-mode engine.
Have a question about this repo?
These articles cover the highlights, but your codebase questions are specific. Give your agent direct access to the source. Share this with your agent to get started:
curl -s "https://instagit.com/install.md" Maintain an open-source project? Get it listed too →