Skip to main content

flydsl-kernel-code-cleanup

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.

설치로 이동

소스 정보

저장소
ROCm/aiter
최근 소스 활동
2026년 9월 15일 01:12
감지된 SKILL.md 언어
영어
스타
562
포크
572

설치 방법

기본적으로 소스를 먼저 확인하는 Prompt가 선택됩니다. 직접 명령으로 전환하거나 로컬 사본을 다운로드할 수도 있습니다.

소스 파일 검토

설치 여부를 결정하기 전에 SKILL.md와 SkillsMP에 표시된 보조 파일을 읽어 보세요.

SKILL.md 표시 중

SKILL.md
소스 지침 · 읽기 전용 미리보기
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. ```python # Before acc = ArithValue(val) + peer lane = ArithValue(tid) % fx.Index(64) cond = arith.unwrap(idx >= limit) off = arith.index_cast(T.index, x) # After acc = val + peer # val already fx.Float32 / fx.Vector lane = tid % fx.Int64(64) cond = (idx >= limit).ir_value() # only if a raw scf.IfOp needs it off = fx.Index(x) # preserve this consumer's index contract ``` 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. ```python # Before 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) # tx already fx.Int32 # After for i in range_constexpr(N): off = base + fx.Int64(4) tile = fx.make_layout(BLOCK, 1) # builders take Python ints 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. ```python # Before (manual offsets — see PA //4 offset bugs) 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) # After 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) # after partitioning tA (§7b: prefer fx.copy) ``` - `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`. ```python # Before (hardcoded address space) p = buffer_ops.create_llvm_ptr(lds_addr, address_space=3) p = mem_ops._create_llvm_ptr(val, address_space=1) # a.k.a. mem_ops.to_llvm_ptr # After p = ptr.llvm_ptr # property on an fx pointer p = fx.to_llvm_ptr(ptr) # equivalent free function; backend resolves the AS ``` - 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. ```python # Before rocdl.s_waitcnt(_encode_waitcnt(lgkmcnt=0)) # per-kernel encoder rocdl.s_waitcnt(0) # raw "wait for everything" _s_waitcnt(0xC07F) # magic LGKMCNT_0_ONLY bitfield # After fx.rocdl.s_waitcnt(lgkmcnt=0) # wait for LDS/SMEM only fx.rocdl.s_waitcnt(vmcnt=0, lgkmcnt=0, expcnt=0) # matches raw s_waitcnt(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**. ```python # Before allocator = SmemAllocator(None, arch=GPU_ARCH, global_sym_name="smem")
GitHub에서 보기
이 SKILL.md는 매우 커서 SkillsMP가 여기에는 첫 섹션만 미리 보여줍니다. GitHub에서 보기