mirror of
https://github.com/tile-ai/tilelang.git
synced 2026-10-04 07:18:17 +08:00
[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:
@@ -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)`.
|
||||
|
||||
@@ -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
|
||||
|
||||
@@ -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
|
||||
|
||||
@@ -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
@@ -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
@@ -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
@@ -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);
|
||||
|
||||
@@ -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
@@ -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>();
|
||||
|
||||
@@ -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
@@ -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));
|
||||
|
||||
|
||||
@@ -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,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()
|
||||
@@ -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",
|
||||
|
||||
@@ -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,
|
||||
)
|
||||
@@ -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)
|
||||
@@ -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,
|
||||
)
|
||||
@@ -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)
|
||||
@@ -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))
|
||||
@@ -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"
|
||||
|
||||
@@ -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]
|
||||
|
||||
|
||||
@@ -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
@@ -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",
|
||||
|
||||
@@ -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,
|
||||
|
||||
@@ -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,
|
||||
)
|
||||
|
||||
@@ -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],
|
||||
|
||||
@@ -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(
|
||||
|
||||
@@ -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
|
||||
|
||||
@@ -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(
|
||||
|
||||
@@ -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
|
||||
|
||||
@@ -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,
|
||||
)
|
||||
@@ -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
|
||||
|
||||
@@ -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)
|
||||
|
||||
@@ -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
|
||||
|
||||
@@ -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
|
||||
|
||||
Reference in New Issue
Block a user