Repository navigation
feat(flydsl): support strided GDN decode state routing - #4573
Merged
Merged
Conversation
Contributor
🏷️ CI GuideRuns automatically on every PR:
Extended tests (opt-in via labels):
|
junna2016
force-pushed
the
xjn_flydsl_decode_stride_indices
branch
from
August 6, 2026 01:34
a9d8f3e to
96c13f7
Compare
xytpai
reviewed
Aug 6, 2026
xytpai
previously approved these changes
Aug 6, 2026
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
force-pushed
the
xjn_flydsl_decode_stride_indices
branch
from
August 6, 2026 05:59
a5f8036 to
38b91eb
Compare
xytpai
approved these changes
Aug 6, 2026
valarLip
approved these changes
Aug 7, 2026
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.
This file contains hidden or bidirectional Unicode text that may be interpreted or compiled differently than what appears below. To review, open the file in an editor that reveals hidden Unicode characters.
Learn more about bidirectional Unicode characters
Sign up for free
to join this conversation on GitHub.
Already have an account?
Sign in to comment
Add this suggestion to a batch that can be applied as a single commit.This suggestion is invalid because no changes were made to the code.Suggestions cannot be applied while the pull request is closed.Suggestions cannot be applied while viewing a subset of changes.Only one suggestion per line can be applied in a batch.Add this suggestion to a batch that can be applied as a single commit.Applying suggestions on deleted lines is not supported.You must change the existing code in this line in order to create a valid suggestion.Outdated suggestions cannot be applied.This suggestion has been applied or marked resolved.Suggestions cannot be applied from pending reviews.Suggestions cannot be applied on multi-line comments.Suggestions cannot be applied while the pull request is queued to merge.Suggestion cannot be applied right now. Please check back later.
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 convertedq,k, andvto contiguous tensors before launching the kernel: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
indicestensor, requiring therecurrent 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:
The kernel constructs
GTensordescriptors using these explicit strides:The kernel can therefore address non-contiguous
q,k, andvviewsdirectly, eliminating the previous
.contiguous()packing operations.The supported layout contract is:
q,k,v,a, andbdo not need to be fully contiguous.qandkmust be contiguous(
stride(-1) == 1) because the kernel uses vectorized loads along the headdimension.
v,a, andbare 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
indicestensor for both loading and storingthe recurrent state:
When a request crossed a cache-block boundary, RTP-LLM had to perform:
The interface now accepts independent
read_indicesandwrite_indices:For each request:
read_indicesidentifies the cache block containing the previous recurrentstate.
write_indicesidentifies the cache block where the updated state is stored.decode kernel.
For example:
This represents:
This design removes the standalone state-copy kernel and its associated launch
and memory-traffic overhead.
Backward compatibility
The existing required
indicesargument and positional argument order arepreserved. When the new arguments are omitted, they fall back to
indices:Existing callers therefore continue to use the original in-place state update
semantics without modification.
Expected Benefits
q,k, andv.Test Plan
The FlyDSL linear-attention tests cover:
torch.split.read_indicesandwrite_indices.