[JAX] Remove unnecessary SWA calculation in _segment_ids_pos_to_seqlens_offsets() - #2201
Merged
KshitijLakhani merged 2 commits intoNov 21, 2025
Conversation
KshitijLakhani
force-pushed
the
klakhani/test/trim-segment_ids_pos_to_seqlens_offsets
branch
from
September 25, 2025 06:00
918e244 to
6a75356
Compare
KshitijLakhani
force-pushed
the
klakhani/test/trim-segment_ids_pos_to_seqlens_offsets
branch
3 times, most recently
from
November 7, 2025 22:40
90c2cbf to
79f92f9
Compare
…ffsets Signed-off-by: Kshitij Lakhani <klakhani@nvidia.com>
KshitijLakhani
force-pushed
the
klakhani/test/trim-segment_ids_pos_to_seqlens_offsets
branch
from
November 21, 2025 02:02
00c779f to
f746bf5
Compare
KshitijLakhani
marked this pull request as ready for review
November 21, 2025 02:02
for more information, see https://pre-commit.ci
Contributor
Greptile OverviewGreptile SummaryRemoved unnecessary sliding window attention (SWA) mask calculation from the Key Changes:
Analysis: Confidence Score: 4/5
Important Files ChangedFile Analysis
Sequence DiagramsequenceDiagram
participant User
participant fused_attn
participant SequenceDescriptor
participant _segment_ids_pos_to_seqlens_offsets
participant _mask_to_seqlens_offset
participant tex.fused_attn_fwd
User->>fused_attn: Call with QKV, mask_type, window_size
fused_attn->>SequenceDescriptor: get_seqlens_and_offsets()
SequenceDescriptor->>_segment_ids_pos_to_seqlens_offsets: segment_ids, segment_pos, attn_mask_type, window_size
Note over _segment_ids_pos_to_seqlens_offsets: Check if fast path applies<br/>(causal without window_size)
alt Fast path: causal without window
_segment_ids_pos_to_seqlens_offsets->>_segment_ids_pos_to_seqlens_offsets: Use fast causal path
else Slow path: BRCM or with window
_segment_ids_pos_to_seqlens_offsets->>_segment_ids_pos_to_seqlens_offsets: Create segment_mask
Note over _segment_ids_pos_to_seqlens_offsets: REMOVED: SWA mask calculation<br/>Window is handled by cuDNN kernel
alt Bottom-right causal
_segment_ids_pos_to_seqlens_offsets->>_segment_ids_pos_to_seqlens_offsets: Apply BRCM mask
else Regular causal
_segment_ids_pos_to_seqlens_offsets->>_segment_ids_pos_to_seqlens_offsets: Apply causal mask
end
_segment_ids_pos_to_seqlens_offsets->>_mask_to_seqlens_offset: attn_mask_with_id
_mask_to_seqlens_offset-->>_segment_ids_pos_to_seqlens_offsets: q_seqlen, kv_seqlen, q_offset, kv_offset
end
_segment_ids_pos_to_seqlens_offsets-->>SequenceDescriptor: seqlens and offsets
SequenceDescriptor-->>fused_attn: seqlens and offsets
fused_attn->>tex.fused_attn_fwd: qkv, seqlens, offsets, window_size
Note over tex.fused_attn_fwd: cuDNN kernel applies<br/>sliding window attention<br/>using window_size parameter
tex.fused_attn_fwd-->>fused_attn: output
fused_attn-->>User: result
|
Collaborator
Author
|
/te-ci jax |
KshitijLakhani
requested review from
jberchtold-nvidia,
mgoldfarb-nvidia,
mingxu1067,
phu0ngng and
zlsh80826
November 21, 2025 07:38
Collaborator
Author
|
CI pipeline passes: 38909687 |
KshitijLakhani
added a commit
that referenced
this pull request
Nov 23, 2025
…ns_offsets() (#2201) * Remove unnecessary SWA calculation from _segment_ids_pos_to_seqlens_offsets Signed-off-by: Kshitij Lakhani <klakhani@nvidia.com> * [pre-commit.ci] auto fixes from pre-commit.com hooks for more information, see https://pre-commit.ci --------- Signed-off-by: Kshitij Lakhani <klakhani@nvidia.com> Co-authored-by: pre-commit-ci[bot] <66853113+pre-commit-ci[bot]@users.noreply.github.com>
KshitijLakhani
added a commit
to KshitijLakhani/TransformerEngine
that referenced
this pull request
Dec 12, 2025
…ns_offsets() (NVIDIA#2201) * Remove unnecessary SWA calculation from _segment_ids_pos_to_seqlens_offsets Signed-off-by: Kshitij Lakhani <klakhani@nvidia.com> * [pre-commit.ci] auto fixes from pre-commit.com hooks for more information, see https://pre-commit.ci --------- Signed-off-by: Kshitij Lakhani <klakhani@nvidia.com> Co-authored-by: pre-commit-ci[bot] <66853113+pre-commit-ci[bot]@users.noreply.github.com>
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.
Description
I noticed that in
_segment_ids_pos_to_seqlens_offsets(), a temp / intermediate mask is created, the purpose of which is only to help with the calculation of seqlens and offsets which is then passed onto the cuDNN backend.This intermediate attn_mask should only need to use the segment ids and pos to create the attn_mask followed by a layer of masking based on the application of the type of mask (causal or brcm) as that would influence the final seqlens and offsets generated.
However, SW mask calculation and application to attn_mask in no way should affect the final seqlens and offsets generated hence this is an unnecessary operation. I tested this out as well and it seems to check out.
I had put in a TODO for this observation and am addressing the same in this PR
Type of change
Checklist: