Repository navigation
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>
|
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>
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>
|
@Oleg-Goncharov , @ksivaman , hello! can i have your assistance here? |
|
@ksivaman kindly |
|
/te-ci pytorch L0 |
This comment has been minimized.
This comment has been minimized.
|
Want your agent to iterate on Greptile's feedback? Try greploops. |
Replace the cold-start test's cache_info assertions with checks for capture-time allocation and release of workspace tensors. Keep output and gradient comparisons across CUDA, CPU and meta initialization. Remove BasicLinear's obsolete cache preallocation and use the existing opaque-object capability guard for compiled execution. Signed-off-by: Matvey Saprykin <mtvey.s@gmail.com>
|
Fixed in f01b35e, on top of the latest main merge. The cold-start test now checks capture-time workspace allocation and release of Python tensor references instead of On H200 with CUDA 13 and PyTorch 2.13, the original |
|
@timmoon10 @ksivaman conflicts have been solved |
|
/te-ci pytorch L0 |
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