How LuisaCompute Uses C++ Template Metaprogramming to Trace Kernel AST Construction
LuisaCompute's embedded DSL leverages expression templates and compile-time type deduction to construct a typed abstract syntax tree by overloading C++ operators and forwarding calls to a singleton FunctionBuilder that records every operation as an AST node.
The luisagroup/luisacompute repository implements a high-performance GPU computing framework that embeds a domain-specific language directly into C++. Instead of parsing strings or using external compilers, the library uses C++ template metaprogramming to trace kernel execution and build an AST at compile time. This architecture enables type-safe GPU programming while maintaining the syntax of standard C++.
Expression Templates and Type-Safe Wrappers
The foundation of the DSL lies in two primary template classes: Expr<T> and Var<T>. In src/dsl/builtin.cpp, Expr<T> serves as a thin wrapper that holds a pointer to an internal detail::Expression node, while Var<T> derives from Expr<T> to represent mutable variables.
When you declare a kernel parameter using Var<float>, the constructor registers a parameter with the current function builder and stores the resulting expression pointer inside the template wrapper. This mechanism allows the C++ type system to track the underlying GPU type while the AST tracks the computational graph.
// From src/dsl/builtin.cpp - Conceptual representation
template<typename T>
class Expr {
detail::Expression* _expression;
public:
explicit Expr(detail::Expression* expr) : _expression(expr) {}
detail::Expression* expression() const noexcept { return _expression; }
};
template<typename T>
class Var : public Expr<T> {
public:
Var() : Expr<T>(FunctionBuilder::current()->declare<T>()) {}
};
The FunctionBuilder Singleton and Node Recording
At the core of the tracing mechanism is the FunctionBuilder class, accessible via FunctionBuilder::current(). This singleton maintains the mutable AST for the kernel currently being defined, storing nodes in a detail::Function object. Every DSL operation ultimately forwards to this builder to append nodes to the graph.
When a kernel lambda executes, it captures the current builder context. As the user writes expressions, the DSL inserts nodes such as detail::AccessExpr, detail::CallExpr, and detail::BinaryExpr into the builder's internal storage. The process concludes with FunctionBuilder::finalize(), which serializes the AST into an intermediate representation consumed by the GPU backends.
// Conceptual usage pattern from src/dsl/soa.cpp
auto kernel = [](Var<float> a, Var<float> b) {
return a + b; // Each operation records nodes via FunctionBuilder::current()
};
// After lambda execution: FunctionBuilder::current()->finalize()
Operator Overloading and AST Construction
The DSL intercepts arithmetic and logical operations through overloaded operators defined for Expr<T>. Rather than executing the operation immediately, these overloads construct AST nodes. The LUISA_EXPR macro expands to extract_expression(std::forward<decltype(v)>(v)), which unwraps Var<T> instances, literals, or sub-expressions to retrieve the underlying Expression* pointer.
In src/dsl/builtin.cpp, operators like operator+ forward to templated helpers such as binary_op<detail::Add, T>(lhs, rhs). The template parameter T ensures type safety, while the operation tag detail::Add determines the node type created. This compile-time dispatch guarantees that ill-typed kernel code fails during C++ compilation rather than at runtime.
// From src/dsl/builtin.cpp - Simplified operator implementation
template<typename T>
Expr<T> operator+(Expr<T> lhs, Expr<T> rhs) {
return binary_op<detail::Add, T>(lhs, rhs);
}
template<typename Op, typename T>
Expr<T> binary_op(Expr<T> lhs, Expr<T> rhs) {
auto expr = FunctionBuilder::current()->create_binary_expression(
extract_expression(lhs),
extract_expression(rhs)
);
return Expr<T>{expr};
}
Compile-Time Dispatch with if constexpr
Complex constructors like make_vector and make_matrix utilize C++17 if constexpr to select implementation paths based on template parameters. Located in src/dsl/builtin.cpp, these functions deduce dimensions and element types at compile time, eliminating runtime branches while ensuring the generated AST nodes match the static type information.
This metaprogramming technique allows a single templated interface to generate specific constructor nodes—such as detail::MatrixCtorExpr—with exact dimensionality encoded in the type system. The compiler unrolls these constructs entirely, leaving only the AST node creation in the compiled binary.
// From src/dsl/builtin.cpp - Template metaprogramming pattern
template<typename T, size_t N>
auto make_vector(/* args */) {
if constexpr (N == 2) {
return Expr<Vector<T, 2>>{FunctionBuilder::current()->create_vector2(/* args */)};
} else if constexpr (N == 3) {
return Expr<Vector<T, 3>>{FunctionBuilder::current()->create_vector3(/* args */)};
}
// Additional compile-time branches...
}
From C++ Source to GPU Intermediate Representation
The tracing process transforms ordinary C++ syntax into a rich, typed AST without parsing. As shown in src/ir/ir2ast.cpp, the finalized AST from FunctionBuilder undergoes translation into the framework's intermediate representation (IR). This IR is then lowered to CUDA, Metal, or Vulkan instructions.
Because the entire pipeline relies on template instantiation rather than string manipulation, the DSL preserves source location information through __builtin_LINE and __builtin_FILE macros. This enables error messages that map directly back to the original C++ kernel code, despite the code executing on the GPU.
Summary
- Expr and Var template classes in
src/dsl/builtin.cppwrap raw expression pointers while preserving compile-time type information. - FunctionBuilder::current() maintains a singleton AST builder that records every operation during kernel lambda execution.
- LUISA_EXPR and extract_expression utilities unwrap DSL objects to feed raw pointers into AST node constructors.
- Operator overloads such as
binary_op<Op, T>convert C++ syntax intodetail::BinaryExprnodes without executing the operations. - if constexpr dispatch in
make_vectorandmake_matrixenables zero-overhead abstraction for complex type constructors. - The pipeline concludes with FunctionBuilder::finalize() and translation through
src/ir/ir2ast.cppto generate GPU-specific code.
Frequently Asked Questions
What are expression templates in the context of LuisaCompute?
Expression templates are a C++ technique where overloaded operators return proxy objects rather than computed values. In LuisaCompute, when you write a + b where a and b are Var<T> types, the operation returns an Expr<T> containing a detail::BinaryExpr node pointer. This defers execution and allows the framework to record the computation structure as an AST node instead of performing the addition immediately.
How does FunctionBuilder maintain AST state during C++ compilation?
FunctionBuilder uses a thread-local singleton pattern accessed via FunctionBuilder::current(). When a kernel lambda executes, the DSL internally sets the current builder context. Each DSL operation—whether variable declaration, arithmetic, or function calls—extracts expression pointers and appends them to the builder's internal graph. Because this happens during the lambda's execution at compile time, the C++ compiler essentially "traces" the kernel logic into a static data structure.
Why does LuisaCompute use C++ templates instead of a separate DSL parser?
Using C++ template metaprogramming eliminates the need for external parsing tools and maintains type safety through the host compiler. By overloading operators and using expression templates, LuisaCompute ensures that type mismatches in GPU code trigger standard C++ compiler errors. Additionally, this approach allows developers to use existing C++ development environments, debuggers, and language features while the framework automatically extracts the computation graph through template instantiation.
How are complex types like matrices constructed within the AST?
Complex constructors utilize if constexpr to handle different dimensions at compile time. Functions like make_matrix<T, dim>() in src/dsl/builtin.cpp use template parameters to determine the exact node type to create—such as detail::MatrixCtorExpr—and recursively extract expressions from each argument. This template-based dispatch ensures that matrix dimensions are encoded in the AST node types themselves, allowing the backend to generate optimal GPU code without runtime dimension checks.
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 →