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:

  1. Capture input buffers (or pass them as kernel arguments)
  2. Create a Callable that contains the $autodiff block to isolate AD logic and enable reuse
  3. Pack inputs into a temporary structure (ArrayFloat<N> or a tuple) if you need more than one differentiable argument
  4. Mark the packed variables with requires_grad
  5. Compute the forward expression, call backward, then read gradients with grad
  6. 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_grad accepts multiple arguments simultaneously (requires_grad(x, y, z)).

Summary

  • LuisaCompute implements reverse-mode AD directly in the DSL through the $autodiff block, eliminating the need for external autodiff libraries.
  • Mark differentiable variables with requires_grad(), trigger differentiation with backward(expr), and retrieve results using grad(var).
  • The underlying implementation emits intrinsic AD ops (e.g., AUTODIFF_REQUIRES_GRADIENT) as defined in src/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:

Share the following with your agent to get started:
curl -s "https://instagit.com/install.md"

Works with
Claude Codex Cursor VS Code OpenClaw Any MCP Client

Maintain an open-source project? Get it listed too →