Skip to content

feat(flydsl): support strided GDN decode state routing - #4573

Merged
xytpai merged 3 commits into
ROCm:mainfrom
junna2016:xjn_flydsl_decode_stride_indices
Aug 8, 2026
Merged

xytpai merged 3 commits into
ROCm:mainfrom
junna2016:xjn_flydsl_decode_stride_indices

Conversation

@junna2016

@junna2016 junna2016 commented Aug 5, 2026 •

Copy link
Copy Markdown
Contributor

Accept production strided q/k/v inputs without contiguous packing and route recurrent state through independent read and write block indices while preserving the legacy indices API. Add ROCm coverage for tuned head shapes and BF16/FP32 state.

Motivation

RTP-LLM's production GDN decode inputs may be strided views produced by
operations such as torch.split. The original AITER FlyDSL wrapper converted
q, k, and v to contiguous tensors before launching the kernel:

query.contiguous()
key.contiguous()
value.contiguous()

Although this ensured correctness, it introduced additional packing kernels,
temporary allocations, and memory traffic on the latency-sensitive decode
path.

The original interface also accepted only one indices tensor, requiring the
recurrent state to be read from and written to the same cache block. When a
request crossed a cache-block boundary, RTP-LLM had to launch a separate kernel
to copy the state before executing GDN decode.

This change allows the FlyDSL kernel to consume production strided inputs
directly and perform cross-block state routing inside the GDN decode kernel.

Technical Details

Strided Q/K/V inputs

The wrapper now passes the tensors' actual strides to the FlyDSL kernel
factory:

query.stride()
key.stride()
value.stride()

The kernel constructs GTensor descriptors using these explicit strides:

q_tensor = GTensor(..., stride=q_strides)
k_tensor = GTensor(..., stride=k_strides)
v_tensor = GTensor(..., stride=v_strides)

The kernel can therefore address non-contiguous q, k, and v views
directly, eliminating the previous .contiguous() packing operations.

The supported layout contract is:

  • q, k, v, a, and b do not need to be fully contiguous.
  • The innermost dimensions of q and k must be contiguous
    (stride(-1) == 1) because the kernel uses vectorized loads along the head
    dimension.
  • Outer dimensions may use arbitrary valid strides.
  • v, a, and b are accessed using their explicitly supplied strides.

The strides are compile-time kernel specialization parameters, so the generated
kernel does not require runtime stride branches.

Independent state read and write indices

The original interface used one indices tensor for both loading and storing
the recurrent state:

state[indices] -> GDN decode -> state[indices]

When a request crossed a cache-block boundary, RTP-LLM had to perform:

state[read_indices] -> copy kernel -> state[write_indices]
state[write_indices] -> GDN decode -> state[write_indices]

The interface now accepts independent read_indices and write_indices:

state[read_indices] -> GDN decode -> state[write_indices]

For each request:

  • read_indices identifies the cache block containing the previous recurrent
    state.
  • write_indices identifies the cache block where the updated state is stored.
  • Equal indices retain the original in-place update behavior.
  • Different indices perform cross-block state migration as part of the GDN
    decode kernel.

For example:

read_indices = [5, 8, 11]
write_indices = [6, 8, 12]

This represents:

request 0: state[5]  -> state[6]
request 1: state[8]  -> state[8]
request 2: state[11] -> state[12]

This design removes the standalone state-copy kernel and its associated launch
and memory-traffic overhead.

Backward compatibility

The existing required indices argument and positional argument order are
preserved. When the new arguments are omitted, they fall back to indices:

read_indices = indices if read_indices is None else read_indices
write_indices = indices if write_indices is None else write_indices

Existing callers therefore continue to use the original in-place state update
semantics without modification.

Expected Benefits

  • Eliminates contiguous packing for production strided q, k, and v.
  • Avoids temporary tensors and redundant memory traffic.
  • Eliminates the separate cross-block state-copy kernel.
  • Reduces decode-path kernel launches.
  • Supports independent state routing for each request in a mixed batch.
  • Preserves compatibility with existing AITER callers.

Test Plan

The FlyDSL linear-attention tests cover:

  • Existing contiguous-input behavior.
  • Production-like strided views created with torch.split.
  • Batch sizes greater than one.
  • Independent read_indices and write_indices.
  • In-place and cross-block state updates.
  • BF16 and FP32 recurrent-state dtypes.
  • Seven tuned GDN head-shape configurations on gfx942.
  • Numerical equivalence with the legacy explicit state-copy workflow.
  • Backward compatibility when the new arguments are omitted.

@junna2016
junna2016 requested a review from a team August 5, 2026 08:31
@github-actions

github-actions Bot commented Aug 5, 2026

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 4573 --add-label <label>

@junna2016
junna2016 force-pushed the xjn_flydsl_decode_stride_indices branch from a9d8f3e to 96c13f7 Compare August 6, 2026 01:34
Comment thread aiter/ops/flydsl/kernels/gdr_decode.py
xytpai
xytpai previously approved these changes Aug 6, 2026
@zufayu
zufayu requested a review from yadaish August 6, 2026 02:34
Accept production strided q/k/v inputs without contiguous packing and route recurrent state through independent read and write block indices while preserving the legacy indices API. Add ROCm coverage for tuned head shapes and BF16/FP32 state.
@junna2016
junna2016 force-pushed the xjn_flydsl_decode_stride_indices branch from a5f8036 to 38b91eb Compare August 6, 2026 05:59
@xytpai xytpai added the ci:atom label Aug 6, 2026
@xytpai
xytpai merged commit 7e08acd into ROCm:main Aug 8, 2026
102 of 119 checks passed
waqahmed-amd-fi added a commit that referenced this pull request Aug 8, 2026
#4609 removed the `vector`/`arith` imports the per-channel branch still used,
and the merge left `r_g_vec` outside the scalar `else`, so building a
`gate_mode="kda"` kernel raised NameError. The benchmark's direct kernel call
also needed #4573's q/k/v strides and read/write index split.

Scalar path bit-identical; 47 tests pass on gfx950.
Sign up for free to join this conversation on GitHub. Already have an account? Sign in to comment

Labels

Projects

None yet

Development

Successfully merging this pull request may close these issues.

3 participants