Skip to content

[Triton/Gluon] [FlyDSL] [CI] feat(flydsl): Add paged-attention Tile kernel - #4332

Merged
coderfeli merged 70 commits into
mainfrom
flydsl-pa-decode-tile
Sep 22, 2026
Merged

coderfeli merged 70 commits into
mainfrom
flydsl-pa-decode-tile

Conversation

@fsx950223

@fsx950223 fsx950223 commented Jul 22, 2026 •

Copy link
Copy Markdown
Contributor

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

  • Add FlyDSL paged-attention decode under aiter.ops.flydsl.
  • Support:
    • FP8 K/V caches
    • BF16/FP16 queries
    • Block sizes 16, 64, and 128
    • Per-tensor and per-token KV scales
    • GQA and multi-token prediction
    • Plain and transposed V-cache layouts
  • Process context in 256-token tiles using four wave64 waves per CTA.
  • Use wide K=128 FP8 MFMA for the tuned MTP3/MTP4 gfx950 path.
  • Split small MTP2/MTP4 grids into one query per CTA; keep larger grids and
    MTP3 fused to share KV/scales.
  • Select the single-M per-token V-prefetch pipeline from one layout/workgroup
    policy; static Hkv2 uses the tuned path while high-grid page128 transposed V
    retains the lower-register schedule.
  • Support explicit partition counts up to 256 and a caller-configurable clamp for
    automatic partition selection.
  • Use the native FlyDSL reducer for partition outputs.

Test Plan

op_tests/test_flydsl_pa_decode.py has one generic parametrized pytest function
with 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.

pytest -q op_tests/test_flydsl_pa_decode.py

HIP_VISIBLE_DEVICES=5 python3 op_tests/test_flydsl_pa_decode.py \
  -d bf16 -b 200 -q 4 -s 16,1,128,200000 \
  --block-size 16 128 --trans-v 0 1 --per-token 1

Latest MI355 validation:

285 passed in 190.50s

Black, Ruff, Python syntax, Bash syntax, and git diff --check also pass.

Performance

Earlier representative cases

MI355/gfx950 CUDA Graph end-to-end latency, including compute and reduce:

Case Block size PA Tile PA Gluon Speedup
small 16 11.79 us 12.24 us 1.04x
small 64 11.75 us 12.07 us 1.03x
batch 16 14.20 us 20.43 us 1.44x
batch 64 13.76 us 20.32 us 1.48x
MTP3 16 22.90 us 41.52 us 1.81x
MTP3 64 22.81 us 41.08 us 1.80x
MTP4-var 16 28.72 us 38.14 us 1.33x
MTP4-var 64 23.91 us 38.53 us 1.61x
long 16 35.93 us 43.69 us 1.22x
long 64 35.24 us 43.52 us 1.23x

PA Tile wins all 10 cases with a 1.37x geometric-mean speedup. Workloads use
Hkv=1, head dimension 128, and transposed V. small is B3/Hq8/Q1/C1027
per-tensor; batch is B81/Hq8/Q1/C1027 per-tensor; MTP3 is
B81/Hq16/Q3/C1027 per-token; MTP4-var is B81/Hq16/Q4/C1027 per-token with
variable KV lengths; and long is B81/Hq16/Q1/C8192 per-token.

MI355: BS=200, MTP4, C=200k

Measured on MI355 (gfx950, 256 CUs) with QL=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 perftest profiler
means (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.

Page V layout FlyDSL auto NP FlyDSL us FlyDSL TB/s Gluon auto NP Gluon us Gluon TB/s Speedup .005
16 rank-4 8 2409.320 4.3898 3 6751.326 1.5666 2.8022x PASS
16 trans-V 8 2372.853 4.4573 3 6719.592 1.5740 2.8319x PASS
128 rank-4 4 2387.181 4.4269 3 6441.091 1.6407 2.6982x PASS
128 trans-V 4 3415.613 3.0940 3 6309.910 1.6748 1.8474x PASS

For 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 .005 check against the reference and also failed their mutual .005
comparison. 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

@fsx950223
fsx950223 requested review from a team and a lite review from Copilot July 22, 2026 08:19
@github-actions

Copy link
Copy Markdown
Contributor

🏷️ CI Guide

Runs automatically on every PR:

  • ✅ Pre-checks (submodule verification, code formatting)
  • ✅ Aiter op tests (gfx942 + gfx950)
  • ✅ Triton tests on MI35X (only when aiter/ops/triton/** or related paths are changed)

Extended tests (opt-in via labels):

Label Tests
ci:triton-300x Run an additional Triton test job on MI300X in PRs; main branch always runs both MI35X and MI300X
ci:sglang SGLang integration tests: DeepSeek-R1-MXFP4 accuracy, Qwen 3.5 accuracy
ci:atom ATOM benchmark: DeepSeek-R1-0528, GPT-OSS-120B
ci:atom_full ATOM accuracy suite for PR and main models from ATOM models_accuracy.json
ci:vllm vLLM benchmark: GPT-OSS-120B, DeepSeek-R1-0528, Kimi-K2.5
ci:all All standard extended tests (excludes ci:atom_full)

Only add ci:atom_full for FlyDSL or Triton upgrades.
Add labels via the sidebar or gh pr edit 4332 --add-label <label>

Copilot AI left a comment

Copy link
Copy Markdown
Contributor

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

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.

Comment thread aiter/ops/flydsl/pa_decode_tile.py Outdated
Comment on lines +1336 to +1345
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)
)
Comment thread aiter/ops/flydsl/pa_decode_tile.py Outdated
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
Comment thread aiter/ops/flydsl/pa_decode_tile.py Outdated
when omitted they are picked/allocated here.
"""
num_seqs = context_lengths.shape[0]
total_q_rows, num_q_heads, head_dim = query.shape
Comment thread aiter/ops/flydsl/pa_decode_tile.py Outdated
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
Comment thread aiter/ops/flydsl/pa_decode_tile.py Outdated
Comment on lines +1370 to +1377
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,
)

Comment on lines +139 to +148
pa_decode_tile(
output,
query,
key_cache,
value_cache,
block_tables,
context_lengths,
key_scale,
value_scale,
)
@zufayu
zufayu requested a review from coderfeli July 23, 2026 01:40
@Bernard-Liu
Bernard-Liu force-pushed the flydsl-pa-decode-tile branch from b9c299c to 94a91cf Compare August 20, 2026 06:39
Copilot AI review requested due to automatic review settings August 20, 2026 06:39
@github-actions github-actions Bot changed the title feat(flydsl): Add paged-attention Tile kernel [FlyDSL] feat(flydsl): Add paged-attention Tile kernel Aug 20, 2026

Copilot AI left a comment

Copy link
Copy Markdown
Contributor

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

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 in op_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 pass device= 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 (uses torch.cuda.* and with torch.cuda.device(dev)), but it never checks query.is_cuda / dev.type == "cuda". Calling this with CPU tensors will fail with a low-level CUDA error instead of a clear ValueError; 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 assert statements for runtime input validation (shapes/dtypes/devices/contiguity). In optimized Python (-O) these checks are stripped entirely, and even without -O they raise AssertionError rather than a stable API error type. Consider replacing these with explicit if ...: raise ValueError/TypeError checks (consistent with other FlyDSL public wrappers like aiter/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_lengths must be int32; key_scale/value_scale must be float32 on 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_tile from aiter.ops.flydsl, but the public export added here is pa_decode. If the intended API is pa_decode, please update the PR description; otherwise consider exporting an alias named pa_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

fsx950223 added a commit that referenced this pull request Aug 21, 2026
…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>
Copilot AI review requested due to automatic review settings August 21, 2026 08:50

Copilot AI left a comment

Copy link
Copy Markdown
Contributor

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

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_tile as an API exported from aiter.ops.flydsl, but this adds only pa_decode; no pa_decode_tile attribute is defined or included in __all__, so from aiter.ops.flydsl import pa_decode_tile fails. 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_gluon uses ps to select the partitioned kernel/reducer, but this drop-in wrapper silently deletes the argument. A caller passing ps=False therefore still receives the PS kernel and C++ PS reduction, with no indication that the requested mode was ignored. Reject ps=False explicitly 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; the f16 query specialization selected in pa_decode.py:174-177 and pa_decode_tile.py:104-108 is 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 output as 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. Add output to 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 adds op_tests/test_flydsl_pa_decode.py and 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-D trans_v=True branch in pa_decode_tile.py:682-687 is 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()

Comment on lines +500 to +502
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)
Comment on lines +553 to +557
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)
Comment on lines +761 to +765
# 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.
Comment thread aiter/ops/flydsl/pa_decode.py Outdated
Comment on lines +338 to +343
_run_compiled(
compiled["launch"],
output,
pmax.view(-1),
psum.view(-1),
pout.view(-1),
Comment thread aiter/ops/flydsl/pa_decode.py Outdated
Comment on lines +139 to +144
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,
Copilot AI review requested due to automatic review settings September 8, 2026 07:06
@github-actions github-actions Bot changed the title [FlyDSL] feat(flydsl): Add paged-attention Tile kernel [Triton/Gluon] [FlyDSL] feat(flydsl): Add paged-attention Tile kernel Sep 8, 2026

Copilot AI left a comment

Copy link
Copy Markdown
Contributor

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

🔵 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

Comment thread aiter/ops/flydsl/pa_decode.py Outdated
Comment thread op_tests/test_flydsl_pa_decode.py Outdated
Comment thread aiter/ops/flydsl/pa_decode.py Outdated
Comment thread aiter/ops/flydsl/kernels/pa_decode_tile.py Outdated
Comment thread aiter/ops/flydsl/kernels/pa_decode_tile.py Outdated
@fsx950223
fsx950223 force-pushed the flydsl-pa-decode-tile branch from 85b2374 to d8b00e0 Compare September 11, 2026 06:14
@github-actions github-actions Bot changed the title [Triton/Gluon] [FlyDSL] feat(flydsl): Add paged-attention Tile kernel [Triton/Gluon] [FlyDSL] [CI] feat(flydsl): Add paged-attention Tile kernel Sep 13, 2026
@github-actions github-actions Bot added the CI label Sep 13, 2026
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.
@github-actions github-actions Bot changed the title [Triton/Gluon] [FlyDSL] [CI] feat(flydsl): Add paged-attention Tile kernel [Triton/Gluon] [HIP] [FlyDSL] feat(flydsl): Add paged-attention Tile kernel Sep 17, 2026
@github-actions github-actions Bot added the HIP label Sep 17, 2026
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>
@github-actions github-actions Bot changed the title [Triton/Gluon] [HIP] [FlyDSL] feat(flydsl): Add paged-attention Tile kernel [Triton/Gluon] [FlyDSL] [CI] feat(flydsl): Add paged-attention Tile kernel Sep 21, 2026
@github-actions github-actions Bot removed the HIP label Sep 21, 2026
@fsx950223 fsx950223 removed the ci:atom label Sep 21, 2026
@coderfeli
coderfeli merged commit 94dca7b into main Sep 22, 2026
70 checks passed
@coderfeli
coderfeli deleted the flydsl-pa-decode-tile branch September 22, 2026 04:49
Sign up for free to join this conversation on GitHub. Already have an account? Sign in to comment

Projects

None yet

Development

Successfully merging this pull request may close these issues.

6 participants