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: