Skip to content

Let the CUDA allocator manage cuBLAS workspaces - #3625

Open
Vtmpas wants to merge 14 commits into
NVIDIA:mainfrom
Vtmpas:fix/cublas-workspace-per-stream
Open

Vtmpas wants to merge 14 commits into
NVIDIA:mainfrom
Vtmpas:fix/cublas-workspace-per-stream

Conversation

@Vtmpas

@Vtmpas Vtmpas commented Oct 5, 2026 •

Copy link
Copy Markdown

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 NT GEMMs 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

  • Documentation change
  • Bug fix (non-breaking change which fixes an issue)
  • New feature
  • Breaking change
  • Infra/Build change
  • Code refactoring

Changes

  • Allocate scratch per invocation instead of caching it in Python.
  • Remove the preallocation workaround for globally cached scratch.
  • Test eager ownership, temporary-stream cleanup, graph-pool isolation across capture generations, concurrent FP32/BF16 graph replay, grouped GEMM scratch reuse, and compiled Linear forward/backward across graph generations.
  • Cover router dgrad NN with 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.
  • Include the workspace regression in the PyTorch L0 launcher.

Validation

  • One NVIDIA H200, CUDA 13.0, PyTorch 2.13.0+cu130.
  • With the exact workspace functions from 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.
  • The integration patch passes all 17 workspace cases, including grouped GEMM, and is idempotent.
  • Rebuilt current main for SM90 and ran 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.
  • Scoped forward/backward timing used benchmarks/linear/benchmark_linear.py::benchmark_linear with 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.
  • Repository pre-commit checks, changed-file pylint and the L0 license checker pass on macOS.
  • Full current-main L0 validation is still pending.

The original workspace factory is identical in main, stable, release_v2.19 and release_v2.20.

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

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>
@github-actions github-actions Bot added the community-contribution PRs from external contributor outside the core maintainers, representing community-driven work. label Oct 5, 2026
@Vtmpas
Vtmpas marked this pull request as ready for review October 5, 2026 10:06
@Vtmpas
Vtmpas requested a review from ksivaman as a code owner October 5, 2026 10:06
@greptile-apps

greptile-apps Bot commented Oct 5, 2026 •

Copy link
Copy Markdown
Contributor

RetriggerConfidence Score: 5/5

[Medium impact] The PR appears safe to merge based on the reviewed changes.

Summary

The PR removes the Python cuBLAS workspace cache so each GEMM invocation obtains allocator-managed scratch, removes Linear’s preallocation workaround, and adds workspace regression coverage. The latest merge also brings in fused MXFP8 RMSNorm changes and tests. Both previously reported workspace-cache issues are addressed; no new actionable issue was established.

Diagram

%%{init: {'theme': 'neutral'}}%%
flowchart LR
  GEMM[GEMM invocation] --> ALLOC[Allocate workspace]
  ALLOC --> EAGER[Eager: caching allocator]
  ALLOC --> GRAPH[Capture: graph-private pool]
  EAGER --> RETURN[Release invocation reference]
  GRAPH --> RETURN
Loading

Reviews (14) · Last reviewed commit: "Merge branch 'main' into fix/cublas-work..." · Reviewed by Greptile

Comment thread transformer_engine/pytorch/cpp_extensions/gemm.py Outdated
Comment thread transformer_engine/pytorch/cpp_extensions/gemm.py Outdated
@Vtmpas
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 Vtmpas changed the title Cache cuBLAS workspace by CUDA stream Let the CUDA allocator manage cuBLAS workspaces Oct 5, 2026
@Vtmpas
Vtmpas marked this pull request as ready for review October 5, 2026 10:44
Vtmpas and others added 3 commits October 5, 2026 14:08
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>
@Vtmpas

Vtmpas commented Oct 6, 2026

Copy link
Copy Markdown
Author

@Oleg-Goncharov , @ksivaman , hello! can i have your assistance here?

@Vtmpas

Vtmpas commented Oct 8, 2026

Copy link
Copy Markdown
Author

@ksivaman kindly

@ksivaman

ksivaman commented Oct 8, 2026

Copy link
Copy Markdown
Member

/te-ci pytorch L0

@itwastony itwastony left a comment •

Copy link
Copy Markdown

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

had the same trouble, thanks

@greptile-apps

This comment has been minimized.

@greptile-apps

greptile-apps Bot commented Oct 8, 2026

Copy link
Copy Markdown
Contributor

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>
@Vtmpas
Vtmpas requested a review from timmoon10 as a code owner October 8, 2026 21:45
@Vtmpas

Vtmpas commented Oct 8, 2026

Copy link
Copy Markdown
Author

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 cache_info(). BasicLinear no longer preallocates scratch for the removed cache. The six initialization/training combinations, fullgraph capture checks, and output/gradient comparisons are preserved.

On H200 with CUDA 13 and PyTorch 2.13, the original AttributeError reproduced before the fix. Afterward, 22 workspace tests, the fake-initialization test, and three BasicLinear CUDA-graph cases passed. The six compiled cold-start cases are capability-skipped because this PyTorch lacks register_custom_class; execution of those cases on a compatible PyTorch build remains unverified. Formatting, scoped pylint, and the license check passed.

@Vtmpas

Vtmpas commented Oct 9, 2026

Copy link
Copy Markdown
Author

@timmoon10 @ksivaman conflicts have been solved
ci seems okay

@ksivaman

ksivaman commented Oct 9, 2026

Copy link
Copy Markdown
Member

/te-ci pytorch L0

This branch has not been deployed

No deployments
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.

3 participants