[Refactor][Language] Move backend-specific op hints into their owning dialects (#3203)

* [Refactor][Language] Move backend-specific op hints into their owning dialects

Common tile ops advertised knobs that only one backend consumes: T.copy took
disable_tma/eviction_policy/prefer_instruction (read only by the CUDA copy
analysis and codegen), T.im2col took eviction_policy (CUDA TMA im2col),
T.gemm took mbar (Blackwell TCGEN5MMA) and k_pack (ROCm MFMA/WMMA),
T.atomic_add took use_tma (sm90+ cp.reduce), T.Parallel took prefer_async and
T.unroll took unroll_factor (both honored by CUDA codegen only). On every
other target these silently did nothing, and every dialect autocompleted
another backend's vocabulary.

Following the T.Kernel dialect pattern (#3186), each dialect now declares the
knobs its backend understands as explicit keyword parameters:

- common keeps the target-neutral core (src/dst, coalesced_width,
  loop_layout, transpose/policy/clear_accum, explicit, ...) plus the
  `annotations` dict as the untyped escape hatch;
- the CUDA dialect (and therefore the default `tilelang.language` facade)
  shadows copy, im2col, gemm, atomic_add, Parallel and unroll/Unroll with the
  CUDA hints;
- the ROCm dialect shadows gemm with k_pack.

Unlike launch annotations, these are performance hints and stay droppable: a
kernel written with the CUDA dialect still compiles for targets that have no
use for the hints (covered by a cpu-compile test); enforcement lives only in
the dialect signatures. The gemm/im2col call protocol is unchanged (the
hint slots stay, fed by the dialect wrappers).

AMD examples and tests that pass k_pack now import the ROCm dialect
(`import tilelang.rocm.language as T`), which is the intended direction for
backend-specific code; the k_pack validation regression test (#3035) moves
with them.

* [Refactor][Language] Carry gemm k_pack as an annotation instead of a call slot

k_pack is a backend lowering knob like use_2cta or sf_layout, so transport it
the same way: as an annotation on the tile-op call, read by the GemmNode
constructor (validated 1/2, default 1). This removes the dummy `1` every
_gemm_impl caller had to thread through the positional protocol; the slots
after it (wg_wait, mbar, C_coords, SFA/SFB, k_start) shift down by one.
GemmNode::kPack_ and its reflection accessor are unchanged, so the ROCm
lowering keeps reading the same field.

* [Refactor][Language] Carry gemm wg_wait as an annotation and drop it from common gemm_sp

wg_wait is a Hopper warpgroup knob: it belongs to the CUDA dialect, not the
common surface or the positional call protocol. Transport it as a tile-op
annotation like k_pack:

- dense gemm: the wg_wait slot is gone (mbar, C_coords, SFA/SFB, k_start
  shift down by one); wgmma_gemm records wg_wait=-1 as an annotation,
  tcgen05_gemm_blockscaled forwards its keyword the same way.
- gemm_sp: the common public signature loses wg_wait (and the dead k_pack
  parameter with its never-consumed protocol slot); the CUDA dialect shadows
  gemm_sp with the wg_wait keyword; wgmma_gemm_sp records wg_wait=-1.

GemmNode::wgWait_ / GemmSPNode::wg_wait and their reflection accessors are
unchanged, so ws_analysis, the auto-scheduler and the WGMMA emitters keep
reading the same fields. Overlaps with the gemm_sp k_pack hunk of #3202;
whichever lands second rebases trivially.

* [Refactor][Language] Finish the tile-op call-protocol cleanup

Apply the slots-are-operands / knobs-are-annotations principle to the rest of
the tile-op protocols:

- im2col: the eviction_policy slot (the last backend knob riding a positional
  slot; read only by the CUDA TMA im2col lowering) moves to the annotations,
  like copy's. The common builder no longer emits a constant 0.
- gemm: the stride_a/stride_b/offset_a/offset_b slots are gone. They were
  parsed into fields nobody read (every lowering consumes the operand
  BufferRegions) and self-documented as deprecated. mbar, C_coords, SFA/SFB
  and k_start shift down accordingly; fields, reflection and the unread
  Python accessors are removed.
- gemm_sp: same for its stride/offset slots, plus the vestigial kPack field
  that nothing has set since the k_pack slot was dropped.
- Stale registration metadata corrected: the copy family takes 2 positional
  inputs (not 5), im2col 8 (not 9), reducer_init/finalize_reducer are
  variadic; constructor doc comments now describe the actual layouts.

This changes the serialized call protocol for out-of-tree builders of
tl.tileop.gemm/gemm_sp/im2col calls, in the same release as the k_pack and
wg_wait transport moves, so external code adapts once.

* [Refactor][Language] Move CUDA-only builtins and math intrinsics into the CUDA dialect

Second batch of dialect ownership, this time at function level. Definitions
stay where they are; only the export surface moves, so the default
`tilelang.language` facade (the CUDA dialect) is unchanged for users.

- reduce_max/min/absmax lose the CUDA-only `nan_propagate` keyword on the
  common surface; the CUDA dialect shadows them with it (transport stays the
  `nan_propagate` annotation, which still hard-errors on targets without
  __hmax_nan/__hmin_nan).
- get_lane_idx / get_warp_idx / get_warp_idx_sync (op registered only in the
  CUDA registry; no ROCm codegen handler exists), no_set_max_nreg (pairing it
  with the already-dialect-owned set_max_nreg), and the mbarrier family
  (barrier_arrive/barrier_wait/mbarrier_*: the HIP codegen would emit
  tl::mbarrier_* templates that do not exist under tl_templates/hip) are now
  exported by the CUDA dialect only.
- math_intrinsics split: the packed-x2 family (add2/sub2/mul2/fma2/max2/min2/
  abs2, lowered on CUDA and ROCm) stays common; the PTX fast-math (__exp,
  __log, ..., fast_rcp) and rounding-mode ieee_* families are CUDA-dialect
  exports.
- sync_global is deleted: the CUDA codegen has FATALed on it ("Global storage
  sync is no longer supported") and it carried a stray debug print; the
  unexported duplicate of loop_break in builtin.py goes with it.
- The CUDA and HIP codegens now reject unknown threadblock-swizzle patterns
  up front, so T.use_swizzle(order="mlx") (Metal-only) fails with a clear
  message instead of a missing-symbol error from nvcc/hipcc.
This commit is contained in:
Lei Wang
2026-09-12 01:52:03 +08:00
committed by GitHub
parent c629fe5e61
commit 66c003c3e7
43 changed files with 1067 additions and 426 deletions
+8 -3
View File
@@ -13,9 +13,13 @@ level, how they map to hardware concepts, and how to use them correctly.
## Data Movement
Use `T.copy(src, dst, *, coalesced_width=None, disable_tma=False, eviction_policy=None, loop_layout=None)`
Use `T.copy(src, dst, *, coalesced_width=None, loop_layout=None)`
to move tiles between memory scopes. It accepts `tir.Buffer`, `BufferLoad`, or
`BufferRegion`; extents are inferred or broadcast when possible.
`BufferRegion`; extents are inferred or broadcast when possible. Backend
dialects add their lowering hints as extra keywords: the CUDA dialect (the
default `tilelang.language` facade) accepts `disable_tma`, `eviction_policy`
and `prefer_instruction`. Hints are recorded on the op and ignored by targets
that have no use for them.
```python
# Global → Shared tiles (extents inferred from dst)
@@ -219,7 +223,8 @@ Warp-match (CUDA sm_70+, not supported on HIP). `mask` defaults to `0xFFFFFFFF`.
> **Note on HIP:** `any_sync`/`all_sync` ignore the mask and call `__any`/`__all` directly. `ballot_sync`, `ballot`, and `activemask` call `__ballot` which returns `uint64` natively on 64-thread wavefronts — no truncation occurs. Shuffle intrinsics lower to `__shfl`/`__shfl_xor`/`__shfl_down`/`__shfl_up` (mask ignored). `syncthreads_count/and/or` have identical signatures on both platforms. `match_any_sync` and `match_all_sync` have no HIP equivalent and will fail to codegen on HIP.
Atomics
- `T.atomic_add(dst, value, memory_order=None, return_prev=False, use_tma=False)`.
- `T.atomic_add(dst, value, memory_order=None, return_prev=False)`; the CUDA
dialect additionally accepts `use_tma=True` (sm90+ TMA `cp.reduce`).
- `T.atomic_addx2(dst, value, return_prev=False)`; `T.atomic_addx4(...)`.
- `T.atomic_max(dst, value, memory_order=None, return_prev=False)`.
- `T.atomic_min(dst, value, memory_order=None, return_prev=False)`.
+1 -1
View File
@@ -2,7 +2,7 @@ import sys
import torch
import torch.nn.functional as F
import tilelang
import tilelang.language as T
import tilelang.rocm.language as T
from tilelang.tileop.base import GemmWarpPolicy
import itertools
import argparse
+1 -1
View File
@@ -2,7 +2,7 @@ import sys
import torch
import torch.nn.functional as F
import tilelang
import tilelang.language as T
import tilelang.rocm.language as T
from tilelang.tileop.base import GemmWarpPolicy
import itertools
import argparse
@@ -1,6 +1,6 @@
import torch
import tilelang
import tilelang.language as T
import tilelang.rocm.language as T
from tilelang.utils.tensor import torch_assert_close
from tilelang.language.fp8 import determine_fp8_type, determine_torch_fp8_type
import itertools
@@ -3,7 +3,7 @@ import itertools
import tilelang
import tilelang.testing
from tilelang import tvm as tvm
import tilelang.language as T
import tilelang.rocm.language as T
from tilelang.tileop.base import GemmWarpPolicy
from tilelang.layout import make_swizzled_layout
from tilelang.rocm.intrinsics.mfma_macro_generator import MatrixCorePreshuffleIntrinEmitter
+8
View File
@@ -5148,6 +5148,14 @@ void CodeGenTileLangCUDA::VisitStmt_(const AttrStmtNode *op) {
}
}
ICHECK(!func_name.empty() && panel_size > 0);
// Only the row/column rasterizations exist in the CUDA device templates;
// e.g. T.use_swizzle(order="mlx") is Metal-only and must fail here
// instead of surfacing as a missing-symbol error from the device
// compiler.
ICHECK(func_name == "rasterization2DRow" ||
func_name == "rasterization2DColumn")
<< "threadblock swizzle pattern `" << func_name
<< "` is not supported by the CUDA backend";
if (this->cluster_dims.has_value()) {
auto [cluster_grid_x_ext, cluster_grid_y_ext, cluster_grid_z_ext] =
this->cluster_dims.value();
+15 -11
View File
@@ -571,8 +571,8 @@ Stmt CopyNode::Lower(const LowerArgs &lower_args,
}
// Constructs an Im2ColOp node from call arguments.
// args: src, dst, nhw_step, c_step, kernel, stride, dilation, padding,
// eviction_policy
// args: src, dst, nhw_step, c_step, kernel, stride, dilation, padding.
// The CUDA-only eviction_policy hint rides in the annotations map.
Im2ColOp::Im2ColOp(Array<PrimExpr> args, Map<String, ObjectRef> annotations) {
ObjectPtr<Im2ColOpNode> node = make_object<Im2ColOpNode>();
auto src_access = NormalizeToAccessRegion(args[0], kAccessRead);
@@ -588,7 +588,11 @@ Im2ColOp::Im2ColOp(Array<PrimExpr> args, Map<String, ObjectRef> annotations) {
node->stride_ = args[5].as<IntImm>().value()->value;
node->dilation_ = args[6].as<IntImm>().value()->value;
node->padding_ = args[7].as<IntImm>().value()->value;
node->eviction_policy_ = args[8].as<IntImm>().value()->value;
if (auto val = annotations.Get("eviction_policy")) {
const auto *int_val = val->as<IntImmNode>();
ICHECK(int_val) << "eviction_policy annotation must be IntImmNode";
node->eviction_policy_ = int_val->value;
}
node->annotations_ = annotations;
data_ = std::move(node);
}
@@ -606,12 +610,12 @@ Stmt Im2ColOpNode::Lower(const LowerArgs &lower_args,
// Register the Copy operation with TVM's TIR system
// This makes the copy operation available for use in TVM programs
// - Takes 5 inputs: src_buffer, dst_buffer, and annotation-driven options.
// - Takes 2 inputs (src_buffer, dst_buffer); options ride in annotations.
// - Marked as opaque since it has side effects (memory writes)
TIR_REGISTER_TL_TILE_OP(Copy, copy)
.set_attr<OpBlockAnnotationHandlerFunc>(kTLOpBlockAnnotationHandler,
ApplyCopyBlockAnnotations)
.set_num_inputs(5)
.set_num_inputs(2)
.set_attr<TCallEffectKind>("TCallEffectKind",
Integer(CallEffectKind::kOpaque));
@@ -627,7 +631,7 @@ TVM_REGISTER_OP("tl.tileop.async_copy")
})
.set_attr<OpBlockAnnotationHandlerFunc>(kTLOpBlockAnnotationHandler,
ApplyCopyBlockAnnotations)
.set_num_inputs(5)
.set_num_inputs(2)
.set_attr<TCallEffectKind>("TCallEffectKind",
Integer(CallEffectKind::kOpaque));
@@ -645,7 +649,7 @@ TVM_REGISTER_OP("tl.tileop.tma_copy")
})
.set_attr<OpBlockAnnotationHandlerFunc>(kTLOpBlockAnnotationHandler,
ApplyCopyBlockAnnotations)
.set_num_inputs(5)
.set_num_inputs(2)
.set_attr<TCallEffectKind>("TCallEffectKind",
Integer(CallEffectKind::kOpaque));
@@ -658,11 +662,11 @@ LayoutMap Im2ColOpNode::InferLayout(const LayoutInferArgs &layout_args,
// Register the Im2Col operation with TVM's TIR system
// This operation performs im2col transformation for 2D convolutions using a
// target-specific lowering.
// - Takes 9 inputs: src_buffer, dst_buffer, nhw_step, c_step, kernel, stride,
// dilation, padding, eviction_policy
// - Takes 8 inputs: src_buffer, dst_buffer, nhw_step, c_step, kernel, stride,
// dilation, padding; the CUDA eviction_policy hint rides in annotations
// - Marked as opaque since it has side effects (memory writes)
TIR_REGISTER_TL_TILE_OP(Im2ColOp, im2col)
.set_num_inputs(9)
.set_num_inputs(8)
.set_attr<TCallEffectKind>("TCallEffectKind",
Integer(CallEffectKind::kOpaque));
@@ -675,7 +679,7 @@ TVM_REGISTER_OP("tl.tileop.c2d_im2col")
Map<String, ObjectRef> annotations) {
return Im2ColOp(args, annotations);
})
.set_num_inputs(9)
.set_num_inputs(8)
.set_attr<TCallEffectKind>("TCallEffectKind",
Integer(CallEffectKind::kOpaque));
+1 -1
View File
@@ -182,7 +182,7 @@ public:
int padding_; // Padding amount
int dilation_; // Dilation factor
int kernel_; // Kernel size
int eviction_policy_; // Cache eviction policy
int eviction_policy_ = 0; // Cache eviction policy (annotation)
PrimExpr nhw_step_; // Step size in NHW dimensions
PrimExpr c_step_; // Step size in channel dimension
Map<String, ObjectRef> annotations_; // Annotations from Call node
+28 -26
View File
@@ -67,17 +67,17 @@ void RegisterGemmImpl(GemmImpl impl) {
*
* Deserializes operator parameters from `args` and resolves buffer references,
* populating an internal GemmNode with buffers, transpose flags, M/N/K,
* warp policy, clear_accum, strides, offsets, optional kPack/wg_wait, and
* optional mbarrier.
* warp policy, clear_accum, an optional mbarrier operand and the C tile
* coordinates.
*
* @param args Positional serialized arguments produced by the TL frontend:
* expected layout is:
* [Aptr, Bptr, Cptr, trans_A (Bool), trans_B (Bool),
* M (Int), N (Int), K (Int), policy (Int), clear_accum (Bool),
* stride_A (Int), stride_B (Int), offset_A (PrimExpr),
* offset_B (PrimExpr),
* (optional) kPack (Int), (optional) internal wg_wait (Int),
* (optional) mbar (BufferLoad), cCoord_y (PrimExpr), cCoord_x (PrimExpr)]
* (optional) mbar (BufferLoad or const-0 placeholder),
* cCoord_y (PrimExpr), cCoord_x (PrimExpr),
* (optional, blockscaled) SFA, SFB regions, k_start (PrimExpr)]
* Backend lowering knobs (k_pack, wg_wait) ride in the annotations map.
*/
Gemm::Gemm(Array<PrimExpr> args, Map<String, ObjectRef> annotations) {
ObjectPtr<GemmNode> node = make_object<GemmNode>();
@@ -101,18 +101,20 @@ Gemm::Gemm(Array<PrimExpr> args, Map<String, ObjectRef> annotations) {
node->k_ = args[7].as<IntImm>().value()->value;
node->policy_ = GemmWarpPolicy(args[8].as<IntImm>().value()->value);
node->clearAccum_ = args[9].as<PrimExpr>().value();
node->strideA_ = args[10].as<IntImm>().value()->value;
node->strideB_ = args[11].as<IntImm>().value()->value;
node->offsetA_ = args[12].as<PrimExpr>().value();
node->offsetB_ = args[13].as<PrimExpr>().value();
if (args.size() > 14) {
node->kPack_ = args[14].as<IntImm>().value()->value;
if (node->kPack_ != 1 && node->kPack_ != 2) {
ICHECK(false) << "kPack must be 1 or 2";
}
// k_pack rides in the annotations (a ROCm MFMA/WMMA lowering knob set by
// the ROCm dialect), not in the positional call protocol.
if (auto val = annotations.Get("k_pack")) {
const auto *int_val = val->as<IntImmNode>();
ICHECK(int_val) << "k_pack annotation must be IntImmNode";
node->kPack_ = int_val->value;
ICHECK(node->kPack_ == 1 || node->kPack_ == 2) << "kPack must be 1 or 2";
}
if (args.size() > 15) {
node->wgWait_ = args[15].as<IntImm>().value()->value;
// wg_wait is a Hopper warpgroup knob set by the CUDA dialect; like k_pack
// it rides in the annotations rather than the positional call protocol.
if (auto val = annotations.Get("wg_wait")) {
const auto *int_val = val->as<IntImmNode>();
ICHECK(int_val) << "wg_wait annotation must be IntImmNode";
node->wgWait_ = int_val->value;
}
if (auto val = annotations.Get("is_wgmma")) {
const auto *int_val = val->as<IntImmNode>();
@@ -124,19 +126,19 @@ Gemm::Gemm(Array<PrimExpr> args, Map<String, ObjectRef> annotations) {
ICHECK(int_val) << "is_tcgen05 annotation must be IntImmNode";
node->isTcgen05_ = int_val->value != 0;
}
if (args.size() > 16 && args[16]->IsInstance<BufferLoadNode>()) {
node->mbar_ = Downcast<BufferLoad>(args[16]);
if (args.size() > 10 && args[10]->IsInstance<BufferLoadNode>()) {
node->mbar_ = Downcast<BufferLoad>(args[10]);
}
node->cCoords_ = Array<PrimExpr>(
{args[17].as<PrimExpr>().value(), args[18].as<PrimExpr>().value()});
if (args.size() > 19) {
node->sfaRegion_ = NormalizeToBufferRegion(args[19]);
{args[11].as<PrimExpr>().value(), args[12].as<PrimExpr>().value()});
if (args.size() > 13) {
node->sfaRegion_ = NormalizeToBufferRegion(args[13]);
}
if (args.size() > 20) {
node->sfbRegion_ = NormalizeToBufferRegion(args[20]);
if (args.size() > 14) {
node->sfbRegion_ = NormalizeToBufferRegion(args[14]);
}
if (args.size() > 21) {
node->sfKStart_ = args[21].as<PrimExpr>().value();
if (args.size() > 15) {
node->sfKStart_ = args[15].as<PrimExpr>().value();
}
node->annotations_ = annotations;
data_ = std::move(node);
-7
View File
@@ -107,9 +107,6 @@ public:
BufferRegion aRegion_, bRegion_, cRegion_;
bool transA_, transB_;
int m_, n_, k_;
int strideA_, strideB_;
// Offsets may be symbolic (e.g. a sliced operand B[:, j*64:...] in a loop).
PrimExpr offsetA_, offsetB_;
PrimExpr clearAccum_ = const_false();
tirx::BufferLoad mbar_; // mbar is optional, only used for TCGEN5MMA
Array<PrimExpr> cCoords_;
@@ -140,10 +137,6 @@ public:
.def_ro("m", &GemmNode::m_)
.def_ro("n", &GemmNode::n_)
.def_ro("k", &GemmNode::k_)
.def_ro("strideA", &GemmNode::strideA_)
.def_ro("strideB", &GemmNode::strideB_)
.def_ro("offsetA", &GemmNode::offsetA_)
.def_ro("offsetB", &GemmNode::offsetB_)
.def_ro("clearAccum", &GemmNode::clearAccum_)
.def_ro("mbar", &GemmNode::mbar_)
.def_ro("cCoords", &GemmNode::cCoords_)
+9 -16
View File
@@ -76,15 +76,14 @@ void RegisterGemmSPImpl(GemmSPImpl impl) {
*
* Deserializes operator parameters from `args` and resolves buffer references,
* populating an internal GemmSPNode with buffers, transpose flags, M/N/K,
* warp policy, clear_accum, strides, offsets, and optional kPack/wg_wait.
* warp policy and clear_accum.
*
* @param args Positional serialized arguments produced by the TL frontend:
* expected layout is:
* [Aptr, Eptr, Bptr, Cptr, trans_A (Bool), trans_E (Bool),
* trans_B (Bool), M (Int), N (Int), K (Int), policy (Int),
* clear_accum (Bool), stride_A (Int), stride_B (Int),
* offset_A (Int), offset_B (Int),
* (optional) kPack (Int), (optional) wg_wait (Int)]
* clear_accum (Bool)]
* Backend lowering knobs (wg_wait) ride in the annotations map.
*/
GemmSP::GemmSP(Array<PrimExpr> args, Map<String, ObjectRef> annotations) {
ObjectPtr<GemmSPNode> node = make_object<GemmSPNode>();
@@ -112,18 +111,12 @@ GemmSP::GemmSP(Array<PrimExpr> args, Map<String, ObjectRef> annotations) {
node->K = args[9].as<IntImm>().value()->value;
node->policy = GemmSPWarpPolicy(args[10].as<IntImm>().value()->value);
node->clear_accum = args[11].as<PrimExpr>().value();
node->stride_A = args[12].as<IntImm>().value()->value;
node->stride_B = args[13].as<IntImm>().value()->value;
node->offset_A = args[14].as<IntImm>().value()->value;
node->offset_B = args[15].as<IntImm>().value()->value;
if (args.size() > 16) {
node->kPack = args[16].as<IntImm>().value()->value;
if (node->kPack != 1 && node->kPack != 2) {
ICHECK(false) << "kPack must be 1 or 2";
}
}
if (args.size() > 17) {
node->wg_wait = args[17].as<IntImm>().value()->value;
// wg_wait is a Hopper warpgroup knob set by the CUDA dialect; it rides in
// the annotations rather than the positional call protocol.
if (auto val = annotations.Get("wg_wait")) {
const auto *int_val = val->as<IntImmNode>();
ICHECK(int_val) << "wg_wait annotation must be IntImmNode";
node->wg_wait = int_val->value;
}
if (auto val = annotations.Get("is_wgmma")) {
const auto *int_val = val->as<IntImmNode>();
-10
View File
@@ -84,12 +84,7 @@ public:
BufferRegion aRegion_, eRegion_, bRegion_, cRegion_;
bool trans_A, trans_B, trans_E;
int M, N, K;
int stride_A, stride_B;
int offset_A, offset_B;
PrimExpr clear_accum = const_false();
// k_pack please ref to bitblas/tl/mfma_macro_generator.py::k_pack
// only will be enabled under cdna mfma instructions
int kPack = 1;
int wg_wait = 0;
bool isWgmma_ = false;
bool isTcgen05_ = false;
@@ -114,12 +109,7 @@ public:
.def_ro("M", &GemmSPNode::M)
.def_ro("N", &GemmSPNode::N)
.def_ro("K", &GemmSPNode::K)
.def_ro("stride_A", &GemmSPNode::stride_A)
.def_ro("stride_B", &GemmSPNode::stride_B)
.def_ro("offset_A", &GemmSPNode::offset_A)
.def_ro("offset_B", &GemmSPNode::offset_B)
.def_ro("clear_accum", &GemmSPNode::clear_accum)
.def_ro("kPack", &GemmSPNode::kPack)
.def_ro("wg_wait", &GemmSPNode::wg_wait)
.def_ro("isWgmma", &GemmSPNode::isWgmma_)
.def_ro("isTcgen05", &GemmSPNode::isTcgen05_)
+5 -2
View File
@@ -406,7 +406,8 @@ TileOperator ReducerInitOpNode::Clone() const {
}
TIR_REGISTER_TL_TILE_OP(ReducerInitOp, reducer_init)
.set_num_inputs(1)
// reducer region plus an optional init value.
.set_num_inputs(-1)
.set_attr<TCallEffectKind>("TCallEffectKind",
Integer(CallEffectKind::kOpaque));
@@ -659,7 +660,9 @@ TileOperator FinalizeReducerOpNode::Clone() const {
}
TIR_REGISTER_TL_TILE_OP(FinalizeReducerOp, finalize_reducer)
.set_num_inputs(1)
// user form: reducer region; materialized form appends the combine-op
// enum and the flattened (reducing_threads, scale) plan pairs.
.set_num_inputs(-1)
.set_attr<TCallEffectKind>("TCallEffectKind",
Integer(CallEffectKind::kOpaque));
+7
View File
@@ -2101,6 +2101,13 @@ void CodeGenTileLangHIP::VisitStmt_(const AttrStmtNode *op) {
ICHECK(!func_name.empty() && panel_size > 0)
<< "threadblock_swizzle_pattern: failed to extract func_name and "
"panel_size";
// Only the row/column rasterizations exist in the HIP device templates;
// e.g. T.use_swizzle(order="mlx") is Metal-only and must fail here
// instead of surfacing as a missing-symbol error from hipcc.
ICHECK(func_name == "rasterization2DRow" ||
func_name == "rasterization2DColumn")
<< "threadblock swizzle pattern `" << func_name
<< "` is not supported by the ROCm backend";
this->stream << "const dim3 blockIdx = tl::" << func_name << "<"
<< panel_size << ">();\n";
this->VisitStmt(op->body);
@@ -2,7 +2,7 @@ import pytest
import torch
import tilelang.testing
from tilelang import tvm as tvm
import tilelang.language as T
import tilelang.rocm.language as T
from tilelang.rocm.intrinsics import make_mfma_swizzle_layout as make_swizzle_layout
from tilelang.rocm.intrinsics.mfma_macro_generator import (
MatrixCoreIntrinEmitter,
@@ -3,7 +3,7 @@ import torch
import tilelang
import tilelang.testing
from tilelang import tvm as tvm
import tilelang.language as T
import tilelang.rocm.language as T
from tilelang.rocm.intrinsics import make_mfma_swizzle_layout as make_swizzle_layout
from tilelang.rocm.intrinsics.mfma_macro_generator import MatrixCorePreshuffleIntrinEmitter
from tilelang.transform import simplify_prim_func
@@ -11,7 +11,7 @@ Two new behaviours introduced in commit dfa63b10:
import pytest
import tilelang as tl
import tilelang.language as T
import tilelang.rocm.language as T
import tilelang.testing
from tilelang.testing import _check_is_gfx950 as _is_gfx950
+1 -1
View File
@@ -1,7 +1,7 @@
import pytest
from tilelang import tvm as tvm
import tilelang as tl
import tilelang.language as T
import tilelang.rocm.language as T
import tilelang.testing
@@ -2,7 +2,7 @@
import pytest
import tilelang.language as T
import tilelang.rocm.language as T
import tilelang.testing
@@ -56,8 +56,15 @@ def test_default_language_is_static_cuda_facade():
def test_common_language_preserves_special_dsl_exports():
assert T_comm.__tilelang_dialect__ == "common"
assert "__log" in T_comm.__all__
assert hasattr(T_comm, "__log")
# Packed-x2 math is target-neutral and stays on the common surface; the
# CUDA-only fast-math family (dunder names like __log) moved to the CUDA
# dialect, whose export machinery must keep supporting them.
assert "add2" in T_comm.__all__
assert "__log" not in T_comm.__all__
from tilelang.cuda import language as cuda_language
assert "__log" in cuda_language.__all__
assert hasattr(cuda_language, "__log")
assert CUDA_ONLY_TIR_EXPORTS.isdisjoint(T_comm.__all__)
assert METAL_ONLY_TIR_EXPORTS.isdisjoint(T_comm.__all__)
assert ROCM_ONLY_TIR_EXPORTS.isdisjoint(T_comm.__all__)
@@ -72,7 +79,11 @@ def test_cuda_language_composes_common_and_cuda_symbols():
from tilelang.cuda import language as T
from tilelang.cuda import debug as cuda_debug
assert T.copy is T_comm.copy
# The CUDA dialect shadows T.copy with a version exposing CUDA hints
# (disable_tma, eviction_policy, prefer_instruction); see
# test_tilelang_language_dialect_op_hints.py for the full contract.
assert T.copy is not T_comm.copy
assert T.async_copy is T_comm.async_copy
assert T.tcgen05_mma is T.tcgen05_gemm
assert T.tcgen05_mma_blockscaled is T.tcgen05_gemm_blockscaled
assert T.wgmma_mma is T.wgmma_gemm
@@ -0,0 +1,214 @@
"""Backend-specific op hints are declared by the owning dialect.
Common tile ops (`T.copy`, `T.gemm`, ...) carry only target-neutral
parameters; each backend dialect shadows them with versions exposing that
backend's knobs as typed keywords (CUDA: TMA/cache hints, ROCm: k_pack, ...).
The hints are recorded as tile-op annotations, so a kernel written with one
dialect still compiles for targets that have no use for them.
"""
import importlib
import inspect
import pytest
import torch
import tilelang
import tilelang.cpu.language as Tcpu
import tilelang.language as T
import tilelang.rocm.language as Trocm
import tilelang.testing
from tvm import tirx
from tvm.tirx.stmt_functor import post_order_visit
def _keywords(func) -> set[str]:
return {
name
for name, p in inspect.signature(func).parameters.items()
if p.kind is inspect.Parameter.KEYWORD_ONLY or p.default is not inspect.Parameter.empty
}
def _tileop_annotations(func, opname: str) -> dict:
found = {}
def visit(node):
if isinstance(node, tirx.Call) and str(getattr(node.op, "name", "")) == opname:
found.update({str(k): v for k, v in node.annotations.items()})
post_order_visit(func.body, visit)
return found
def test_each_dialect_declares_its_own_op_hints():
common = importlib.import_module("tilelang.language.common")
cuda = importlib.import_module("tilelang.cuda.language")
expected_cuda_extra = {
"copy": {"disable_tma", "eviction_policy", "prefer_instruction"},
"im2col": {"eviction_policy"},
"gemm": {"mbar"},
"gemm_sp": {"wg_wait"},
"atomic_add": {"use_tma"},
"Parallel": {"prefer_async"},
"reduce_max": {"nan_propagate"},
"reduce_min": {"nan_propagate"},
"reduce_absmax": {"nan_propagate"},
"unroll": {"unroll_factor"},
"Unroll": {"unroll_factor"},
}
for name, extra in expected_cuda_extra.items():
assert _keywords(getattr(cuda, name)) - _keywords(getattr(common, name)) == extra, name
# The default facade is the CUDA dialect.
assert getattr(T, name) is getattr(cuda, name), name
# The CPU dialect keeps the neutral surface.
assert getattr(Tcpu, name) is getattr(common, name), name
assert _keywords(Trocm.gemm) - _keywords(common.gemm) == {"k_pack"}
assert Trocm.copy is common.copy
def test_cuda_copy_hints_are_recorded_as_annotations():
@T.prim_func
def main(A: T.Tensor((128,), "float16"), B: T.Tensor((128,), "float16")):
with T.Kernel(1, threads=128):
S = T.alloc_shared((128,), "float16")
T.copy(A, S, disable_tma=True, eviction_policy="evict_last", prefer_instruction="cp_async")
T.copy(S, B)
ann = _tileop_annotations(main, "tl.tileop.copy")
assert int(ann["disable_tma"]) == 1
assert int(ann["eviction_policy"]) == 2
assert isinstance(ann["prefer_instruction"], tirx.StringImm) and ann["prefer_instruction"].value == "cp_async"
def test_cuda_atomic_add_use_tma_is_recorded_as_annotation():
@T.prim_func
def main(A: T.Tensor((64,), "float32"), B: T.Tensor((64,), "float32")):
with T.Kernel(1, threads=64):
S = T.alloc_shared((64,), "float32")
T.copy(A, S)
T.atomic_add(B, S, use_tma=True)
ann = _tileop_annotations(main, "tl.tileop.atomicadd")
assert int(ann["use_tma"]) == 1
def test_cuda_loop_hints_are_recorded_as_loop_annotations():
@T.prim_func
def main(A: T.Tensor((64,), "float32"), B: T.Tensor((64,), "float32")):
with T.Kernel(1, threads=64):
for i in T.Parallel(64, prefer_async=True):
B[i] = A[i]
for j in T.unroll(4, unroll_factor=2):
B[j] = A[j]
annotations = {}
def visit(node):
if isinstance(node, tirx.For):
annotations.update({str(k): v for k, v in node.annotations.items()})
post_order_visit(main.body, visit)
assert bool(annotations["parallel_prefer_async"])
assert int(annotations["pragma_unroll_factor"]) == 2
def test_cuda_wgmma_gemm_records_wg_wait_annotation():
"""wgmma_gemm defers the warpgroup wait: wg_wait=-1 must ride the tile-op
annotations now that the positional slot is gone."""
@T.prim_func
def main(A: T.Tensor((64, 64), "float16"), B: T.Tensor((64, 64), "float16"), C: T.Tensor((64, 64), "float32")):
with T.Kernel(1, threads=128):
a = T.alloc_shared((64, 64), "float16")
b = T.alloc_shared((64, 64), "float16")
c = T.alloc_fragment((64, 64), "float32")
T.copy(A, a)
T.copy(B, b)
T.clear(c)
T.wgmma_gemm(a, b, c)
T.copy(c, C)
ann = _tileop_annotations(main, "tl.tileop.wgmma_gemm")
assert int(ann["wg_wait"]) == -1
def test_rocm_gemm_k_pack_traces_and_validates():
@Trocm.prim_func
def main(A: Trocm.Tensor((64, 64), "float16"), B: Trocm.Tensor((64, 64), "float16"), C: Trocm.Tensor((64, 64), "float32")):
with Trocm.Kernel(1, threads=256):
a = Trocm.alloc_shared((64, 64), "float16")
b = Trocm.alloc_shared((64, 64), "float16")
c = Trocm.alloc_fragment((64, 64), "float32")
Trocm.copy(A, a)
Trocm.copy(B, b)
Trocm.clear(c)
Trocm.gemm(a, b, c, k_pack=2)
Trocm.copy(c, C)
assert main is not None
ann = _tileop_annotations(main, "tl.tileop.gemm")
assert int(ann["k_pack"]) == 2
with pytest.raises(ValueError, match="k_pack must be an int equal to 1 or 2"):
@Trocm.prim_func
def bad(A: Trocm.Tensor((64, 64), "float16")):
with Trocm.Kernel(1):
Trocm.gemm(A, A, A, k_pack=3)
def test_neutral_dialects_reject_backend_hints():
with pytest.raises(TypeError, match="unexpected keyword argument 'disable_tma'"):
@Tcpu.prim_func
def cpu_copy(A: Tcpu.Tensor((16,), "float32"), B: Tcpu.Tensor((16,), "float32")):
with Tcpu.Kernel(1):
Tcpu.copy(A, B, disable_tma=True)
with pytest.raises(TypeError, match="unexpected keyword argument 'k_pack'"):
@T.prim_func
def cuda_gemm(A: T.Tensor((64, 64), "float16")):
with T.Kernel(1):
T.gemm(A, A, A, k_pack=2)
def test_cuda_hints_do_not_block_cpu_compilation():
"""Hints are droppable: a kernel authored with the CUDA dialect still
compiles for a target that has no use for them."""
@T.prim_func
def main(A: T.Tensor((64,), "float32"), B: T.Tensor((64,), "float32")):
with T.Kernel(1):
T.copy(A, B, disable_tma=True, eviction_policy="evict_first")
mod = tilelang.compile(main, target="c", out_idx=[-1])
x = torch.arange(64, dtype=torch.float32)
torch.testing.assert_close(mod(x), x)
@tilelang.testing.requires_cuda
def test_cuda_hinted_copy_compiles_and_runs():
@tilelang.jit(out_idx=[-1])
def add(N):
@T.prim_func
def main(A: T.Tensor((N,), "float32"), B: T.Tensor((N,), "float32")):
with T.Kernel(T.ceildiv(N, 128), threads=128) as bx:
s = T.alloc_shared((128,), "float32")
T.copy(A[bx * 128], s, disable_tma=True, eviction_policy="evict_first")
for i in T.Parallel(128):
s[i] = s[i] + 1.0
T.copy(s, B[bx * 128])
return main
k = add(256)
a = torch.randn(256, device="cuda")
torch.testing.assert_close(k(a), a + 1)
if __name__ == "__main__":
tilelang.testing.main()
+78 -2
View File
@@ -8,6 +8,16 @@ from tilelang.language.allocate import alloc_cluster_barrier, alloc_descriptor,
from tilelang.language.annotations import annotate_l2_hit_ratio, annotate_min_blocks_per_sm # noqa: F401
from tilelang.language.builtin import ( # noqa: F401
annotate_consumer_reg_alloc,
barrier_arrive,
barrier_wait,
get_lane_idx,
get_warp_idx,
get_warp_idx_sync,
mbarrier_arrive,
mbarrier_arrive_expect_tx,
mbarrier_expect_tx,
mbarrier_wait_parity,
no_set_max_nreg,
annotate_producer_reg_dealloc,
create_tma_descriptor,
deallocate_tmem,
@@ -46,9 +56,37 @@ 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).
# The CUDA dialect shadows a handful of common constructs with versions that
# expose CUDA's knobs as typed keywords: T.Kernel (threads, prelude,
# cluster_dims), T.copy / T.im2col (TMA and cache hints), T.gemm (mbar),
# T.atomic_add (use_tma), T.Parallel (prefer_async) and T.unroll
# (unroll_factor). Semantics match the common versions; the extra keywords
# are recorded on the op and consumed by the CUDA pipeline.
from .kernel import Kernel # noqa: F401
from .reduce_op import reduce_absmax, reduce_max, reduce_min # noqa: F401
from tilelang.language.math_intrinsics import ( # noqa: F401
__cos,
__exp,
__exp10,
__log,
__log2,
__log10,
__sin,
__tan,
fast_rcp,
ieee_add,
ieee_fdiv,
ieee_fmaf,
ieee_frcp,
ieee_frsqrt,
ieee_fsqrt,
ieee_mul,
ieee_sub,
)
from .copy_op import copy, im2col # noqa: F401
from .gemm_op import gemm, gemm_sp # noqa: F401
from .atomic import atomic_add # noqa: F401
from .loop import Parallel, Unroll, unroll # noqa: F401
from .cluster import * # noqa: F401,F403
from .cluster import __all__ as _CLUSTER_ALL
from .intrinsics import * # noqa: F401,F403
@@ -70,6 +108,44 @@ _CUDA_API_ALL = (
"ClusterKernel",
"CUDASourceCodeKernel",
"Kernel",
"__cos",
"__exp",
"__exp10",
"__log",
"__log2",
"__log10",
"__sin",
"__tan",
"barrier_arrive",
"barrier_wait",
"fast_rcp",
"get_lane_idx",
"get_warp_idx",
"get_warp_idx_sync",
"ieee_add",
"ieee_fdiv",
"ieee_fmaf",
"ieee_frcp",
"ieee_frsqrt",
"ieee_fsqrt",
"ieee_mul",
"ieee_sub",
"mbarrier_arrive",
"mbarrier_arrive_expect_tx",
"mbarrier_expect_tx",
"mbarrier_wait_parity",
"no_set_max_nreg",
"reduce_absmax",
"reduce_max",
"reduce_min",
"Parallel",
"Unroll",
"atomic_add",
"copy",
"gemm",
"gemm_sp",
"im2col",
"unroll",
"alloc_cluster_barrier",
"alloc_descriptor",
"alloc_tmem",
+50
View File
@@ -0,0 +1,50 @@
"""CUDA dialect of the atomic operators: the common ops plus CUDA knobs."""
from __future__ import annotations
from tvm.tirx import Buffer, PrimExpr
from tilelang.language.atomic import atomic_add as _common_atomic_add
__all__ = ["atomic_add"]
def atomic_add(
dst: Buffer,
value: PrimExpr,
memory_order: str | None = None,
return_prev: bool = False,
use_tma: bool = False,
annotations: dict | None = None,
) -> PrimExpr:
"""Atomically add ``value`` into ``dst``, with CUDA lowering knobs.
Same semantics as the common :func:`tilelang.language.atomic.atomic_add`.
``use_tma`` selects the sm90+ TMA ``cp.reduce`` lowering for the
tile-region path; targets without TMA reject it at compile time.
Parameters:
dst (Buffer): Destination buffer/address to apply the atomic add.
value (PrimExpr): Value to add atomically.
memory_order (Optional[str]): Memory-order name controlling the atomic
operation's ordering ("relaxed", "consume", "acquire", "release",
"acq_rel", "seq_cst").
return_prev (bool): Return the previous value (scalar path only).
use_tma (bool): If True, lower the tile-region atomic add through TMA
``cp.reduce``. Available on sm90+ only (default False).
annotations (Optional[dict]): Extra annotations for the tile-region
path; values in it take precedence over the individual keywords.
Returns:
PrimExpr: A handle to the atomic operation.
"""
ann: dict = dict(annotations) if annotations is not None else {}
if use_tma:
ann.setdefault("use_tma", 1)
return _common_atomic_add(
dst,
value,
memory_order=memory_order,
return_prev=return_prev,
annotations=ann or None,
)
+94
View File
@@ -0,0 +1,94 @@
"""CUDA dialect of the copy operators: the common ops plus CUDA copy hints."""
from __future__ import annotations
from typing import Any, Literal
from tvm import tirx
from tilelang._typing import BufferLikeType
from tilelang.language.copy_op import (
EVICTION_POLICY_IDS,
copy as _common_copy,
im2col_impl,
)
__all__ = ["copy", "im2col"]
def copy(
src: BufferLikeType,
dst: BufferLikeType,
*,
coalesced_width: int | None = None,
disable_tma: bool = False,
eviction_policy: Literal["evict_normal", "evict_first", "evict_last"] | None = None,
prefer_instruction: str | None = None,
annotations: dict | None = None,
loop_layout: Any | None = None,
) -> tirx.PrimExpr | tirx.Stmt:
"""Copy data between memory regions, with CUDA lowering hints.
Same semantics as the common :func:`tilelang.language.copy_op.copy`; the
extra keywords steer how the CUDA backend lowers the copy. They are
performance hints recorded on the tile op: compiling the same kernel for a
target that has no use for them leaves the result unchanged.
Args:
src: Source memory region (Buffer, BufferLoad or BufferRegion).
dst: Destination memory region.
coalesced_width (Optional[int], keyword-only): Width for coalesced
memory access. Defaults to None.
disable_tma (bool, keyword-only): Never lower this copy through TMA
even when the shape and scopes qualify. Defaults to False.
eviction_policy (Optional[str], keyword-only): L2 cache eviction
priority for the generated load/store or TMA instruction, one of
``"evict_normal"``, ``"evict_first"``, ``"evict_last"``.
prefer_instruction (Optional[str], keyword-only): Preferred lowering
instruction category: ``"tma"``, ``"cp_async"`` or ``"sync"``. For
``"tma"``, T.copy keeps synchronous copy semantics; global ->
shared copies lower through TMA with an automatically allocated
barrier and wait when constraints are satisfied.
annotations (Optional[dict], keyword-only): Additional annotations
dict; values in it take precedence over the individual keywords.
loop_layout (Optional[Fragment], keyword-only): Parallel loop layout
hint for the SIMT copy path.
Returns:
tirx.Call: A handle to the copy operation.
"""
ann: dict = dict(annotations) if annotations is not None else {}
if "disable_tma" not in ann and disable_tma:
ann["disable_tma"] = disable_tma
if "eviction_policy" not in ann and eviction_policy is not None:
ann["eviction_policy"] = EVICTION_POLICY_IDS[eviction_policy]
if "prefer_instruction" not in ann and prefer_instruction is not None:
ann["prefer_instruction"] = tirx.StringImm(prefer_instruction)
return _common_copy(
src,
dst,
coalesced_width=coalesced_width,
annotations=ann or None,
loop_layout=loop_layout,
)
def im2col(
img: BufferLikeType,
col: BufferLikeType,
nhw_step: tirx.PrimExpr,
c_step: tirx.PrimExpr,
kernel: int,
stride: int,
dilation: int,
pad: int,
eviction_policy: Literal["evict_normal", "evict_first", "evict_last"] | None = None,
annotations: dict | None = None,
) -> tirx.PrimExpr:
"""Perform im2col transformation for 2D convolution, with CUDA hints.
Same semantics as the common :func:`tilelang.language.copy_op.im2col`;
``eviction_policy`` is the L2 cache hint consumed by the CUDA TMA im2col
lowering (ignored by the generic SIMT fallback other targets use).
"""
return im2col_impl(img, col, nhw_step, c_step, kernel, stride, dilation, pad, eviction_policy=eviction_policy, annotations=annotations)
+117
View File
@@ -0,0 +1,117 @@
"""CUDA dialect of ``T.gemm``: the common GEMM plus CUDA-specific knobs."""
from __future__ import annotations
from tvm import tirx
from tilelang._typing import BufferLikeType
from tilelang.language.experimental.gemm_sp_op import _gemm_sp_impl
from tilelang.language.gemm_op import BarrierType, GemmWarpPolicy, _gemm_impl
__all__ = ["gemm", "gemm_sp"]
def gemm(
A: BufferLikeType,
B: BufferLikeType,
C: BufferLikeType,
transpose_A: bool = False,
transpose_B: bool = False,
policy: GemmWarpPolicy = GemmWarpPolicy.Square,
clear_accum: bool = False,
mbar: BarrierType | None = None,
annotations: dict | None = None,
) -> tirx.PrimExpr:
"""TileLang GEMM operator for CUDA.
Same semantics as the common :func:`tilelang.language.gemm_op.gemm`: the
default synchronous GEMM. On Hopper, if the compiler selects WGMMA
lowering, TileLang inserts the corresponding wait implicitly. On Blackwell
TCGEN5MMA, TileLang inserts the corresponding
``mbarrier_wait_parity(...)`` implicitly after issue.
For manual asynchronous scheduling, use ``T.wgmma_gemm(...)`` with
``T.wait_wgmma(...)`` on Hopper, or ``T.tcgen05_gemm(...)`` with
``T.mbarrier_wait_parity(...)`` on Blackwell.
Args:
A (BufferLikeType, i.e. Buffer | BufferLoad | BufferRegion, or Var): Input buffer A.
B (BufferLikeType): Input buffer B.
C (BufferLikeType): Output buffer C.
transpose_A (bool): Whether to transpose A. Defaults to False.
transpose_B (bool): Whether to transpose B. Defaults to False.
policy (GemmWarpPolicy): GEMM warp partition policy.
clear_accum (bool): Whether to clear the accumulator.
mbar (BarrierType, i.e. Buffer | BufferLoad, or Var, optional): Mbarrier in Blackwell.
Required when this GEMM lowers to TCGEN5MMA. Defaults to None.
annotations (Optional[dict]): Additional annotations.
Returns:
tirx.Call: A handle to the GEMM operation.
"""
return _gemm_impl(
"tl.tileop.gemm",
A,
B,
C,
transpose_A,
transpose_B,
policy,
clear_accum,
mbar,
annotations=annotations,
)
def gemm_sp(
A_sparse: BufferLikeType | tirx.Var,
E: BufferLikeType | tirx.Var,
B: BufferLikeType | tirx.Var,
C: BufferLikeType | tirx.Var,
transpose_A: bool = False,
transpose_E: bool = False,
transpose_B: bool = False,
policy: GemmWarpPolicy = GemmWarpPolicy.Square,
clear_accum: bool = False,
wg_wait: int = 0,
annotations: dict | None = None,
) -> tirx.Call:
"""Sparse GEMM (2:4 structured sparsity) for CUDA.
Same semantics as the common :func:`tilelang.language.experimental.gemm_sp_op.gemm_sp`.
``wg_wait`` is the Hopper warpgroup wait count consumed when the WGMMA SP
lowering is selected (``-1`` defers the wait to an explicit
``T.wait_wgmma``); it rides in the tile-op annotations.
Args:
A_sparse: Compressed sparse matrix containing only non-zero elements.
E: Metadata tensor encoding the sparsity pattern of A.
B: Dense input matrix.
C: Output accumulator matrix.
transpose_A: Whether to transpose A. Defaults to False.
transpose_E: Whether to transpose E. Defaults to False.
transpose_B: Whether to transpose B. Defaults to False.
policy: Warp partition policy. Defaults to GemmWarpPolicy.Square.
clear_accum: Whether to zero the accumulator before computation. Defaults to False.
wg_wait: Warp group wait count. Defaults to 0.
annotations: Additional annotations; values in it take precedence.
Returns:
tirx.Call: A handle to the sparse GEMM operation.
"""
ann = dict(annotations) if annotations is not None else {}
if wg_wait != 0:
ann.setdefault("wg_wait", wg_wait)
return _gemm_sp_impl(
"tl.tileop.gemm_sp",
A_sparse,
E,
B,
C,
transpose_A,
transpose_E,
transpose_B,
policy,
clear_accum,
annotations=ann or None,
)
+101
View File
@@ -0,0 +1,101 @@
"""CUDA dialect of the loop constructs: the common loops plus CUDA hints."""
from __future__ import annotations
from typing import Any
from tvm import tirx
from tvm.tirx.script.builder import frame
from tilelang.language.loop import (
Parallel as _common_Parallel,
unroll as _common_unroll,
)
__all__ = ["Parallel", "Unroll", "unroll"]
def Parallel(
*extents: int | tirx.PrimExpr,
coalesced_width: int | None = None,
loop_layout: Any | None = None,
prefer_async: bool | None = None,
annotations: dict[str, Any] | None = None,
) -> frame.ForFrame:
"""Construct a nested parallel loop, with CUDA lowering hints.
Same semantics as the common :func:`tilelang.language.loop.Parallel`.
``prefer_async`` requests the PTX cp.async rewrite for copies in this loop
subtree even outside pipelined loops; it is a performance hint ignored by
targets without async copy.
Parameters
----------
extents : int | PrimExpr
Extents of the parallel loop nest.
coalesced_width : Optional[int]
Width for coalesced memory access.
loop_layout : Optional[Fragment]
Layout annotation for the parallel loop nest.
prefer_async : Optional[bool]
When True, requests cp.async injection for this subtree; when False,
forbids it. Lowered as the ``"parallel_prefer_async"`` annotation.
annotations : Optional[Dict[str, Any]]
Additional loop annotations; values in it take precedence.
"""
ann: dict[str, Any] = dict(annotations) if annotations is not None else {}
if prefer_async is not None:
ann.setdefault("parallel_prefer_async", prefer_async)
return _common_Parallel(
*extents,
coalesced_width=coalesced_width,
loop_layout=loop_layout,
annotations=ann or None,
)
def unroll(
start: tirx.PrimExpr,
stop: tirx.PrimExpr | None = None,
step: tirx.PrimExpr | None = None,
*,
explicit: bool = False,
unroll_factor: int | None = None,
annotations: dict[str, Any] | None = None,
) -> frame.ForFrame:
"""The unrolled For statement, with the CUDA unroll-factor pragma.
Same semantics as the common :func:`tilelang.language.loop.unroll`.
``unroll_factor`` emits ``#pragma unroll N``, which only the CUDA codegen
honors; it is mutually exclusive with ``explicit``.
Parameters
----------
start, stop, step : PrimExpr
Iteration range.
explicit : bool
Whether to explicitly unroll the loop at compile time.
unroll_factor : Optional[int]
Partial unroll factor, lowered as the ``"pragma_unroll_factor"``
annotation.
annotations : Optional[Dict[str, Any]]
Additional loop annotations; values in it take precedence.
"""
ann: dict[str, Any] = dict(annotations) if annotations is not None else {}
if unroll_factor is not None:
ann.setdefault("pragma_unroll_factor", unroll_factor)
return _common_unroll(start, stop, step, explicit=explicit, annotations=ann or None)
def Unroll(
start: tirx.PrimExpr,
stop: tirx.PrimExpr | None = None,
step: tirx.PrimExpr | None = None,
*,
explicit: bool = False,
unroll_factor: int | None = None,
annotations: dict[str, Any] | None = None,
) -> frame.ForFrame:
"""Alias of the CUDA dialect's :func:`unroll`."""
return unroll(start, stop, step, explicit=explicit, unroll_factor=unroll_factor, annotations=annotations)
+73
View File
@@ -0,0 +1,73 @@
"""CUDA dialect of the reduction operators: the common ops plus CUDA knobs."""
from __future__ import annotations
from tvm import tirx
from tilelang.language.reduce_op import (
reduce_absmax as _common_reduce_absmax,
reduce_max as _common_reduce_max,
reduce_min as _common_reduce_min,
)
__all__ = ["reduce_absmax", "reduce_max", "reduce_min"]
def _with_nan_propagate(annotations: dict | None, nan_propagate: bool) -> dict | None:
if not nan_propagate:
return annotations
ann = dict(annotations) if annotations is not None else {}
ann.setdefault("nan_propagate", True)
return ann
def reduce_max(
buffer: tirx.Buffer,
out: tirx.Buffer,
dim: int = -1,
clear: bool = True,
batch: int = 1,
nan_propagate: bool = False,
annotations: dict | None = None,
) -> None:
"""Perform reduce max, with the CUDA NaN-propagation knob.
Same semantics as the common :func:`tilelang.language.reduce_op.reduce_max`.
``nan_propagate`` is meaningful for float16/bfloat16 only: when True the
reduction lowers to ``__hmax_nan`` so NaNs propagate; when False (default)
``__hmax`` returns the non-NaN operand. Targets without these intrinsics
reject the annotation at compile time.
"""
_common_reduce_max(buffer, out, dim, clear, batch=batch, annotations=_with_nan_propagate(annotations, nan_propagate))
def reduce_min(
buffer: tirx.Buffer,
out: tirx.Buffer,
dim: int = -1,
clear: bool = True,
batch: int = 1,
nan_propagate: bool = False,
annotations: dict | None = None,
) -> None:
"""Perform reduce min, with the CUDA NaN-propagation knob.
See :func:`reduce_max`; this lowers to ``__hmin_nan``/``__hmin``.
"""
_common_reduce_min(buffer, out, dim, clear, batch=batch, annotations=_with_nan_propagate(annotations, nan_propagate))
def reduce_absmax(
buffer: tirx.Buffer,
out: tirx.Buffer,
dim: int = -1,
clear: bool = True,
batch: int = 1,
nan_propagate: bool = False,
annotations: dict | None = None,
) -> None:
"""Perform reduce absolute max, with the CUDA NaN-propagation knob.
See :func:`reduce_max`.
"""
_common_reduce_absmax(buffer, out, dim, clear, batch=batch, annotations=_with_nan_propagate(annotations, nan_propagate))
+16 -3
View File
@@ -11,8 +11,21 @@ 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
# Imported by name so static type checkers resolve the CUDA-typed signatures
# through this facade (they cannot evaluate the dynamic __all__).
from tilelang.cuda.language import ( # noqa: F401
Kernel,
Parallel,
Unroll,
atomic_add,
copy,
gemm,
gemm_sp,
im2col,
reduce_absmax,
reduce_max,
reduce_min,
unroll,
)
__tilelang_dialect__ = "cuda"
+3 -4
View File
@@ -213,7 +213,6 @@ def atomic_add(
value: PrimExpr,
memory_order: str | None = None,
return_prev: bool = False,
use_tma: bool = False,
annotations: dict | None = None,
) -> PrimExpr:
"""
@@ -226,7 +225,9 @@ def atomic_add(
value (PrimExpr): Value to add atomically.
memory_order (Optional[str]): Optional memory-order name controlling the atomic operation's ordering.
return_prev (bool): If True, return the previous value; if False, return handle (default False).
use_tma (bool): If True, use TMA (cp.reduce) to perform the atomic add. This is available only for sm90+ (default False).
annotations (Optional[dict]): Extra annotations for the tile-region path. Backend
hints ride through this dict; the CUDA dialect (``tilelang.cuda.language.atomic_add``)
exposes ``use_tma`` (sm90+ TMA cp.reduce) as a typed keyword instead.
Returns:
PrimExpr: A handle representing the atomic addition operation, or the previous value if return_prev is True.
@@ -304,8 +305,6 @@ def atomic_add(
raise NotImplementedError("return_prev is not supported for tile-region-based atomic operations")
# Build annotations dict
if use_tma:
ann["use_tma"] = 1
if memory_order is not None:
ann["memory_order"] = _MEMORY_ORDER_ID_MAP[memory_order]
-15
View File
@@ -7,7 +7,6 @@ from tilelang._typing import BufferLikeType, BufferLikeTypeTuple, BarrierType, D
from tilelang import tvm as tvm
from tilelang.language.common import ptx_arrive_barrier, evaluate
from tilelang.language.eager.builder import macro
from tilelang.language.kernel import get_thread_bindings, get_block_extents
from tvm import DataType, DataTypeCode, tirx
from tvm.runtime import convert
from tvm.tirx import PrimExpr, Var, Call, BufferLoad, BufferRegion
@@ -1179,15 +1178,6 @@ def match_all_sync(
return tirx.call_intrin("uint32", tirx.op.Op.get("tl.match_all_sync"), _as_uint32_mask(mask), value)
def sync_global():
"""Synchronize all threads in the entire grid."""
tx, ty, tz = get_thread_bindings()
ex, ey, ez = get_block_extents()
print(tx, ty, tz, ex, ey, ez)
args = ["global", tx == 0 and ty == 0 and tz == 0, ex * ey * ez]
return evaluate(tirx.Call("handle", "tirx.tvm_storage_sync", args))
def sync_grid():
"""Synchronize all threads in a grid."""
return tirx.call_intrin("handle", tirx.op.Op.get("tl.sync_grid"))
@@ -1393,11 +1383,6 @@ def cooperative_tensor_multiply_accumulate(
)
def loop_break():
"""Break out of the innermost loop."""
return tirx.call_intrin("handle", tirx.op.Op.get("tl.loop_break"))
def cp_async_barrier_noinc(barrier: BarrierType):
"""Perform a ptx async copy barrier using cp.async.mbarrier.arrive.noinc."""
barrier = _mbar_to_buffer_load(barrier)
+10 -24
View File
@@ -30,7 +30,15 @@ from .loop import (
Vectorized, # noqa: F401
)
from .frame import has_let_value, get_let_value # noqa: F401
from .math_intrinsics import * # noqa: F401,F403
from .math_intrinsics import ( # noqa: F401
abs2,
add2,
fma2,
max2,
min2,
mul2,
sub2,
)
from .kernel import (
Kernel, # noqa: F401
KernelLaunchFrame, # noqa: F401
@@ -118,21 +126,10 @@ from .builtin import ( # noqa: F401
any_sync,
ballot,
ballot_sync,
barrier_arrive,
barrier_wait,
get_lane_idx,
get_warp_idx,
get_warp_idx_sync,
mbarrier_arrive,
mbarrier_arrive_expect_tx,
mbarrier_expect_tx,
mbarrier_wait_parity,
no_set_max_nreg,
shfl_down,
shfl_sync,
shfl_up,
shfl_xor,
sync_global,
sync_grid,
sync_threads,
sync_warp,
@@ -189,7 +186,7 @@ def import_source(source: str | None = None):
from .tir.common import __all__ as _TIR_COMMON_ALL # noqa: E402
from .eager import __all__ as _EAGER_ALL # noqa: E402
from .tir.ir import __all__ as _TIR_IR_ALL # noqa: E402
from .math_intrinsics import __all__ as _MATH_ALL # noqa: E402
from .math_intrinsics import COMMON_MATH_INTRINSICS as _MATH_ALL # noqa: E402
_LOCAL_EXPORTS = (
"BaseTileScheduler",
@@ -251,8 +248,6 @@ _LOCAL_EXPORTS = (
"atomic_store",
"ballot",
"ballot_sync",
"barrier_arrive",
"barrier_wait",
"c2d_im2col",
"clamp",
"clear",
@@ -277,14 +272,11 @@ _LOCAL_EXPORTS = (
"get_cluster_id",
"get_cluster_ids",
"get_cluster_size",
"get_lane_idx",
"get_let_value",
"get_thread_binding",
"get_thread_bindings",
"get_thread_extent",
"get_thread_extents",
"get_warp_idx",
"get_warp_idx_sync",
"has_let_value",
"im2col",
"import_source",
@@ -293,12 +285,7 @@ _LOCAL_EXPORTS = (
"loop_break",
"make_tensor",
"make_tensor_from_addr",
"mbarrier_arrive",
"mbarrier_arrive_expect_tx",
"mbarrier_expect_tx",
"mbarrier_wait_parity",
"meta_class",
"no_set_max_nreg",
"reduce",
"reduce_absmax",
"reduce_abssum",
@@ -313,7 +300,6 @@ _LOCAL_EXPORTS = (
"shfl_sync",
"shfl_up",
"shfl_xor",
"sync_global",
"sync_grid",
"sync_threads",
"sync_warp",
+54 -46
View File
@@ -51,14 +51,16 @@ def _normalize_copy_regions(
return src, dst
# Cache eviction priority names -> integer ids used in the tile-op call
# protocol (consumed by CUDA codegen; see the CUDA dialect's copy/im2col).
EVICTION_POLICY_IDS = {"evict_normal": 0, "evict_first": 1, "evict_last": 2}
def copy(
src: BufferLikeType,
dst: BufferLikeType,
*,
coalesced_width: int | None = None,
disable_tma: bool = False,
eviction_policy: Literal["evict_normal", "evict_first", "evict_last"] | None = None,
prefer_instruction: str | None = None,
annotations: dict | None = None,
loop_layout: Any | None = None,
) -> tirx.PrimExpr | tirx.Stmt:
@@ -68,17 +70,12 @@ def copy(
src (Union[tirx.Buffer, tirx.BufferLoad, tirx.BufferRegion]): Source memory region
dst (Union[tirx.Buffer, tirx.BufferLoad, tirx.BufferRegion]): Destination memory region
coalesced_width (Optional[int], keyword-only): Width for coalesced memory access. Defaults to None.
disable_tma (bool, keyword-only): Whether to disable TMA acceleration. Defaults to False.
eviction_policy (Optional[str], keyword-only): Cache eviction policy. Defaults to None.
prefer_instruction (Optional[str], keyword-only): Backend-specific preferred lowering
instruction category. For CUDA, recognized values include "tma", "cp_async", and
"sync". For "tma", T.copy keeps synchronous copy semantics; global -> shared copies
lower through TMA with an automatically allocated barrier and wait when constraints
are satisfied.
annotations (Optional[dict], keyword-only): Additional annotations dict. If provided,
coalesced_width, disable_tma, eviction_policy, and prefer_instruction can also
be specified here.
Values in annotations take precedence over individual arguments.
coalesced_width can also be specified here. Values in annotations take precedence
over individual arguments. Backend-specific copy hints ride through this dict;
the backend dialects expose them as typed keywords instead
(``tilelang.cuda.language.copy`` adds ``disable_tma``, ``eviction_policy`` and
``prefer_instruction``). Hints a target does not understand are ignored.
loop_layout (Optional[Fragment], keyword-only): A parallel loop layout hint for the SIMT copy
(only valid for normal SIMT copy; incompatible with TMA/LDSM/STSM/TMem). When provided,
it is attached to the outermost parallel loop generated by this copy.
@@ -116,13 +113,6 @@ def copy(
# Individual arguments take lower precedence than annotations
if "coalesced_width" not in ann and coalesced_width is not None:
ann["coalesced_width"] = coalesced_width
if "disable_tma" not in ann and disable_tma:
ann["disable_tma"] = disable_tma
if "eviction_policy" not in ann and eviction_policy is not None:
eviction_policy_map = {"evict_normal": 0, "evict_first": 1, "evict_last": 2}
ann["eviction_policy"] = eviction_policy_map[eviction_policy]
if "prefer_instruction" not in ann and prefer_instruction is not None:
ann["prefer_instruction"] = tirx.StringImm(prefer_instruction)
# Parallel loop layout hint (Fragment). Mirrors T.Parallel(loop_layout=...)
if loop_layout is not None and "parallel_loop_layout" not in ann:
@@ -186,8 +176,7 @@ def copy_cluster(
if "barrier" not in ann and remote_barrier is not None:
ann["barrier"] = remote_barrier
if "eviction_policy" not in ann and eviction_policy is not None:
eviction_policy_map = {"evict_normal": 0, "evict_first": 1, "evict_last": 2}
ann["eviction_policy"] = eviction_policy_map[eviction_policy]
ann["eviction_policy"] = EVICTION_POLICY_IDS[eviction_policy]
if "coalesced_width" not in ann and coalesced_width is not None:
ann["coalesced_width"] = coalesced_width
if loop_layout is not None and "parallel_loop_layout" not in ann:
@@ -330,8 +319,7 @@ def tma_copy(
ann["leader_scope_threads"] = leader_scope_threads
if "eviction_policy" not in ann and eviction_policy is not None:
eviction_policy_map = {"evict_normal": 0, "evict_first": 1, "evict_last": 2}
ann["eviction_policy"] = eviction_policy_map[eviction_policy]
ann["eviction_policy"] = EVICTION_POLICY_IDS[eviction_policy]
return tirx.call_intrin("handle", tirx.op.Op.get("tl.tileop.tma_copy"), src, dst, annotations=ann)
@@ -574,7 +562,7 @@ def transpose(
)
def im2col(
def im2col_impl(
img: BufferLikeType,
col: BufferLikeType,
nhw_step: tirx.PrimExpr,
@@ -586,26 +574,14 @@ def im2col(
eviction_policy: Literal["evict_normal", "evict_first", "evict_last"] | None = None,
annotations: dict | None = None,
) -> tirx.PrimExpr:
"""Perform im2col transformation for 2D convolution.
"""Shared im2col implementation behind the common and dialect wrappers.
Args:
img (tirx.Buffer): Input image buffer
col (tirx.Buffer): Output column buffer
nhw_step (tirx.PrimExpr): Step size for batch and spatial dimensions
c_step (tirx.PrimExpr): Step size for channel dimension
kernel (int): Kernel size
stride (int): Stride of the convolution
dilation (int): Dilation rate
pad (int): Padding size
annotations: Optional annotations to attach to the call
Returns:
tirx.Call: A handle to the im2col operation
``eviction_policy`` rides in the tile-op annotations; only the CUDA TMA
lowering reads it, so only the CUDA dialect exposes it.
"""
if eviction_policy is None:
eviction_policy = 0
else:
eviction_policy = {"evict_normal": 0, "evict_first": 1, "evict_last": 2}[eviction_policy]
ann = _normalize_annotations(annotations)
if eviction_policy is not None and "eviction_policy" not in ann:
ann["eviction_policy"] = EVICTION_POLICY_IDS[eviction_policy]
img_region = to_buffer_region(img)
col_region = to_buffer_region(col)
img_extents = [r.extent for r in img_region.region]
@@ -623,11 +599,43 @@ def im2col(
stride,
dilation,
pad,
eviction_policy,
annotations=_normalize_annotations(annotations),
annotations=ann,
)
def im2col(
img: BufferLikeType,
col: BufferLikeType,
nhw_step: tirx.PrimExpr,
c_step: tirx.PrimExpr,
kernel: int,
stride: int,
dilation: int,
pad: int,
annotations: dict | None = None,
) -> tirx.PrimExpr:
"""Perform im2col transformation for 2D convolution.
Args:
img (tirx.Buffer): Input image buffer
col (tirx.Buffer): Output column buffer
nhw_step (tirx.PrimExpr): Step size for batch and spatial dimensions
c_step (tirx.PrimExpr): Step size for channel dimension
kernel (int): Kernel size
stride (int): Stride of the convolution
dilation (int): Dilation rate
pad (int): Padding size
annotations: Optional annotations to attach to the call
The CUDA dialect (``tilelang.cuda.language.im2col``) additionally accepts
``eviction_policy``, a cache hint for the TMA lowering.
Returns:
tirx.Call: A handle to the im2col operation
"""
return im2col_impl(img, col, nhw_step, c_step, kernel, stride, dilation, pad, annotations=annotations)
@deprecated("T.c2d_im2col", "T.im2col", "0.14.0")
def c2d_im2col(
img: BufferLikeType,
@@ -647,7 +655,7 @@ def c2d_im2col(
Use :func:`im2col` instead. This alias is scheduled for removal in
TileLang 0.14.0.
"""
return im2col(
return im2col_impl(
img,
col,
nhw_step,
+11 -33
View File
@@ -7,8 +7,6 @@ from tvm import tirx
from tilelang.utils.language import (
to_buffer_region,
retrieve_shape,
retrieve_stride,
retrieve_offset,
prim_expr_equal,
)
from tilelang.language.utils import (
@@ -29,12 +27,12 @@ def _gemm_sp_impl(
transpose_B: bool = False,
policy: GemmWarpPolicy = GemmWarpPolicy.Square,
clear_accum: bool = False,
wg_wait: int = 0,
annotations: dict | None = None,
) -> tirx.Call:
"""Shared sparse GEMM implementation.
Returns a call_intrin handle for the given op key.
Returns a call_intrin handle for the given op key. Backend lowering knobs
such as ``wg_wait`` ride in ``annotations``.
"""
def legalize_arguments(arg: BufferLikeType | tirx.Var) -> BufferLikeType:
@@ -59,9 +57,6 @@ def _gemm_sp_impl(
B_shape = retrieve_shape(B)
C_shape = retrieve_shape(C)
A_stride = retrieve_stride(A_sparse)
B_stride = retrieve_stride(B)
assert len(C_shape) == 2, "current only support C as a 2D tensor"
assert len(A_shape) >= 2, "current only support A as a 2D or higher-order tensor"
assert len(B_shape) >= 2, "current only support B as a 2D or higher-order tensor"
@@ -85,16 +80,6 @@ def _gemm_sp_impl(
if not isinstance(dim, tirx.IntImm):
raise ValueError(f"T.gemm_sp requires static tile dimensions, but {name} is symbolic: {dim}")
stride_a = A_stride[-2]
stride_b = B_stride[-2]
A_offset = retrieve_offset(A_sparse)
B_offset = retrieve_offset(B)
assert A_offset[-2] == 0, "The offset of the first dimension of A must be 0"
assert B_offset[-2] == 0, "The offset of the first dimension of B must be 0"
offset_a = A_offset[-1]
offset_b = B_offset[-1]
A_arg = buffer_region_to_tile_region(A_region, "r", [r for r in A_shape])
E_arg = buffer_region_to_tile_region(E_region, "r", [r for r in E_shape])
B_arg = buffer_region_to_tile_region(B_region, "r", [r for r in B_shape])
@@ -114,15 +99,6 @@ def _gemm_sp_impl(
K,
policy,
clear_accum,
stride_a,
stride_b,
offset_a,
offset_b,
# k_pack call slot: parsed and validated on the C++ side but never
# consumed by any sparse-GEMM lowering; kept at 1 for protocol
# stability.
1,
wg_wait,
annotations=annotations,
)
@@ -137,7 +113,6 @@ def gemm_sp(
transpose_B: bool = False,
policy: GemmWarpPolicy = GemmWarpPolicy.Square,
clear_accum: bool = False,
wg_wait: int = 0,
annotations: dict | None = None,
) -> tirx.Call:
"""TileLang sparse GEMM operator.
@@ -159,8 +134,11 @@ def gemm_sp(
transpose_B: Whether to transpose B. Defaults to False.
policy: Warp partition policy. Defaults to GemmSPWarpPolicy.Square.
clear_accum: Whether to zero the accumulator before computation. Defaults to False.
wg_wait: Warp group wait count. Defaults to 0.
annotations: Additional annotations.
annotations: Additional annotations. The CUDA dialect
(``tilelang.cuda.language.gemm_sp``) additionally exposes
``wg_wait`` (Hopper warpgroup wait count) as a typed keyword. The
former ``k_pack`` parameter was parsed but never consumed by any
sparse-GEMM lowering and has been removed.
Returns:
tirx.Call: A handle to the sparse GEMM operation.
@@ -176,7 +154,6 @@ def gemm_sp(
transpose_B,
policy,
clear_accum,
wg_wait,
annotations=annotations,
)
@@ -218,6 +195,9 @@ def wgmma_gemm_sp(
Returns:
tirx.Call: A handle to the sparse GEMM operation.
"""
ann = dict(annotations) if annotations is not None else {}
# Explicit async WGMMA SP: never auto-emit the warpgroup wait.
ann.setdefault("wg_wait", -1)
return _gemm_sp_impl(
"tl.tileop.wgmma_gemm_sp",
A_sparse,
@@ -229,8 +209,7 @@ def wgmma_gemm_sp(
transpose_B,
policy,
clear_accum,
-1,
annotations=annotations,
annotations=ann,
)
@@ -284,6 +263,5 @@ def tcgen05_gemm_sp(
transpose_B,
policy,
clear_accum,
0,
annotations=annotations,
)
+20 -75
View File
@@ -10,8 +10,6 @@ from tvm import tirx
from tilelang.utils.language import (
to_buffer_region,
retrieve_shape,
retrieve_stride,
retrieve_offset,
prim_expr_equal,
)
from tilelang.language.utils import (
@@ -29,17 +27,15 @@ def _gemm_impl(
transpose_B: bool = False,
policy: GemmWarpPolicy = GemmWarpPolicy.Square,
clear_accum: bool = False,
k_pack: int = 1,
wg_wait: int = 0,
mbar: BarrierType | None = None,
annotations: dict | None = None,
) -> tirx.PrimExpr:
"""Shared GEMM implementation.
Returns a call_intrin handle for the given op key.
Returns a call_intrin handle for the given op key. Backend lowering knobs
such as ``k_pack`` and ``wg_wait`` ride in ``annotations``; the dialect
wrappers and the CUDA gemm variants put them there.
"""
if not (isinstance(k_pack, int) and not isinstance(k_pack, bool) and k_pack in (1, 2)):
raise ValueError(f"T.gemm k_pack must be an int equal to 1 or 2, got {k_pack!r}")
def legalize_arguments(arg: BufferLikeType | tirx.Var) -> BufferLikeType:
"""Convert let-bound variables to their corresponding buffers.
@@ -97,20 +93,6 @@ def _gemm_impl(
if not isinstance(dim, tirx.IntImm):
raise ValueError(f"T.gemm requires static tile dimensions, but {name} is symbolic: {dim}")
# Deprecated: every lowering consumes the complete operand BufferRegions,
# so the serialized per-axis strides and final-axis offsets below are no
# longer read in-tree and are NOT validated (the historic
# ``A_offset[-2] == 0`` assertions are gone). They are kept in the call
# protocol only for out-of-tree consumers of the GemmNode fields.
A_stride = retrieve_stride(A_region)
B_stride = retrieve_stride(B_region)
stride_a = A_stride[-2]
stride_b = B_stride[-2]
A_offset = retrieve_offset(A_region)
B_offset = retrieve_offset(B_region)
offset_a = A_offset[-1]
offset_b = B_offset[-1]
if mbar is not None:
assert isinstance(mbar, (tirx.Buffer, tirx.BufferLoad)), (
f"mbar for tcgen5mma must be a tirx.Buffer or tirx.BufferLoad, but got {type(mbar)}"
@@ -121,9 +103,9 @@ def _gemm_impl(
A_arg = buffer_region_to_tile_region(A_region, "r", [r for r in A_shape])
B_arg = buffer_region_to_tile_region(B_region, "r", [r for r in B_shape])
C_arg = buffer_region_to_tile_region(C_region, "rw", [r for r in C_shape])
# When mbar is None, pass a placeholder constant (0).
# The C++ side checks if arg 16 is a BufferLoadNode before using it,
# so a non-BufferLoad value will be correctly ignored.
# When mbar is None, pass a placeholder constant (0). The C++ side only
# accepts the mbar slot when it is a BufferLoadNode, so the placeholder is
# correctly ignored.
mbar_arg = mbar if mbar is not None else tirx.const(0, dtype="int32")
return tirx.call_intrin(
"handle",
@@ -138,12 +120,6 @@ def _gemm_impl(
K,
policy,
clear_accum,
stride_a,
stride_b,
offset_a,
offset_b,
k_pack,
wg_wait,
mbar_arg,
C_coords[0],
C_coords[1],
@@ -159,8 +135,6 @@ def gemm(
transpose_B: bool = False,
policy: GemmWarpPolicy = GemmWarpPolicy.Square,
clear_accum: bool = False,
k_pack: int = 1,
mbar: BarrierType | None = None,
annotations: dict | None = None,
) -> tirx.PrimExpr:
"""TileLang GEMM operator.
@@ -182,11 +156,12 @@ def gemm(
transpose_B (bool): Whether to transpose B. Defaults to False.
policy (GemmWarpPolicy): GEMM warp partition policy.
clear_accum (bool): Whether to clear the accumulator.
k_pack (int): Number of packed matrix cores, for ROCm only. Must be 1 or 2. Defaults to 1.
mbar (BarrierType, i.e. Buffer | BufferLoad, or Var, optional): Mbarrier in Blackwell.
Required when this GEMM lowers to TCGEN5MMA. Defaults to None.
annotations (Optional[dict]): Additional annotations.
Backend dialects extend this signature with their hardware's knobs:
``tilelang.cuda.language.gemm`` adds ``mbar`` (Blackwell TCGEN5MMA
barrier), ``tilelang.rocm.language.gemm`` adds ``k_pack`` (packed MFMA).
Returns:
tirx.Call: A handle to the GEMM operation.
"""
@@ -199,9 +174,7 @@ def gemm(
transpose_B,
policy,
clear_accum,
k_pack,
0,
mbar,
None,
annotations=annotations,
)
@@ -227,6 +200,9 @@ def wgmma_gemm(
compilation fails instead of silently falling back to MMA.
"""
ann = _normalize_annotations(annotations)
# Explicit async WGMMA: never auto-emit the warpgroup wait.
ann.setdefault("wg_wait", -1)
return _gemm_impl(
"tl.tileop.wgmma_gemm",
A,
@@ -236,10 +212,8 @@ def wgmma_gemm(
transpose_B,
policy,
clear_accum,
1,
-1,
None,
annotations=annotations,
annotations=ann,
)
@@ -287,8 +261,6 @@ def tcgen05_gemm(
transpose_B,
policy,
clear_accum,
1,
0,
mbar,
annotations=ann,
)
@@ -354,6 +326,8 @@ def tcgen05_gemm_blockscaled(
ann = {} if ann is None else dict(ann)
ann["sf_a_granularity_k"] = int(sf_a_granularity_k)
ann["sf_b_granularity_k"] = int(sf_b_granularity_k)
if wg_wait != 0:
ann["wg_wait"] = wg_wait
# Re-read normalized regions below after let legalization.
@@ -398,15 +372,6 @@ def tcgen05_gemm_blockscaled(
# Deprecated: kept in the call protocol only for out-of-tree consumers;
# not read or validated in-tree.
A_stride = retrieve_stride(A_region)
B_stride = retrieve_stride(B_region)
stride_a = A_stride[-2]
stride_b = B_stride[-2]
A_offset = retrieve_offset(A_region)
B_offset = retrieve_offset(B_region)
offset_a = A_offset[-1]
offset_b = B_offset[-1]
if mbar is not None:
assert isinstance(mbar, (tirx.Buffer, tirx.BufferLoad)), (
f"mbar for tcgen5mma must be a tirx.Buffer or tirx.BufferLoad, but got {type(mbar)}"
@@ -443,18 +408,12 @@ def tcgen05_gemm_blockscaled(
K,
policy,
clear_accum,
stride_a,
stride_b,
offset_a,
offset_b,
1, # k_pack
wg_wait,
mbar,
C_coords[0],
C_coords[1],
SFA_arg, # arg 19
SFB_arg, # arg 20
k_start, # arg 21
SFA_arg,
SFB_arg,
k_start,
annotations=ann,
)
@@ -529,14 +488,6 @@ def mma_gemm_blockscaled(
assert prim_expr_equal(K, K_B), f"T.mma_gemm_blockscaled K shape check failed: K_A = {K}, K_B = {K_B}"
assert prim_expr_equal(N_B, N), f"T.mma_gemm_blockscaled N shape check failed: N_B = {N_B}, N_C = {N}"
A_stride = retrieve_stride(A_region)
B_stride = retrieve_stride(B_region)
A_offset = retrieve_offset(A_region)
B_offset = retrieve_offset(B_region)
stride_a = A_stride[-2]
stride_b = B_stride[-2]
offset_a = A_offset[-1]
offset_b = B_offset[-1]
C_coords = [r.min for r in C_region.region]
A_arg = buffer_region_to_tile_region(A_region, "r", [r for r in A_shape])
@@ -561,12 +512,6 @@ def mma_gemm_blockscaled(
K,
policy,
clear_accum,
stride_a,
stride_b,
offset_a,
offset_b,
1, # k_pack
0, # wg_wait
tirx.const(0, dtype="int32"), # no mbarrier for synchronous mma.sync
C_coords[0],
C_coords[1],
+8 -22
View File
@@ -14,7 +14,6 @@ def Parallel(
*extents: int | tirx.PrimExpr,
coalesced_width: int | None = None,
loop_layout: Any | None = None,
prefer_async: bool | None = None,
annotations: dict[str, Any] | None = None,
) -> frame.ForFrame:
"""Tools to construct nested parallel for loop.
@@ -35,16 +34,11 @@ def Parallel(
For a k-dimensional ``T.Parallel(...)`` nest, the fragment's
``InputDim`` must equal ``k``.
prefer_async : Optional[bool]
Optional hint for PTX async-copy rewrite in this parallel loop subtree.
When set to ``True``, it requests cp.async injection even outside
pipelined loops. ``False``/``None`` keeps default behavior.
Internally lowered as loop annotation ``"parallel_prefer_async"``.
annotations : Optional[Dict[str, Any]]
Optional user-provided loop annotations attached to the outermost
generated parallel loop. For example:
``{"parallel_async_without_async_commit_wait": True}``.
generated parallel loop. Backend hints ride through this dict; the
CUDA dialect (``tilelang.cuda.language.Parallel``) exposes
``prefer_async`` (PTX cp.async rewrite) as a typed keyword instead.
Notes on layout constraints
---------------------------
@@ -82,8 +76,6 @@ def Parallel(
# Pass through to C++ as the standard parallel loop layout key.
# The builder will attach it only on the outermost parallel loop.
merged_annotations["parallel_loop_layout"] = loop_layout
if prefer_async is not None:
merged_annotations["parallel_prefer_async"] = prefer_async
return _ffi_api.Parallel(extents, merged_annotations) # type: ignore[attr-defined] # pylint: disable=no-member
@@ -233,7 +225,6 @@ def unroll(
step: tirx.PrimExpr | None = None,
*,
explicit: bool = False,
unroll_factor: int | None = None,
annotations: dict[str, Any] | None = None,
) -> frame.ForFrame:
"""The unrolled For statement.
@@ -252,11 +243,10 @@ def unroll(
explicit : bool
Whether to explicitly unroll the loop.
unroll_factor : int
The unroll factor of the loop.
annotations : Dict[str, Any]
The optional annotations of the For statement.
The optional annotations of the For statement. The CUDA dialect
(``tilelang.cuda.language.unroll``) additionally exposes
``unroll_factor`` (``#pragma unroll N``, honored by CUDA codegen only).
Returns
-------
@@ -282,10 +272,7 @@ def unroll(
else:
explicit = annotations.get("pragma_unroll_explicit", False)
if unroll_factor is not None:
annotations["pragma_unroll_factor"] = unroll_factor
else:
unroll_factor = annotations.get("pragma_unroll_factor")
unroll_factor = annotations.get("pragma_unroll_factor")
if explicit and unroll_factor is not None:
raise ValueError("T.unroll's explicit and unroll_factor params are mutually exclusive.")
@@ -317,12 +304,11 @@ def Unroll(
step: tirx.PrimExpr | None = None,
*,
explicit: bool = False,
unroll_factor: int | None = None,
annotations: dict[str, Any] | None = None,
) -> frame.ForFrame:
"""Alias of T.unroll."""
return unroll(start, stop, step, explicit=explicit, unroll_factor=unroll_factor, annotations=annotations)
return unroll(start, stop, step, explicit=explicit, annotations=annotations)
def vectorized(
+35 -25
View File
@@ -427,29 +427,39 @@ def abs2(x: PrimExpr) -> PrimExpr:
return tirx.call_intrin(x.dtype, tirx.op.Op.get("tl.abs2"), x)
__all__ = [
"__log", # noqa: F401
"__log2", # noqa: F401
"__log10", # noqa: F401
"__tan", # noqa: F401
"__cos", # noqa: F401
"__sin", # noqa: F401
"__exp10", # noqa: F401
"__exp", # noqa: F401
"fast_rcp", # noqa: F401
"ieee_add", # noqa: F401
"ieee_sub", # noqa: F401
"ieee_mul", # noqa: F401
"ieee_fmaf", # noqa: F401
"ieee_frcp", # noqa: F401
"ieee_fsqrt", # noqa: F401
"ieee_frsqrt", # noqa: F401
"ieee_fdiv", # noqa: F401
"add2", # noqa: F401
"sub2", # noqa: F401
"mul2", # noqa: F401
"fma2", # noqa: F401
"max2", # noqa: F401
"min2", # noqa: F401
"abs2", # noqa: F401
# Packed x2 element-wise math: registered target-neutrally and lowered on
# both CUDA and ROCm (ROCm currently supports the float32x2 flavors).
COMMON_MATH_INTRINSICS = [
"add2",
"sub2",
"mul2",
"fma2",
"max2",
"min2",
"abs2",
]
# CUDA-only intrinsics: PTX fast-math approximations and IEEE ops with an
# explicit rounding mode. Registered and lowered only by the CUDA backend,
# so only the CUDA dialect exports them.
CUDA_MATH_INTRINSICS = [
"__log",
"__log2",
"__log10",
"__tan",
"__cos",
"__sin",
"__exp10",
"__exp",
"fast_rcp",
"ieee_add",
"ieee_sub",
"ieee_mul",
"ieee_fmaf",
"ieee_frcp",
"ieee_fsqrt",
"ieee_frsqrt",
"ieee_fdiv",
]
__all__ = CUDA_MATH_INTRINSICS + COMMON_MATH_INTRINSICS
+7 -24
View File
@@ -29,7 +29,6 @@ def reduce(
dim: int,
clear: bool,
batch: int = 1,
nan_propagate: bool = False,
annotations: dict | None = None,
) -> None:
"""Perform a reduction operation on a buffer along a specified dimension.
@@ -45,11 +44,10 @@ def reduce(
compiler emits ceil(N/batch) batched AllReduce calls each sharing
a single pair of barriers, reducing total barrier count by batch×.
batch must evenly divide the per-thread output element count N.
nan_propagate (bool): Only meaningful for max/min/absmax on
float16/bfloat16. When True, lower to CUDA __hmax_nan/__hmin_nan so
NaNs propagate through the reduction. When False (default), use
__hmax/__hmin which return the non-NaN operand. CUDA-only.
annotations (dict, optional): Additional lowering controls. On CUDA
annotations (dict, optional): Additional lowering controls. The CUDA
dialect exposes ``nan_propagate`` on reduce_max/min/absmax as a
typed keyword (lowering to __hmax_nan/__hmin_nan); it rides here
as the ``{"nan_propagate": True}`` annotation. On CUDA
SM100+, FP32 sum/abssum reductions accept
``{"enable_fadd2": False}`` to keep the reducer scalar. Packed
FP32x2 reduction remains enabled by default, and can be disabled
@@ -72,8 +70,6 @@ def reduce(
annotations = _normalize_annotations(annotations)
if batch > 1:
annotations["batch"] = batch
if nan_propagate:
annotations["nan_propagate"] = True
# Emit local reductions before macro expansion so alloc_var retains its
# underlying Buffer rather than becoming a scalar expression.
@@ -172,7 +168,6 @@ def reduce_max(
dim: int = -1,
clear: bool = True,
batch: int = 1,
nan_propagate: bool = False,
annotations: dict | None = None,
) -> None:
"""Perform reduce max on input buffer, store the result to output buffer
@@ -189,16 +184,12 @@ def reduce_max(
If set to True, the output buffer will first be initialized to -inf.
batch : int
Number of output elements per batched AllReduce call (default 1).
nan_propagate : bool
For float16/bfloat16 only. When True, NaN inputs propagate through the
reduction (CUDA __hmax_nan). When False (default), NaN inputs are
ignored in favor of the other operand (CUDA __hmax). CUDA-only.
Returns
-------
handle : PrimExpr
"""
dim = _legalize_dim(buffer, dim)
reduce(buffer, out, "max", dim, clear, batch=batch, nan_propagate=nan_propagate, annotations=annotations)
reduce(buffer, out, "max", dim, clear, batch=batch, annotations=annotations)
def reduce_min(
@@ -207,7 +198,6 @@ def reduce_min(
dim: int = -1,
clear: bool = True,
batch: int = 1,
nan_propagate: bool = False,
annotations: dict | None = None,
) -> None:
"""Perform reduce min on input buffer, store the result to output buffer.
@@ -218,15 +208,12 @@ def reduce_min(
dim (int): The dimension to perform reduce on
clear (bool, optional): If True, output buffer will be initialized to inf. Defaults to True.
batch (int): Number of output elements per batched AllReduce call (default 1).
nan_propagate (bool, optional): For float16/bfloat16 only. When True,
NaN inputs propagate (CUDA __hmin_nan). When False (default), NaNs
are ignored (CUDA __hmin). CUDA-only.
Returns:
tirx.Call: Handle to the reduction operation
"""
dim = _legalize_dim(buffer, dim)
reduce(buffer, out, "min", dim, clear, batch=batch, nan_propagate=nan_propagate, annotations=annotations)
reduce(buffer, out, "min", dim, clear, batch=batch, annotations=annotations)
def reduce_sum(
@@ -287,7 +274,6 @@ def reduce_absmax(
dim: int = -1,
clear: bool = True,
batch: int = 1,
nan_propagate: bool = False,
annotations: dict | None = None,
) -> None:
"""Perform reduce absolute max on input buffer, store the result to output buffer.
@@ -297,15 +283,12 @@ def reduce_absmax(
out (tirx.Buffer): The output buffer
dim (int): The dimension to perform reduce on
batch (int): Number of output elements per batched AllReduce call (default 1).
nan_propagate (bool, optional): For float16/bfloat16 only. When True,
NaN inputs propagate (CUDA __hmax_nan). When False (default), NaNs
are ignored. CUDA-only.
Returns:
tirx.Call: Handle to the reduction operation
"""
dim = _legalize_dim(buffer, dim)
reduce(buffer, out, "absmax", dim, clear, batch=batch, nan_propagate=nan_propagate, annotations=annotations)
reduce(buffer, out, "absmax", dim, clear, batch=batch, annotations=annotations)
def reduce_bitand(
+8 -3
View File
@@ -10,7 +10,12 @@ 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, *_KERNEL_ALL)))
# The ROCm dialect's T.gemm shadows the common one: same semantics, plus the
# k_pack knob of the MFMA/WMMA lowering.
from .gemm_op import * # noqa: F401,F403
from .gemm_op import __all__ as _GEMM_ALL
del _COMMON_ALL, _ROCM_ALL, _KERNEL_ALL
__tilelang_dialect__ = "rocm"
__all__ = tuple(dict.fromkeys((*_COMMON_ALL, *_ROCM_ALL, *_KERNEL_ALL, *_GEMM_ALL)))
del _COMMON_ALL, _ROCM_ALL, _KERNEL_ALL, _GEMM_ALL
+63
View File
@@ -0,0 +1,63 @@
"""ROCm dialect of ``T.gemm``: the common GEMM plus ROCm-specific knobs."""
from __future__ import annotations
from tvm import tirx
from tilelang._typing import BufferLikeType
from tilelang.language.gemm_op import GemmWarpPolicy, _gemm_impl
from tilelang.language.utils import _normalize_annotations
__all__ = ["gemm"]
def gemm(
A: BufferLikeType,
B: BufferLikeType,
C: BufferLikeType,
transpose_A: bool = False,
transpose_B: bool = False,
policy: GemmWarpPolicy = GemmWarpPolicy.Square,
clear_accum: bool = False,
k_pack: int = 1,
annotations: dict | None = None,
) -> tirx.PrimExpr:
"""TileLang GEMM operator for ROCm.
Same semantics as the common :func:`tilelang.language.gemm_op.gemm`.
``k_pack`` packs multiple matrix-core operations along K in the MFMA/WMMA
lowering (CDNA/RDNA); it is a performance knob with no counterpart on
other targets.
Args:
A (BufferLikeType, i.e. Buffer | BufferLoad | BufferRegion, or Var): Input buffer A.
B (BufferLikeType): Input buffer B.
C (BufferLikeType): Output buffer C.
transpose_A (bool): Whether to transpose A. Defaults to False.
transpose_B (bool): Whether to transpose B. Defaults to False.
policy (GemmWarpPolicy): GEMM warp partition policy.
clear_accum (bool): Whether to clear the accumulator.
k_pack (int): Number of packed matrix cores along K. Must be 1 or 2.
Defaults to 1.
annotations (Optional[dict]): Additional annotations.
Returns:
tirx.Call: A handle to the GEMM operation.
"""
if not (isinstance(k_pack, int) and not isinstance(k_pack, bool) and k_pack in (1, 2)):
raise ValueError(f"T.gemm k_pack must be an int equal to 1 or 2, got {k_pack!r}")
ann = _normalize_annotations(annotations)
if k_pack != 1:
ann.setdefault("k_pack", k_pack)
return _gemm_impl(
"tl.tileop.gemm",
A,
B,
C,
transpose_A,
transpose_B,
policy,
clear_accum,
None,
annotations=ann,
)
+1 -17
View File
@@ -33,7 +33,7 @@ def gemm_lower(
class Gemm(Node, Scriptable):
# FFI fields (LLVM/MLIR-style lowerCamel via reflection):
# a, b, c, aPtr, bPtr, cPtr, m, n, k, transA, transB,
# strideA, strideB, offsetA, offsetB, clearAccum, kPack, wgWait, policy
# clearAccum, kPack, wgWait, policy
#
# Backward-compat alias properties are provided below to support old names.
@@ -82,22 +82,6 @@ class Gemm(Node, Scriptable):
def trans_B(self):
return self.transB
@property
def stride_A(self):
return self.strideA
@property
def stride_B(self):
return self.strideB
@property
def offset_A(self):
return self.offsetA
@property
def offset_B(self):
return self.offsetB
@property
def clear_accum(self):
return self.clearAccum
-16
View File
@@ -132,22 +132,6 @@ class GemmBase:
def CRegion(self):
return getattr(self.gemm_node, "cRegion", None)
@property
def stride_A(self) -> int:
return getattr(self.gemm_node, "strideA", None)
@property
def stride_B(self) -> int:
return getattr(self.gemm_node, "strideB", None)
@property
def offset_A(self) -> int:
return getattr(self.gemm_node, "offsetA", None)
@property
def offset_B(self) -> int:
return getattr(self.gemm_node, "offsetB", None)
@property
def clear_accum(self) -> PrimExpr:
return getattr(self.gemm_node, "clearAccum", None)
-9
View File
@@ -30,19 +30,10 @@ class GemmSP(Node, Scriptable):
trans_B: bool
trans_E: bool
stride_A: int
stride_B: int
offset_A: int
offset_B: int
clear_accum: bool
kPack: int
wg_wait: int
policy: GemmSPWarpPolicy
@property
def k_pack(self):
return self.kPack
@tvm_ffi.register_global_func("tl.gemm_sp.infer_layout")
def gemm_sp_infer_layout(self, target: Target, thread_bounds: Range):
thread_nums = thread_bounds.extent
-20
View File
@@ -104,30 +104,10 @@ class GemmSPBase:
def CRegion(self) -> tirx.PrimExpr:
return self.gemm_sp_node.cRegion
@property
def stride_A(self) -> int:
return self.gemm_sp_node.stride_A
@property
def stride_B(self) -> int:
return self.gemm_sp_node.stride_B
@property
def offset_A(self) -> int:
return self.gemm_sp_node.offset_A
@property
def offset_B(self) -> int:
return self.gemm_sp_node.offset_B
@property
def clear_accum(self) -> bool:
return self.gemm_sp_node.clear_accum
@property
def k_pack(self) -> int:
return self.gemm_sp_node.k_pack
@property
def wg_wait(self) -> int:
return self.gemm_sp_node.wg_wait