From da52081af27499683298a55007250f12fa41d753 Mon Sep 17 00:00:00 2001 From: Matvey Saprykin Date: Mon, 5 Oct 2026 12:38:39 +0300 Subject: [PATCH] Cache cuBLAS workspace by CUDA stream Keep scratch buffers separate for concurrent GEMMs while preserving reuse on each stream. Add ownership checks for dense, userbuffer and grouped workspaces, plus exact FP32/BF16 numerical coverage for concurrent large reductions. Signed-off-by: Matvey Saprykin --- qa/L0_pytorch_unittest/test.sh | 1 + tests/pytorch/test_gemm_workspace.py | 55 +++++++++++++++++++ .../pytorch/cpp_extensions/gemm.py | 12 +++- 3 files changed, 66 insertions(+), 2 deletions(-) create mode 100644 tests/pytorch/test_gemm_workspace.py diff --git a/qa/L0_pytorch_unittest/test.sh b/qa/L0_pytorch_unittest/test.sh index b3b6ccacac7..457fb5b798d 100644 --- a/qa/L0_pytorch_unittest/test.sh +++ b/qa/L0_pytorch_unittest/test.sh @@ -31,6 +31,7 @@ pip3 install pytest==8.2.1 || error_exit "Failed to install pytest" python3 -m pytest --tb=auto --junitxml=$XML_LOG_DIR/pytest_test_optional_flash_attn_import.xml $TE_PATH/tests/pytorch/test_optional_flash_attn_import.py || test_fail "test_optional_flash_attn_import.py" NVTE_GROUPED_LINEAR_SINGLE_PARAM=1 python3 -m pytest --tb=auto --junitxml=$XML_LOG_DIR/pytest_test_sanity.xml $TE_PATH/tests/pytorch/test_sanity.py || test_fail "test_sanity.py" +python3 -m pytest --tb=auto --junitxml=$XML_LOG_DIR/pytest_test_gemm_workspace.xml $TE_PATH/tests/pytorch/test_gemm_workspace.py || test_fail "test_gemm_workspace.py" python3 -m pytest --tb=auto --junitxml=$XML_LOG_DIR/pytest_test_recipe.xml $TE_PATH/tests/pytorch/test_recipe.py || test_fail "test_recipe.py" python3 -m pytest --tb=auto --junitxml=$XML_LOG_DIR/pytest_test_custom_recipe.xml $TE_PATH/tests/pytorch/test_custom_recipe.py || test_fail "test_custom_recipe.py" python3 -m pytest --tb=auto --junitxml=$XML_LOG_DIR/pytest_test_deferred_init.xml $TE_PATH/tests/pytorch/test_deferred_init.py || test_fail "test_deferred_init.py" diff --git a/tests/pytorch/test_gemm_workspace.py b/tests/pytorch/test_gemm_workspace.py new file mode 100644 index 00000000000..28e265c696e --- /dev/null +++ b/tests/pytorch/test_gemm_workspace.py @@ -0,0 +1,55 @@ +# Copyright (c) 2022-2026, NVIDIA CORPORATION & AFFILIATES. All rights reserved. +# +# See LICENSE for license information. + +"""cuBLAS workspace ownership and numerical checks for independent CUDA streams.""" + +import pytest +import torch + +from transformer_engine.pytorch.cpp_extensions.gemm import general_gemm, get_cublas_workspace + + +@pytest.mark.parametrize("ub,grouped_gemm", [(False, False), (True, False), (False, True)]) +def test_workspace_is_reused_only_on_its_stream(ub, grouped_gemm): + """A workspace remains reusable without aliasing another stream's active scratch.""" + device = torch.cuda.current_device() + streams = [torch.cuda.Stream(), torch.cuda.Stream()] + workspaces = [] + for stream in streams: + with torch.cuda.stream(stream): + first = get_cublas_workspace(device, ub, grouped_gemm) + assert get_cublas_workspace(device, ub, grouped_gemm) is first + workspaces.append(first if grouped_gemm else [first]) + assert len(workspaces[0]) == len(workspaces[1]) + assert {w.data_ptr() for w in workspaces[0]}.isdisjoint(w.data_ptr() for w in workspaces[1]) + + +@pytest.mark.parametrize("dtype", [torch.float32, torch.bfloat16]) +def test_concurrent_gemms_match_exact_reference(dtype): + """Large reductions exercise cuBLAS algorithms that use scratch for partial sums.""" + streams = [torch.cuda.Stream(), torch.cuda.Stream()] + reduction, hidden, experts = 16384, 4096, 128 + inputs = [ + torch.full((reduction, hidden), value, dtype=dtype, device="cuda") for value in (1.0, 2.0) + ] + gradients = [ + torch.full((reduction, experts), value, dtype=dtype, device="cuda") for value in (1.0, 3.0) + ] + ready = torch.cuda.Event() + ready.record() + for stream in streams: + stream.wait_event(ready) + + outputs = [] + for _ in range(8): + for stream, inp, grad in zip(streams, inputs, gradients): + with torch.cuda.stream(stream): + out, *_ = general_gemm(inp, grad, dtype, layout="NT", grad=True) + outputs.append(out) + for stream in streams: + stream.synchronize() + # Both operands and the result are exactly representable in either dtype. + for index, out in enumerate(outputs): + expected = torch.full_like(out, reduction * (1 if index % 2 == 0 else 6)) + torch.testing.assert_close(out, expected, rtol=0, atol=0) diff --git a/transformer_engine/pytorch/cpp_extensions/gemm.py b/transformer_engine/pytorch/cpp_extensions/gemm.py index 6939847c8dc..56f23b872f5 100644 --- a/transformer_engine/pytorch/cpp_extensions/gemm.py +++ b/transformer_engine/pytorch/cpp_extensions/gemm.py @@ -48,9 +48,17 @@ def get_cublas_workspace_size_bytes() -> None: return 4_194_304 -@functools.lru_cache(maxsize=None) def get_cublas_workspace(device: int, ub: bool, grouped_gemm: bool) -> torch.Tensor: - """Returns workspace for cublas GEMM.""" + """Returns workspace for cublas GEMM on the current CUDA stream.""" + stream = torch.cuda.current_stream(device).cuda_stream + return _get_cublas_workspace_for_stream(device, ub, grouped_gemm, stream) + + +@functools.lru_cache(maxsize=None) +def _get_cublas_workspace_for_stream( + device: int, ub: bool, grouped_gemm: bool, _stream: int +) -> torch.Tensor: + """Cache by stream so concurrent GEMMs do not overwrite each other's scratch space.""" assert not (ub and grouped_gemm), "UB is unsupported for grouped GEMM." if ub: