| name | flydsl-kernel-code-cleanup |
| description | Clean up aiter FlyDSL kernels and shared helpers while preserving numerical behavior and performance. Use for raw-IR migrations, helper deduplication, dead-code removal, and reviews of those changes against the pinned FlyDSL API.
|
| allowed-tools | Read Edit Bash Grep Glob Agent |
FlyDSL Kernel Code Cleanup
Maps legacy kernel constructs to the fx.* surface available in aiter's pinned
FlyDSL version. See the in-tree flydsl-kernel-authoring skill for API guidance.
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.
Before applying the recipes, check requirements.txt and the imported FlyDSL
version and module path. A sibling FlyDSL checkout can expose newer APIs than
aiter supports. Kernel cleanup does not require a dependency bump or edits to
FlyDSL itself; honor the task's explicit version, architecture and single-/multi-
GPU scope.
Aiter layout
| Location | Role |
|---|
aiter/ops/flydsl/kernels/ | @flyc.kernel device kernels and shared helpers (tensor_shim.py, kernels_common.py, …) |
aiter/ops/flydsl/*.py | Launch wrappers, compile helpers, public op entry points |
op_tests/test_flydsl_*.py | Top-level FlyDSL correctness / perf tests |
op_tests/flydsl_tests/ | Additional FlyDSL kernel tests |
Prefer from aiter.ops.flydsl.kernels.tensor_shim import _run_compiled, ptr_arg.
Reuse tensor_shim.py for compilation, cached dispatch and failure recovery.
moe_kernels._run_compiled(exe, args) is a tuple-argument adapter for existing
MoE/AOT callers; preserve that contract when consolidating launch code.
Reuse existing owners before adding helpers: act.py for activations;
tensor_shim.py for pointer/base/dtype extraction;
kernels_common.py for LOG2E and host or wrapping integer ceildiv; family common modules
for reductions and specialized memory operations. Prefer typed fx.min/fx.max
and fx.ceildiv where their signedness, NaN and overflow semantics match.
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. Compare before/after
numerics and generated code with
FLYDSL_RUNTIME_ENABLE_CACHE=0; use a fresh
dump directory per specialization so one shape cannot overwrite another.
- Raw boundaries are semantic. Retain exact scope/ordering, volatile and
alias metadata, raw SSA contracts, or integer widths unsupported by the pinned
API. Record the specific reason; do not merely hide a dialect in a new facade.
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 |
| index-typed arithmetic counters | fx.Int64(...) or fx.Int32(...) when the consumer permits a fixed-width integer |
arith.index_cast(T.index, v) at an index-typed boundary | fx.Index(v) |
fx.Index maps to MLIR index. Prefer explicit-width fx.Int64/fx.Int32 for
arithmetic, choosing width and signedness deliberately. Keep fx.Index where a
launch, layout, loop or other API requires the index type; replacing it merely
to remove the name can change the IR contract. Do not widen i32 counters or
narrow an index without checking the consumer and supported bounds.
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.Index(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.
- aiter still ships
aiter/ops/flydsl/kernels/buffer_ops.py for legacy kernels;
migrate off it when touching load/store paths.
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.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) |
Check the rewriter's loop contract in the pinned version. In FlyDSL 0.3.2:
range_constexpr requests Python unrolling.
range(..., init=[...]) emits an scf.for with explicit carried state and
converts its bounds to index, including Python integer bounds. It does not
discard init merely because the bounds are static.
- Ordinary
range without init uses automatic carried-state dispatch and
expects i32 bounds; index bounds are converted to i32, and i64 is rejected.
Preserve the selected loop form, supported bounds and state types. See §5 for
runtime branches inside helper functions.
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. Use existing
intrinsic wrappers where equivalent; keep unsupported boundaries local to
aiter rather than extending FlyDSL as part of a cleanup.
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.
- Preserve byte versus element GEPs and alignment provenance. An equal numeric
address alone does not guarantee equal memory instructions; compare the
generated loads and stores when replacing an epilog pointer path.
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, expcnt=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.
- Use the public
fx.rocdl.sched_barrier / fx.rocdl.sched_group_barrier
wrappers when exposed by the pinned version. The legacy wait form remains
available as positional fx.rocdl.s_waitcnt(bitfield) for a boundary the
keyword form cannot 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")