| name | kernel-code-cleanup |
| description | Modernize FlyDSL kernels: replace raw MLIR dialects (arith, scf, vector, llvm, memref, math), ArithValue, redundant fx.* wrapping, fx.Index, buffer_ops, SmemPtr/SmemAllocator, copy_atom_call/mma_atom_call (loop or single atom), and raw rocdl.mfma_* with the current fx.* surface (fx types, Python control flow, make_buffer_tensor, SharedAllocator, fx.copy/fx.gemm, make_layout_tv/make_tiled_copy TV layouts, to_llvm_ptr, arch-dispatched fx.rocdl.s_waitcnt, local @flyc.jit if/else). Also trims comments and dead code and applies the _run_compiled fast launch path. Use when reviewing, cleaning, or migrating existing kernels.
|
| allowed-tools | Read Edit Bash Grep Glob Agent |
FlyDSL Kernel Code Cleanup
Maps legacy/deprecated kernel constructs to the current fx.* surface. Companion
to flydsl-kernel-authoring (API reference) and flydsl-tile-programming
(authoring wizard).
Golden rule: in @flyc.kernel / @flyc.jit bodies, use fx.* and Python
operators first. Drop to a raw dialect only at a hard boundary with no wrapper,
and localize it.
Cautions
- Surgical, behavior-preserving. Migration is a refactor: minimal diffs, match
local style.
- Don't mass-rewrite heavily-legacy kernels unless asked (e.g.
flash_attn_gfx950.py uses _scf.IfOp/_raw pervasively). Clean what the task
touches.
- Verify. Offset/type/SSA changes can shift results — run the kernel test
before/after; clear cache with
FLYDSL_RUNTIME_ENABLE_CACHE=0 if unsure.
expr/ stays target-neutral: no rocdl/llvm/buffer imports in
python/flydsl/expr/ top-level (guarded by test_expr_optional_rocdl.py).
1. ArithValue and index helpers (deprecated in expr/arith.py)
| Deprecated | Replacement |
|---|
ArithValue(x) (wrap for operators) | fx.Int32/Int64/Float32/Vector — already overload + - * / % << >> == < > |
arith.unwrap(v) / arith._to_raw(v) | v.ir_value(), only where a raw ir.Value is needed |
arith.index(n) / arith.index_cast(T.index, v) / fx.Index(n) | fx.Int64(...) (or fx.Int32(...)) |
fx.Index maps to MLIR index — platform-defined width, ambiguous, and forces
implicit casts. Prefer explicit-width fx.Int64/fx.Int32; pick the width on
purpose (don't widen counters that must stay i32).
acc = ArithValue(val) + peer
lane = ArithValue(tid) % fx.Index(64)
cond = arith.unwrap(idx >= limit)
off = arith.index_cast(T.index, x)
acc = val + peer
lane = tid % fx.Int64(64)
cond = (idx >= limit).ir_value()
off = fx.Int64(x)
If an operand is a raw ir.Value, wrap it once at the source (fx.Float32(v)),
not with ArithValue per use. Keep an explicit arith.*FOp only for non-default
fastmath.
1b. Drop redundant fx.* wraps
Wrap only to introduce a type (Python literal / raw ir.Value) or change one.
Re-wrapping an already-typed value is noise; double-wrapping is dead.
for i in range_constexpr(fx.Int32(N)):
off = fx.Int64(fx.Int64(base) + fx.Int64(4))
tile = fx.make_layout(fx.Int32(BLOCK), fx.Int32(1))
idx = fx.Int32(tx)
for i in range_constexpr(N):
off = base + fx.Int64(4)
tile = fx.make_layout(BLOCK, 1)
idx = tx
- Compile-time shapes/strides/bounds (
make_layout, make_shape,
range_constexpr, Constexpr) take plain Python ints.
- Wrap a runtime value once, at first typed use.
- A real cast (
fx.Int64(i32) widen, fx.Int32(index) narrow) is not redundant —
it replaces arith.index_cast.
2. buffer_ops → make_buffer_tensor + copy atoms
create_buffer_resource + manual offsets is legacy. Build a buffer-resource view
with fx.rocdl.make_buffer_tensor(), then use layout ops + fx.copy (§7b);
the OOB-checked V# descriptor is built for you.
rsrc = buffer_ops.create_buffer_resource(A, max_size=True)
data = buffer_ops.buffer_load(rsrc, row * K + k, vec_width=4, dtype=fx.Float32)
buffer_ops.buffer_store(data, rsrc, row * N + col)
bufA = fx.rocdl.make_buffer_tensor(A)
tA = fx.make_view(fx.get_iter(bufA), fx.make_layout((M, K), (K, 1)))
copy = fx.make_copy_atom(fx.rocdl.BufferCopy128b(), fx.Float32)
fx.copy(copy, fx.slice(tA, (None, tid)), rA)
make_buffer_tensor(tensor, max_size=True) mirrors create_buffer_resource;
pass num_records_bytes= for a const byte count, or max_size=False to derive
from the layout.
- gfx1250 TDM uses a different atom —
fx.rocdl.make_tdm_atom (raw VA, not a
buffer resource).
- A scalar-base + per-thread-offset load with no layout form may stay on
buffer_ops — note it. buffer_load/store offset is in elements (×
sizeof(dtype) internally) — a classic bug.
3. Raw upstream dialects → fx.* and Python
arith
| Raw | Preferred |
|---|
arith.constant(42, index=True) | fx.Int64(42) |
arith.mulf/addf(a,b) | a * b / a + b |
arith.trunc_f(ty, v) / ext_f | v.to(fx.BFloat16) |
arith.index_cast(T.i32, v) | fx.Int32(v) |
arith.select(cond, t, f) | cond.select(t, f) |
arith.cmpi(slt, a, b) | a < b |
arith.maximumf/minimumf(a,b) | fx.max(a, b) / fx.min(a, b) |
arith.maxsi/maxui/minsi/minui(a,b) | fx.max(a, b) / fx.min(a, b) |
arith.maxnumf(a,b) | fx.maxnumf(a, b) — different NaN semantics from fx.max |
arith.minnumf(a,b) | fx.minnumf(a, b) — different NaN semantics from fx.min |
arith.ceildivsi/ceildivui(a,b) | fx.ceildiv(a, b) |
Keep arith.cmpf / explicit *FOp only where no operator exists or fastmath is
needed.
scf
| Raw | Preferred |
|---|
scf.ForOp | range_constexpr(N) (unrolled) or range(lo, hi, step, init=[...]) (runtime, loop-carried) |
scf.IfOp(_raw(cond)) | Python if cond: (runtime) / if const_expr(flag): (compile-time) |
Runtime bounds must be typed (fx.Int64) or the rewriter unrolls and drops
init=. See §5 for branches the rewriter can't express.
vector
| Raw | Preferred |
|---|
vector.extract(v, static_position=[i]) | fx.Vector(v)[i] |
vector.bitcast(ty, v) | fx.Vector(v).bitcast(fx.Float32) |
vector.splat / const vector | fx.Vector.filled(width, val, fx.Float32) |
| build from scalars | fx.Vector.from_elements(...) |
| reg-memref load/store | fx.memref_load_vec(r) / fx.memref_store_vec(v, r) |
llvm / memref / math
llvm.* ptr math / load/store / const → layout views (fx.make_view,
fx.get_iter), fx.Array + SharedAllocator, fx constants. Real intrinsics
go in expr/rocdl/inline_asm.py or rocdl wrappers.
memref.* → layout tensors/views + copy atoms.
math.* → fx math helpers (expr/math.py); keep math_dialect.fma etc. only
where no wrapper exists.
3b. fly.ptr → !llvm.ptr (backend-resolved address space)
When you hold an fx pointer (fly.ptr) and need a raw !llvm.ptr at a hard
boundary, use the DSL primitive — it maps the pointer's semantic address space to
the backend's LLVM address-space number for you. Don't hand-build one with a
hardcoded <1> / <3> via IntToPtrOp.
p = buffer_ops.create_llvm_ptr(lds_addr, address_space=3)
p = mem_ops._create_llvm_ptr(val, address_space=1)
p = ptr.llvm_ptr
p = fx.to_llvm_ptr(ptr)
- Applies only when you already have a
fly.ptr. A raw int/index address (e.g. an
LDS byte offset with no pointer form) still needs manual construction — note it.
mem_ops.get_llvm_ptr / element_ptr also fold in + offset*dtype_bytes
arithmetic; keep the offset math (layout views / get_element_ptr) and only swap
the final ptr cast for .llvm_ptr.
3c. Manual s_waitcnt bitfields → fx.rocdl.s_waitcnt(vmcnt=/lgkmcnt=/expcnt=)
Hand-encoding a wait-counter bitfield (or calling rocdl.s_waitcnt(magic) with a
raw number) is arch-fragile — the field widths differ per arch (CDNA3 lgkmcnt
max 15 vs RDNA 63). The keyword form of fx.rocdl.s_waitcnt
(expr/rocdl/universal.py) is arch-dispatched across gfx942/gfx950/gfx11xx/gfx120x
and packs the correct bitfield for you.
rocdl.s_waitcnt(_encode_waitcnt(lgkmcnt=0))
rocdl.s_waitcnt(0)
_s_waitcnt(0xC07F)
fx.rocdl.s_waitcnt(lgkmcnt=0)
fx.rocdl.s_waitcnt(vmcnt=0, lgkmcnt=0)
fx.rocdl.s_waitcnt(lgkmcnt=0)
- Unset fields default to "no wait" (their per-arch max) — name only the counters
you need.
- Delete the now-unused per-kernel
_encode_waitcnt / _s_waitcnt shims and magic
*CNT_* constants once your changes make them dead.
sched_barrier / sched_group_barrier have no keyword fx wrapper — keep raw
rocdl.sched_barrier(...). The legacy raw form stays available as positional
fx.rocdl.s_waitcnt(bitfield) for a boundary the keyword form can't express;
localize it.
- Scheduler-sensitive.
s_waitcnt placement drives hot-loop pipelining in
tuned attention/GEMM kernels — an op-identical swap can still shift the schedule.
Verify perf (median-based), not just correctness, and don't mass-migrate
pervasively-tuned kernels (e.g. flash_attn_gfx950.py, mla_fwd_decode_*).
4. SmemAllocator / SmemPtr → SharedAllocator
Legacy LDS path uses a manual base pointer, byte offsets, and finalize(). New
kernels declare an @fx.struct of fx.Array fields and allocate via
fx.SharedAllocator — the compiler sizes the LDS global; no finalize.
allocator = SmemAllocator(None, arch=GPU_ARCH, global_sym_name="smem")
base = allocator.get_base()
smem_a = SmemPtr(base, 0, dtype_, shape=(BLOCK_M * BLOCK_K,))
smem_b = SmemPtr(base, a_bytes, dtype_, shape=(BLOCK_K * BLOCK_N,))
allocator.finalize()
@fx.struct
class SharedStorage:
a: fx.Array[fx.Float16, BLOCK_M * BLOCK_K]
b: fx.Array[fx.Float16, BLOCK_K * BLOCK_N]
lds = fx.SharedAllocator().allocate(SharedStorage).peek()
lds_a = lds.a.view(fx.make_layout((BLOCK_M, BLOCK_K), (BLOCK_K, 1)))
lds_b = lds.b.view(fx.make_layout((BLOCK_K, BLOCK_N), (BLOCK_N, 1)))
- Default
static=True leaves launch(smem=...) unset; only static=False
auto-infers smem from allocated_bytes.
SmemPtr.get() caches its view — reusing it in an epilogue after a scf.for
causes a dominance error. SharedAllocator avoids this (view taken per use); for
legacy code, clear ptr._view_cache = None.
- Structural change — migrate a kernel's whole LDS at once and re-run its test.
5. Runtime if/else with side effects → local @flyc.jit