[Refactor][Language] Make T.Kernel target-neutral with dialect-owned launch annotations (#3186)

* [Refactor][Language] Make T.Kernel target-neutral with dialect-owned launch annotations

T.Kernel is traced before the Target is known, yet it used to emit
threadIdx.* thread_binding loops with a hard-coded default of 128 threads.
That is a SIMT launch baked into the frontend: CPU lowering silently pinned
the thread index to 0 (so per-thread code computed 1/128 of the work), and
every backend inherited CUDA's launch vocabulary.

Split the launch into the part every target shares and the parts a backend
interprets:

- The frontend records the grid as blockIdx.* thread_binding loops (the
  program-index space all targets share) and reserves thread identities as
  `tx/ty/tz = tl.launch_thread_idx(axis)` Bind placeholders. `threads=`
  becomes the `tl.launch_threads` annotation on the tilelang_root block;
  the frontend no longer chooses a default.
- MaterializeKernelLaunch(lower_thread_binding, default_threads,
  unsupported_annotations) is where the target-dependent meaning is decided.
  SIMT backends rebind each placeholder as a threadIdx.* thread_extent over
  the same Var (extent from the annotation, else the backend default; the
  CUDA/ROCm/Metal/WebGPU pipelines pass 128). Backends without SIMT drop the
  placeholders and reject a body that references a thread index instead of
  miscomputing. Launch annotations a backend cannot honor (cluster_dims on
  cpu/rocm/metal/webgpu) are rejected here rather than ignored downstream.
- KernelLaunchFrame exposes grid_vars/grid_extents/thread_vars/
  thread_extents explicitly; get_thread_extent() raises at trace time when
  threads= was omitted, since only the backend knows the extent.
- Each language dialect owns its `Kernel` with explicit keyword parameters:
  cuda offers threads/prelude/cluster_dims, rocm/metal threads/prelude,
  webgpu threads, cpu prelude, and the common `Kernel(*grid)` none. All of
  them funnel into `tilelang.language.kernel.launch_kernel`. This is what
  editors show on hover and what type checkers verify; unknown keywords are
  rejected by Python itself. ClusterKernel keeps its interface.
- Cluster identity: `T.get_cluster_id/ids/dims/size/extent(s)` are
  program-space arithmetic (`bx // cluster_dims`) computed at trace time and
  therefore target-neutral, while `T.block_rank_in_cluster()` stays the
  CUDA hardware intrinsic. The two agree under the cluster-aware threadblock
  swizzle, which permutes the grid at cluster granularity.

Tests that applied passes directly to a traced function now materialize the
launch first (thread indices otherwise carry no extent), fixtures that wrote
per-thread code and compiled it for the `c` target are rewritten at the tile
level, and the span-preservation check asserts that no statement loses its
span rather than that the spanned count never decreases.

* [Fix][Language] Unwind kernel frames after tracing errors

Pass the active exception to the pending frame exits when tracing fails, so
the Python launch state (KernelLaunchFrame stack) is cleared without
finalizing incomplete TIR. Previously the next trace observed the failed
kernel's launch frame as KernelLaunchFrame.Current(). Covered for both the
prim_func and eager JIT tracing paths.

* [Fix][Language] Recognize every dialect's Kernel as a launch in the eager rewriter

Eager JIT phase 1 infers the signature with symbolic T.const values and must
not execute the kernel body; the AST rewriter guards `with T.Kernel(...)`
with `if __tb.skip_kernel_ctx(): return`. It found the launch by object
identity against the common `Kernel`, which no longer holds now that each
dialect (and the `tilelang.language` facade) exports its own `Kernel`: the
guard was not injected, phase 1 ran the body, and e.g. `T.gemm` on a tile
whose extent comes from `T.const` failed its static-shape check.

Mark every launch factory (`Kernel` of each dialect, `ClusterKernel`) with
`kernel_launch_factory` and let the rewriter test that mark instead. It now
resolves plain names and attribute chains only, so aliases such as
`K = T.Kernel` are recognized and no factory expression is executed at
rewrite time. Covered by a rewriter-level test over all dialects and an
end-to-end eager gemm whose tile extent depends on a const.
This commit is contained in:
Lei Wang
2026-09-10 17:33:50 +08:00
committed by GitHub
parent 752312c2a9
commit 85fd8fc2d3
33 changed files with 1388 additions and 205 deletions
+29 -5
View File
@@ -85,16 +85,40 @@ Notes
## 2. Launching Work with `T.Kernel`
`with T.Kernel(...)` declares a launch context and creates block/thread
bindings. For GPU backends, specify a grid and threads per block.
`with T.Kernel(...)` declares a grid of tile programs. The positional
arguments give the grid extent along each axis and the returned variables are
the program indices along those axes. This is the part of a launch every
target shares: on CUDA a program is a thread block and `bx`/`by` are
`blockIdx.x`/`blockIdx.y`; on CPU the grid becomes the outer loop.
```python
with T.Kernel(grid_x, grid_y, threads=128) as (bx, by):
... # bx/by are blockIdx.x/y
... # bx/by are the program indices (blockIdx.x/y on CUDA)
```
You rarely need raw thread indices; most kernels use structured loops
(`T.serial`, `T.unroll`, `T.Parallel`, `T.Pipelined`) inside a `T.Kernel`.
Keyword arguments are launch annotations that the backend interprets once the
target is known. Each language dialect's `Kernel` declares the annotations its
backend understands as explicit keyword parameters, so hovering or
autocompleting `T.Kernel` shows exactly those and anything else is rejected:
`tilelang.language` (the CUDA dialect) offers `threads`, `prelude` and
`cluster_dims`; `tilelang.rocm.language` / `tilelang.metal.language` offer
`threads` and `prelude`; `tilelang.cpu.language` offers only `prelude`.
`threads` is the SIMT one: how many threads run each tile program on
GPU-style backends. Those backends pick a default (128) when it is omitted;
a kernel written with the CUDA dialect still compiles for CPU, which ignores
the thread count. Code inside `T.Kernel` operates at the tile-program level,
so you rarely need raw thread indices; most kernels use structured loops
(`T.serial`, `T.unroll`, `T.Parallel`, `T.Pipelined`) that the compiler maps
onto threads. `T.get_thread_binding()` exposes the thread index for
thread-level code on SIMT targets; a kernel that uses it is rejected when
compiled for a target without SIMT threads.
`T.ClusterKernel(..., cluster_dims=...)` adds the CUDA thread-block-cluster
annotation (SM90+). A cluster is a `cluster_dims`-shaped tile of the grid, so
`T.get_cluster_id(axis)` is plain program-index arithmetic
(`bx // cluster_dims[axis]`) and works on every target, while
`T.block_rank_in_cluster()` reads the hardware rank and is CUDA-only. Targets
without clusters reject `cluster_dims` at compile time.
## 3. Loops and Control Flow
+74 -18
View File
@@ -27,10 +27,10 @@ using namespace script::ir_builder::tirx;
using namespace ffi;
// Build a ForFrame that emits a target-neutral kThreadBinding loop for one
// kernel-launch dimension. The launch nest is materialized into the
// target-specific form (thread_extent AttrStmt on GPU, serial For on CPU) by
// the tl.MaterializeKernelLaunch pass once the Target is known at compile
// time.
// grid (program index) axis of a kernel launch. The launch nest is
// materialized into the target-specific form (thread_extent AttrStmt on GPU,
// serial For on CPU) by the tl.MaterializeKernelLaunch pass once the Target is
// known at compile time.
static ForFrame MakeThreadBindingFrame(const std::string &name,
const String &thread_tag,
const PrimExpr &extent) {
@@ -216,6 +216,38 @@ ForFrame PersistentFor(const Array<PrimExpr> &domain, const PrimExpr &wave_size,
return ForFrame(n);
}
// Build a frame whose exit prefixes the body with
// `tx = tl.launch_thread_idx(0); ty = ...; tz = ...` Bind statements. The
// launch nest is traced before the Target is known, so the thread indices are
// only placeholders here: the Vars keep their identity through
// tl.MaterializeKernelLaunch, which rebinds them as threadIdx.* thread_extent
// scopes on SIMT backends and drops them elsewhere.
static ForFrame MakeLaunchThreadFrame() {
using namespace tvm::tirx;
static const char *kThreadVarNames[3] = {"tx", "ty", "tz"};
DataType dtype = DataType::Int(32);
ObjectPtr<ForFrameNode> n = make_object<ForFrameNode>();
for (int axis = 0; axis < 3; axis++) {
n->vars.push_back(Var(kThreadVarNames[axis], dtype));
// The extent is decided by the backend at materialization; this dom only
// keeps the ForFrame invariants satisfied.
n->doms.push_back(Range(make_const(dtype, 0), make_const(dtype, 1)));
}
n->f_make_for_loop = [](const Array<Var> &vars, const Array<Range> &doms,
const Array<Optional<PrimExpr>> &steps,
Stmt body) -> Stmt {
Array<Stmt> seq;
for (int axis = 0; axis < static_cast<int>(vars.size()); axis++) {
PrimExpr thread_idx = Call(vars[axis]->dtype, launch_thread_idx(),
{IntImm(DataType::Int(32), axis)});
seq.push_back(tvm::tirx::Bind(vars[axis], thread_idx));
}
seq.push_back(body);
return SeqStmt::Flatten(seq);
};
return ForFrame(n);
}
/*!
* \brief A frame that represents a kernel launch.
*
@@ -223,12 +255,26 @@ ForFrame PersistentFor(const Array<PrimExpr> &domain, const PrimExpr &wave_size,
*/
class KernelLaunchFrameNode : public TIRFrameNode {
public:
/*! \brief Grid loops, thread placeholders and the root block, outer to
* inner. */
Array<TIRFrame> frames;
/*! \brief Program (grid) index vars, one per launch axis. */
Array<tvm::tirx::Var> grid_vars;
/*! \brief Grid extents, one per launch axis. */
Array<PrimExpr> grid_extents;
/*! \brief Placeholder thread index vars for the x, y and z axes. */
Array<tvm::tirx::Var> thread_vars;
/*! \brief Requested SIMT thread-block extents, when threads= was given. */
Optional<Array<PrimExpr>> thread_extents;
static void RegisterReflection() {
namespace refl = reflection;
refl::ObjectDef<KernelLaunchFrameNode>().def_ro(
"frames", &KernelLaunchFrameNode::frames);
refl::ObjectDef<KernelLaunchFrameNode>()
.def_ro("frames", &KernelLaunchFrameNode::frames)
.def_ro("grid_vars", &KernelLaunchFrameNode::grid_vars)
.def_ro("grid_extents", &KernelLaunchFrameNode::grid_extents)
.def_ro("thread_vars", &KernelLaunchFrameNode::thread_vars)
.def_ro("thread_extents", &KernelLaunchFrameNode::thread_extents);
}
TVM_FFI_DECLARE_OBJECT_INFO_FINAL("tl.KernelLaunchFrame",
@@ -270,30 +316,40 @@ KernelLaunchFrame KernelLaunch(const Array<PrimExpr> &grid_size,
const Map<String, Any> &attrs) {
ObjectPtr<KernelLaunchFrameNode> n = make_object<KernelLaunchFrameNode>();
auto block_size = block_size_opt.value_or(Array<PrimExpr>());
ICHECK(grid_size.size() <= 3);
ICHECK(block_size.size() <= 3);
static const char *kBlockVarNames[3] = {"bx", "by", "bz"};
static const char *kBlockTags[3] = {"blockIdx.x", "blockIdx.y", "blockIdx.z"};
static const char *kThreadVarNames[3] = {"tx", "ty", "tz"};
static const char *kThreadTags[3] = {"threadIdx.x", "threadIdx.y",
"threadIdx.z"};
for (size_t i = 0; i < grid_size.size(); i++) {
n->frames.push_back(
MakeThreadBindingFrame(kBlockVarNames[i], kBlockTags[i], grid_size[i]));
ForFrame frame =
MakeThreadBindingFrame(kBlockVarNames[i], kBlockTags[i], grid_size[i]);
n->grid_vars.push_back(frame->vars[0]);
n->grid_extents.push_back(grid_size[i]);
n->frames.push_back(frame);
}
for (size_t i = 0; i < block_size.size(); i++) {
n->frames.push_back(MakeThreadBindingFrame(kThreadVarNames[i],
kThreadTags[i], block_size[i]));
// Thread placeholders are always emitted so the body may reference a thread
// index regardless of whether threads= was given; the backend decides what
// they mean.
ForFrame thread_frame = MakeLaunchThreadFrame();
n->thread_vars = thread_frame->vars;
n->frames.push_back(thread_frame);
Map<String, Any> block_annotations =
attrs.defined() ? attrs : Map<String, Any>{};
if (block_size_opt.defined()) {
Array<PrimExpr> block_size = block_size_opt.value();
ICHECK(block_size.size() <= 3);
while (block_size.size() < 3) {
block_size.push_back(IntImm(DataType::Int(32), 1));
}
n->thread_extents = block_size;
block_annotations.Set(attr::kLaunchThreads, block_size);
}
auto empty_block = tvm::script::ir_builder::tirx::Block(DeviceMainBlockName);
empty_block->reads = Array<tvm::tirx::BufferRegion>();
empty_block->writes = Array<tvm::tirx::BufferRegion>();
Map<String, Any> block_annotations =
attrs.defined() ? attrs : Map<String, Any>{};
empty_block->annotations = block_annotations;
n->frames.push_back(empty_block);
+5
View File
@@ -57,6 +57,11 @@ TIR_DEFINE_TL_BUILTIN(access_ptr)
TIR_DEFINE_TL_BUILTIN(region).set_num_inputs(-1).set_attr<TCallEffectKind>(
"TCallEffectKind", Integer(CallEffectKind::kPure));
TIR_DEFINE_TL_BUILTIN(launch_thread_idx)
.set_num_inputs(1)
.set_attr<TCallEffectKind>("TCallEffectKind",
Integer(CallEffectKind::kOpaque));
TIR_DEFINE_TL_BUILTIN(add2).set_num_inputs(2).set_attr<TCallEffectKind>(
"TCallEffectKind", Integer(CallEffectKind::kPure));
+19
View File
@@ -33,6 +33,12 @@ static constexpr const char *kLocalVarInit = "tl.local_var_init";
static constexpr const char *kNonRestrictParams = "tl.non_restrict_params";
static constexpr const char *kLexicalAllocScope = "lexical_alloc_scope";
// Annotation on the tilelang_root block recording the SIMT thread-block
// extents requested by T.Kernel(threads=...). It is a launch hint: SIMT
// backends materialize it as threadIdx.* thread_extent scopes, other
// backends ignore it.
static constexpr const char *kLaunchThreads = "tl.launch_threads";
} // namespace attr
inline ffi::Optional<PrimExpr> GetAnnotatedMbarPhaseExpr(
@@ -185,6 +191,19 @@ TVM_DLL const Op &access_ptr();
*/
TVM_DLL const Op &region();
/*!
* \brief Placeholder for the thread index along one launch axis.
*
* T.Kernel binds each thread variable as `LetStmt(tx, launch_thread_idx(axis))`
* so the kernel body can reference a thread index before the target is known.
* tl.MaterializeKernelLaunch replaces the binding with a real threadIdx.*
* thread_extent scope on SIMT backends and rejects any use on backends
* without SIMT. It must never reach codegen.
*
* int32 launch_thread_idx(axis)
*/
TVM_DLL const Op &launch_thread_idx();
// Packed x2 element-wise math (float32x2, bfloat16x2, float16x2)
TVM_DLL const Op &add2();
TVM_DLL const Op &sub2();
+224 -47
View File
@@ -3,40 +3,65 @@
* \brief Materialize the target-neutral kernel launch nest emitted by
* T.Kernel into a backend-specific form.
*
* T.Kernel traces into a nest of For loops with ForKind::kThreadBinding
* tagged blockIdx.* / threadIdx.*. This pass runs right after BindTarget
* and rewrites the nest according to `lower_thread_binding`, which each
* backend pipeline chooses for itself (no target dispatch happens here):
* T.Kernel traces into
*
* for bx in thread_binding(blockIdx.x): # one per grid axis
* tx = tl.launch_thread_idx(0) # x, y, z placeholders
* ty = tl.launch_thread_idx(1)
* tz = tl.launch_thread_idx(2)
* block tilelang_root { annotations: tl.launch_threads?, ... }
*
* The grid loops are the program-index space every backend shares. The
* thread placeholders only reserve Var identities the body may reference;
* what they mean, and how many threads run, is decided here once the Target
* is bound. Each backend pipeline chooses the mode for itself (no target
* dispatch happens in this pass):
* - lower_thread_binding = true (SIMT backends, e.g. CUDA/ROCm/Metal):
* each launch loop becomes an AttrStmt thread_extent, reusing the loop
* variable so body references stay valid.
* grid loops become AttrStmt thread_extent scopes; each thread placeholder
* is rebound as a threadIdx.* thread_extent over the same Var, with the
* extent taken from the `tl.launch_threads` annotation (T.Kernel
* threads=...) or, failing that, from `default_threads`.
* - lower_thread_binding = false (backends without SIMT, e.g. CPU):
* blockIdx.* loops become serial For loops over the grid extent;
* threadIdx.* loops are ignored; they become unit serial loops so the
* loop variable stays defined (pinned to 0) while the requested thread
* extent (e.g. the default threads=128) is dropped.
* grid loops become serial For loops; thread placeholders are dropped and
* `tl.launch_threads` is ignored. A body that references a thread index
* has no meaning on such a target and is rejected.
* Launch annotations listed in `unsupported_annotations` (e.g. `cluster_dims`
* on a target without thread block clusters) are rejected instead of being
* silently dropped further down the pipeline.
*
* Only the outermost contiguous launch nest is converted; thread_binding
* loops deeper inside the kernel body (separated by the tilelang_root
* block) are left for LowerOpaqueBlock to handle at its usual stage.
*/
#include "../op/builtin.h"
#include "common/attr.h"
#include "support/check.h"
#include <tvm/ir/transform.h>
#include <tvm/runtime/logging.h>
#include <tvm/target/target.h>
#include <tvm/tirx/analysis.h>
#include <tvm/tirx/op.h>
#include <tvm/tirx/stmt_functor.h>
#include <tvm/tirx/transform.h>
#include <utility>
#include <vector>
namespace tvm {
namespace tl {
using namespace tirx;
using ffi::Array;
using ffi::GetRef;
using ffi::Optional;
namespace {
constexpr int kNumThreadAxes = 3;
constexpr const char *kThreadTags[kNumThreadAxes] = {
"threadIdx.x", "threadIdx.y", "threadIdx.z"};
bool IsBlockBinding(const ForNode *op) {
if (op->kind != ForKind::kThreadBinding || !op->thread_binding.defined())
return false;
@@ -44,69 +69,221 @@ bool IsBlockBinding(const ForNode *op) {
return tag.rfind("blockIdx.", 0) == 0;
}
bool IsThreadBinding(const ForNode *op) {
if (op->kind != ForKind::kThreadBinding || !op->thread_binding.defined())
// `v = tl.launch_thread_idx(axis)` emitted by T.Kernel.
bool IsLaunchThreadPlaceholder(const Stmt &stmt) {
const BindNode *bind = stmt.as<BindNode>();
if (!bind)
return false;
std::string tag = op->thread_binding.value()->thread_tag;
return tag.rfind("threadIdx.", 0) == 0;
const CallNode *call = bind->value.as<CallNode>();
return call && call->op.same_as(launch_thread_idx());
}
bool IsLaunchBinding(const ForNode *op) {
return IsBlockBinding(op) || IsThreadBinding(op);
int ThreadAxisOf(const BindNode *bind) {
const CallNode *call = bind->value.as<CallNode>();
ICHECK(call && call->args.size() == 1);
const int64_t *axis = as_const_int(call->args[0]);
ICHECK(axis && *axis >= 0 && *axis < kNumThreadAxes)
<< "tl.launch_thread_idx expects a constant axis in [0, 3), got "
<< call->args[0];
return static_cast<int>(*axis);
}
// The tilelang_root block that carries the launch annotations, if `body` is
// the kernel body directly below the launch nest.
const SBlockNode *GetLaunchBlock(const Stmt &body) {
const SBlockRealizeNode *realize = body.as<SBlockRealizeNode>();
if (!realize || !IsDeviceMainBlock(realize->block.get()))
return nullptr;
return realize->block.get();
}
Optional<Array<PrimExpr>> GetLaunchThreads(const Stmt &body) {
const SBlockNode *block = GetLaunchBlock(body);
if (!block)
return std::nullopt;
if (auto threads = block->annotations.Get(attr::kLaunchThreads)) {
if (auto arr = threads.value().try_cast<Array<PrimExpr>>())
return arr.value();
LOG(FATAL) << "Expected `" << attr::kLaunchThreads
<< "` to be an Array<PrimExpr>, but got "
<< threads.value().GetTypeKey();
}
return std::nullopt;
}
class KernelLaunchMaterializer : public StmtMutator {
public:
explicit KernelLaunchMaterializer(bool lower_thread_binding)
: lower_thread_binding_(lower_thread_binding) {}
KernelLaunchMaterializer(bool lower_thread_binding,
Optional<Array<PrimExpr>> default_threads,
Array<ffi::String> unsupported_annotations,
ffi::String target_name)
: lower_thread_binding_(lower_thread_binding),
default_threads_(std::move(default_threads)),
unsupported_annotations_(std::move(unsupported_annotations)),
target_name_(std::move(target_name)) {}
Stmt VisitStmt_(const ForNode *op) final {
if (IsLaunchBinding(op)) {
return ConvertNest(op);
if (IsBlockBinding(op)) {
return ConvertNest(GetRef<Stmt>(op));
}
return StmtMutator::VisitStmt_(op);
}
// A launch without grid axes starts directly at the thread placeholders.
Stmt VisitStmt_(const SeqStmtNode *op) final {
if (op->size() > 0 && IsLaunchThreadPlaceholder(op->seq[0])) {
return ConvertNest(GetRef<Stmt>(op));
}
return StmtMutator::VisitStmt_(op);
}
private:
// Peel the contiguous launch nest rooted at `op` without descending into
// the kernel body below it.
Stmt ConvertNest(const ForNode *op) {
Stmt body;
if (const ForNode *inner = op->body.as<ForNode>();
inner && IsLaunchBinding(inner)) {
body = ConvertNest(inner);
} else {
body = op->body;
// Peel the contiguous launch nest rooted at `root` without descending into
// the kernel body below it, then rebuild it in the backend's form.
Stmt ConvertNest(const Stmt &root) {
std::vector<const ForNode *> grid_loops;
Stmt body = root;
while (const ForNode *loop = body.as<ForNode>()) {
if (!IsBlockBinding(loop))
break;
grid_loops.push_back(loop);
body = loop->body;
}
if (lower_thread_binding_) {
ffi::String tag = op->thread_binding.value()->thread_tag;
IterVar iter_var(Range::FromMinExtent(op->min, op->extent), op->loop_var,
IterVarType::kThreadIndex, tag);
return AttrStmt(std::move(iter_var), tirx::attr::thread_extent,
op->extent, std::move(body), op->span);
std::vector<const BindNode *> thread_binds;
body = PeelThreadPlaceholders(body, &thread_binds);
RejectUnsupportedAnnotations(body);
body = lower_thread_binding_ ? BindThreads(thread_binds, body)
: DropThreads(thread_binds, body);
for (auto it = grid_loops.rbegin(); it != grid_loops.rend(); ++it) {
const ForNode *loop = *it;
if (lower_thread_binding_) {
ffi::String tag = loop->thread_binding.value()->thread_tag;
IterVar iter_var(Range::FromMinExtent(loop->min, loop->extent),
loop->loop_var, IterVarType::kThreadIndex, tag);
body = AttrStmt(std::move(iter_var), tirx::attr::thread_extent,
loop->extent, std::move(body), loop->span);
} else {
body = For(loop->loop_var, loop->min, loop->extent, ForKind::kSerial,
std::move(body),
/*thread_binding=*/std::nullopt, loop->annotations,
loop->step, loop->span);
}
}
// No SIMT: grid dims run as plain serial loops; thread dims are ignored
// (a unit loop keeps the loop variable defined and pinned to 0).
PrimExpr extent = IsThreadBinding(op)
? PrimExpr(IntImm(op->extent.dtype(), 1))
: op->extent;
return For(op->loop_var, op->min, std::move(extent), ForKind::kSerial,
std::move(body),
/*thread_binding=*/std::nullopt, op->annotations, op->step,
op->span);
return body;
}
// Split the leading `v = tl.launch_thread_idx(axis)` binds off `stmt` and
// return what follows them.
static Stmt PeelThreadPlaceholders(const Stmt &stmt,
std::vector<const BindNode *> *binds) {
const SeqStmtNode *seq = stmt.as<SeqStmtNode>();
if (!seq)
return stmt;
size_t i = 0;
while (i < seq->size() && IsLaunchThreadPlaceholder(seq->seq[i])) {
binds->push_back(seq->seq[i].as<BindNode>());
++i;
}
if (i == 0)
return stmt;
ICHECK_LT(i, seq->size())
<< "T.Kernel launch has thread placeholders but no body";
if (i + 1 == seq->size())
return seq->seq[i];
return SeqStmt(Array<Stmt>(seq->seq.begin() + i, seq->seq.end()));
}
// SIMT: every placeholder becomes a threadIdx.* thread_extent scope over
// the same Var so body references stay valid.
Stmt BindThreads(const std::vector<const BindNode *> &thread_binds,
Stmt body) {
if (thread_binds.empty())
return body;
Array<PrimExpr> extents = ResolveThreadExtents(body);
for (auto it = thread_binds.rbegin(); it != thread_binds.rend(); ++it) {
const BindNode *bind = *it;
int axis = ThreadAxisOf(bind);
PrimExpr extent = extents[axis];
IterVar iter_var(
Range::FromMinExtent(make_zero(bind->var.dtype()), extent), bind->var,
IterVarType::kThreadIndex, kThreadTags[axis]);
body = AttrStmt(std::move(iter_var), tirx::attr::thread_extent, extent,
std::move(body), bind->span);
}
return body;
}
// No SIMT: thread placeholders carry no meaning, so they are removed. A body
// that reads one would otherwise silently run as a single thread.
Stmt DropThreads(const std::vector<const BindNode *> &thread_binds,
Stmt body) {
for (const BindNode *bind : thread_binds) {
const VarNode *var = bind->var.get();
if (UsesVar(body, [var](const VarNode *v) { return v == var; })) {
LOG(FATAL) << "T.Kernel body references thread index `"
<< bind->var->name_hint << "`, but target `" << target_name_
<< "` has no SIMT threads. Express the computation with "
"tile-level operators (T.Parallel, T.copy, ...) instead "
"of per-thread indexing.";
}
}
return body;
}
void RejectUnsupportedAnnotations(const Stmt &body) const {
const SBlockNode *block = GetLaunchBlock(body);
if (!block)
return;
for (const ffi::String &key : unsupported_annotations_) {
if (block->annotations.count(key)) {
LOG(FATAL) << "T.Kernel launch annotation `" << key
<< "` is not supported on target `" << target_name_ << "`";
}
}
}
Array<PrimExpr> ResolveThreadExtents(const Stmt &body) {
Optional<Array<PrimExpr>> threads = GetLaunchThreads(body);
if (!threads.defined())
threads = default_threads_;
ICHECK(threads.defined())
<< "T.Kernel did not specify threads= and target `" << target_name_
<< "` provides no default thread-block size";
Array<PrimExpr> extents = threads.value();
ICHECK_LE(extents.size(), static_cast<size_t>(kNumThreadAxes));
while (extents.size() < static_cast<size_t>(kNumThreadAxes)) {
extents.push_back(IntImm(DataType::Int(32), 1));
}
return extents;
}
bool lower_thread_binding_;
Optional<Array<PrimExpr>> default_threads_;
Array<ffi::String> unsupported_annotations_;
ffi::String target_name_;
};
} // namespace
tvm::transform::Pass MaterializeKernelLaunch(bool lower_thread_binding) {
tvm::transform::Pass
MaterializeKernelLaunch(bool lower_thread_binding,
Optional<Array<PrimExpr>> default_threads,
Optional<Array<ffi::String>> unsupported_annotations) {
using namespace tirx::transform;
auto pass_func = [lower_thread_binding](
Array<ffi::String> unsupported =
unsupported_annotations.value_or(Array<ffi::String>());
auto pass_func = [lower_thread_binding, default_threads, unsupported](
PrimFunc func, const IRModule &mod,
const tvm::transform::PassContext &ctx) -> PrimFunc {
KernelLaunchMaterializer mutator(lower_thread_binding);
ffi::String target_name = "<unbound>";
if (auto target = func->GetAttr<Target>(tvm::attr::kTarget)) {
target_name = target.value()->kind->name;
}
KernelLaunchMaterializer mutator(lower_thread_binding, default_threads,
unsupported, target_name);
func.CopyOnWrite()->body = mutator(func->body);
return func;
};
+12 -12
View File
@@ -49,9 +49,9 @@ def _simple_program():
@T.prim_func
def program(A: T.Tensor((128,), "float32"), B: T.Tensor((128,), "float32")):
with T.Kernel(threads=128):
tid = T.get_thread_binding()
B[tid] = A[tid] + 1.0
with T.Kernel(1):
for i in T.serial(128):
B[i] = A[i] + 1.0
return program
@@ -186,9 +186,9 @@ def test_multiple_pipelines_share_one_compile_session(monkeypatch, tmp_path):
@T.prim_func
def tiny(A: T.Tensor((32,), "float32"), B: T.Tensor((32,), "float32")):
with T.Kernel(32):
tid = T.get_thread_binding()
B[tid] = A[tid] + 1.0
with T.Kernel(1):
for i in T.serial(32):
B[i] = A[i] + 1.0
mod = tvm.IRModule({"main": tiny})
context = create_backend_context("c", "c", "cython")
@@ -279,9 +279,9 @@ def test_no_skipped_phantom_records(monkeypatch, tmp_path):
@T.prim_func
def tiny(A: T.Tensor((32,), "float32"), B: T.Tensor((32,), "float32")):
with T.Kernel(32):
tid = T.get_thread_binding()
B[tid] = A[tid] + 1.0
with T.Kernel(1):
for i in T.serial(32):
B[i] = A[i] + 1.0
tilelang.lower(tiny, target="c")
@@ -325,9 +325,9 @@ def test_terminal_mode_no_html(monkeypatch, tmp_path):
@T.prim_func
def tiny(A: T.Tensor((32,), "float32"), B: T.Tensor((32,), "float32")):
with T.Kernel(32):
tid = T.get_thread_binding()
B[tid] = A[tid] + 1.0
with T.Kernel(1):
for i in T.serial(32):
B[i] = A[i] + 1.0
tilelang.lower(tiny, target="c")
@@ -44,6 +44,22 @@ def get_cta_rank_in_cluster(cluster_size=4):
return main
@tilelang.jit(out_idx=-1)
def get_cluster_id_kernel(cluster_size=4):
assert 128 % cluster_size == 0
@T.prim_func
def main(A: T.Tensor((128, 2), T.int32)):
with T.ClusterKernel(128, cluster_dims=(cluster_size, 1, 1)) as bx:
if T.get_thread_binding() == 0:
A[bx, 0] = T.get_cluster_id()
# Program-space cluster id and the hardware rank must agree on
# which programs form a cluster.
A[bx, 1] = T.get_cluster_id() * T.get_cluster_size() + T.block_rank_in_cluster()
return main
@tilelang.jit(out_idx=-1)
def barrier_kernel():
@T.prim_func
@@ -110,6 +126,15 @@ def test_cluster_launch_intrinsics(cluster_size=4):
assert torch.all(result == ref)
@tilelang.testing.requires_cuda
@tilelang.testing.requires_cuda_compute_version_ge(9, 0)
def test_cluster_id_matches_hardware_rank(cluster_size=4):
result = get_cluster_id_kernel(cluster_size)()
bx = torch.arange(128, dtype=torch.int32, device="cuda")
assert torch.all(result[:, 0] == bx // cluster_size)
assert torch.all(result[:, 1] == bx)
@tilelang.testing.requires_cuda
@tilelang.testing.requires_cuda_compute_version_ge(9, 0)
def test_cluster_barrier():
@@ -238,5 +238,83 @@ def test_jit2_compile_with_consts():
transpose.compile(M=1024, N=1024, block_M=64, block_N=64)
def _kernel_body_is_skipped_in_phase1(func) -> bool:
"""Whether the eager rewriter guards the launch body with skip_kernel_ctx,
which is what keeps phase-1 signature inference from executing it."""
from tilelang.language.eager.ast import mutate
return "skip_kernel_ctx" in mutate(func).source
def test_jit2_recognizes_launch_from_every_dialect():
from tilelang.cpu import language as Tcpu
from tilelang.cuda import language as Tcuda
from tilelang.rocm import language as Trocm
Launch = T.Kernel
def default_facade(A):
with T.Kernel(1):
pass
def cuda_dialect(A):
with Tcuda.Kernel(1, threads=128):
pass
def rocm_dialect(A):
with Trocm.Kernel(1, threads=64):
pass
def cpu_dialect(A):
with Tcpu.Kernel(1):
pass
def cluster(A):
with T.ClusterKernel(2, cluster_dims=2):
pass
def aliased(A):
with Launch(1):
pass
def not_a_launch(A):
with T.ws(0):
pass
for func in (default_facade, cuda_dialect, rocm_dialect, cpu_dialect, cluster, aliased):
assert _kernel_body_is_skipped_in_phase1(func), func.__name__
assert not _kernel_body_is_skipped_in_phase1(not_a_launch)
@tilelang.testing.requires_cuda
def test_jit2_phase1_does_not_execute_kernel_body():
"""Phase 1 infers the signature with symbolic T.const values, so the launch
body must not run then: a gemm tile that depends on a const is only valid
once the values are bound in phase 2."""
@tilelang.jit
def gemm_full_n(A, B, block_M, block_K):
M, N, K = T.const("M, N, K")
A: T.Tensor[[M, K], T.float16]
B: T.Tensor[[K, N], T.float16]
C = T.empty((M, N), T.float16)
with T.Kernel(T.ceildiv(M, block_M), threads=128) as bx:
A_shared = T.alloc_shared((block_M, block_K), T.float16)
B_shared = T.alloc_shared((block_K, N), T.float16)
C_local = T.alloc_fragment((block_M, N), T.float32)
T.clear(C_local)
for k in T.Pipelined(T.ceildiv(K, block_K), num_stages=2):
T.copy(A[bx * block_M, k * block_K], A_shared)
T.copy(B[k * block_K, 0], B_shared)
T.gemm(A_shared, B_shared, C_local)
T.copy(C_local, C[bx * block_M, 0])
return C
a = torch.randn(256, 128, device="cuda", dtype=torch.float16)
b = torch.randn(128, 64, device="cuda", dtype=torch.float16)
c = gemm_full_n(a, b, 64, 32)
torch.testing.assert_close(c, (a.float() @ b.float()).half(), rtol=1e-2, atol=1e-2)
if __name__ == "__main__":
tilelang.testing.main()
@@ -20,7 +20,6 @@ def test_normalize_threads_rejects_non_positive(threads):
@pytest.mark.parametrize(
"threads, expected",
[
(None, [128, 1, 1]),
(256, [256, 1, 1]),
([32, 4], [32, 4, 1]),
((32, 2, 2), [32, 2, 2]),
@@ -31,5 +30,11 @@ def test_normalize_threads_accepts_positive(threads, expected):
assert _normalize_threads(threads) == expected
def test_normalize_threads_leaves_default_to_backend():
"""No threads= means no SIMT hint: the backend picks its default when it
materializes the launch, the frontend does not guess one."""
assert _normalize_threads(None) is None
if __name__ == "__main__":
tilelang.testing.main()
@@ -26,6 +26,15 @@ def _strip_block_reads_writes(stmt, strip_annotations: bool = False):
return ir_transform(stmt, None, _postorder)
def _materialize_launch(func):
"""Run the launch materialization the pipeline performs before
LegalizeSafeMemoryAccess, so thread indices carry their extents."""
mod = tvm.IRModule({func.attrs["global_symbol"]: func})
mod = tvm.tirx.transform.BindTarget(tvm.target.Target("cuda"))(mod)
mod = tl.transform.MaterializeKernelLaunch()(mod)
return mod[func.attrs["global_symbol"]]
def _collect_call_nodes(stmt, op_names):
if isinstance(op_names, str):
op_names = {op_names}
@@ -115,6 +124,7 @@ def vectorize_access_legalize(M: int = 64, N: int = 64, M_offset: int = 2, N_off
def assert_vectorize_access(M: int = 64, N: int = 64):
func, expected = vectorize_access_legalize(M, N)
func, expected = _materialize_launch(func), _materialize_launch(expected)
mod = tvm.IRModule({func.attrs["global_symbol"]: func})
transformed = tl.transform.LegalizeSafeMemoryAccess()(mod)
@@ -160,6 +170,7 @@ def vectorize_access_with_atmoic_add_legalize(M: int = 64, N: int = 64, M_offset
def assert_vectorize_access_with_atmoic_add(M: int = 64, N: int = 64):
func, expected = vectorize_access_with_atmoic_add_legalize(M, N)
func, expected = _materialize_launch(func), _materialize_launch(expected)
mod = tvm.IRModule({func.attrs["global_symbol"]: func})
transformed = tl.transform.LegalizeSafeMemoryAccess()(mod)
print(transformed)
@@ -198,6 +209,7 @@ def oob_store_legalize(M: int = 64, N: int = 64, M_offset: int = 2, N_offset: in
def assert_oob_store_legalize(M: int = 64, N: int = 64):
func, expected = oob_store_legalize(M, N)
func, expected = _materialize_launch(func), _materialize_launch(expected)
mod = tvm.IRModule({func.attrs["global_symbol"]: func})
transformed = tl.transform.LegalizeSafeMemoryAccess()(mod)
tvm.ir.assert_structural_equal(
@@ -0,0 +1,337 @@
"""Tests for the target-neutral T.Kernel encoding and MaterializeKernelLaunch.
T.Kernel is traced before the Target is known, so it only records the grid
loops, thread-index placeholders and launch annotations; the backend pipeline
decides what the thread placeholders mean when it runs MaterializeKernelLaunch.
"""
import importlib
import inspect
import pytest
import tilelang as tl
import tilelang.language as T
import tilelang
import tilelang.testing
from tilelang import tvm
from tvm.tirx.stmt_functor import post_order_visit
def _collect(root, kind):
found = []
def _visit(node):
if isinstance(node, kind):
found.append(node)
post_order_visit(root.body if hasattr(root, "body") else root, _visit)
return found
def _launch_placeholders(func):
return [
stmt
for stmt in _collect(func, tvm.tirx.Bind)
if isinstance(stmt.value, tvm.tirx.Call) and str(stmt.value.op.name) == "tl.launch_thread_idx"
]
def _thread_extents(func):
"""{thread_tag: extent} of every thread_extent AttrStmt in `func`."""
extents = {}
for attr in _collect(func, tvm.tirx.AttrStmt):
if attr.attr_key == "thread_extent":
extents[str(attr.node.thread_tag)] = int(attr.value)
return extents
def _root_block(func):
blocks = [b for b in _collect(func, tvm.tirx.SBlock) if b.name_hint == "tilelang_root"]
assert len(blocks) == 1
return blocks[0]
def _materialize(func, target: str, **kwargs):
mod = tvm.IRModule.from_expr(func)
mod = tvm.tirx.transform.BindTarget(tvm.target.Target(target))(mod)
mod = tl.transform.MaterializeKernelLaunch(**kwargs)(mod)
return mod[func.attrs["global_symbol"]]
def _parallel_kernel(threads=None):
@T.prim_func
def main(A: T.Tensor((256,), "float32"), B: T.Tensor((256,), "float32")):
with T.Kernel(2, threads=threads) as bx:
for i in T.Parallel(128):
B[bx * 128 + i] = A[bx * 128 + i] + 1.0
return main
def _thread_indexed_kernel():
@T.prim_func
def main(A: T.Tensor((256,), "float32"), B: T.Tensor((256,), "float32")):
with T.Kernel(2, threads=128) as bx:
tx = T.get_thread_binding()
B[bx * 128 + tx] = A[bx * 128 + tx] + 1.0
return main
def test_traced_launch_records_grid_and_thread_placeholders():
func = _parallel_kernel()
grid = [f for f in _collect(func, tvm.tirx.For) if f.kind == tvm.tirx.ForKind.THREAD_BINDING]
assert [str(f.thread_binding.thread_tag) for f in grid] == ["blockIdx.x"]
placeholders = _launch_placeholders(func)
assert [p.var.name for p in placeholders] == ["tx", "ty", "tz"]
assert [int(p.value.args[0]) for p in placeholders] == [0, 1, 2]
# No threads= means no SIMT hint is recorded; the backend picks.
assert "tl.launch_threads" not in _root_block(func).annotations
assert _thread_extents(func) == {}
def test_traced_launch_records_requested_threads_as_annotation():
func = _parallel_kernel(threads=(64, 2))
threads = _root_block(func).annotations["tl.launch_threads"]
assert [int(x) for x in threads] == [64, 2, 1]
# Still only placeholders: the frontend does not bind threadIdx itself.
assert _thread_extents(func) == {}
def test_kernel_launch_annotations_are_recorded_on_the_root_block():
@T.prim_func
def main(A: T.Tensor((16,), "int32")):
with T.Kernel(1, threads=64, prelude="// hi", cluster_dims=2):
A[0] = 0
annotations = _root_block(main).annotations
assert [int(x) for x in annotations["tl.launch_threads"]] == [64, 1, 1]
assert [int(x) for x in annotations["cluster_dims"]] == [2, 1, 1]
assert str(annotations["pragma_import_c"]) == "// hi"
def test_kernel_rejects_unknown_launch_annotation():
"""A dialect's Kernel declares its launch annotations as explicit keyword
parameters, so a misspelled or foreign key fails at trace time."""
with pytest.raises(TypeError, match="unexpected keyword argument 'thread'"):
@T.prim_func
def typo(A: T.Tensor((16,), "int32")):
with T.Kernel(1, thread=128):
A[0] = 0
with pytest.raises(TypeError, match="unexpected keyword argument 'core_type'"):
@T.prim_func
def foreign(A: T.Tensor((16,), "int32")):
with T.Kernel(1, core_type="aiv"):
A[0] = 0
def _launch_annotations(kernel) -> set[str]:
return {name for name, p in inspect.signature(kernel).parameters.items() if p.kind is inspect.Parameter.KEYWORD_ONLY}
def test_each_dialect_declares_its_own_launch_annotations():
expected = {
"tilelang.language.common": set(),
"tilelang.cuda.language": {"threads", "prelude", "cluster_dims"},
"tilelang.rocm.language": {"threads", "prelude"},
"tilelang.metal.language": {"threads", "prelude"},
"tilelang.webgpu.language": {"threads"},
"tilelang.cpu.language": {"prelude"},
}
for module, keys in expected.items():
dialect = importlib.import_module(module)
assert _launch_annotations(dialect.Kernel) == keys, module
assert dialect.Kernel.__module__.startswith(module.removesuffix(".common")), module
# The default facade is the CUDA dialect.
assert T.Kernel is importlib.import_module("tilelang.cuda.language").Kernel
def test_cpu_dialect_kernel_has_no_threads():
from tilelang.cpu import language as Tcpu
with pytest.raises(TypeError, match="unexpected keyword argument 'threads'"):
@Tcpu.prim_func
def main(A: Tcpu.Tensor((16,), "int32")):
with Tcpu.Kernel(1, threads=128):
A[0] = 0
@Tcpu.prim_func
def ok(A: Tcpu.Tensor((16,), "int32")):
with Tcpu.Kernel(1, prelude="// cpu"):
A[0] = 0
assert str(_root_block(ok).annotations["pragma_import_c"]) == "// cpu"
def test_simt_binds_requested_threads():
func = _materialize(_parallel_kernel(threads=64), "cuda")
assert _thread_extents(func) == {"blockIdx.x": 2, "threadIdx.x": 64, "threadIdx.y": 1, "threadIdx.z": 1}
assert _launch_placeholders(func) == []
def test_simt_uses_backend_default_when_threads_omitted():
func = _materialize(_parallel_kernel(), "cuda", default_threads=256)
assert _thread_extents(func)["threadIdx.x"] == 256
func = _materialize(_parallel_kernel(), "cuda")
assert _thread_extents(func)["threadIdx.x"] == tl.transform.DEFAULT_SIMT_THREADS
def test_simt_requires_threads_when_backend_has_no_default():
with pytest.raises(Exception, match="did not specify threads="):
_materialize(_parallel_kernel(), "cuda", default_threads=None)
def test_simt_preserves_thread_var_identity():
"""The Var handed out by T.get_thread_binding() at trace time must be the
Var bound by the threadIdx.x thread_extent after materialization."""
func = _thread_indexed_kernel()
(placeholder,) = [p for p in _launch_placeholders(func) if p.var.name == "tx"]
lowered = _materialize(func, "cuda")
(attr,) = [a for a in _collect(lowered, tvm.tirx.AttrStmt) if str(a.node.thread_tag) == "threadIdx.x"]
assert attr.node.var.same_as(placeholder.var)
body_vars = [v for v in _collect(lowered, tvm.tirx.Var) if v.name == "tx"]
assert body_vars and all(v.same_as(placeholder.var) for v in body_vars)
def test_non_simt_drops_thread_placeholders():
func = _materialize(_parallel_kernel(threads=128), "c", lower_thread_binding=False)
assert _thread_extents(func) == {}
assert _launch_placeholders(func) == []
grid = [f for f in _collect(func, tvm.tirx.For) if f.loop_var.name == "bx"]
assert len(grid) == 1 and grid[0].kind == tvm.tirx.ForKind.SERIAL and int(grid[0].extent) == 2
assert not any(v.name in ("tx", "ty", "tz") for v in _collect(func, tvm.tirx.Var))
def test_non_simt_rejects_thread_index_use():
with pytest.raises(Exception, match="references thread index `tx`"):
_materialize(_thread_indexed_kernel(), "c", lower_thread_binding=False)
def test_get_thread_extent_requires_threads_at_trace_time():
with pytest.raises(ValueError, match="not known at trace time"):
@T.prim_func
def main(A: T.Tensor((16,), "int32")):
with T.Kernel(1):
A[0] = T.get_thread_extent()
def test_get_thread_extent_with_threads_at_trace_time():
@T.prim_func
def main(A: T.Tensor((16,), "int32")):
with T.Kernel(1, threads=(32, 4)):
A[0] = T.get_thread_extent(0) * T.get_thread_extent(1)
(store,) = _collect(main, tvm.tirx.BufferStore)
assert int(store.value) == 128
class _TraceFailure(Exception):
pass
def test_failed_trace_unwinds_launch_frames():
"""An exception inside T.Kernel must leave no stale launch frame behind,
otherwise the next trace sees the previous kernel's KernelLaunchFrame."""
with pytest.raises(_TraceFailure):
@T.prim_func
def failing(A: T.Tensor((16,), "int32")):
with T.Kernel(1, threads=128):
raise _TraceFailure()
assert T.KernelLaunchFrame.Current() is None
@tilelang.jit
def failing_jit(A):
A: T.Tensor[[16], T.int32]
with T.Kernel(1, threads=128):
raise _TraceFailure()
import torch
with pytest.raises(_TraceFailure):
failing_jit.get_tir(torch.zeros(16, dtype=torch.int32))
assert T.KernelLaunchFrame.Current() is None
def _cluster_kernel():
@T.prim_func
def main(A: T.Tensor((16,), "int32")):
with T.ClusterKernel(8, 4, threads=128, cluster_dims=2) as (bx, by):
A[0] = 0
return main
def test_cluster_dims_is_a_launch_annotation():
dims = _root_block(_cluster_kernel()).annotations["cluster_dims"]
assert [int(d) for d in dims] == [2, 1, 1]
def test_cluster_id_is_program_space_arithmetic():
"""Cluster identity is derived from the program index and cluster_dims at
trace time, so it needs no target-specific intrinsic."""
captured = {}
@T.prim_func
def main(A: T.Tensor((16,), "int32")):
with T.ClusterKernel(8, 4, threads=128, cluster_dims=2) as (bx, by):
captured["bx"], captured["by"] = bx, by
captured["ids"] = T.get_cluster_ids()
captured["dims"] = T.get_cluster_dims()
captured["size"] = T.get_cluster_size()
captured["extents"] = T.get_cluster_extents()
A[0] = T.get_cluster_id(0)
cx, cy = captured["ids"]
assert isinstance(cx, tvm.tirx.FloorDiv) and cx.a.same_as(captured["bx"]) and int(cx.b) == 2
# A unit cluster axis is the program index itself.
assert cy.same_as(captured["by"])
assert captured["dims"] == [2, 1, 1]
assert captured["size"] == 2
assert captured["extents"] == [4, 4, 1]
(store,) = _collect(main, tvm.tirx.BufferStore)
assert isinstance(store.value, tvm.tirx.FloorDiv)
def test_cluster_id_without_clusters_is_the_program_index():
captured = {}
@T.prim_func
def main(A: T.Tensor((16,), "int32")):
with T.Kernel(8) as bx:
captured["bx"] = bx
captured["id"] = T.get_cluster_id()
captured["dims"] = T.get_cluster_dims()
captured["extents"] = T.get_cluster_extents()
A[0] = 0
assert captured["id"].same_as(captured["bx"])
assert captured["dims"] == [1, 1, 1]
# Axes beyond the launched grid have a single cluster.
assert captured["extents"] == [8, 1, 1]
def test_cluster_dims_accepted_by_default_and_rejected_when_unsupported():
func = _materialize(_cluster_kernel(), "cuda")
assert "cluster_dims" in _root_block(func).annotations
with pytest.raises(Exception, match="`cluster_dims` is not supported on target `c`"):
_materialize(_cluster_kernel(), "c", lower_thread_binding=False, unsupported_annotations=["cluster_dims"])
# A launch without the annotation is unaffected by the rejection list.
_materialize(_parallel_kernel(), "c", lower_thread_binding=False, unsupported_annotations=["cluster_dims"])
if __name__ == "__main__":
tilelang.testing.main()
@@ -58,9 +58,9 @@ def _marker_line(marker: str) -> int:
def _make_vector_add():
@T.prim_func
def main(A: T.Tensor((1024,), "float32"), B: T.Tensor((1024,), "float32")):
with T.Kernel(1024):
tid = T.get_thread_binding()
B[tid] = A[tid] + 1.0 # span_marker_vadd_store
with T.Kernel(8) as bx:
for i in T.Parallel(128):
B[bx * 128 + i] = A[bx * 128 + i] + 1.0 # span_marker_vadd_store
return main
@@ -96,8 +96,9 @@ def _lower_with_recorder(func, target: str) -> _SpanCoverageRecorder:
return recorder
# Passes whose span propagation was fixed; they must never drop a span
# (statement *deletion* is fine — it also reduces the total).
# Passes whose span propagation was fixed; they must never strip a span from
# a statement or introduce a new statement without one (statement *deletion*
# is fine — it reduces the spanned and total counts alike).
_SPAN_SAFE_PASSES = {
"tl.MaterializeKernelLaunch",
"tl.AddWrapperForSingleBufStore",
@@ -112,11 +113,13 @@ _SPAN_SAFE_PASSES = {
def _assert_no_span_loss(recorder: _SpanCoverageRecorder):
prev_w = None
for name, w, _t in recorder.rows:
if prev_w is not None and name in _SPAN_SAFE_PASSES:
assert w >= prev_w, f"pass {name} dropped spans: {prev_w} -> {w}"
prev_w = w
prev = None
for name, w, t in recorder.rows:
if prev is not None and name in _SPAN_SAFE_PASSES:
prev_w, prev_t = prev
unspanned, prev_unspanned = t - w, prev_t - prev_w
assert unspanned <= prev_unspanned, f"pass {name} dropped spans: {prev_w}/{prev_t} spanned -> {w}/{t} spanned"
prev = (w, t)
def test_span_survives_lowering_cpu():
+6 -3
View File
@@ -5,7 +5,10 @@ from __future__ import annotations
from tilelang.language.common import * # noqa: F401,F403
from tilelang.language.common import __all__ as _COMMON_ALL
__tilelang_dialect__ = "cpu"
__all__ = tuple(_COMMON_ALL)
from .kernel import * # noqa: F401,F403
from .kernel import __all__ as _KERNEL_ALL
del _COMMON_ALL
__tilelang_dialect__ = "cpu"
__all__ = tuple(dict.fromkeys((*_COMMON_ALL, *_KERNEL_ALL)))
del _COMMON_ALL, _KERNEL_ALL
+41
View File
@@ -0,0 +1,41 @@
"""CPU dialect of ``T.Kernel``: the common launch plus CPU launch annotations."""
from __future__ import annotations
from tvm import tirx
from tilelang.language.kernel import KernelLaunchFrame, kernel_launch_factory, launch_kernel
__all__ = ["Kernel"]
@kernel_launch_factory
def Kernel(
*blocks: int | tirx.PrimExpr,
prelude: str | None = None,
) -> KernelLaunchFrame:
"""Construct a kernel launch frame for CPU: a grid of tile programs.
The grid becomes the outer loop nest of the generated function and each
tile program runs as a plain serial body; ``T.Parallel`` loops are lowered
to serial loops. There are no SIMT threads, so this dialect has no
``threads`` and ``T.get_thread_binding()`` is rejected at compile time.
Parameters
----------
*blocks : int | PrimExpr
Grid extent along each axis (1-3 dimensions). The launch yields one
program index per axis.
prelude : str, optional
C source injected before the generated kernel, e.g. ``#include`` lines
or helper functions.
Examples
--------
.. code-block:: python
with T.Kernel(T.ceildiv(N, 128)) as bx:
for i in T.Parallel(128):
...
"""
return launch_kernel(blocks, prelude=prelude)
+1 -1
View File
@@ -14,7 +14,7 @@ from tilelang.backend.pass_pipeline.pipeline_utils import (
def CPUPassPipelineBody(mod: IRModule, target: Target) -> IRModule:
mod = tirx.transform.BindTarget(target)(mod)
mod = tilelang.transform.MaterializeKernelLaunch(lower_thread_binding=False)(mod)
mod = tilelang.transform.MaterializeKernelLaunch(lower_thread_binding=False, unsupported_annotations=["cluster_dims"])(mod)
pass_ctx = tilelang.transform.get_pass_context()
if should_force_let_inline():
+4
View File
@@ -46,6 +46,9 @@ from tilelang.language.builtin import ( # noqa: F401
from tilelang.language.copy_op import copy_cluster, tma_copy, tma_gather4, tma_gather4_bytes, tma_scatter4 # noqa: F401
from tilelang.language.kernel import ClusterKernel, CUDASourceCodeKernel # noqa: F401
# The CUDA dialect's T.Kernel shadows the target-neutral one from common: same
# launch, plus the CUDA launch annotations (threads, prelude, cluster_dims).
from .kernel import Kernel # noqa: F401
from .cluster import * # noqa: F401,F403
from .cluster import __all__ as _CLUSTER_ALL
from .intrinsics import * # noqa: F401,F403
@@ -66,6 +69,7 @@ from .warpgroup import __all__ as _WARPGROUP_ALL
_CUDA_API_ALL = (
"ClusterKernel",
"CUDASourceCodeKernel",
"Kernel",
"alloc_cluster_barrier",
"alloc_descriptor",
"alloc_tmem",
+55
View File
@@ -0,0 +1,55 @@
"""CUDA dialect of ``T.Kernel``: the common launch plus CUDA launch annotations."""
from __future__ import annotations
from tvm import tirx
from tilelang.language.kernel import KernelLaunchFrame, kernel_launch_factory, launch_kernel
__all__ = ["Kernel"]
@kernel_launch_factory
def Kernel(
*blocks: int | tirx.PrimExpr,
threads: int | list[int] | tuple[int, ...] | None = None,
prelude: str | None = None,
cluster_dims: int | tuple[int, int, int] | list[int] | None = None,
) -> KernelLaunchFrame:
"""Construct a kernel launch frame for CUDA: a grid of thread blocks.
Code inside the launch operates at the block level: ``T.Parallel``,
``T.copy`` and friends are mapped onto threads by the compiler.
``T.get_thread_binding()`` exposes ``threadIdx`` for thread-level code.
The keyword arguments are recorded at trace time and materialized by the
CUDA pipeline once the target is known.
Parameters
----------
*blocks : int | PrimExpr
Grid extent along each axis (1-3 dimensions, ``gridDim.(x|y|z)``). The
launch yields one block index per axis (``blockIdx.(x|y|z)``).
threads : int | list[int] | tuple[int, ...], optional
Threads per block: a count for ``blockDim.x`` or up to three
per-dimension extents for ``blockDim.(x|y|z)``. Defaults to 128 when
omitted.
prelude : str, optional
CUDA source injected before the generated kernel, e.g. ``#include``
lines or helper functions.
cluster_dims : int | tuple[int, int, int] | list[int], optional
Thread block cluster shape (SM90+). ``2`` or ``(2, 1, 1)`` launches
2-CTA clusters via ``cudaLaunchKernelEx``. ``T.ClusterKernel`` is the
same launch with a required ``cluster_dims``.
Examples
--------
.. code-block:: python
with T.Kernel(T.ceildiv(N, 128), threads=128) as bx:
...
with T.Kernel(grid_x, grid_y, threads=(64, 2)) as (bx, by):
tx, ty = T.get_thread_bindings()
...
"""
return launch_kernel(blocks, threads=threads, prelude=prelude, cluster_dims=cluster_dims)
+1 -1
View File
@@ -67,7 +67,7 @@ def _module_has_shared_barrier(mod: IRModule) -> bool:
def CUDAPassPipelineBodyPrologue(mod: IRModule, target: Target) -> IRModule:
mod = tirx.transform.BindTarget(target)(mod)
mod = tilelang.transform.MaterializeKernelLaunch()(mod)
mod = tilelang.transform.MaterializeKernelLaunch(lower_thread_binding=True, default_threads=128)(mod)
# Record body-bound global bases before optional let inlining obscures
# their provenance. CopyAnalysis consumes this marker for every lowering
# path, independently of whether warp specialization is enabled.
+4
View File
@@ -11,4 +11,8 @@ from __future__ import annotations
from tilelang.cuda.language import * # noqa: F401,F403
from tilelang.cuda.language import __all__ as __all__ # noqa: F401
# Imported by name so static type checkers resolve the CUDA-typed launch
# signature through this facade (they cannot evaluate the dynamic __all__).
from tilelang.cuda.language import Kernel # noqa: F401
__tilelang_dialect__ = "cuda"
+12
View File
@@ -42,6 +42,12 @@ from .kernel import (
get_block_bindings, # noqa: F401
get_block_extent, # noqa: F401
get_block_extents, # noqa: F401
get_cluster_dims, # noqa: F401
get_cluster_size, # noqa: F401
get_cluster_id, # noqa: F401
get_cluster_ids, # noqa: F401
get_cluster_extent, # noqa: F401
get_cluster_extents, # noqa: F401
)
from .allocate import (
alloc_var, # noqa: F401
@@ -265,6 +271,12 @@ _LOCAL_EXPORTS = (
"get_block_bindings",
"get_block_extent",
"get_block_extents",
"get_cluster_dims",
"get_cluster_extent",
"get_cluster_extents",
"get_cluster_id",
"get_cluster_ids",
"get_cluster_size",
"get_lane_idx",
"get_let_value",
"get_thread_binding",
+14 -4
View File
@@ -621,11 +621,21 @@ class DSLMutator(ast.NodeTransformer):
is_kernel_ctx = False
for expr in node.items:
cexpr = expr.context_expr
if isinstance(cexpr, ast.Call) and isinstance(cexpr.func, ast.Attribute) and cexpr.func.attr in ("Kernel", "ClusterKernel"):
eval_res = self._try_eval(cexpr.func)
from tilelang.language.kernel import ClusterKernel, Kernel
if isinstance(cexpr, ast.Call) and isinstance(cexpr.func, (ast.Attribute, ast.Name)):
# Only resolve plain names and module attribute chains, so no
# factory expression such as make_scope().context() is executed
# at rewrite time.
root = cexpr.func
while isinstance(root, ast.Attribute):
root = root.value
if not isinstance(root, ast.Name):
continue
# Every dialect's Kernel (and ClusterKernel) is marked as a launch
# factory; identity against one implementation would miss the
# others, and aliases such as `K = T.Kernel`.
from tilelang.language.kernel import is_kernel_launch_factory
if eval_res is Kernel or eval_res is ClusterKernel:
if is_kernel_launch_factory(self._try_eval(cexpr.func)):
is_kernel_ctx = True
break
node = self.generic_visit(node)
+7 -1
View File
@@ -388,7 +388,13 @@ class Builder(BaseBuilder):
self.current_file,
self.current_line,
)
yield self.enter_frame(frame)
try:
yield self.enter_frame(frame)
except BaseException as exc:
# Unwind Python frame state without finalizing incomplete TIR.
while len(self.frames) > pop_idx:
self.frames.pop().__exit__(type(exc), exc, exc.__traceback__)
raise
if self._spans_enabled:
# Flush leaf stmts of the frame being exited before its __exit__
# moves them into the produced node.
+220 -83
View File
@@ -3,6 +3,7 @@
from __future__ import annotations
from collections import deque
import os
from typing import Any
from tvm import tirx
from tvm.tirx import Var
from tvm.tirx.script.builder import evaluate as T_evaluate
@@ -97,21 +98,25 @@ def _normalize_bindings(bindings: list[Var]) -> Var | list[Var]:
def _normalize_threads(
threads: int | list[int] | tuple | None,
) -> list[int]:
) -> list[int] | None:
"""Normalize a thread-block specification into a 3-D extent list.
Args:
threads: A thread count, a per-dimension extent list/tuple, or None for the default.
threads: A thread count, a per-dimension extent list/tuple, or None to
leave the choice to the backend.
Returns:
The extents as ``[x, y, z]``, padding missing dimensions with 1.
The extents as ``[x, y, z]``, padding missing dimensions with 1, or
None when no thread count was requested. The frontend does not pick a
default: the thread count is a SIMT launch hint whose default (if any)
belongs to the backend that materializes the launch.
Raises:
ValueError: If ``threads`` has an unsupported type, or any concrete extent
is not positive.
"""
if threads is None:
threads = 128 # default thread number
return None
if isinstance(threads, int):
normalized = [threads, 1, 1]
@@ -151,13 +156,18 @@ class KernelLaunchFrame(TIRFrame):
"""
KernelLaunchFrame is a custom TIRFrame that manages block/thread indices
and handles the entry and exit of the kernel launch scope.
Grid (program index) vars are bound by the frame itself. Thread vars are
placeholders: they have an identity so the body can reference them, but
their extent is only known once a backend materializes the launch. Thread
extents are therefore available at trace time only when ``threads=`` was
passed to :func:`Kernel`.
"""
def __enter__(self) -> Var | list[Var]:
"""
Enters the KernelLaunchFrame scope and pushes this frame onto the stack.
Returns one Var if we detect exactly 5 frames (meaning there is a single
block dimension), or a list of Vars otherwise.
Returns one Var for a single grid dimension, or a list of Vars otherwise.
"""
super().__enter__()
_get_current_stack().push(self)
@@ -165,9 +175,7 @@ class KernelLaunchFrame(TIRFrame):
last_block_frame = self.frames[-1]
assert isinstance(last_block_frame, SBlockFrame), f"Last frame must be a block frame, got {last_block_frame}"
# Return a list of grid loop vars (excluding the last 4 frames:
# threadIdx.x, threadIdx.y, threadIdx.z and the block frame with attributes).
return _normalize_bindings([frame.vars[0] for frame in self.frames[0:-4]])
return _normalize_bindings(list(self.grid_vars))
def __exit__(self, ptype, value, trace):
"""
@@ -192,9 +200,11 @@ class KernelLaunchFrame(TIRFrame):
"""
Returns the block extent for the given dimension.
dim=0 corresponds to blockIdx.x, dim=1 to blockIdx.y, and dim=2 to blockIdx.z.
Grid axes that were not launched have extent 1.
"""
iter_var = self.frames[dim].doms[0]
return int(iter_var.extent)
if dim >= len(self.grid_extents):
return 1
return int(self.grid_extents[dim])
def get_block_extents(self) -> list[int]:
"""
@@ -206,9 +216,17 @@ class KernelLaunchFrame(TIRFrame):
"""
Returns the thread extent for the given dimension.
dim=0 corresponds to threadIdx.x, dim=1 to threadIdx.y, and dim=2 to threadIdx.z.
Raises:
ValueError: If the kernel was launched without ``threads=``. The
extent is then chosen by the backend and is not known at trace time.
"""
iter_var = self.frames[-4 + dim].doms[0]
return int(iter_var.extent)
if self.thread_extents is None:
raise ValueError(
"The thread extent is not known at trace time: T.Kernel(...) was called without "
"threads=. Pass threads= explicitly when the kernel body needs the thread-block size."
)
return int(self.thread_extents[dim])
def get_thread_extents(self) -> list[int]:
"""
@@ -221,14 +239,14 @@ class KernelLaunchFrame(TIRFrame):
Returns the thread binding for the given dimension.
dim=0 corresponds to threadIdx.x, dim=1 to threadIdx.y, and dim=2 to threadIdx.z.
"""
return self.frames[-4 + dim].vars[0]
return self.thread_vars[dim]
def get_thread_bindings(self) -> list[Var]:
"""
Returns the thread binding for the given dimension.
dim=0 corresponds to threadIdx.x, dim=1 to threadIdx.y, and dim=2 to threadIdx.z.
"""
return [frame.vars[0] for frame in self.frames[-4:-1]]
return list(self.thread_vars)
def get_num_threads(self) -> int:
"""
@@ -244,27 +262,94 @@ class KernelLaunchFrame(TIRFrame):
Returns the block binding for the given dimension.
dim=0 corresponds to blockIdx.x, dim=1 to blockIdx.y, and dim=2 to blockIdx.z.
"""
return self.frames[dim].vars[0]
return self.grid_vars[dim]
def get_block_bindings(self) -> list[Var]:
"""
Returns all three block bindings.
"""
return [frame.vars[0] for frame in self.frames[0:-4]]
return list(self.grid_vars)
def get_launch_annotation(self, key: str, default=None):
"""
Returns the launch annotation ``key`` recorded by T.Kernel (e.g. ``cluster_dims``),
or ``default`` when it was not given.
"""
annotations = self.frames[-1].annotations
if annotations is None or key not in annotations:
return default
return annotations[key]
def get_cluster_dims(self) -> list[int]:
"""
Returns the cluster dimensions as ``[x, y, z]``. A launch without
``cluster_dims`` has clusters of a single program, i.e. ``[1, 1, 1]``.
"""
dims = self.get_launch_annotation("cluster_dims")
if dims is None:
return [1, 1, 1]
dims = [int(d) for d in dims]
return dims + [1] * (3 - len(dims))
def get_cluster_size(self) -> int:
"""
Returns the number of programs per cluster (product of the cluster dimensions).
"""
size = 1
for dim in self.get_cluster_dims():
size *= dim
return size
def get_cluster_id(self, dim: int = 0) -> Var | tirx.PrimExpr:
"""
Returns the index of the cluster the current program belongs to along
``dim``, in program-space arithmetic: ``block_id // cluster_dims[dim]``.
A cluster is a ``cluster_dims``-shaped tile of the grid, so this is the
same on every target (clusterIdx on CUDA, a group of consecutive
programs elsewhere) and stays consistent with threadblock swizzling,
which permutes the grid at cluster granularity.
"""
if dim >= len(self.grid_vars):
return tirx.IntImm("int32", 0)
block = self.grid_vars[dim]
cluster_dim = self.get_cluster_dims()[dim]
if cluster_dim == 1:
return block
return tirx.floordiv(block, tirx.IntImm(block.dtype, cluster_dim))
def get_cluster_ids(self) -> list[Var | tirx.PrimExpr]:
"""
Returns the cluster index along every launched grid axis.
"""
return [self.get_cluster_id(dim) for dim in range(len(self.grid_vars))]
def get_cluster_extent(self, dim: int = 0) -> int:
"""
Returns the number of clusters along ``dim``: ``ceil(grid_extent / cluster_dims[dim])``.
"""
cluster_dim = self.get_cluster_dims()[dim]
return -(-self.get_block_extent(dim) // cluster_dim)
def get_cluster_extents(self) -> list[int]:
"""
Returns the number of clusters along all three dimensions.
"""
return [self.get_cluster_extent(dim) for dim in range(3)]
@property
def blocks(self) -> list[Var]:
"""
Returns the block indices from the topmost frame.
"""
return [frame.vars[0] for frame in self.frames[0:-4]]
return list(self.grid_vars)
@property
def threads(self) -> list[Var]:
"""
Returns the thread indices from the topmost frame.
"""
return [frame.vars[0] for frame in self.frames[-4:-1]]
return list(self.thread_vars)
@property
def num_threads(self) -> int:
@@ -274,53 +359,51 @@ class KernelLaunchFrame(TIRFrame):
return self.get_num_threads()
def Kernel(
*blocks: int | tirx.PrimExpr,
threads: int | list[int] | tuple | None = None,
# ---------------------------------------------------------------------------
# Launch annotations
#
# T.Kernel(*grid) is the launch every target shares. Everything a backend may
# additionally need (thread count, clusters, ...) is a *launch annotation*: the
# frontend records it verbatim and the backend interprets it once the target
# is known (MaterializeKernelLaunch). Which annotations exist is declared per
# language dialect as the explicit keyword parameters of its own `Kernel`
# (see tilelang/<backend>/language/kernel.py), so `tilelang.cuda.language.Kernel`
# shows, autocompletes and accepts exactly what CUDA understands. Every
# dialect's `Kernel` funnels into `launch_kernel` below.
# ---------------------------------------------------------------------------
_KERNEL_LAUNCH_FACTORY_ATTR = "__tilelang_kernel_launch__"
def kernel_launch_factory(func):
"""Mark ``func`` as a launch factory: a callable used as ``with func(...)``
to open a kernel launch. Every dialect's ``Kernel`` (and ``ClusterKernel``)
carries this mark so the eager JIT rewriter can find the launch regardless
of which dialect or alias the user went through."""
setattr(func, _KERNEL_LAUNCH_FACTORY_ATTR, True)
return func
def is_kernel_launch_factory(obj) -> bool:
return getattr(obj, _KERNEL_LAUNCH_FACTORY_ATTR, False) is True
def launch_kernel(
blocks: tuple[int | tirx.PrimExpr, ...],
*,
threads: int | list[int] | tuple[int, ...] | None = None,
prelude: str | None = None,
):
"""Tools to quickly construct a kernel launch frame.
cluster_dims: int | tuple[int, int, int] | list[int] | None = None,
**annotations: Any,
) -> KernelLaunchFrame:
"""Shared implementation behind every dialect's ``T.Kernel``.
The launch nest is emitted in a target-neutral form (thread_binding
For loops); each backend pipeline materializes it via
MaterializeKernelLaunch. Backends without SIMT (e.g. CPU) simply
ignore the thread extents at compile time, so the same kernel can be
compiled for any target.
Parameters
----------
blocks : int
A list of extent, can be 1-3 dimension, representing gridDim.(x|y|z)
threads : int
A integer representing blockDim.x
Or a list of integers representing blockDim.(x|y|z)
if the value is -1, we skip the threadIdx.x binding.
prelude : str
The import c code of the kernel,
will be injected before the generated kernel code.
Returns
-------
res : Tuple[frame.LaunchThreadFrame]
The result LaunchThreadFrame.
Examples
--------
Create a 1-D CUDA kernel launch and unpack the single block index:
.. code-block:: python
with T.Kernel(T.ceildiv(N, 128), threads=128) as bx:
# bx is the blockIdx.x binding (also iterable as (bx,))
...
Launch a 2-D grid while requesting two thread dimensions:
.. code-block:: python
with T.Kernel(grid_x, grid_y, threads=(64, 2)) as (bx, by):
tx, ty = T.get_thread_bindings()
...
The well-known launch annotations are normalized here; any other keyword
a dialect forwards is recorded verbatim on the launch block for that
backend's pipeline to consume. Dialects, not this function, decide which
keywords exist: they only forward what their own ``Kernel`` signature
declares.
"""
# In eager mode, we construct AST directly without prim_func,
# so there must be a Builder available. If not, this function
@@ -332,14 +415,46 @@ def Kernel(
raise JITNoBuilderError("T.Kernel() can only be used inside @tilelang.jit or @T.prim_func context. No Builder is available.")
attrs: dict = {}
threads = _normalize_threads(threads)
if prelude is not None:
attrs["pragma_import_c"] = prelude
cluster_dims = _normalize_cluster_dims(cluster_dims)
if cluster_dims is not None:
attrs["cluster_dims"] = cluster_dims
for key, value in annotations.items():
if value is not None:
attrs[key] = value
return _ffi_api.KernelLaunch(blocks, threads, attrs)
return _ffi_api.KernelLaunch(blocks, _normalize_threads(threads), attrs)
@kernel_launch_factory
def Kernel(*blocks: int | tirx.PrimExpr) -> KernelLaunchFrame:
"""Construct a kernel launch frame: a grid of tile programs.
This is the target-neutral launch: the part every backend shares. Backend
dialects offer their own ``T.Kernel`` with the launch annotations that
backend understands, e.g. ``tilelang.cuda.language.Kernel(..., threads=128)``;
``tilelang.language`` is the CUDA dialect.
Parameters
----------
*blocks : int | PrimExpr
Extent of the grid along each axis (1-3 dimensions). The launch yields
one program index per axis (``blockIdx`` on CUDA, the outer loop on
CPU, the core index on an NPU).
Examples
--------
.. code-block:: python
with T.Kernel(T.ceildiv(N, 128)) as bx:
# bx is the program index along x; also iterable as (bx,)
...
"""
return launch_kernel(blocks)
@kernel_launch_factory
def ClusterKernel(
*blocks: int | tirx.PrimExpr,
cluster_dims: int | tuple[int, int, int] | list[int],
@@ -375,22 +490,7 @@ def ClusterKernel(
with T.ClusterKernel(grid_x, grid_y, cluster_dims=2, threads=128) as (bx, by):
...
"""
from tilelang.language.eager.builder import Builder
if Builder.current() is None:
raise JITNoBuilderError("T.ClusterKernel() can only be used inside @tilelang.jit or @T.prim_func context. No Builder is available.")
attrs: dict = {}
threads = _normalize_threads(threads)
if prelude is not None:
attrs["pragma_import_c"] = prelude
cluster_dims = _normalize_cluster_dims(cluster_dims)
if cluster_dims is not None:
attrs["cluster_dims"] = cluster_dims
return _ffi_api.KernelLaunch(blocks, threads, attrs)
return launch_kernel(blocks, threads=threads, prelude=prelude, cluster_dims=cluster_dims)
# For CUDA source kernels, we need to load the source code from a file or string.
@@ -530,3 +630,40 @@ def get_block_extents() -> list[int]:
"""Returns all three block extents."""
assert KernelLaunchFrame.Current() is not None, "KernelLaunchFrame is not initialized"
return KernelLaunchFrame.Current().get_block_extents()
def get_cluster_dims() -> list[int]:
"""Returns the cluster dimensions ``[x, y, z]`` of the current launch (``[1, 1, 1]`` without clusters)."""
assert KernelLaunchFrame.Current() is not None, "KernelLaunchFrame is not initialized"
return KernelLaunchFrame.Current().get_cluster_dims()
def get_cluster_size() -> int:
"""Returns the number of programs per cluster of the current launch."""
assert KernelLaunchFrame.Current() is not None, "KernelLaunchFrame is not initialized"
return KernelLaunchFrame.Current().get_cluster_size()
def get_cluster_id(dim: int = 0) -> Var | tirx.PrimExpr:
"""Returns the cluster index of the current program along ``dim``
(``block_id // cluster_dims[dim]``). See :meth:`KernelLaunchFrame.get_cluster_id`."""
assert KernelLaunchFrame.Current() is not None, "KernelLaunchFrame is not initialized"
return KernelLaunchFrame.Current().get_cluster_id(dim)
def get_cluster_ids() -> list[Var | tirx.PrimExpr]:
"""Returns the cluster index along every launched grid axis."""
assert KernelLaunchFrame.Current() is not None, "KernelLaunchFrame is not initialized"
return KernelLaunchFrame.Current().get_cluster_ids()
def get_cluster_extent(dim: int = 0) -> int:
"""Returns the number of clusters along ``dim``."""
assert KernelLaunchFrame.Current() is not None, "KernelLaunchFrame is not initialized"
return KernelLaunchFrame.Current().get_cluster_extent(dim)
def get_cluster_extents() -> list[int]:
"""Returns the number of clusters along all three dimensions."""
assert KernelLaunchFrame.Current() is not None, "KernelLaunchFrame is not initialized"
return KernelLaunchFrame.Current().get_cluster_extents()
+4 -1
View File
@@ -11,6 +11,8 @@ from tilelang.language.builtin import ( # noqa: F401
cooperative_tensor_store,
)
from .kernel import * # noqa: F401,F403
from .kernel import __all__ as _KERNEL_ALL
from .tir import * # noqa: F401,F403
from .tir import __all__ as _TIR_ALL
@@ -19,6 +21,7 @@ __all__ = tuple(
dict.fromkeys(
(
*_COMMON_ALL,
*_KERNEL_ALL,
*_TIR_ALL,
"cooperative_tensor_fill",
"cooperative_tensor_load",
@@ -28,4 +31,4 @@ __all__ = tuple(
)
)
del _COMMON_ALL, _TIR_ALL
del _COMMON_ALL, _KERNEL_ALL, _TIR_ALL
+45
View File
@@ -0,0 +1,45 @@
"""Metal dialect of ``T.Kernel``: the common launch plus Metal launch annotations."""
from __future__ import annotations
from tvm import tirx
from tilelang.language.kernel import KernelLaunchFrame, kernel_launch_factory, launch_kernel
__all__ = ["Kernel"]
@kernel_launch_factory
def Kernel(
*blocks: int | tirx.PrimExpr,
threads: int | list[int] | tuple[int, ...] | None = None,
prelude: str | None = None,
) -> KernelLaunchFrame:
"""Construct a kernel launch frame for Metal: a grid of threadgroups.
Code inside the launch operates at the threadgroup level: ``T.Parallel``,
``T.copy`` and friends are mapped onto threads by the compiler.
``T.get_thread_binding()`` exposes the thread index for thread-level code.
The keyword arguments are recorded at trace time and materialized by the
Metal pipeline once the target is known.
Parameters
----------
*blocks : int | PrimExpr
Grid extent along each axis (1-3 dimensions). The launch yields one
threadgroup index per axis.
threads : int | list[int] | tuple[int, ...], optional
Threads per threadgroup: a count or up to three per-dimension extents.
Defaults to 128 when omitted.
prelude : str, optional
Source injected before the generated kernel, e.g. ``#include`` lines or
helper functions.
Examples
--------
.. code-block:: python
with T.Kernel(T.ceildiv(N, 128), threads=128) as bx:
...
"""
return launch_kernel(blocks, threads=threads, prelude=prelude)
+3 -1
View File
@@ -17,7 +17,9 @@ from tilelang.metal.transform import MetalFragmentToSimdgroup
def MetalPassPipelineBody(mod: IRModule, target: Target) -> IRModule:
mod = tirx.transform.BindTarget(target)(mod)
mod = tilelang.transform.MaterializeKernelLaunch()(mod)
mod = tilelang.transform.MaterializeKernelLaunch(
lower_thread_binding=True, default_threads=128, unsupported_annotations=["cluster_dims"]
)(mod)
pass_ctx = tilelang.transform.get_pass_context()
if should_force_let_inline():
+4 -2
View File
@@ -7,8 +7,10 @@ from tilelang.language.common import __all__ as _COMMON_ALL
from .intrinsics import * # noqa: F401,F403
from .intrinsics import __all__ as _ROCM_ALL
from .kernel import * # noqa: F401,F403
from .kernel import __all__ as _KERNEL_ALL
__tilelang_dialect__ = "rocm"
__all__ = tuple(dict.fromkeys((*_COMMON_ALL, *_ROCM_ALL)))
__all__ = tuple(dict.fromkeys((*_COMMON_ALL, *_ROCM_ALL, *_KERNEL_ALL)))
del _COMMON_ALL, _ROCM_ALL
del _COMMON_ALL, _ROCM_ALL, _KERNEL_ALL
+45
View File
@@ -0,0 +1,45 @@
"""ROCm dialect of ``T.Kernel``: the common launch plus ROCm launch annotations."""
from __future__ import annotations
from tvm import tirx
from tilelang.language.kernel import KernelLaunchFrame, kernel_launch_factory, launch_kernel
__all__ = ["Kernel"]
@kernel_launch_factory
def Kernel(
*blocks: int | tirx.PrimExpr,
threads: int | list[int] | tuple[int, ...] | None = None,
prelude: str | None = None,
) -> KernelLaunchFrame:
"""Construct a kernel launch frame for ROCm: a grid of workgroups.
Code inside the launch operates at the workgroup level: ``T.Parallel``,
``T.copy`` and friends are mapped onto threads by the compiler.
``T.get_thread_binding()`` exposes the thread index for thread-level code.
The keyword arguments are recorded at trace time and materialized by the
ROCm pipeline once the target is known.
Parameters
----------
*blocks : int | PrimExpr
Grid extent along each axis (1-3 dimensions). The launch yields one
workgroup index per axis.
threads : int | list[int] | tuple[int, ...], optional
Threads per workgroup: a count or up to three per-dimension extents.
Defaults to 128 when omitted.
prelude : str, optional
Source injected before the generated kernel, e.g. ``#include`` lines or
helper functions.
Examples
--------
.. code-block:: python
with T.Kernel(T.ceildiv(N, 128), threads=128) as bx:
...
"""
return launch_kernel(blocks, threads=threads, prelude=prelude)
+3 -1
View File
@@ -16,7 +16,9 @@ from tilelang.backend.pass_pipeline.pipeline_utils import (
def ROCMPassPipelineBody(mod: IRModule, target: Target) -> IRModule:
mod = tirx.transform.BindTarget(target)(mod)
mod = tilelang.transform.MaterializeKernelLaunch()(mod)
mod = tilelang.transform.MaterializeKernelLaunch(
lower_thread_binding=True, default_threads=128, unsupported_annotations=["cluster_dims"]
)(mod)
pass_ctx = tilelang.transform.get_pass_context()
if should_force_let_inline():
+35 -10
View File
@@ -198,26 +198,51 @@ def MakePackedAPI():
return _ffi_api.MakePackedAPI() # type: ignore
def MaterializeKernelLaunch(lower_thread_binding: bool = True):
"""Materialize the target-neutral kernel launch nest (thread_binding
For loops emitted by T.Kernel) into a backend-specific form. Each
backend pipeline decides the mode for itself:
DEFAULT_SIMT_THREADS = 128
def MaterializeKernelLaunch(
lower_thread_binding: bool = True,
default_threads: int | list[int] | tuple | None = DEFAULT_SIMT_THREADS,
unsupported_annotations: list[str] | tuple[str, ...] | None = None,
):
"""Materialize the target-neutral kernel launch nest emitted by T.Kernel
into a backend-specific form. Each backend pipeline decides the mode for
itself; this is where the target-dependent parts of a launch (whether
threads exist and how many run by default) are decided.
Parameters
----------
lower_thread_binding : bool
If True (SIMT backends, e.g. CUDA/ROCm/Metal), lower the
blockIdx.*/threadIdx.* loops into thread_extent AttrStmts.
If False (backends without SIMT, e.g. CPU), lower blockIdx.*
loops into plain serial For loops and ignore threadIdx.* loops
(their extents are dropped; the loop vars are pinned to 0).
If True (SIMT backends, e.g. CUDA/ROCm/Metal), lower the blockIdx.*
grid loops into thread_extent AttrStmts and bind the thread
placeholders as threadIdx.* thread_extent scopes.
If False (backends without SIMT, e.g. CPU), lower blockIdx.* loops
into plain serial For loops and drop the thread placeholders. A body
that references a thread index is rejected on such targets.
default_threads : int | list[int] | tuple | None
Thread-block extents used by SIMT backends when T.Kernel was called
without ``threads=``. Ignored when ``lower_thread_binding`` is False.
None means the backend has no default and ``threads=`` is required.
unsupported_annotations : list[str] | None
Launch annotations (keys on the ``tilelang_root`` block, e.g.
``cluster_dims``) that have no meaning on this backend. A launch
carrying one is rejected here instead of being silently ignored by
later passes.
Returns
-------
fpass : tvm.transform.Pass
The result pass
"""
return _ffi_api.MaterializeKernelLaunch(lower_thread_binding) # type: ignore
if default_threads is not None:
if isinstance(default_threads, int):
default_threads = [default_threads, 1, 1]
else:
default_threads = list(default_threads) + [1] * (3 - len(default_threads))
if unsupported_annotations is not None:
unsupported_annotations = list(unsupported_annotations)
return _ffi_api.MaterializeKernelLaunch(lower_thread_binding, default_threads, unsupported_annotations) # type: ignore
def AnnotateDeviceRegions():
+6 -3
View File
@@ -5,7 +5,10 @@ from __future__ import annotations
from tilelang.language.common import * # noqa: F401,F403
from tilelang.language.common import __all__ as _COMMON_ALL
__tilelang_dialect__ = "webgpu"
__all__ = tuple(_COMMON_ALL)
from .kernel import * # noqa: F401,F403
from .kernel import __all__ as _KERNEL_ALL
del _COMMON_ALL
__tilelang_dialect__ = "webgpu"
__all__ = tuple(dict.fromkeys((*_COMMON_ALL, *_KERNEL_ALL)))
del _COMMON_ALL, _KERNEL_ALL
+41
View File
@@ -0,0 +1,41 @@
"""WebGPU dialect of ``T.Kernel``: the common launch plus WebGPU launch annotations."""
from __future__ import annotations
from tvm import tirx
from tilelang.language.kernel import KernelLaunchFrame, kernel_launch_factory, launch_kernel
__all__ = ["Kernel"]
@kernel_launch_factory
def Kernel(
*blocks: int | tirx.PrimExpr,
threads: int | list[int] | tuple[int, ...] | None = None,
) -> KernelLaunchFrame:
"""Construct a kernel launch frame for WebGPU: a grid of workgroups.
Code inside the launch operates at the workgroup level: ``T.Parallel``,
``T.copy`` and friends are mapped onto invocations by the compiler.
``T.get_thread_binding()`` exposes the invocation index for thread-level
code. ``threads`` is recorded at trace time and materialized by the WebGPU
pipeline once the target is known.
Parameters
----------
*blocks : int | PrimExpr
Grid extent along each axis (1-3 dimensions). The launch yields one
workgroup index per axis.
threads : int | list[int] | tuple[int, ...], optional
Invocations per workgroup: a count or up to three per-dimension
extents. Defaults to 128 when omitted.
Examples
--------
.. code-block:: python
with T.Kernel(T.ceildiv(N, 128), threads=128) as bx:
...
"""
return launch_kernel(blocks, threads=threads)
+3 -1
View File
@@ -18,7 +18,9 @@ def WebGPUPassPipelineBody(mod: IRModule, target: Target) -> IRModule:
mod = tirx.transform.BindTarget(target)(mod)
# WebGPU is a SIMT backend: lower the launch nest to thread_extent
# bindings for codegen.
mod = tilelang.transform.MaterializeKernelLaunch()(mod)
mod = tilelang.transform.MaterializeKernelLaunch(
lower_thread_binding=True, default_threads=128, unsupported_annotations=["cluster_dims"]
)(mod)
pass_ctx = tilelang.transform.get_pass_context()
if should_force_let_inline():