mirror of
https://github.com/tile-ai/tilelang.git
synced 2026-10-02 06:34:36 +08:00
* [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.