Repository navigation
[HIP] [OPUS] [JIT] [GFX950]OPUS PA MQA Logits MXFP4 - #5332
Conversation
🏷️ CI GuideRuns automatically on every PR:
Extended tests (opt-in via labels):
PR title tags & labels: |
There was a problem hiding this comment.
Remaining comments which cannot be posted as a review comment to avoid GitHub Rate Limit
ruff
Function definition does not bind loop variable inp
aiter/op_tests/test_pa_mqa_logits_opus.py
Line 565 in 1135f24
Undefined name inp
aiter/op_tests/test_pa_mqa_logits_opus.py
Line 565 in 1135f24
Function definition does not bind loop variable inp
aiter/op_tests/test_pa_mqa_logits_opus.py
Line 566 in 1135f24
Undefined name inp
aiter/op_tests/test_pa_mqa_logits_opus.py
Line 566 in 1135f24
Function definition does not bind loop variable le
aiter/op_tests/test_pa_mqa_logits_opus.py
Line 566 in 1135f24
Undefined name inp
aiter/op_tests/test_pa_mqa_logits_opus.py
Line 566 in 1135f24
Function definition does not bind loop variable inp
aiter/op_tests/test_pa_mqa_logits_opus.py
Line 566 in 1135f24
Function definition does not bind loop variable next_n
aiter/op_tests/test_pa_mqa_logits_opus.py
Line 566 in 1135f24
Undefined name inp
aiter/op_tests/test_pa_mqa_logits_opus.py
Line 567 in 1135f24
Function definition does not bind loop variable inp
aiter/op_tests/test_pa_mqa_logits_opus.py
Line 567 in 1135f24
Function definition does not bind loop variable block_k
aiter/op_tests/test_pa_mqa_logits_opus.py
Line 568 in 1135f24
Undefined name out
aiter/op_tests/test_pa_mqa_logits_opus.py
Line 568 in 1135f24
Function definition does not bind loop variable out
aiter/op_tests/test_pa_mqa_logits_opus.py
Line 568 in 1135f24
There was a problem hiding this comment.
🟡 Changes recommended
The new test and wrapper/launcher code has concrete shape/contiguity validation issues that can cause runtime errors or silent misbehavior.
Once you've addressed the issues Copilot identified, you can request another Copilot review.
Pull request overview
Adds a new gfx950-only OPUS (HIP) implementation of an MXFP4 paged MQA-logits kernel, along with Python wrappers, pybind/JIT build wiring, and targeted correctness/perf tests for both prefill and decode entry points.
Changes:
- Introduces the device kernel (
pa_mqa_logits_mxfp4_opus.h) plus HIP launchers and a device-side per-row window builder. - Adds a Python wrapper module (
aiter.ops.opus.pa_mqa_logits_opus) and exports it viaaiter.ops.opus. - Adds new OP tests for correctness/perf and for detecting a known store-alignment failure mode.
File summaries
| File | Description |
|---|---|
| op_tests/test_pa_mqa_logits_opus.py | New correctness + perf harness comparing against dequantized reference (and FlyDSL when available). |
| op_tests/test_fp4_store_alignment.py | New focused repro for potential window-base store alignment issues via local_start sweep. |
| csrc/pybind/pa_mqa_logits_mxfp4_pybind.cu | New pybind module entry for the MXFP4 logits op. |
| csrc/py_itfs_cu/pa_mqa_logits_mxfp4_kernels.cu | New HIP host launchers and device-side window build kernel. |
| csrc/include/rocm_ops.hpp | Adds pybind macro bindings for the new entry points. |
| csrc/include/pa_mqa_logits_mxfp4_opus.h | New OPUS kernel + ABI/traits definitions and device implementation. |
| aiter/ops/opus/pa_mqa_logits_opus.py | New Python API + JIT stubs for prefill/decode/window-builder. |
| aiter/ops/opus/init.py | Exposes the new OPUS MXFP4 logits APIs (and stubs on unsupported arch). |
| aiter/jit/optCompilerConfig.json | Registers the new JIT module build configuration and sources. |
Review details
Suppressed comments (1)
aiter/ops/opus/pa_mqa_logits_opus.py:333
pa_mqa_logits_mxfp4_decode()accepts a caller-providedoutbut does not validate its shape/dtype/device against(total_q, max_seq_len)andtorch.float32. The C++ side checks dtype/dims, but a mismatched width can still lead to out-of-bounds accesses whenlocal_endsexceeds the actual allocation width.
if out is None:
out = torch.full(
(total_q, max_seq_len),
float("-inf"),
dtype=torch.float32,
device=q_fp4.device,
)
pa_mqa_logits_mxfp4_fwd_decode(
q_fp4,
q_scale,
kv_cache,
kv_scale,
block_tables,
weights,
cu_seq_q,
local_ends.to(torch.int32).contiguous(),
out,
int(batch),
int(next_n_max),
int(split_kv),
float(weight_scale),
block_k,
int(kv_block_size),
int(max_seq_len),
)
- Files reviewed: 9/9 changed files
- Comments generated: 4
- Review effort level: Lite
💡 Add a code-review agent skill or configure MCP servers for context-aware, tailored reviews. Learn more in the docs.
efe944c to
1135f24
Compare
There was a problem hiding this comment.
Remaining comments which cannot be posted as a review comment to avoid GitHub Rate Limit
ruff
Function definition does not bind loop variable inp
aiter/op_tests/test_pa_mqa_logits_opus.py
Line 565 in 1135f24
Undefined name inp
aiter/op_tests/test_pa_mqa_logits_opus.py
Line 565 in 1135f24
Function definition does not bind loop variable inp
aiter/op_tests/test_pa_mqa_logits_opus.py
Line 566 in 1135f24
Undefined name inp
aiter/op_tests/test_pa_mqa_logits_opus.py
Line 566 in 1135f24
Function definition does not bind loop variable le
aiter/op_tests/test_pa_mqa_logits_opus.py
Line 566 in 1135f24
Undefined name inp
aiter/op_tests/test_pa_mqa_logits_opus.py
Line 566 in 1135f24
Function definition does not bind loop variable inp
aiter/op_tests/test_pa_mqa_logits_opus.py
Line 566 in 1135f24
Function definition does not bind loop variable next_n
aiter/op_tests/test_pa_mqa_logits_opus.py
Line 566 in 1135f24
Undefined name inp
aiter/op_tests/test_pa_mqa_logits_opus.py
Line 567 in 1135f24
Function definition does not bind loop variable inp
aiter/op_tests/test_pa_mqa_logits_opus.py
Line 567 in 1135f24
Function definition does not bind loop variable block_k
aiter/op_tests/test_pa_mqa_logits_opus.py
Line 568 in 1135f24
Undefined name out
aiter/op_tests/test_pa_mqa_logits_opus.py
Line 568 in 1135f24
Function definition does not bind loop variable out
aiter/op_tests/test_pa_mqa_logits_opus.py
Line 568 in 1135f24
There was a problem hiding this comment.
🟡 Changes recommended
There is a confirmed test bug in max_err() and the kernel currently trusts device-provided window bounds without clamping, which can lead to invalid pointers/table OOB behavior for malformed local_start/local_end inputs.
Once you've addressed the issues Copilot identified, you can request another Copilot review.
Review details
Suppressed comments (1)
op_tests/test_pa_mqa_logits_opus.py:263
max_err()computes a relative error but the denominator is a 0-dim tensor (vals.abs().max().clamp(...)), so the division produces a tensor andmax(err, ...)will raise (or behave unexpectedly) when comparing a float with a tensor. Convert the denominator to a Python float (or call.item()on the final ratio) before feeding it tomax().
err = max(err, (got - vals).abs().max().item() / vals.abs().max().clamp(min=1e-6))
- Files reviewed: 9/9 changed files
- Comments generated: 4
- Review effort level: Lite
There was a problem hiding this comment.
🟡 Changes recommended
Several launch/window-build paths lack required contiguity/size/device validations (and one Python path can divide by zero), which can lead to undefined behavior or runtime errors.
Once you've addressed the issues Copilot identified, you can request another Copilot review.
Review details
Suppressed comments (5)
csrc/py_itfs_cu/pa_mqa_logits_mxfp4_kernels.cu:7
- The file header comment says decode derives windows in-kernel, but the decode path actually reads per-row window information from the
local_endsinput (seekargs.ptr_local_endsand the kernel indexingp_local_ends[row_id]). This comment is misleading for users trying to understand the ABI.
// the input ABI. Prefill launches a 1D grid with per-row window arrays, decode a 3D grid
// with the windows derived in-kernel; both are schedule-free and cudagraph-safe.
csrc/py_itfs_cu/pa_mqa_logits_mxfp4_kernels.cu:90
pa_mqa_logits_mxfp4_check_shapesonly enforces last-dim contiguity, but the kernel uses fixed pointer arithmetic for q/weights/q_scale/kv_* and ignores the runtime strides. Non-contiguous tensors (or tensors with padded strides) can silently produce wrong results or OOB reads.
AITER_CHECK(out.dtype() == AITER_DTYPE_fp32, "out must be fp32");
AITER_CHECK(q.stride(2) == 1 && weights.stride(1) == 1 && out.stride(1) == 1,
"q / weights / out must be contiguous along their last dim");
csrc/py_itfs_cu/pa_mqa_logits_mxfp4_kernels.cu:198
- The decode launch validates
grid_xandnext_n_maxagainst the 65535 grid-dimension limit but doesn’t validatesplit_kv(grid.z). Ifsplit_kvis ever increased beyond 65535, the launch becomes invalid.
AITER_CHECK(grid_x <= 65535 && next_n_max <= 65535,
"decode launch: padded batch / next_n_max exceed grid.x/.y limit (65535)");
csrc/py_itfs_cu/pa_mqa_logits_mxfp4_kernels.cu:360
pa_mqa_logits_mxfp4_prefill_windowsdoesn’t validate that output buffers (row_to_batch,local_starts,local_ends) are int32, contiguous, on the same device, and sized fortotal_q. The kernel writestotal_qelements unconditionally, so a mismatch can cause OOB writes.
const int B = static_cast<int>(context_lens.size(0));
AITER_CHECK(cu_seq_q.dtype() == AITER_DTYPE_i32 && context_lens.dtype() == AITER_DTYPE_i32,
"cu_seq_q / context_lens must be int32");
AITER_CHECK(cu_seq_q.size(0) == B + 1, "cu_seq_q must have length B+1");
aiter/ops/opus/pa_mqa_logits_opus.py:152
compute_prefill_windowscan be called on non-gfx950 devices (even thoughaiter.ops.opusis importable on other supported arches) and accepts anouttuple without validating dtype/device/contiguity/size. That can lead to confusing JIT failures on other arches and potential OOB writes if undersized output tensors are passed through.
dev = cu_seq_q.device
cu = cu_seq_q.to(torch.int32).contiguous()
ctx = context_lens.to(torch.int32).contiguous()
if out is None:
- Files reviewed: 8/8 changed files
- Comments generated: 3
- Review effort level: Lite
There was a problem hiding this comment.
🟡 Changes recommended
The decode launcher enforces an incorrect grid.x limit (65535), which can reject otherwise valid launches and is inconsistent with established grid.x bounds used elsewhere in the repo.
Once you've addressed the issues Copilot identified, you can request another Copilot review.
Review details
- Files reviewed: 8/8 changed files
- Comments generated: 2
- Review effort level: Lite
There was a problem hiding this comment.
🟡 Changes recommended
Critical correctness, memory-safety, validation, and architecture-gating issues remain unresolved.
Once you've addressed the issues Copilot identified, you can request another Copilot review.
Review details
Suppressed comments (2)
csrc/py_itfs_cu/pa_mqa_logits_mxfp4_kernels.cu:74
max_seq_lenis passed in the launch arguments but never used by the kernel; the bounded store is sized fromlocal_end - local_start. Thus a caller-provided window such as[0, 1025)for an output with width 1024 writes past the row, and a negative start can moveout_basebefore the row. The public wrapper documents only0 <= start <= end, so this is an unchecked memory-safety precondition. Author must reject or device-side guard every window against the output extent before issuing stores.
// The kernel bounds its store by the WINDOW, not by max_seq_len -- `max_seq_len` is
// carried in kargs and never read, only `stride_out_row` is. So these two are the only
// place an undersized `out` can be caught at all; a `local_ends` entry past out.size(1)
// is the caller's contract (see the wrapper docstring).
op_tests/test_pa_mqa_logits_opus.py:354
- All value assertions are fed only
min(8, len(nonempty))randomly selected rows.window_is_writtenandoob_is_neginfcheck only finite/-inf state, so a wrong logit in any unselected row—especially in the 16K-row perf cases—still reports a pass. Keep sampling for the expensive perf sweep, but make the small correctness cases compare every row or add a deterministic all-row check.
def sample_rows(total, le, n=N_COS_SAMPLE, seed=0):
nonempty = torch.nonzero(le > 0).flatten().tolist()
if not nonempty:
return []
rng = random.Random(seed)
- Files reviewed: 8/8 changed files
- Comments generated: 9
- Review effort level: Lite
There was a problem hiding this comment.
🟡 Changes recommended
Unresolved API, bounds-validation, architecture-guard, mutation-metadata, and device-validation issues block approval.
Get a fresh assessment by requesting another Copilot review.
Review details
Suppressed comments (3)
aiter/ops/opus/init.py:55
- [verified]
_arch_okis true for gfx942 and gfx1250, so this import exposespa_mqa_logits_mxfp4_build_schedon those GPUs, but onlypa_mqa_logits_mxfp4_schedcalls_require_gfx950. Calling the builder on a non-gfx950 device therefore enters the gfx950-only JIT/module path instead of the documented unsupported-architecture stub, which can fail at first compilation or launch. Author must apply the exact-arch guard to the builder as well, or install an unsupported stub for it.
# gfx950-only, unlike the gemm entries above: _SUPPORTED is wider than this
# kernel, so the wrappers re-check at call time rather than at import.
from .pa_mqa_logits_opus import (
pa_mqa_logits_mxfp4_build_sched,
pa_mqa_logits_mxfp4_sched,
pa_mqa_logits_mxfp4_sched_buffer_ints,
pa_mqa_logits_mxfp4_sched_slots,
aiter/ops/opus/pa_mqa_logits_opus.py:101
- [verified]
compile_opswraps this binding withtorch_compile_guard's defaultmutates_args="unknown", so schema inference marks every Tensor parameter as mutated even though this launch only writesout. That can force auto-functionalized writebacks for the read-only Q/KV/weight inputs and inflate or break compiled graph consumers; the existing batched-GEMM wrapper documents this exact failure mode. Author must provide precise mutation metadata for this raw binding (andcta_infofor the builder).
@compile_ops(MD_NAME_MXFP4, develop=True)
csrc/py_itfs_cu/pa_mqa_logits_mxfp4_kernels.cu:43
- [verified] This only checks that
weightshas two dimensions and enough rows; the kernel uses a fixedT::W_ROW_ELEMS == 64stride for every row. An input such asweights.shape == [T, 1]therefore passes validation and makes the remaining 63 loads read past the allocation. Author must validateweights.size(1) == q.size(1)before launching.
AITER_CHECK(weights.dim() == 2, "weights must be 2-D [T, H], got ndim=", weights.dim());
- Files reviewed: 8/8 changed files
- Comments generated: 7
- Review effort level: Lite
There was a problem hiding this comment.
🔵 Needs a closer look
Unresolved moderate findings remain around the API contract, device safety, output bounds, and test failure coverage.
Review details
Suppressed comments (9)
aiter/ops/opus/init.py:19
- The public surface added here exports only
build_schedand the singleschedlauncher; there are no..._prefillor..._decodeentry points. This forces every caller to allocate/reusecta_infoand run the device-side schedule builder, contradicting the PR's stated schedule-free 1D/3D APIs and “no device-side schedule” guarantee. Please either implement/export those entry points or update the contract before merging.
from .pa_mqa_logits_opus import (
pa_mqa_logits_mxfp4_build_sched,
pa_mqa_logits_mxfp4_sched,
pa_mqa_logits_mxfp4_sched_buffer_ints,
pa_mqa_logits_mxfp4_sched_slots,
aiter/ops/opus/pa_mqa_logits_opus.py:170
- The PR description promises schedule-free
..._prefill/..._decodelaunches with no persistent schedule buffer, but this public path instead builds a device-side per-row table and requires callers to retain/passcta_infoand its scratch allocation. That is a different caller contract and removes the claimed schedule-free behavior. Author must either implement the described entry points or update the description and API contract to document the schedule-table requirement.
"""Build the per-row schedule, for either entry point. Device-side, cudagraph-safe, no sync.
Call this ONCE PER FORWARD, not once per layer: it depends only on ``local_ends``, a
per-forward quantity, while the kernel runs per CSA layer. A caller inside a CUDAGraph
capture builds it in its metadata builder and hands the same buffer to the capture.
aiter/ops/opus/pa_mqa_logits_opus.py:102
- This raw entry point is exported without calling
_require_gfx950, unlike the two wrappers around it. On gfx942, the header selects the empty kernel stub because__gfx950__is false, so a direct caller receives no-op output (or stale values in a reusedout) instead of the documented gfx950-only error. Author must hide this raw symbol or wrap it with the same architecture guard before exporting it.
@compile_ops(MD_NAME_MXFP4, develop=True)
def pa_mqa_logits_mxfp4_fwd_sched(
aiter/ops/opus/pa_mqa_logits_opus.py:198
- Because this is a
develop=Truepybind call,compile_opssets the extension's thread-local HIP stream fromtorch.cuda.current_device(), while the C++ builder switches tolocal_ends.device_id. Calling this helper with metadata on a non-current GPU can therefore enqueue on a stream belonging to the wrong device; wrap the call inwith torch.cuda.device(local_ends.device)(asaiter/ops/opus/gemm_op_a16w16.py:72does) or otherwise select the stream for that device.
_pa_mqa_logits_mxfp4_build_sched_raw(
local_starts if local_starts is not None else empty,
local_ends.to(torch.int32).contiguous(),
row_to_batch if row_to_batch is not None else empty,
cta_info,
aiter/ops/opus/pa_mqa_logits_opus.py:245
- The same stream-selection mismatch exists for the forward launch:
compile_opssets a stream for the current device, while the C++ function guards toq.device_id. A caller using a non-current GPU can launch this kernel on an incompatible stream. Enter the q device context before this raw call, or make the develop wrapper choose the tensor's device.
pa_mqa_logits_mxfp4_fwd_sched(
q_fp4,
q_scale,
kv_cache,
kv_scale,
block_tables,
csrc/py_itfs_cu/pa_mqa_logits_mxfp4_kernels.cu:58
local_endis copied into each schedule record and later used to size the bounded output descriptor, but this validation only checksout.size(1) >= max_seq_len; it never enforceslocal_end <= max_seq_len. For example,local_end=65,max_seq_len=64, andout.shape=(T,64)passes these checks and permits the kernel to write column 64 past the output allocation. Author must enforce the window/output bound before launching or clamp it in a way that cannot expose memory beyondout.
// The kernel bounds its store by the WINDOW, not by max_seq_len -- `max_seq_len` is
// carried in kargs and never read, only `stride_out_row` is. So these two are the only
// place an undersized `out` can be caught at all; a `local_ends` entry past out.size(1)
// is the caller's contract (see the wrapper docstring).
op_tests/test_pa_mqa_logits_opus.py:494
- The pass condition treats
fly_err=NaNas success, while the FlyDSL helpers returnNoneafter broad import/launch exceptions and the caller then omits that candidate. A FlyDSL compile/ABI failure can therefore turn off the independent scale-layout oracle and still report a passing correctness run, allowing a kernel/layout regression to be missed. Author must make unexpected baseline failures fail or be reported as an explicit skip rather than a successful check.
ok = err == 0 and oob and wr and (math.isnan(fly_err) or fly_err == 0)
op_tests/test_pa_mqa_logits_opus.py:783
- The 16,384-row prefill sweep is the only case that exercises the multi-workgroup
build_schedpath, but this code samples only eight rows and does not check full-output/window invariants. A bug inmqa_logits_build_sched_emit/finishcould leave most rows wrong while the performance table still passes. Add a dedicated >4096-row correctness case that validates all in-window writes and representative logits.
ref = ref_rows(inp, sample_rows(total_q, le, seed=seed), rb, ls, le)
ret = {}
for name, us in times.items():
ret[f"{name} us"] = round(us, 2)
ret[f"{name} TFLOPS"] = round(flops / us / 1e6, 1)
op_tests/test_pa_mqa_logits_opus.py:785
- The perf sweeps compute
check_rowserrors, but those values are only placed in the returned markdown and never affect the process exit status:main()uses onlyrun_corner()andrun_nan_scale()to decide success. A regression that appears only on a large prefill/decode shape can therefore report a nonzero*_errand still pass CI. Propagate a nonzero sampled-reference error into the test result (or raise here) before returning.
ret[f"{name} err"] = check_rows(outs[name], ref, name)
- Files reviewed: 8/8 changed files
- Comments generated: 0 new
- Review effort level: Lite
An MXFP4 paged MQA-logits kernel for DeepSeek-V4-style sparse-attention indexers on gfx1250, in the manner of the gfx950 op in #5332. Prefill and decode run through ONE launch over a per-tile schedule table built on device. wave32 and `v_wmma_scale_f32_32x16x128_f4`, where K = 128 is the whole head_dim in one instruction, with TDM for HBM->LDS. A CTA is 4 waves x 32 lanes sharing one 128-token KV tile, each wave owning one of the tile's up-to-4 query rows. 589 VGPR, occupancy 1, LDS 19456, no spill. The per-row `[local_start, local_end)` windows are the caller's, so any window rule works, including a CSA-compressed cache's `floor((pos + 1) / R)`. The loop bound is the tile's UNION window and is CTA-uniform; the store mask is each wave's OWN row. `pa_mqa_logits_mxfp4_plan` builds the table device-side, once per FORWARD against a per-layer `pa_mqa_logits_mxfp4`, off caller-allocated buffers and a grid that is a function of the static shapes, so the launch stays cudagraph-safe. Both dispatch on the arch of the tensor's own device, so the gfx950 op lands behind the same two entry points. All five inputs take their natural layout, there being no FlyDSL fp4 kernel on this target to be byte-compatible with. gfx950 permutes three of them, and every fp4 scale layout has the same BYTE COUNT, so the C++ `numel` checks accept the other target's arrays and return plausible wrong logits; the shapes differ, and the python entry points are the only place one is seen before the pointer is taken.
* [GFX1250] OPUS PA MQA Logits MXFP4 An MXFP4 paged MQA-logits kernel for DeepSeek-V4-style sparse-attention indexers on gfx1250, in the manner of the gfx950 op in #5332. Prefill and decode run through ONE launch over a per-tile schedule table built on device. wave32 and `v_wmma_scale_f32_32x16x128_f4`, where K = 128 is the whole head_dim in one instruction, with TDM for HBM->LDS. A CTA is 4 waves x 32 lanes sharing one 128-token KV tile, each wave owning one of the tile's up-to-4 query rows. 589 VGPR, occupancy 1, LDS 19456, no spill. The per-row `[local_start, local_end)` windows are the caller's, so any window rule works, including a CSA-compressed cache's `floor((pos + 1) / R)`. The loop bound is the tile's UNION window and is CTA-uniform; the store mask is each wave's OWN row. `pa_mqa_logits_mxfp4_plan` builds the table device-side, once per FORWARD against a per-layer `pa_mqa_logits_mxfp4`, off caller-allocated buffers and a grid that is a function of the static shapes, so the launch stays cudagraph-safe. Both dispatch on the arch of the tensor's own device, so the gfx950 op lands behind the same two entry points. All five inputs take their natural layout, there being no FlyDSL fp4 kernel on this target to be byte-compatible with. gfx950 permutes three of them, and every fp4 scale layout has the same BYTE COUNT, so the C++ `numel` checks accept the other target's arrays and return plausible wrong logits; the shapes differ, and the python entry points are the only place one is seen before the pointer is taken. * spell the dtypes opus's way: bf16 from opus/dtypes.hpp, D_ACC in the wmma * hand the plan's buffers and grid to the caller, and drop cta_target * take the op test to the ubench contract, and add a decode sweep * check the weights row width, and take the reference's E8M0 decode from fp4_utils * bound row_id against num_rows, and tighten the launcher and plan checks * narrow the KV tile to one page, and move the resident CTA count with it * let the caller name the kernel instance, and stop inferring it from the shape --------- Co-authored-by: la <junchen2@amd.com>
…0 row guard, trimmed ABI Launcher and ABI - One check_shapes/launch_sched template for both arches (traits share KV_PAGE_BYTES / KVS_PAGE_BYTES / PAGES_PER_TILE / READS_ROW_WINDOWS). - The plan's row count (num_rows) reaches build_sched and fwd_sched (private ABI); every per-row array is checked against it at any q_per_block, and fwd_sched checks it against q / weights / out on both arches. Replaces the num_tiles-based length checks, which used the wrong unit above q_per_block == 1, and the gfx950-only Python row bound. - Arch cached per device (SynchronizedCache); a single device guard per launch. - build_sched requires cta_resident >= 1. - Drop the six kargs fields no kernel reads (144 -> 112 B). Kernels - gfx950: CTA-uniform `row_id >= num_rows` guard, as gfx1250 has; SCHED defaults to Table and is static_asserted; the unused Prefill/Decode arms are removed; traits and kernel renamed *_mfma_*. Schedule builder - Skip the row_to_batch load on empty tiles (read one past the end on surplus tiles); guard general_emit's divide; drop __restrict__ on the cta_info read-back; __syncthreads() instead of a hand-written copy; remove dead helpers. Tables are byte-identical for qlen1_kv64. Python - MqaLogitsVariant gains a trailing, defaulted `arch`; cross-arch guards compare it instead of relying on value equality. block_table_width does no device probe for explicit variant objects. MqaLogitsPlan gains a defaulted num_rows. - qlen1_kv256 cta_resident back to 1024. - Document the empty-window contract for rows outside any live sequence, and that `out` is only written inside each row's window. Tests - FlyDSL cross-check on gfx950, a real NaN control, host-raise and raw row-guard probes on both arches, ATOM-shaped padded decode, kv256 perf rows. Comments condensed throughout; no code change from that part (ISA identical). Validated on gfx950: corner 36/36, NaN 8/8, err 0 on all perf tables; MFMA ISA unchanged apart from the kargs and row-guard prologue; perf vs pre-PR main within +/-3% on #5332 and #5656 shapes. Co-Authored-By: Claude Opus 5.5 (1M context) <noreply@anthropic.com>
… one module, schedule and API (#5761) Fold the gfx950 (#5332) and gfx1250 (#5656) OPUS MXFP4 paged MQA-logits ops into one implementation. The device math of both bodies is unchanged. - One JIT module, module_pa_mqa_logits_mxfp4_opus. Sources move to csrc/kernels/opus_mqa_logits/pa_mqa_logits_mxfp4/: - _opus.h: ABI, kargs, traits; - _sched.cuh: shared builder; - _gfx950.cuh / _gfx1250.cuh: arch bodies, each an empty stub on the other arch; - _kernels.cu: launcher. Namespaces are opus_logits::gfx950 / ::gfx1250, with generic names (pa_mqa_logits_mxfp4_traits and _kernel). The per-arch sources and module_pa_mqa_logits_mxfp4_gfx1250_opus are deleted. - One schedule: gfx1250's build_tiles + build_sched serve both arches. build_tiles is skipped at q_per_block == 1, where the cut is the identity, which saves one launch per forward on every gfx950 call and on gfx1250's MTP=1 decode. - One runtime-arch dispatch: fwd_sched probes the arch once per process, then dispatches on (q_per_block, block_k). An unmatched config raises. - Host-visible bounds on both arches: - num_rows is checked against q, weights and out, and against the lengths of local_ends, local_starts and row_to_batch; - block_tables width is checked in KV tiles; - both kernels drop records with row_id >= num_rows. The caller contract (four conditions) is documented in the module docstring and the header. - One Python API (aiter.ops.opus.pa_mqa_logits_mxfp4: plan_buffers / plan / pa_mqa_logits_mxfp4), with per-arch MqaLogitsVariant instances. gfx1250: qlen4_kv64 / qlen1_kv64. gfx950: qlen1_kv64 / qlen1_kv256, cta_resident 1024. - Retired: #5332's gfx950 entry points (pa_mqa_logits_mxfp4_sched / _build_sched / _sched_slots / _sched_buffer_ints). - Kept: #5656's public gfx1250 API, except the gfx950-only plan(block_k=). - -mllvm -enable-post-misched=1 is applied on gfx950 only. It is load-bearing there; on gfx1250 it slowed chunked prefill. - kargs drops six unread fields (144 -> 112 B). - One op test for both arches, op_tests/test_pa_mqa_logits_mxfp4_opus.py. It adds ATOM-shaped padded decode, host-raise and raw row-guard cases, a real NaN control, and on gfx950 a FlyDSL cross-check.
* [Config] Add gfx942 a8w8 blockscale GEMM tunings for Qwen3/Qwen3.5/GLM/DSV4 shapes (#5839) MI325X (gfx942, 304 CUs): 286 (M, N, K) rows over 28 weight shapes that no existing model_configs table covers, tuned on main with gemm_a8w8_blockscale_tune.py --libtype ck --splitK. Keys already present in model_configs and rows slower than the heuristic default are left out. The rows go into each model's existing model_configs table; Qwen3-14B and Qwen3.5-122B-A10B get new tables. A shape two models share is kept in one table only, since the config merge rejects duplicate keys across files. Co-authored-by: Cursor <cursoragent@cursor.com> (cherry picked from commit 8260de6) * [HIP] [OPUS] [FlyDSL] gfx950 MXFP8 e8m0 GEMM + DeepSeek-V4/V4.1 tuned configs (#5896) * feat(flydsl): gfx950 MXFP8 e8m0 GEMM on FlyDSL + DeepSeek-V4/V4.1 configs - New FlyDSL gfx950 MXFP8 batched GEMM (bmm_a8w8_mxscale_gfx950): scaled MFMA 16x16x128, async LDS DMA pipeline, 32x32 / 128x128 e8m0 blocks, split-K with same-XCD last-arrival reduction, B direct to registers, XCD tile order, non-temporal B, per-stage or preloaded scale panels, K % 64 tail, column-major (blockscale) x_scale. Scale rows that are not whole dwords (e.g. 128-wide blocks at K = 384 / 768) are copied a byte per LDS slot with exact buffer bounds. check_bmm_config is the single legality source. - Arch-neutral front door flydsl.batched_gemm_a8w8 dispatching by (arch, w_scale block, x_scale layout); gfx950 module with kernelName parsing, untuned-shape heuristic and a compiled-launcher cache (host 71 -> 21 us). - gemm_a8w8_blockscale_bpreshuffle on gfx950 routes e8m0 x_scale + w_scale to it through the tuned table AITER_CONFIG_GEMM_A8W8_BLOCKSCALE_MXSCALE_ BPRESHUFFLE; FP32 scales keep CK / asm. gemm_a8w8_blockscale with isBpreshuffled forwards native group32 operands there. The V4 wo_a batched GEMM uses the same kernel. - AOT precompiles kernelId=bmm rows (both x_scale layouts for 128x128). - Tuned gfx950 configs: DeepSeek-V4-Pro 128x128 linears (980 rows, TP1/4/8, 1.30x vs CK/asm), DeepSeek-V4.1 group32 (1224 rows, 1.22x vs Triton), V4 wo_a batched (680 rows). - Also carries the other local changes in the tree: inverse_rope_group_quant, opus policy / bmm tune, gfx1250 bmm wrapper, tensor_shim helpers. * refactor(flydsl): tidy gfx950 MXFP8 bmm kernel per FlyDSL cleanup guide - Use fx.copy instead of fx.copy_atom_call for the single-atom copies. - Build the async-LDS DMA destination from the LDS pointer with fx.add_offset + fx.to_llvm_ptr instead of ptrtoint/inttoptr with a hand-picked address space. The raw buffer_load_async_lds stays: the BufferLoadAsyncLDS atom has no 1-byte size and no cache-policy operand. - Factor the compile-and-run + leaked ir.Context recovery out of tensor_shim._run_compiled into _compile_and_run, and use it for the bmm wrapper's per-config compile cache too. Generated ISA is byte-identical on 9 reference configs. (cherry picked from commit 40d524b) * [Triton/Gluon] [ASM] [HIP] Mha v4: adds bf16 sparse, LSE support, KV varlen, fixes, etc (#5798) Motivation: adds log-sum-exp output to MHA v4 gfx950 kernels (only dense variants now), so it can run under ring / context parallelism, which merges per-rank partials via LSE. Also extends block-sparse to the BF16 recipes, canonicalizes the MXFP4 rows, and improves K/V quantization accuracy. # Kernels: - New BF16 and BF16FP8 sorted-sparse kernels. - Dense LSE epilogue on all ten hd128 recipes, gated at runtime on s_lse. - f8f6, f6f4 and mxfp4 sparse rows moved to FP6-P V, matching their dense siblings. - Dense mxfp4 claims the canonical FP6-P V order; the duplicate f4f4 row is disabled. # Host: - bugfix: per-channel V amax clamped so empty heads cannot quantize to NaN. - Optional lse output on mha_v4 / mha_v4_packed, plumbed through asm_mha_v4_fwd.cu. + ABI unchanged: ptr_lse / s_lse / s_lse_Hs were already reserved in the kernarg. - _LSE_CAPABLE_QV gates the supported format pairs; sorted-sparse still raises. - K mean is subtracted before quantizing (K-smoothing), fused into the MX quantizer kernels. - Dense MHA v4 accepts per-batch key lengths (ragged seqlen_k). # Minor: - Retired MXFP4 Q/K + FP8 V and the deprecated mha_v4_mxfp8 alias. - bench_sage.py: improve input distributions, diffusion-calibrated default, BF16 sparse modes. - Split the block-sparse cases into op_tests/test_mha_v4_sparse.py. (cherry picked from commit 105615e) * [tuner] Fail a task as soon as its worker process exits (#5841) A GPU memory fault aborts the worker process, so its task never returns a result. mp_tuner only noticed at the task timeout (1800s by default) or never without one. Record which worker started each task and, when that process is gone, fail the task and restart the pool right away, like the existing accelerator-error path. Fixes #5840. Co-authored-by: Cursor <cursoragent@cursor.com> (cherry picked from commit bcb56d9) * [CI] Update Aiter artifact downloads to v8.0.1 (#5908) (cherry picked from commit c03689f) * [MLA v4 nm] Test fix _run_one_point reading packed BF16 as FP32 partials (#5905) * [MLA v4 nm] Fix _run_one_point reading packed BF16 as FP32 partials Since #4311 the v4 nm dispatcher derives out_16_nosplit from num_kv_splits and ignores the caller's value, so a single-split launch writes packed BF16 into the logits buffer. _run_one_point still read logits_buf[:, 0] as FP32 for num_kv_splits == 1, which compared every other BF16 element (plus the never-written tail of the buffer) against the reference and printed spurious `fp8_dequant_ref vs asm` checkAllclose failures in the script-mode sweep. Read output_buf, which holds the final result for every split count. Co-authored-by: Cursor <cursoragent@cursor.com> * [MLA v4 nm] Gate accuracy checks so a failed! fails the test checkAllclose only raises on a catastrophic delta; otherwise it logs `failed!` and returns the mismatch fraction. Four of the six accuracy checks in test_mla_v4_nm.py dropped that return value, which is how the packed-BF16 readback bug fixed in the previous commit passed both pytest and the script-mode CI run. Route them through _gated_allclose, which asserts the mismatch fraction against the same tol_err_ratio checkAllclose uses for `failed!`. The script-mode sweep keeps going past a failing shape, prints a summary, and exits non-zero so aiter_test.sh reports the file as failed. Co-authored-by: Cursor <cursoragent@cursor.com> --------- Co-authored-by: Cursor <cursoragent@cursor.com> (cherry picked from commit 049fae4) * [AOT] Inline the FlyDSL FP8 FMHA head shapes, drop the config CSVs (#5901) Follow-up to #5796. The AOT job list for the gfx950 FlyDSL FP8 flash attention was driven by a header-only aiter/configs/fmha_fp8_aot.csv merged with aiter/configs/model_configs/*_fmha_fp8_aot.csv. A single row does not justify the CSV plumbing, so the head shapes now live in a DEFAULT_SHAPES table in the module, the way mega_moe.py already does it, and the comments introduced by #5796 are trimmed. - fmha_fp8.py: DEFAULT_SHAPES replaces parse_csv()/DEFAULT_CSVS; default_jobs() replaces the CSV walk; --csv is gone (--shape still overrides). The cu_num column is dropped with the CSVs: the kernel is gfx950-only and AOT_ARCH is fixed, so every non-gfx950 row was warned about and skipped anyway. - common.py: FMHA_FP8 returns default_jobs() next to MEGA_MOE instead of going through collect_aot_jobs(). - jit/core.py: drop AITER_CONFIG_FMHA_FP8_AOT and its config-file property, now unused. - README.md: document the table instead of the CSVs. _variant_space()/jobs_for_shape() are unchanged, so coverage is unchanged: --list still emits the same 92 kernel names for Kimi-K3 TP8 (12:12:192:128 varlen_cross) as the CSV-driven version, and setup.py source builds still compile them via run_aot(). Co-authored-by: Claude Opus 5 <noreply@anthropic.com> (cherry picked from commit 71a31be) * [FlyDSL] gfx942 fp8_mqa_logits: let _auto_variant choose rows_per_block (#4963) * [FlyDSL] gfx942 fp8_mqa_logits: let _auto_variant choose rows_per_block _auto_variant returned f"mfma_r2_w{wpb}", so rows_per_block was pinned at 2 and seven of the nine registered variants -- including every member of the r4 family -- could never be selected. r4 amortizes each KV tile load over twice as many query rows and is faster from seq_len 8 upward. The gap is widest exactly where it costs most: vLLM chunks indexer prefill to fit VLLM_SPARSE_INDEXER_MAX_LOGITS_MB (512 MB), which caps seq_len at 1024 when seq_len_kv is 131072, so the existing seq_len >= 2048 branch cannot fire at long context and every such call took mfma_r2_w4 -- the median of the nine by speed, with the best 1.65x faster. Measured on MI325X (gfx942), seq_len_kv 131072, best variant vs the r2 pick: seq_len 1 r2 25.2 us r4 79.1 us r4 3.1x worse seq_len 4 r2 24.3 us r4 31.5 us r4 1.3x worse seq_len 8 r2 34.0 us r4 32.9 us r4 1.03x better seq_len 16 r2 56.5 us r4 47.5 us r4 1.19x better seq_len 1024 r2 2564.3 us r4 1520.3 us r4 1.69x better Below seq_len 8 the host padding of seq_len up to a multiple of RPB dominates -- at seq_len 1 an r4 kernel computes 4 rows to obtain 1 -- so r2 is kept there and behaviour for those shapes is unchanged. Logits are bitwise identical across all variants at every shape tested, so this is purely a blocking/occupancy change. End to end on 8x MI325X, TP8, GLM-5.2-FP8, 131072 in / 1024 out, concurrency 8: median TPOT improves 7.22% and output throughput 6.60%. This kernel is 16.8% of GPU time at that point. Signed-off-by: Jin Tao <jin.tao@amd.com> * [FlyDSL] gfx942 fp8_mqa_logits: pick RPB on element count, and keep it a divisor Refines the previous commit's rule after a 2-D sweep. That rule keyed RPB off seq_len alone with a crossover measured only at seq_len_kv=131072; sweeping the other contexts shows the crossover is not a seq_len threshold at all, and that a second effect was being read as one. RPB tracks the logits element count. Over seq_len 1..8192 x seq_len_kv 1024..262144 on MI325X, the boundaries land on the same element count at every context: RPB=1 wins below 2**19 elements (27/27 shapes), RPB=2 at 2**19 (6/6), RPB=4 from 2**21 up (38/38), with 2**20 a transition band split 3/3. Keying off seq_len instead put the previous rule on the wrong side at low context: at seq_len 16, seq_len_kv 1024 it chose RPB=4 and ran 1.26x slower than RPB=1. RPB must also divide seq_len. When it does not, the launcher pads with four torch.cat calls; that is a flat ~44 us of host-side overhead, independent of seq_len_kv, and it is the whole of the "small seq_len" penalty the previous commit attributed to wasted rows. At seq_len 1, seq_len_kv 131072: RPB=1 23.1 us, RPB=2 67.8 us, of which the four cats are 44.1 us and pre-padding by hand recovers all of it (21.9 us). So the penalty is not proportional to the padding -- 1 wasted row of 2 costs the same as 3 of 4 -- and it applies to every odd seq_len, which the old rule sent to RPB=2 unconditionally. Stepping down to a divisor is only right while the kernel is cheap relative to that fixed cost, so it is gated to the same 2**21 elements: at seq_len 1025, seq_len_kv 131072 the dividing RPB=1 takes 3880 us against 2601 us for a padded RPB=2. Measured on MI325X, no FLYDSL_FP8_MQA_LOGITS_VARIANT set, old pick vs new: seq_len seq_len_kv old new speedup 1024 131072 r2_w4 2598.7 r4_w4 1623.5 1.60x 1025 131072 r2_w4 2609.3 r4_w4 1594.0 1.64x 512 131072 r2_w4 1190.1 r4_w4 795.3 1.50x 700 50000 r2_w4 661.6 r4_w4 440.3 1.50x 333 12000 r2_w4 118.4 r4_w4 97.1 1.22x 16 131072 r2_w4 58.3 r4_w4 47.4 1.23x 3 131072 r2_w4 66.4 r1_w4 28.1 2.36x 1 131072 r2_w4 69.4 r1_w4 23.2 2.99x 1 1024 r2_w4 66.2 r1_w4 22.2 2.98x No shape measured regresses; the smallest gain is 1.04x. Against the best of the nine variants at each shape, pooled over held-out data (non-power-of-two shapes, a fine seq_len sweep, and head counts 16 and 64), the geometric mean cost falls from 1.45x to 1.03x and the worst case from 3.17x to 1.41x. Logits are bitwise identical across all nine variants at all 180 shapes swept (1620 timings), so this remains purely a blocking/occupancy change. WPB is deliberately left alone. It is worth a few percent at most here, and unlike RPB its optimum moves with the head count -- at 64 heads the current WPB rule costs 1.61x worst case where a fixed WPB=4 costs 1.07x -- so it needs its own sweep rather than a change fitted to one head count. Signed-off-by: Jin Tao <jin.tao@amd.com> Co-authored-by: Cursor <cursoragent@cursor.com> * [FlyDSL] gfx942 fp8_mqa_logits: trim _auto_variant comments per review Move RPB element-count thresholds into _auto_variant and shorten the docstring; logic unchanged. Signed-off-by: Jin Tao <jin.tao@amd.com> Co-authored-by: Cursor <cursoragent@cursor.com> * [FlyDSL] gfx942 fp8_mqa_logits: scope the _auto_variant step-down note The docstring stated the divisor step-down as an unconditional rule, but it only applies in the middle band: above the top threshold RPB stays 4 and the padding is accepted. Say which band it applies to, and why the top band is exempt. No logic change. Co-authored-by: Cursor <cursoragent@cursor.com> * [FlyDSL] gfx942 fp8_mqa_logits: unit-test the variant selector The shape sweep in test_flydsl_fp8_mqa_logits.py never reaches r4 -- its largest default shape is 1024 x 1560, below the 2**21 threshold -- and nothing asserted _auto_variant directly, so a regression could re-pin RPB to 2 with every correctness test still green. Cover the RPB bands at both thresholds, that the edges track seq_len * seq_len_kv rather than seq_len alone, the middle-band step-down on odd seq_len (and that it stops above the top threshold), the production 1024 x 131072 shape, the unchanged WPB rule, and _resolve_variant precedence so the auto path is confirmed to be the default. Pure shape arithmetic, so no kernel launch. Co-authored-by: Cursor <cursoragent@cursor.com> * [FlyDSL] gfx942 fp8_mqa_logits: put selector tests on the CI path CI only discovers op_tests/test_*.py, so the r1/r2/r4 selector coverage never ran. Move it there, run pytest from __main__, and add two r4 shapes to the GPU sweep so auto-selected padding is launched. Signed-off-by: Jin Tao <jin.tao@amd.com> Co-authored-by: Cursor <cursoragent@cursor.com> * [FlyDSL] gfx942 fp8_mqa_logits: fold the selector pins into the op test aiter op tests are plain scripts, not pytest, so drop test_flydsl_fp8_mqa_logits_variant.py and pin the gfx942 auto-selected variant in test_flydsl_fp8_mqa_logits.py instead: a host-only verify_auto_variant table at both RPB band edges, either side of the odd-seq_len step-down, the long-context prefill shape and the WPB switch. It runs in the default verify scenario on gfx942, and its failures count toward the same exit code as the kernel sweep. The selector import sits in the file's existing ImportError guard. Signed-off-by: Jin Tao <jin.tao@amd.com> Co-authored-by: Cursor <cursoragent@cursor.com> --------- Signed-off-by: Jin Tao <jin.tao@amd.com> Co-authored-by: Cursor <cursoragent@cursor.com> Co-authored-by: jin.tao@amd.com <jin.tao@amd.com@tus1-p15-g43.tus.tensorwave.lan> Co-authored-by: Felix Li <felix.li@amd.com> (cherry picked from commit bbaccdf) * [Triton] [MHA] Add a tuned gfx1101 config, split small_head/default (#4493) gfx1101 (RDNA3, e.g. RX 7800 XT) ships no MHA config, so `_get_config` in `_triton_kernels/attention/mha.py` finds no `configs/gfx1101/triton/attention/mha/DEFAULT.json` and every call to `aiter.ops.triton.attention.mha.flash_attn_func` fails before a kernel runs. Of the architectures in RDNA_ARCHS, only gfx1151 ships one today. Nine of the eleven entries are taken verbatim from the gfx1151 donor (RDNA3.5, added in #3423, tuned in #3560), which is the nearest tuned architecture. Two forward entries are tuned on gfx1101 instead of inherited, and they differ from each other only in `num_warps` and `num_stages`: fwd/default BLOCK_M 128, num_warps 8, num_stages 3 (donor: 64 / 4 / 2) fwd/small_head BLOCK_M 128, num_warps 4, num_stages 1 The split uses the `small_head` bucket added in #4414, which is opt-in per architecture by the mere presence of the key, so this stays a data-only change. It is needed because one `fwd/default` cannot serve both halves on this card: measured against the donor, `M128 w4 s1` is 0.935x on head_dim 64 but 1.269x on head_dim 128, while `M128 w8 s3` is 0.914x on head_dim 128 but 1.062x on head_dim 64. Measured on Windows native ROCm, triton 3.8.0, fp16, 10 independent repeats of 20 iterations after 5 warmups, `torch.cuda.synchronize()` per iteration; a result counts only when the [min, max] intervals across repeats are disjoint. Ratios are against the gfx1151 donor entry, i.e. against what a donor-inherited config would do. head_dim 128, `default` Flux joint 0.914x / 0.891x / 0.892x on three independent torch+ROCm stacks (2.11/7.15, 2.11/10.1, 2.15/10.1), all pinned to the same triton; llama3-8B 0.882x, mixtral-7B 0.880x, kimik25-tp4 0.849x at seqlen 16384 head_dim <= 64, `small_head` SDXL self-attn 0.935x; deepseek-V3 0.658x and glm47fp8-tp4 0.857x at seqlen 16384 The LLM shapes come from `op_tests/op_benchmarks/triton/utils/model_configs.json` (prefill, causal, GQA, batch 1); the sequence lengths are not in that file and are chosen here. Not covered: batch > 1, varlen/thd, sliding window, decode. Note that the `small_head` comment in `_get_config` does not describe gfx1101. It states that 16 < d <= 64 suffers a num_stages=1 pipelining pathology which num_stages=3 cures. On this card the ordering is the opposite -- on the tuned M128/N32/w4 tile, num_stages 1/2/3 measured 2.300 / 2.399 / 2.492 ms. The bucket is still the right mechanism here, for a different reason: the two head_dim ranges want a different `num_warps`, not a different `num_stages`. Signed-off-by: Martin Domanský <ragua@email.cz> Co-authored-by: Mehmet Cagri <mehmet.kaymak@amd.com> (cherry picked from commit f1f95ce) * [Triton/Gluon] Consolidate tuning harnesses (#5874) * [Triton/Gluon] Consolidate tuning harnesses * [Triton/Gluon] Simplify tuning harness (cherry picked from commit af3514a) * [Bugfix][Gluon][MLA] Follow-up: fix stale comments + test hardening (#5860) Addresses review feedback on #5648 (kept separate to not disturb the approved PR): - mla_gluon.py: the >2GB global_load calls now carry a bounds mask + other=0.0, so the old "No mask needed / in-bounds" comments above them were stale and gave the opposite (unsafe) guidance. Replace with an accurate one-liner. - test_mla.py: allocate the output with device=q_nope.device instead of relying on the module default device; call torch.cuda.empty_cache() before the mem_get_info() free-memory gate to avoid allocator-state flakiness; and assert exact parity (atol=rtol=0) between the >2GB and <2GB paths, which read identical KV and must be bit-identical (a loose tolerance could hide a regression). Signed-off-by: Rohan138 <rohanpotdar138@gmail.com> (cherry picked from commit a966245) * [Triton/Gluon] [Config] gemm_a16w16: gfx950 tuned per-shape configs (#5830) * [Triton] gfx950: tuned defaults for gemm_a16w16, gmm, MHA fwd and PA decode Tuned on MI355X (gfx950) and validated for correctness and no regressions on broader shape sets than tuned. All changes are gated to gfx950 configs/arch. - gemm_a16w16: per-shape configs for N,K = 2048/2048, 4096/4096, 8192/8192, 10240/8192, 57344/8192, 8192/28672 (non-persistent and persistent). Only the tuned M bucket differs from DEFAULT.json. 1.07-2.5x. - gmm: new "large_kn" config (8 warps), selected for K, N >= 4096 with >= 256 rows per group. 1.14-1.26x there; smaller problems keep "default". - MHA fwd (Triton): new "mid_head" config (128x128, 2 stages) for bf16/fp16 with 64 < head_dim <= 128. 1.03-1.16x; d<=64 and d>128 unchanged. - PA decode (Triton): use v2 whenever there is more than one partition for bf16/fp16 KV. 1.9-5.3x for the batch/context sizes that previously hit v1. Unit tests: test_gemm_a16w16 360 passed, test_gmm 48 passed, test_pa_decode 816 passed, test_mha 2060 passed. * [Triton] pa_decode: sort imports (ruff I001) * Move gmm, MHA fwd and PA decode changes to their own PRs Per review, each kernel gets its own PR; this PR keeps only the gemm_a16w16 tuned configs. The removed changes are on branches gfx950-tuned-gmm, gfx950-tuned-mha-fwd and gfx950-pa-decode-v2. * [Triton] gemm_a16w16_persistent: gfx950 tuned config for N=K=2048 (cherry picked from commit a651db0) * replacing the fp32 mfma with two bf16 mfma. fn is split inside the kernel into hi = bf16(fn) and lo = bf16(fn -hi). a simple bf16 downcast for fn inside the kernel drops accuracy more and performs worse for large M. also change the concatenation of the four streams from column major to stream major (k = stream * TILE_K + column). retuned the gfx950 configs. (#5885) (cherry picked from commit f10cd2a) * perf(fused-moe): add tuned DSV4.1 TP4 configs (#5919) (cherry picked from commit 977ae79) * [Config] Add Qwen3.8-27B TP1 a8w8 blockscale GEMM tunings for gfx942 (#5585) * [Config] Add Qwen3.8-27B TP1 a8w8 blockscale GEMM tunings for gfx942 #3324 covered this model family at TP=2/4/8 only, so the five widths Qwen3.8-27B drives at TP1 have no tuned entries for gfx942/cu_num=304 and every call falls back to the default -- 275 "not found tuned config" messages per profile run, now 4. This adds 641 rows in one file, 249 decode (M <= 512) and 392 prefill (M 907-65536), from a shape list extracted out of real vLLM server logs rather than generated as an M ladder. Decode, against AITER's own default, measured in a real vLLM serving run -- Qwen3.8-27B-FP8 at 128K context, max-concurrency 1, on one MI325X, byte-identical images, the config bind-mount the only variable, three interleaved repeats per arm on an idle node: decode step 15.444 ms -> 12.449 ms 1.24x Per-shape decode GEMM vs the default: down 2.33x, out_proj 1.51x, in_proj 1.12x, qkv 1.00x, gate_up 1.00x. No decode shape regresses. Decode-step spread across repeats was 0.037 ms and 0.012 ms. The 392 prefill rows ship and are tuned, but their end-to-end effect is currently unmeasurable on this workload: gemm_a8w8_blockscale silently stores only M*N mod 2^31 output elements once M*N reaches 2^31, and at a 65,536-token prefill chunk the 34,816-wide gate_up projection crosses that limit. The truncation depends only on M*N and not on which kernel the config selects -- AITER's default truncates identically -- so it is neither introduced nor worsened by this change. Tuned with gemm_a8w8_blockscale_tune.py --libtype ck --splitK, decode rows then re-timed interleaved because the tuner's min() over short samples is noisy on near-ties; only rows beating the default by >=3% were kept, following #5421. test_csv_validation (15 tests) and test_config_shape_collision (17) both pass, there are zero duplicate (gfx, cu_num, M, N, K) keys against the 12 other sources in the runtime merge set on main, and max relative Frobenius error is 4.11e-3 against bf16's 2^-8 = 3.9e-3 rounding floor. Scope: only the decode rows have in-situ evidence, and MI300X shares the gfx942/304 key but is untested here. Two rows pin AITER's default kernel (qkv and gate_up at M=1) after the trace measured the tuned picks slower -- pinned, not deleted, because getPaddedM(1,N,K,0) == 16. Six decode near-ties ship no row. gate_up ships 72 of 80 prefill rows; the 8 missing have M*N > INT32_MAX and fault during tuning, but inherit the tuned M=8192 row at runtime via the same padding collapse. Signed-off-by: Pham Binh <phamhuuthanh.binh@amd.com> Co-authored-by: Claude Opus 5 <noreply@anthropic.com> Co-authored-by: Cursor <cursoragent@cursor.com> * [Config] Thin Qwen3.8-27B TP1 GEMM rows to the padded M ladder Lookup already rounds M through getPaddedM, so a row per traced M is redundant. Keep the power-of-two M list used by the other gfx942 blockscale tables, and omit shapes that main already ships. --------- Signed-off-by: Pham Binh <phamhuuthanh.binh@amd.com> Co-authored-by: Claude Opus 5 <noreply@anthropic.com> Co-authored-by: Cursor <cursoragent@cursor.com> Co-authored-by: akii96 <aakif.nawaz@amd.com> (cherry picked from commit 330f127) * [HIP] [OPUS] [JIT] Unify gfx950 + gfx1250 MXFP4 paged MQA-logits into one module, schedule and API (#5761) Fold the gfx950 (#5332) and gfx1250 (#5656) OPUS MXFP4 paged MQA-logits ops into one implementation. The device math of both bodies is unchanged. - One JIT module, module_pa_mqa_logits_mxfp4_opus. Sources move to csrc/kernels/opus_mqa_logits/pa_mqa_logits_mxfp4/: - _opus.h: ABI, kargs, traits; - _sched.cuh: shared builder; - _gfx950.cuh / _gfx1250.cuh: arch bodies, each an empty stub on the other arch; - _kernels.cu: launcher. Namespaces are opus_logits::gfx950 / ::gfx1250, with generic names (pa_mqa_logits_mxfp4_traits and _kernel). The per-arch sources and module_pa_mqa_logits_mxfp4_gfx1250_opus are deleted. - One schedule: gfx1250's build_tiles + build_sched serve both arches. build_tiles is skipped at q_per_block == 1, where the cut is the identity, which saves one launch per forward on every gfx950 call and on gfx1250's MTP=1 decode. - One runtime-arch dispatch: fwd_sched probes the arch once per process, then dispatches on (q_per_block, block_k). An unmatched config raises. - Host-visible bounds on both arches: - num_rows is checked against q, weights and out, and against the lengths of local_ends, local_starts and row_to_batch; - block_tables width is checked in KV tiles; - both kernels drop records with row_id >= num_rows. The caller contract (four conditions) is documented in the module docstring and the header. - One Python API (aiter.ops.opus.pa_mqa_logits_mxfp4: plan_buffers / plan / pa_mqa_logits_mxfp4), with per-arch MqaLogitsVariant instances. gfx1250: qlen4_kv64 / qlen1_kv64. gfx950: qlen1_kv64 / qlen1_kv256, cta_resident 1024. - Retired: #5332's gfx950 entry points (pa_mqa_logits_mxfp4_sched / _build_sched / _sched_slots / _sched_buffer_ints). - Kept: #5656's public gfx1250 API, except the gfx950-only plan(block_k=). - -mllvm -enable-post-misched=1 is applied on gfx950 only. It is load-bearing there; on gfx1250 it slowed chunked prefill. - kargs drops six unread fields (144 -> 112 B). - One op test for both arches, op_tests/test_pa_mqa_logits_mxfp4_opus.py. It adds ATOM-shaped padded decode, host-raise and raw row-guard cases, a real NaN control, and on gfx950 a FlyDSL cross-check. (cherry picked from commit 9b3885f) * [FlyDSL] feat(mega_moe/gfx1250): bind mori tokoff-ext allocator on the mori dispatch path (#5810) * feat(mega_moe/gfx1250): bind mori tokoff-ext allocator on the mori dispatch path mori's op layer builds the tokoff-ext slot allocator in EpDispatchCombineOpHip.__init__; driving EpDispatchPlan directly bypasses that and leaves EpArgs.tokOffPeers null, keeping dispatch on the serializing cco-window atomic. Build it here from mori's TokOffExt, gated exactly like mori (default on; MORI_EP_TOKOFF_EXT=0 opts out), pass its peer-pointer array to plan.launch, and free it in close(). Needs mori with a public TokOffExt (dispatch_combine_v2.hip_backend). Co-Authored-By: Claude Opus 4.8 (1M context) <noreply@anthropic.com> * style(mega_moe/gfx1250): black-format the tok_off_peers ternary Co-Authored-By: Claude Opus 5.5 (1M context) <noreply@anthropic.com> --------- Co-authored-by: Claude Opus 4.8 (1M context) <noreply@anthropic.com> Co-authored-by: HaonanWang98 <hwang@amd.com> (cherry picked from commit 475cf0f) * [Triton] Do not write segment softmax state at NUM_SEGMENTS_PER_SEQ == 1 (#5935) mla_decode_fwd splits each sequence's KV into NUM_SEGMENTS_PER_SEQ segments and lets a reduce kernel merge the per-segment softmax max and expsum. At one segment there is nothing to merge, so the host skips the reduce kernel and hands both scratch pointers the output buffer itself: else: segm_output = out segm_max = out # dummy ptr segm_expsum = out # dummy ptr The decode kernels store M and L unconditionally, so at one segment they write the softmax state straight over the attention output. Nothing faults and nothing warns; the result is simply wrong. Guard both stores. NUM_SEGMENTS_PER_SEQ is a constexpr, so the branch costs nothing when it is greater than one. Only the gfx1250 Gluon kernel reaches this today: select_3d_config ends its gfx12 branch at max(1, ...), while the other branch floors at MIN_SEGMENTS >= 8. The segment count falls as batch x heads grows, so on gfx1250 it reaches one at large batch -- a DeepSeek-R1 TP2 serving run at --max-running-requests 256 sits in that range throughout, and its decode is silently corrupted. The plain Triton kernel shares both the unguarded store and the host-side aliasing, so guard it as well, before a future tuning change makes it reachable. test_mla_decode_fwd stays green either way, which is why this went unnoticed. Its grid does reach one segment for the larger head count, but the corruption lands only at out.flat[token * num_query_heads + head] -- one element in kv_lora_rank, measured at 0.04% to 0.12% of the output -- while the assertion allows a tol_err_ratio of 0.01. Catching it needs a case with no error budget. Checked on gfx1250. DeepSeek-R1 TP2 at page size 64, gsm8k over 2000 questions, scores 0.945 with the fix. A standalone sweep over batch 1 to 256, bf16 and fp8 e4m3 caches, sequence lengths that are not page multiples and pages scattered through the pool gives a worst relative error of 0.0033 for decode and prefill alike. Compared against a torch reference with no error budget, a one-segment batch differs at a 0.08% ratio before this change -- softmax maxima around 100 sitting where attention outputs near 0.02 belong -- and matches to the last element after it. The plain Triton kernel, forced on by disabling the gfx12 branch, matches the same reference at batch 1 to 256, confirming the guard leaves the multi-segment reduce path intact. test_mla.py is unchanged: its 128 pre-existing failures are all shuffled_kv_cache=True on the pipelined kernel, identical before and after. Signed-off-by: Lin, Soga <soga.lin@amd.com> (cherry picked from commit 2b6ff3d) * [Triton/Gluon] Enable unified-attention skip-mask for gfx950 hd256 FP8 prefill (#5598) Long FP8 prefill at head size 256 matches `D_GEQ_256` (`BLOCK_M=16`, no `SPLIT_UNMASKED_LOOP`) because D outranks Q in the lookup, so it never reaches the hd128 skip key. Add the same `Q_GEQ_256` family under `D_GEQ_256`, with `SW` and `SHUF` siblings left skip-off, so the existing kernel fast path fires only for non-windowed, non-shuffled prefill. (cherry picked from commit e289613) * [Triton/Gluon] [Kimi-K3][ROCm] Add merged MoE front (#5321) * feat(kimi-k3): add minimal large-M MoE front Signed-off-by: jiacao-amd <jiahui.cao@amd.com> * feat(kimi-k3): enable merged front for M=7 decode Signed-off-by: jiacao-amd <jiahui.cao@amd.com> * feat(kimi-k3): tune merged front for M=14 decode Signed-off-by: jiacao-amd <jiahui.cao@amd.com> * perf(kimi-k3): tune merged front decode bucket Signed-off-by: jiacao-amd <jiahui.cao@amd.com> * test(kimi-k3): validate the merged front across the MTP decode bucket test_decode_full_front_matches_reference stopped at m=16, so the shapes an MTP server actually decodes at were never checked against the torch.mm reference: with num_speculative_tokens=3 a pure-decode step submits num_seqs * 4 tokens, which is 32 and up for any server past four concurrent sequences. Extend the parametrization to 32, 48, 64, 80, 96, 112, 128 and 192 -- the shapes the companion vLLM change enables, and for which kimik3_bf16_tuned_gemm.csv already ships tuned solutions at N=6016, K=7168, gfx950, cu_num=256. No kernel or config change is needed; this only closes the correctness gap left by the old parametrization. gfx950: 9 passed -> 17 passed. Signed-off-by: jiacao-amd <jiahui.cao@amd.com> * style(kimi-k3): satisfy the pinned ruff on this PR's own files `ruff==0.16.0`, the version .github/workflows/pre-checks.yaml pins, reports two errors in files this PR introduces: RUF022 on the `__all__` list in kimi_k3_moe_front.py and I001 on the import block in its test. Both are autofixes and neither changes behaviour -- `__all__` ordering only affects `import *`, and every symbol here is imported by name. The two EXE001 findings that remain in the repo are in unrelated flydsl files and predate this PR. Signed-off-by: jiacao-amd <jiahui.cao@amd.com> * perf(kimi-k3): retune the M=16 merged-front FP32 row The decode shape is 4 * concurrency, and getPaddedM(gl=0) rounds every M <= 256 up to 16, so this row serves both the C1 (M=4) and C4 (M=16) decode buckets -- about 40% of iterations in the 8k/1k replay. It was tuned against a single weight buffer. The merged front's weight is 6016x7168 BF16 = 86.2 MB and MI355X has 256 MB of LLC, so that regime serves most of the GEMM from cache; the shipped 16.005 us is faster than anything the shape can reach when the weight is actually streamed. In the real model 92 distinct MoE layers stream 7.9 GB per decode step and nothing is resident. Re-searched all 681 hipblaslt solutions with the weight chained across 604 MB of distinct buffers (past LLC) and timed inside a CUDA graph, which is how vLLM runs decode: solidx 443935 (shipped) 22.364 us 3857 GB/s solidx 443486 (this) 19.279 us 4474 GB/s -13.8% 443486 is not a new solution -- the table already carries it at M=7. The right kernel was present and keyed to the wrong M. Max relative error vs a torch.mm FP32 reference is 4.9e-07, so this is a dispatch change only. The us/tflops/bw columns are the streamed measurement and are therefore not comparable to the cache-hot numbers in neighbouring rows. * perf(kimi-k3): size the front-GEMM config cache to the whitelist _KIMI_K3_MERGED_FRONT_TOKEN_COUNTS admits 28 distinct token counts but the config cache held 16, so the M values a mixed decode/prefill server cycles through evicted each other. Minor: the tuned table is cached upstream, so a miss costs a few get_padded_m() extension calls rather than a CSV parse, and in the graphed decode path it is paid at capture instead of replay. This just makes the cache cover the set it was sized for. * refactor(kimi-k3): separate Triton epilogue layers Move the launchable epilogue kernel under _triton_kernels/moe, keep only the public wrapper in ops/triton/moe, and add a config-aware kernel repr and CUDA Graph benchmark. Remove Kimi-specific weight packing and hipBLASLt orchestration from the Triton module. Signed-off-by: jiacao-amd <jiahui.cao@amd.com> * refactor(moe): generalize the SiTU epilogue Rename the op to describe the SiTU epilogue it actually performs rather than implying that it projects the incoming activation. Make all branch widths caller-provided, mask partial tiles for arbitrary shapes, and reuse the shared Triton tanh helper. Update the numerical and CUDA Graph coverage plus the benchmark for the generic API. Signed-off-by: jiacao-amd <jiahui.cao@amd.com> * test(triton): use Triton arch helper for SiTU epilogue Signed-off-by: jiacao-amd <jiahui.cao@amd.com> * fix(triton): avoid branch-local mask type collisions Signed-off-by: jiacao-amd <jiahui.cao@amd.com> --------- Signed-off-by: jiacao-amd <jiahui.cao@amd.com> Co-authored-by: Jiahui Cao <jiacao@crs-m2m-cpu-spur-014.us-east2-a.compute.internal> (cherry picked from commit aa84815) * [Triton/Gluon] Add SonicMoE pure-Triton grouped GEMM MoE (#5725) * [Triton] Add SonicMoE pure-Triton grouped GEMM MoE Add a pure-Triton, expert-major grouped GEMM MoE with full autograd support (SonicMoE), covering top-k and general routing, GLU/elementwise activations, and blockwise FP8 scales for the grouped GEMM. Includes gfx942/gfx950 tuned configs, unit tests, and a benchmark script. Public entry points: moe_TC_softmax_topk_layer, moe_general_routing_inputs, moe_pre_routed_inputs (aiter.ops.triton.sonicmoe). This is the first of two PRs splitting the SonicMoE contribution; a follow-up PR stacks the hipBLASLt/multistream grouped GEMM backend on top of this pure-Triton implementation. * [Triton] Drop unused SonicMoE stream_id dels and restating comments Keep the positional stream argument for caller compatibility as _stream_id, and remove comments that only restated the next GEMM call. * [Triton] Move SonicMoE autograd API out of _triton_kernels Keep @triton.jit kernels under _triton_kernels/moe/sonicmoe and put the autograd wrappers plus public entry points in aiter.ops.triton.sonicmoe. * [Triton] Use AMD copyright headers on SonicMoE kernels Replace third-party author banners and drop external source-link comments so new files match aiter's SPDX header. * Address SonicMoE review feedback * [Triton] Move SonicMoE routing kernels into moe_routing * [Triton] Collapse SonicMoE host wrappers into one module Keep the public API and tests in a single file instead of a set of sibling wrappers. (cherry picked from commit 2e62094) * [Triton/Gluon] Move the KDA_DECODE configs into the nested config layout (#5941) (cherry picked from commit 92192fd) * [Triton/Gluon] Move _triton_kernels/gated_delta_rule/ to _triton_kernels/gated_delta_net/ (#5943) (cherry picked from commit cf89ffd) * [Triton/Gluon] Move _gluon_kernels/gfx1250/norm/ to _gluon_kernels/gfx1250/normalization/ (#5944) (cherry picked from commit 9ef0d08) * [CI] Select impacted Triton and Gluon unit tests (#5878) Select impacted Triton and Gluon unit tests (cherry picked from commit 32b1cb0) * [Triton/Gluon] [Config] Drop no-op kpack from RDNA GEMM configs (#5917) * [Triton] [Config] Drop no-op kpack from RDNA GEMM configs * [Triton] [Config] Drop redundant matrix_instr_nonkdim from selected RDNA GEMM configs * [Triton] [Config] Extend RDNA matrix_instr_nonkdim cleanup * [Triton] [Config] Retain matrix_instr_nonkdim for shared GEMM consumers (cherry picked from commit 9ef3e41) * mxfp8 gemm cga update, update (#5055) Co-authored-by: Satya Nikhil Kodukula <nikhil.kodukula@gmail.com> (cherry picked from commit 693701c) * stop fused_bmm_rope_kv_cache from using batched_gemm_a8w8_a_per_token_group_prequant_w_per_batched_tensor_quant configs (#5877) (cherry picked from commit a2650cd) * [Triton/Gluon] Add Triton-based Conv3D kernels (#5952) * conv3d implementation on Triton * Remove unnecessary convolution compatibility aliases * Update the README file. * Add Conv2D weight-pack cache cleanup and updated the tests. * Simplify the Conv2D Winograd test guard * Move Conv3D device validation out of shape helper * Add diagnostics to convolution test assertions * Removing benchmark related tests from unit test * Add context to Conv3D numerical assertions * Deduplicate Conv2D and Conv3D blocked-layout kernels * Simplified benchmark and fixed formatting issue on test file * Split Conv2D and Conv3D cache-clear tests into their owning suites * Deduplicate Conv2D and Conv3D Winograd transforms * Deduplicate the Conv2D and Conv3D Winograd filter transform * Consolidate Conv2D Winograd launch paths * Deduplicate Conv2D and Conv3D prepack helpers and cache wrappers * Black formatting * Replace convolution test helper prints with logging * Use safe configs for Conv3D on CDNA (cherry picked from commit c8325e0) --------- Signed-off-by: Jin Tao <jin.tao@amd.com> Signed-off-by: Martin Domanský <ragua@email.cz> Signed-off-by: Rohan138 <rohanpotdar138@gmail.com> Signed-off-by: Pham Binh <phamhuuthanh.binh@amd.com> Signed-off-by: Lin, Soga <soga.lin@amd.com> Signed-off-by: jiacao-amd <jiahui.cao@amd.com> Co-authored-by: siliangchen-amd <SiLiang.Chen@amd.com> Co-authored-by: Cursor <cursoragent@cursor.com> Co-authored-by: Lingpeng Jin <103567126+valarLip@users.noreply.github.com> Co-authored-by: Jesús Carabaño <jcaraban@users.noreply.github.com> Co-authored-by: Leo <drleonid@amd.com> Co-authored-by: liyjiang <liying.jiang@amd.com> Co-authored-by: gbyu-amd <Guanbao.Yu@amd.com> Co-authored-by: Claude Opus 5 <noreply@anthropic.com> Co-authored-by: Jin Tao <jintao12@amd.com> Co-authored-by: jin.tao@amd.com <jin.tao@amd.com@tus1-p15-g43.tus.tensorwave.lan> Co-authored-by: Felix Li <felix.li@amd.com> Co-authored-by: Martin Domanský <8312516+Ragua1@users.noreply.github.com> Co-authored-by: Mehmet Cagri <mehmet.kaymak@amd.com> Co-authored-by: Satya Nikhil Kodukula <nikhil.kodukula@gmail.com> Co-authored-by: Rohan Potdar <rohanpotdar138@gmail.com> Co-authored-by: Nimit Patel <61071220+NimitPtl@users.noreply.github.com> Co-authored-by: Muhammad Ahmed <mm.ahmed2202@gmail.com> Co-authored-by: yifehuan-amd <yifehuan@amd.com> Co-authored-by: Pham Binh <phamhuuthanh.binh@amd.com> Co-authored-by: akii96 <aakif.nawaz@amd.com> Co-authored-by: li,xiangxiang <xiangxli@amd.com> Co-authored-by: jhchouuu <jiahzhou@amd.com> Co-authored-by: HaonanWang98 <hwang@amd.com> Co-authored-by: sogalin_codegen <39478626+sogalin@users.noreply.github.com> Co-authored-by: vorapolsiloai <115975949+vorapolsiloai@users.noreply.github.com> Co-authored-by: jiacao-amd <jiahui.cao@amd.com> Co-authored-by: Jiahui Cao <jiacao@crs-m2m-cpu-spur-014.us-east2-a.compute.internal> Co-authored-by: WuLei-AMD <leiwu@amd.com> Co-authored-by: Hyunjune Kim <132782704+hyjuunn@users.noreply.github.com> Co-authored-by: Shao-Chun Lee <Shao-Chun.Lee@amd.com> Co-authored-by: Saeid Rostami <123997133+saeid-rostami@users.noreply.github.com>
Motivation
Adds an OPUS (HIP) MXFP4 paged MQA-logits kernel for DeepSeek-V4-style sparse-attention
indexers on gfx950, as an alternative to the existing FlyDSL fp4 path.
Two things motivate it beyond raw speed:
nover
floor((pos + 1) / ratio)rows, which noctx - (next_n - 1 - n)derivation canexpress. Both entry points take
local_endsper query row, so any window rule works.persistent
cta_infobuffer and no per-forward schedule kernel to be CUDA-graph safe.Technical Details
Two entry points in
aiter.ops.opus.pa_mqa_logits_opus, both schedule-free, cudagraph-safe,gfx950only:..._prefill— 1D grid, one CTA per row, window[local_starts[r], local_ends[r])...._decode— 3D grid(batch, next_n_max, split_kv), packed rowcu_seq_q[b] + nover[0, local_ends[row]).compute_prefill_windowsbuilds tail-causal windows device-side for either.block_kpicksone of two compiled variants (256 → 4 waves/CTA, 64 → 1 wave, default 64);
H=64,D=128,PAGE=64are compile-time.q [T,H,D/2],weights [T,H]bf16 andkv_cache [num_blocks,4,PAGE,16]are natural /standard-paged,
kv_cachebyte-identical to FlyDSL's. Onlyq_scale [T,2,32,4]andkv_scale [num_blocks,2,32,4]carry a 32x32-MFMA E8M0 permutation.Decode pads
grid.xto a whole number of XCDs so a batch's MTP rows —gridDim.xapart inworkgroup id — share one XCD's L2; the padding CTAs return immediately.
VGPR 223 / 224, occupancy 2, LDS 0, no spill.
Test Plan
gfx950, both compiled variants (
block_k64 and 256), random data.Performance at the shapes ATOM's DeepSeek-V4 indexer runs (CSA ratio 4), OPUS at
block_k=64against FlyDSL at its shippingblock_k=256:total_q=16384— fresh (bs 2/4/8) and chunked into a long sequence(
bs 1/2/4, 25000 committed compressed rows). Geometry follows the reported prefillregression, so the two are directly comparable.
product(max_ctx {1024, 8192}, next_n {1,4,8}, batch {1..128}).Both sides share one set of
q/kv_cache/weights, each gets the scale layout itreads, and each is timed by one
run_perftest(iters=50, warmup=10)— torch-profiler GPUkernel time, identical for both. FlyDSL is handed the precomputed
cta_infoATOM gives it,so neither side builds a schedule inside the timed region. Agreement is checked after
timing, masked to the window.
Test Result
Correctness — all pass.
errischeckAllclose's mismatch fraction, so 0 is exactagreement within
rtol=2e-5:test_pa_mqa_logits_opus.py,block_k=64test_pa_mqa_logits_opus.py,block_k=256vs flydslisnanand omitted from the count on the four decode cases — FlyDSL's decode ophas a different ABI and is compared in the perf harness instead — and on the two
unaligned-start cases, for the reason in the Test Plan.
Resources, clean build:
Performance — faster on every shape measured, agreement 8.6e-08 .. 1.3e-07:
Two caveats worth carrying with the numbers. Decode's small end is a scheduling artifact:
43 of the 48 shapes fall under 9 us once compression shortens the windows, and down there
the margin is mostly FlyDSL's 512-CTA floor under ATOM's
FP4_MQA_PARALLEL_UNIT_NUM. Thethroughput end is the largest shape — 1024 rows over ~2000 compressed columns, 15.15 vs
26.32 us, 2164 vs 1245 TFLOPS. And prefill's fresh column is partly the
block_kchoice,which was already known; the new result is chunked prefill, where the reported long-context
deficit (−8.1 .. −9.7%) is now +7.9 .. +8.4%, i.e. ~3108 against ~2855 TFLOPS.
Each figure is the third of three passes; passes agree to ~1 point on prefill and to a 1.0%
median per-shape drift on decode.
Performance
Prefill
(n+1)//4. Windows 0 -> qlen/4.this chunk is its tail, so row n sees
K - (qlen-1-n)//4. K = 25000 compressed rows isroughly a 100k-raw-token context.
fresh:
chunked, kvlen = 25000 compressed rows:
Decode
next_nnext_n_max = next_n(3D grid)block_tables[bs, pages][total_q, pages]local_endscontext_lenscta_infoprecomputed atP = max(512, total_q)block_kFP4_MQA_BLOCK_K)Submission Checklist