[Triton/Gluon] [FlyDSL] [CI] feat(flydsl): Add paged-attention Tile kernel - #4332
Conversation
🏷️ CI GuideRuns automatically on every PR:
Extended tests (opt-in via labels):
|
There was a problem hiding this comment.
Pull request overview
This PR adds a new FlyDSL implementation of paged-attention FP8 decode using tile-programming, plus supporting FlyDSL kernel utilities and a basic correctness test.
Changes:
- Added
aiter.ops.flydsl.pa_decode_tile.pa_decode_tile(and its FlyDSL kernel generator) for FP8 paged-attention decode. - Introduced consolidated FlyDSL memory helpers (
mem_ops.py) and expanded kernel utilities (e.g., DPP xor helper, utils re-exports). - Added a correctness test for the new PA decode tile kernel on supported ROCm GPU architectures.
Reviewed changes
Copilot reviewed 6 out of 6 changed files in this pull request and generated 6 comments.
Show a summary per file
| File | Description |
|---|---|
| op_tests/flydsl_tests/test_pa_decode_tile.py | Adds a correctness test for pa_decode_tile (currently covers per-tensor scale path). |
| aiter/ops/flydsl/pa_decode_tile.py | Implements the FlyDSL paged-attention decode tile kernel and the Python host wrapper pa_decode_tile(). |
| aiter/ops/flydsl/kernels/utils.py | Adds kernel utility helpers and re-exports memory helpers from mem_ops for backward compatibility. |
| aiter/ops/flydsl/kernels/mem_ops.py | Introduces a consolidated module for pointer arithmetic, global load/store, and atomic helpers. |
| aiter/ops/flydsl/kernels/dpp_utils.py | Adds dpp_xor_f32 helper for 16-lane XOR DPP patterns. |
| aiter/ops/flydsl/init.py | Exposes pa_decode_tile as part of the FlyDSL public API when FlyDSL is available. |
💡 Add Copilot custom instructions for smarter, more guided reviews. Learn how to get started.
| key_scale_t = ( | ||
| key_scale | ||
| if isinstance(key_scale, torch.Tensor) | ||
| else torch.tensor([float(key_scale)], device=dev) | ||
| ) | ||
| value_scale_t = ( | ||
| value_scale | ||
| if isinstance(value_scale, torch.Tensor) | ||
| else torch.tensor([float(value_scale)], device=dev) | ||
| ) |
| assert ( | ||
| context_lengths.dtype == torch.int32 | ||
| ), f"context_lengths must be int32, got {context_lengths.dtype}" | ||
| query_group_size = num_q_heads // num_kv_heads |
| when omitted they are picked/allocated here. | ||
| """ | ||
| num_seqs = context_lengths.shape[0] | ||
| total_q_rows, num_q_heads, head_dim = query.shape |
| query_length = total_q_rows // num_seqs | ||
| _, num_kv_heads, num_hgroups, block_size, hgroup_width = key_cache.shape | ||
|
|
||
| assert num_hgroups == head_dim // 16 and hgroup_width == 16 |
| if not num_partitions: | ||
| blocks_per_partition = KV_COMPUTE_BLOCK // block_size | ||
| num_partitions = get_recommended_splits( | ||
| num_seqs, | ||
| num_kv_heads, | ||
| split_kv_blocks=blocks_per_partition, | ||
| ) | ||
|
|
| pa_decode_tile( | ||
| output, | ||
| query, | ||
| key_cache, | ||
| value_cache, | ||
| block_tables, | ||
| context_lengths, | ||
| key_scale, | ||
| value_scale, | ||
| ) |
b9c299c to
94a91cf
Compare
There was a problem hiding this comment.
Pull request overview
Copilot reviewed 7 out of 7 changed files in this pull request and generated no new comments.
Suppressed comments (5)
op_tests/flydsl_tests/test_pa_decode_tile.py:22
torch.set_default_device("cuda")at module scope leaks global state across the entire pytest session (see the rationale/mitigation inop_tests/triton_tests/attention/test_mha_dao_ai.py:22-41). Consider using a local context manager/fixture to set + restore the default device, and passdevice=explicitly for tensors created in this module so test ordering can’t change behavior.
torch.set_default_device("cuda")
aiter/ops/flydsl/pa_decode.py:104
pa_decode()assumes CUDA/HIP tensors (usestorch.cuda.*andwith torch.cuda.device(dev)), but it never checksquery.is_cuda/dev.type == "cuda". Calling this with CPU tensors will fail with a low-level CUDA error instead of a clearValueError; please add an early device check similar to other FlyDSL public APIs (e.g.flydsl_flash_attn_func).
arch = get_gfx_runtime()
expected_fp8_dtype = {
"gfx942": torch.float8_e4m3fnuz,
"gfx950": torch.float8_e4m3fn,
}.get(arch)
if expected_fp8_dtype is None:
raise NotImplementedError(
f"pa_decode only supports gfx942 and gfx950, got {arch}"
)
aiter/ops/flydsl/pa_decode.py:132
- This function uses many
assertstatements for runtime input validation (shapes/dtypes/devices/contiguity). In optimized Python (-O) these checks are stripped entirely, and even without-Othey raiseAssertionErrorrather than a stable API error type. Consider replacing these with explicitif ...: raise ValueError/TypeErrorchecks (consistent with other FlyDSL public wrappers likeaiter/ops/flydsl/fmha_kernels.py).
num_seqs = context_lengths.shape[0]
total_q_rows, num_q_heads, head_dim = query.shape
assert total_q_rows == num_seqs * query_length, (
f"query.shape[0] ({total_q_rows}) must equal "
f"context_lengths.shape[0] * query_length ({num_seqs} * {query_length})"
)
assert output.shape == query.shape, (
f"output shape {tuple(output.shape)} must match "
f"query shape {tuple(query.shape)}"
)
aiter/ops/attention.py:125
_can_use_flydsl_pa_decode()doesn’t validate some preconditions that the FlyDSL implementation enforces (e.g.block_tables/context_lengthsmust beint32;key_scale/value_scalemust befloat32on the right device and either scalar or[num_blocks,num_kv_heads,block_size{,1}]). As written, callers can satisfy the guard but still hit an error inside the FlyDSL path even though the fallback would work. Consider extending the guard to check these dtypes/shapes (or normalize/cast scales before deciding).
and query.shape[-1] % 64 == 0
and key_cache.shape[1] > 0
and query.shape[1] % key_cache.shape[1] == 0
and context_lengths.dim() == 1
and context_lengths.is_contiguous()
and block_tables.dim() == 2
and block_tables.is_contiguous()
and query.shape[0] == context_lengths.shape[0] * query_length
and 1 <= max_context_partition_num <= 64
and context_partition_size == 256
and compute_type == runtime_fp8_dtype
and query_scale is None
and (key_scale is None) == (value_scale is None)
and scratch_buffers_match()
and alibi_slopes is None
and sinks is None
and sliding_window == 0
)
aiter/ops/flydsl/init.py:63
- The PR description says to add/export
pa_decode_tilefromaiter.ops.flydsl, but the public export added here ispa_decode. If the intended API ispa_decode, please update the PR description; otherwise consider exporting an alias namedpa_decode_tile(or adjusting the module naming) so the code matches the stated deliverable.
from .mla_reduce_kernels import flydsl_mla_reduce_v1
from .moe_kernels import flydsl_moe_stage1, flydsl_moe_stage2
from .pa_decode import pa_decode
…ockers Renames the public entry point added by #4332 and addresses the correctness items raised in review. Dispatch - `pa_decode_gluon` in aiter/ops/attention.py goes back to main's plain import + registration; the FlyDSL kernel is registered separately as `torch.ops.aiter.pa_decode_flydsl`. The eligibility gate (`_can_use_flydsl_pa_decode`) and the Gluon fallback are removed, so nothing is silently rerouted and no input can be captured into FlyDSL and then rejected by an assert. Kernel - Reject head_dim whose head_dim//16 is above 8 and not a multiple of 8 (192/320/448/...). Q staging fetches each lane's chunk in <=8-wide pieces, so those widths silently dropped the tail: measured 7.8% error at 192 and 97.6% at 320 against a torch reference. - Fold the per-row score scale in before the -inf mask in the phase-split path. An all-zero query row quantises to scale 0, and `-inf * 0` reached exp2 and turned the whole output into NaN (4096/4096 at ql=4). max() commutes with a non-negative scale, so the row max is unchanged. - Restrict the per-token pv_max reduction to context-visible tokens, so an extreme scale on an unwritten slot cannot shrink valid probabilities to zero in the fp8 pack. Uses context_len rather than the per-row causal bound, which keeps pv_max independent of the query row. - Pin block-table entries past the sequence's extent to block 0. They hold whatever the caller left there, often a stale id into another sequence's live page, and the PV matmul still multiplies those V bytes because a zero probability is not an annihilator (0 * NaN == NaN). - Widen the K/V element offsets to 64-bit once a cache tensor passes 2 GiB, where the i32 page-stride product wraps; reproduced as a GPU memory-access fault at 2.00 GiB, correct afterwards. Gated on the cache size because the wider math costs ~18% at block_size=64. - Drop the device_index cache key. Move arith.select, arith.constant and arith.constant_vector to the fx surface, and replace the per-op arith.FastMathFlags.contract plumbing with an ambient `CompilationContext.compile_hints({"fastmath": "contract"})` around kernel emission -- byte-identical IR (103 contract / 24 nnan flags on both the single- and multi-M-tile paths). `fm_nnan` stays explicit, since an explicit fastmath arg wins over the ambient hint. Tests - Move op_tests/flydsl_tests/test_pa_decode_tile.py to op_tests/test_flydsl_pa_decode.py so `find op_tests -maxdepth 1` in split_tests.sh collects it; neither the aiter nor the triton shard saw the old path. - Add adversarial cases for each fix above, driven through torch.ops.aiter.pa_decode_flydsl, with explicit isfinite assertions. - Scope the module-level torch.set_default_device("cuda") into a fixture now that this file shares a pytest session. Verified on gfx942 with FlyDSL 0.3.1: 22 passed, 1 skipped. Co-Authored-By: Claude Opus 5 (1M context) <noreply@anthropic.com>
There was a problem hiding this comment.
Pull request overview
Copilot reviewed 7 out of 7 changed files in this pull request and generated 6 comments.
Suppressed comments (6)
aiter/ops/flydsl/init.py:62
- The PR description specifies
pa_decode_tileas an API exported fromaiter.ops.flydsl, but this adds onlypa_decode; nopa_decode_tileattribute is defined or included in__all__, sofrom aiter.ops.flydsl import pa_decode_tilefails. Export an alias if that is the intended public name, or update the stated API contract and tests.
from .pa_decode import pa_decode
aiter/ops/flydsl/pa_decode.py:121
pa_decode_gluonusespsto select the partitioned kernel/reducer, but this drop-in wrapper silently deletes the argument. A caller passingps=Falsetherefore still receives the PS kernel and C++ PS reduction, with no indication that the requested mode was ignored. Rejectps=Falseexplicitly or implement and forward the non-PS path.
# ``ps`` is retained for drop-in API compatibility. Both partitioning
# policies are represented by the caller-provided partition count.
del ps
op_tests/test_flydsl_pa_decode.py:430
- [verified] The only pytest correctness case hardcodes
dtype=dtypes.bf16; thef16query specialization selected inpa_decode.py:174-177andpa_decode_tile.py:104-108is not exercised by the automated test suite. A bad FP16 load or conversion path would therefore go undetected. Author must parameterize the correctness test over both BF16 and FP16, or add an equivalent FP16 test.
result = run_pa_decode_tile_case(
batch_size=3,
num_query_heads=8,
num_kv_heads=1,
head_dim=128,
context_length=257,
block_size=block_size,
dtype=dtypes.bf16,
aiter/ops/flydsl/pa_decode.py:213
- This wrapper passes
outputas a raw pointer to the kernel, but the validation loop below does not require it to be contiguous. The kernel has no output strides, so a non-contiguous output is written with the wrong layout; with multiple partitions,output.reshape(...)can instead create a copy and the reducer updates that copy, leaving the caller's tensor unchanged. Addoutputto the contiguity checks (or implement a strided/copy-back path).
for name, tensor in (
("key_cache", key_cache),
("value_cache", value_cache),
("block_tables", block_tables),
("context_lengths", context_lengths),
op_tests/test_flydsl_pa_decode.py:4
- The PR description points its test command at
op_tests/flydsl_tests/test_pa_decode_tile.py, but this change addsop_tests/test_flydsl_pa_decode.pyand no file exists at the documented path. As written, the advertised pytest command cannot discover these tests; please update the description or move the test to the stated path.
"""Correctness and performance sweep for FlyDSL paged-attention Tile."""
op_tests/test_flydsl_pa_decode.py:327
- [verified] Every numerical case here passes
value_quant.contiguous()as a 4-D plain V cache, so the 5-Dtrans_v=Truebranch inpa_decode_tile.py:682-687is never exercised. The PR's performance claim specifically uses transposed V, so an incorrect transposed offset can pass all current correctness tests. Author must add a correctness case with the 5-D transposed layout and an independent reference.
value_cache = value_quant.contiguous()
| num_tiles_m1 = num_tiles - 1 | ||
| start_safe = (part_start < num_tiles).select(part_start, num_tiles_m1) | ||
| k_pf0, phys_vec0 = _k_ops_flat(start_safe) |
| for i in range_constexpr(n // 4): | ||
| b = i * 4 | ||
| lo = fx.rocdl.cvt_pk_fp8_f32(T.i32, vf32[b], vf32[b + 1], 0, False) | ||
| words.append( | ||
| fx.rocdl.cvt_pk_fp8_f32(T.i32, vf32[b + 2], vf32[b + 3], lo, True) |
| # scales still reach the PV MFMA / the fp8 normalization, so both | ||
| # have to be neutralised here. `context_len` (not the per-row causal | ||
| # bound) is the right cutoff: causally-masked tokens inside the | ||
| # context hold real data, and using it keeps both masks independent | ||
| # of the query row. |
| _run_compiled( | ||
| compiled["launch"], | ||
| output, | ||
| pmax.view(-1), | ||
| psum.view(-1), | ||
| pout.view(-1), |
| q_chunk = head_dim // 16 | ||
| if q_chunk > 8 and q_chunk % 8 != 0: | ||
| raise NotImplementedError( | ||
| f"pa_decode does not support head_dim={head_dim}: head_dim//16 " | ||
| f"({q_chunk}) must be <=8 or a multiple of 8" | ||
| ) |
| query_length: int, | ||
| max_context_partition_num: int, | ||
| context_partition_size: int = 256, | ||
| compute_type: torch.dtype = torch.bfloat16, |
There was a problem hiding this comment.
🔵 Needs a closer look
It introduces a large, low-level GPU kernel and changes reduction dispatch plumbing, which warrants careful human validation beyond automated review.
Review details
- Files reviewed: 8/8 changed files
- Comments generated: 2
- Review effort level: Lite
85b2374 to
d8b00e0
Compare
Share input generation, the FP32 reference, and launch helpers with the CLI while retaining sliding-window, sinks, plan, and graph-replay coverage. Validated on dzm-mi355-gpu34-dev: 72 pytest cases and 10 CLI smoke configurations passed.
Select query splitting and V prefetching in the compiler entry point while caching only the effective specialization. Clarify FP8 operand-pack versus MFMA instruction K and verify scheduling boundaries and cache reuse in the generic decode test.
Remove ps from the core decode API and forward work_plan through the Python wrapper. Place the wrapper's ignored ps argument before sinks and align the shared unit-test launch helper with that order. Update the API documentation. Validation: - Ruff and Black pass for the changed Python files. - Eight CPU helper checks preserve tensor, sink, window, and plan forwarding. - MI355 pytest remains blocked during collection: Torch infer_schema rejects the wrapper's PADecodePlan | None parameter. No GPU numerical tests passed for this API revision.
Torch schema inference rejects PADecodePlan and prevents importing aiter. Keep work_plan on the Python PA decode API and omit it when registering the static torch.ops entry point. Add an opt-in registration option for optional trailing Python arguments. Infer the schema from a separate signature while preserving annotation globals, defaults, and the existing dispatcher implementation. Validation: - 72 PA decode tests passed on dzm-mi355-gpu34-dev (gfx950, Torch 2.10, FlyDSL 0.3.2), including static/planned execution and graph replay. - Torch 2.8 CPU schema/dispatcher checks passed. - Actual aiter import passed with a simulated FlyDSL backend import failure. - Ruff, Black, and git diff --check passed for the changed files.
Derive planned specializations from capacity and visible tiles, and widen the M1 and batch-first selectors using continuous ranges. Simplify compiler caching and returned metadata, refine planned scratch and reduction paths, expand the generic unit test, and remove the obsolete plan document. Validation: 176 pytest cases passed on dzm-mi355-gpu34-dev (gfx950). Ruff and staged diff checks passed. Committed Python sources match the GPU-tested source archive byte-for-byte. Known performance caveat: B12/W8192 batch-first remains about 13.3% slower than flat for decode+reduce in the repeat measurement; plan refresh is excluded from those timings.
Match Black 26.5.1 used by the Checks workflow. Only assertion layout and line wrapping change; the Python AST is unchanged. Validation: Black --check --diff across the worktree, Ruff check for the changed file, and git diff --check passed.
Generate the existing 88 cases from grouped parameter tables and keep one shared correctness/contract flow for static and planned decode. Remove fake-CU selector/cache mirrors and retain only explicit query split/address-width overrides. Preserve the reference, plan oracle, sinks/empty-row checks, poisoned scratch, graph replays and CLI behavior. Validation: all 176 pytest cases passed on dzm-mi355-gpu34-dev (gfx950); the CLI help smoke test, Black 26.5.1 and Ruff checks passed. Independent CPU checks confirmed all case parameters and common helpers are unchanged.
Default the per-sequence plan limit to the context device's CU count and preserve explicit smaller limits when refreshing a plan. Validate planned decode and reduction against this device limit while retaining the static 256-partition bound. Extend the generic PA decode unit test with default-CU plans, out-of-range limits, refresh reuse, and long-context reduction coverage. Verified 180 test points on gfx950 and two numerical cases with a simulated host CU304 limit on the same real GPU. Signed-off-by: fsx950223 <fsx950223@outlook.com>
Condense repetitive kernel, planner, reducer, API, and test documentation while retaining numerical, masking, synchronization, and calling-contract notes. No executable logic or test parameters changed. Signed-off-by: fsx950223 <fsx950223@outlook.com>
Motivation
Add a FlyDSL paged-attention decode implementation to AITER for FP8 KV caches.
The kernel provides explicit control over MFMA execution, KV-page loading, LDS
layout, softmax, partition scheduling, and native FlyDSL partition reduction.
Technical Details
aiter.ops.flydsl.MTP3 fused to share KV/scales.
policy; static Hkv2 uses the tuned path while high-grid page128 transposed V
retains the lower-register schedule.
automatic partition selection.
Test Plan
op_tests/test_flydsl_pa_decode.pyhas one generic parametrized pytest functionwith 14 real-kernel/reference configurations covering QL1/2/3/4, page16/64/128,
plain/transposed V, BF16/FP16, scalar/per-token scales, automatic/explicit
partitions, Hkv2 prefetch policy boundaries, C200k, and a BS200 smoke case. The
same file also covers dynamic planning, graph replay, high-partition reducers,
and provides a direct correctness + performance CLI for larger sweeps.
Latest MI355 validation:
Black, Ruff, Python syntax, Bash syntax, and
git diff --checkalso pass.Performance
Earlier representative cases
MI355/gfx950 CUDA Graph end-to-end latency, including compute and reduce:
PA Tile wins all 10 cases with a 1.37x geometric-mean speedup. Workloads use
Hkv=1, head dimension 128, and transposed V.
smallis B3/Hq8/Q1/C1027per-tensor;
batchis B81/Hq8/Q1/C1027 per-tensor;MTP3isB81/Hq16/Q3/C1027 per-token;
MTP4-varis B81/Hq16/Q4/C1027 per-token withvariable KV lengths; and
longis B81/Hq16/Q1/C8192 per-token.MI355: BS=200, MTP4, C=200k
Measured on MI355 (
gfx950, 256 CUs) withQL=4,Hq/Hkv=16/1,D=128,BF16 Q/O, FP8 E4M3FN KV, FP32 per-token KV scales, dense attention, and the
native FlyDSL reducer. Each value is the median of three
perftestprofilermeans (101 iterations, 2 warmups). Bandwidth is decimal logical TB/s; KV and
scales are counted once per PA call, not multiplied by MTP4.
Equal-length context, backend auto NP
Every sequence has context length 200,000. FlyDSL and Gluon use the same input,
block table, and reference within each row. Each backend uses its own automatic
partition selection: FlyDSL chooses NP8 for page16 and NP4 for page128, while
Gluon chooses NP3.
.005For page128 rank-4 V, a separate controlled FlyDSL NP3/4/5/8 tuning sweep
measured 3.4866/4.4269/5.1639/4.4958 TB/s. Explicit NP5 is faster than auto
NP4 for that layout, but it is a manual tuning result and is not used in the
auto-NP comparison table above.
The earlier page-16 variable-length measurements showed 2.1779x (rank-4 V) and
2.1964x (trans-V) speedups, but both implementations failed the additional
strict
.005check against the reference and also failed their mutual.005comparison. Those rows are excluded from the performance claims above. Their
total context was 19,500,374 tokens, versus 40,000,000 for the equal-length rows.
Submission Checklist