Add better ordering enforcment to split_overlap_rs gemms. - #2056
Conversation
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>
for more information, see https://pre-commit.ci
|
@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. |
|
/te-ci L1 |
Greptile SummaryThis PR fixes GEMM reordering under CUDA Graph execution in Confidence Score: 4/5Safe 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
Sequence DiagramsequenceDiagram
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
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>>>(); } |
There was a problem hiding this comment.
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.
| 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()); | |
| } |
| 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)); | ||
| } |
There was a problem hiding this comment.
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.
* 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>
* 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>
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.
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
Changes
This PR adds:
This changes the kernel ordering in the table above to:
Checklist: