Skip to content

Add better ordering enforcment to split_overlap_rs gemms. - #2056

Merged
ptrendx merged 3 commits into
NVIDIA:mainfrom
chaseblock:main
Apr 22, 2026
Merged

ptrendx merged 3 commits into
NVIDIA:mainfrom
chaseblock:main

Conversation

@chaseblock

Copy link
Copy Markdown
Contributor

This adds a short delay kernel to the split_overlap_rs function, which ensures that the gemms are properly ordered when run with cuda graphs.

Description

In some situations, such as when running a TP4 config with comm/gemm overlap and cuda graphs, CG will reorder the GEMM kernels in the CommOverlapP2PBase::split_overlap_rs function, but will not reorder the relevant communication. This leads to exposed communication (see below for a simplified illustration). This is because the GEMMs are issued on (in this case) three separate streams, without dependencies that enforce the proper order.

Stream T1 T2 T3 T4 T5 T6
1 GemmChunk1 GemmChunk4
2 GemmChunk2
3 GemmChunk3
4 Comm1 Comm 2 Comm 3

We do not want to specify dependencies between each of these GEMM kernels, since that would prevent their prologs/epilogs from overlapping. Additionally, CUBLAS does not support LaunchCompletionEvents. This PR takes a different approach (explained below).

Type of change

  • Documentation change (change only to the documentation, either a fix or a new content)
  • Bug fix (non-breaking change which fixes an issue)
  • New feature (non-breaking change which adds functionality)
  • Breaking change (fix or feature that would cause existing functionality to not work as expected)
  • Infra/Build change
  • Code refactoring

Changes

This PR adds:

  1. Dependencies between every other GEMM in split_overlap_rs (i.e., 1->3, 2->4) to partially order these operations.
  2. A small delay kernel, which is issued in parallel with the first gemm chunk, and which has negligible runtime. This way, we can create a dependency between this delay and the second GEMM chunk. This ensures that the first gemm chunk actually issues first. The kernel itself is empty, but launching this empty kernel results in a 2-4 us delay.

This changes the kernel ordering in the table above to:

Stream T1 T2 T3 T4
1 GemmChunk1 GemmChunk4
2 GemmChunk2
3 GemmChunk3
4 TinyDelayKernel Comm1 Comm2 Comm3

Checklist:

  • I have read and followed the contributing guidelines
  • The functionality is complete
  • I have commented my code, particularly in hard-to-understand areas
  • I have made corresponding changes to the documentation
  • My changes generate no new warnings
  • I have added tests that prove my fix is effective or that my feature works
  • New and existing unit tests pass locally with my changes

This adds a short delay kernel to the split_overlap_rs function, which
ensures that the gemms are properly ordered when run with cuda graphs.

Signed-off-by: Chase Block <cblock@nvidia.com>
@ptrendx

ptrendx commented Aug 26, 2025

Copy link
Copy Markdown
Member

@chaseblock One question - what is preventing the CG from reordering the GemmChunk1 with the GemmChunk2 in the new scheme? If I understand correctly there is no dependency between those and that reordering could still lead to (albeit smaller than the original) exposed communication.

@chaseblock

Copy link
Copy Markdown
Contributor Author

@chaseblock One question - what is preventing the CG from reordering the GemmChunk1 with the GemmChunk2 in the new scheme? If I understand correctly there is no dependency between those and that reordering could still lead to (albeit smaller than the original) exposed communication.

@ptrendx You're correct that there's no dependency between these two. This is necessary in order to allow neighboring gemm kernels to partially overlap during their prolog/epilog. The ordering between GemmChunk1 and GemmChunk2 in this solution is instead enforced via the "TinyDelayKernel". This TinyDelayKernel and GemmChunk1 both become ready to launch at the same time (i.e., share the same dependencies). However, because GemmChunk2 depends on TinyDelayKernel, it cannot launch immediately, whereas GemmChunk1 can, so GemmChunk1 launches first.

I admit that this is a bit of a hack, but since cublas doesn't provide an API to record launchCompletionEvents, I think this is the best we can do for now.

In practice, this seems to fix the problem, and I see perf uplift.

@nvMelissa nvMelissa added the community-contribution PRs from external contributor outside the core maintainers, representing community-driven work. label Oct 9, 2025
@ptrendx

ptrendx commented Apr 21, 2026

Copy link
Copy Markdown
Member

/te-ci L1

@greptile-apps

greptile-apps Bot commented Apr 21, 2026

Copy link
Copy Markdown
Contributor

Greptile Summary

This PR fixes GEMM reordering under CUDA Graph execution in CommOverlapP2PBase::split_overlap_rs by introducing (1) an empty tiny_delay_kern launched on _stream_send[0] to ensure GEMM chunk 0 is always issued before chunk 1, and (2) skip-one event dependencies (0→2, 1→3, …) that partially order the per-stream GEMMs without fully serializing them.

Confidence Score: 4/5

Safe to merge after addressing the missing kernel-launch error check; the ordering logic is correct under standard CUDA event semantics.

One P1 finding: the tiny_delay_kern launch in userbuffers_tiny_delay has no cudaGetLastError() check, so a failed launch would leave downstream cudaStreamWaitEvent calls waiting on an unrecorded event. All other findings are P2 style/documentation concerns. The core ordering logic is sound.

transformer_engine/common/comm_gemm_overlap/userbuffers/userbuffers.cu — missing error check on kernel launch

Important Files Changed

Filename Overview
transformer_engine/common/comm_gemm_overlap/comm_gemm_overlap.cpp Adds tiny-delay launch + multi-event ordering logic (record/wait pairs for i>=1) to split_overlap_rs; correctness relies on CUDA CPU-call-time event-capture semantics
transformer_engine/common/comm_gemm_overlap/userbuffers/userbuffers.cu Adds empty tiny_delay_kern GPU kernel and userbuffers_tiny_delay host wrapper; kernel launch is missing a CUDA error check
transformer_engine/common/comm_gemm_overlap/userbuffers/userbuffers.h Adds declaration for userbuffers_tiny_delay; straightforward header change

Sequence Diagram

sequenceDiagram
    participant SM as stream_main
    participant SS0 as stream_send[0]
    participant SC0 as stream_compute[0]
    participant SC1 as stream_compute[1]
    participant SC2 as stream_compute[2]

    SM->>SS0: EventRecord(_start_compute)
    SM->>SC0: StreamWaitEvent(_start_compute)
    SM->>SC1: StreamWaitEvent(_start_compute)
    SM->>SC2: StreamWaitEvent(_start_compute)

    SS0->>SS0: tiny_delay_kern (empty)
    SS0->>SS0: EventRecord(_start_compute)

    SC0->>SC0: GEMM chunk 0
    Note over SC1: i=1: StreamWaitEvent(tiny_delay)
    SC1->>SC1: GEMM chunk 1

    Note over SC2: i=2: EventRecord(_start_compute on SC0) / StreamWaitEvent - waits for GEMM 0
    SC2->>SC2: GEMM chunk 2

    Note over SC0: i=3: EventRecord(_start_compute on SC1) / StreamWaitEvent - waits for GEMM 1
    SC0->>SC0: GEMM chunk 3

    SC0->>SS0: Comm 1 (via _start_comm)
    SC1->>SS0: Comm 2
    SC2->>SS0: Comm 3
Loading

Reviews (1): Last reviewed commit: "Merge branch 'main' into main" | Re-trigger Greptile

// an empty kernel seems to work well enough.
}

void userbuffers_tiny_delay(cudaStream_t stream) { tiny_delay_kern<<<1, 1, 0, stream>>>(); }

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.

P1 Missing CUDA error check after kernel launch

tiny_delay_kern launch errors are silently dropped. A misconfigured launch (e.g., bad stream handle) would leave _start_compute unrecorded, causing every subsequent cudaStreamWaitEvent that waits on it to hang indefinitely. Add a cudaGetLastError() check immediately after the launch.

Suggested change
void userbuffers_tiny_delay(cudaStream_t stream) { tiny_delay_kern<<<1, 1, 0, stream>>>(); }
void userbuffers_tiny_delay(cudaStream_t stream) {
tiny_delay_kern<<<1, 1, 0, stream>>>();
NVTE_CHECK_CUDA(cudaGetLastError());
}

Comment on lines +1176 to +1182
if (i == 1) {
NVTE_CHECK_CUDA(cudaStreamWaitEvent(_stream_compute[stream_id], _start_compute));
} else if (i > 1) {
NVTE_CHECK_CUDA(
cudaEventRecord(_start_compute, _stream_compute[(i - 2) % _stream_compute.size()]));
NVTE_CHECK_CUDA(cudaStreamWaitEvent(_stream_compute[stream_id], _start_compute));
}

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.

P2 Correctness relies on CPU-enqueue-time event capture semantics

The loop mutates _start_compute in-place: for i==1 the wait captures the tiny-delay recording, but for every i>1 the event is re-recorded on a different stream before the wait is enqueued. This is correct only because CUDA's cudaStreamWaitEvent captures the event's completion generation at the CPU call site (not at GPU execution time). If the GPU were to re-read the latest generation when it executes the wait instruction, i==1's wait could end up waiting on the i==3 recording (which is enqueued on _stream_compute[1] itself, after its own GEMM), creating a deadlock.

The logic is sound under the standard CUDA model, but a short comment at the i==1 branch explaining "the wait is bound to the tiny-delay recording, not a later re-recording" would prevent future readers from accidentally breaking this assumption.

@ptrendx
ptrendx merged commit f2ed86b into NVIDIA:main Apr 22, 2026
51 of 53 checks passed
YigongQin pushed a commit to YigongQin/TransformerEngine that referenced this pull request Apr 23, 2026
* Add better ordering enforcment to split_overlap_rs gemms.

This adds a short delay kernel to the split_overlap_rs function, which
ensures that the gemms are properly ordered when run with cuda graphs.

Signed-off-by: Chase Block <cblock@nvidia.com>

* [pre-commit.ci] auto fixes from pre-commit.com hooks

for more information, see https://pre-commit.ci

---------

Signed-off-by: Chase Block <cblock@nvidia.com>
Co-authored-by: pre-commit-ci[bot] <66853113+pre-commit-ci[bot]@users.noreply.github.com>
Co-authored-by: Przemyslaw Tredak <ptredak@nvidia.com>
faradawn pushed a commit to faradawn/TransformerEngine that referenced this pull request May 14, 2026
* Add better ordering enforcment to split_overlap_rs gemms.

This adds a short delay kernel to the split_overlap_rs function, which
ensures that the gemms are properly ordered when run with cuda graphs.

Signed-off-by: Chase Block <cblock@nvidia.com>

* [pre-commit.ci] auto fixes from pre-commit.com hooks

for more information, see https://pre-commit.ci

---------

Signed-off-by: Chase Block <cblock@nvidia.com>
Co-authored-by: pre-commit-ci[bot] <66853113+pre-commit-ci[bot]@users.noreply.github.com>
Co-authored-by: Przemyslaw Tredak <ptredak@nvidia.com>
Sign up for free to join this conversation on GitHub. Already have an account? Sign in to comment

Labels

community-contribution PRs from external contributor outside the core maintainers, representing community-driven work.

Projects

None yet

Development

Successfully merging this pull request may close these issues.

4 participants