Conversation
The workspace cache currently ignores the calling CUDA stream. Overlapping GEMMs can therefore overwrite the same scratch buffer and return incorrect results. Include the stream in the cache key, keep reuse on each stream, and add ownership and FP32/BF16 numerical regression tests. Signed-off-by: Matvey Saprykin <mtvey.s@gmail.com>
Vtmpas
marked this pull request as ready for review
October 5, 2026 10:06
Contributor
|
Vtmpas
marked this pull request as draft
October 5, 2026 10:30
Keep scratch local to each invocation so graph captures use their own private pools and temporary streams do not retain tensors in Python. Remove Linear's preallocation workaround for the global workspace cache. Cover concurrent graph replays, graph generations, scratch reclamation, grouped GEMM stream joins, and compiled Linear forward and backward. Signed-off-by: Matvey Saprykin <mtvey.s@gmail.com>
Vtmpas
marked this pull request as ready for review
October 5, 2026 10:44
Cover NN router dgrad with FP32 and BF16 operands in concurrent eager execution and independent CUDA graphs. Preserve NT wgrad coverage and exact references while retaining only one pair of large eager outputs. All 22 workspace cases pass on native SM90 H200 with CUDA 13 and PyTorch 2.13, including all four new NN cases. Signed-off-by: Matvey Saprykin <mtvey.s@gmail.com>
This branch has not been deployed
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
Concurrent GEMMs on independent CUDA streams receive the same cuBLAS workspace from the device/mode cache and can overwrite each other's partial sums. On one H200, two FP32
NTGEMMs with a reduction dimension of 16384 produced incorrect results in 35 of 48 calls; separate scratch produced zero mismatches against serial execution.Remove the Python workspace cache and let the PyTorch caching allocator manage scratch for each invocation. Eager allocations are reused safely on their stream. Captured allocations belong to their graph's private pool, so a later capture cannot borrow a buffer from an earlier graph generation. Scratch is returned to the allocator when the invocation ends, without retaining tensors for temporary streams.
This also removes Linear's preallocation workaround for the process-global cache. The workspace signature and sizes, GEMM dtypes and algorithm selection are unchanged. No new device synchronization is added.
Type of change
Changes
weight[128,4096],grad_output[8192,128], output[8192,4096], and reduction 128 against exact analytic references in FP32 and BF16. Preserve NT coverage and check every eager invocation while retaining only one pair of large outputs.Validation
main(d0b4b32) loaded into native TE 2.19.0.dev0+b5599209: the first 13 regression cases fail before the fix and pass after it. The previous per-stream-cache revision fails 11 of those 13 cases, including both concurrent graph-replay numerics cases.NVTE_ALLOW_NONDETERMINISTIC_ALGO=0 python -m pytest -v tests/pytorch/test_gemm_workspace.py: all 22 cases pass in 15.98s, including all four new FP32/BF16 NN eager/independent-graph cases, the existing NT cases, and compiled GEMM graph generations. The compiled Linear cases are capability-skipped because this PyTorch build lacks the required opaque-object API.benchmarks/linear/benchmark_linear.py::benchmark_linearwith BF16/FP8 block scaling on shapes(8192,4096,128)and(128,4096,16384). The repeated timings were unstable, so they do not establish a performance improvement or a reliable regression bound.The original workspace factory is identical in
main,stable,release_v2.19andrelease_v2.20.Checklist