From 93a80c7908eea41426ad2c6044f246bb76d3fd04 Mon Sep 17 00:00:00 2001 From: Matvey Saprykin Date: Mon, 5 Oct 2026 12:59:10 +0300 Subject: [PATCH 1/3] Cache cuBLAS workspace by CUDA stream 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 --- 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: From e5b55e918ffdbc1118a7dd0e1effbe731c55cbbd Mon Sep 17 00:00:00 2001 From: Matvey Saprykin Date: Mon, 5 Oct 2026 13:41:11 +0300 Subject: [PATCH 2/3] Let the CUDA allocator own cuBLAS workspaces 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 --- tests/pytorch/test_gemm_workspace.py | 207 +++++++++++++++++- tests/pytorch/test_torch_compile.py | 43 ++++ .../pytorch/cpp_extensions/gemm.py | 15 +- transformer_engine/pytorch/module/linear.py | 35 --- 4 files changed, 246 insertions(+), 54 deletions(-) diff --git a/tests/pytorch/test_gemm_workspace.py b/tests/pytorch/test_gemm_workspace.py index 28e265c696e..013c126ea41 100644 --- a/tests/pytorch/test_gemm_workspace.py +++ b/tests/pytorch/test_gemm_workspace.py @@ -4,29 +4,109 @@ """cuBLAS workspace ownership and numerical checks for independent CUDA streams.""" +import gc +import weakref + import pytest import torch +from transformer_engine.pytorch.cpp_extensions import gemm -from transformer_engine.pytorch.cpp_extensions.gemm import general_gemm, get_cublas_workspace +from transformer_engine.pytorch.cpp_extensions.gemm import ( + general_gemm, + general_grouped_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.""" +def test_live_workspaces_do_not_alias(ub, grouped_gemm): + """Live invocations own distinct scratch, even on the same capture stream.""" 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]) + second = get_cublas_workspace(device, ub, grouped_gemm) + workspaces.append(second if grouped_gemm else [second]) 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]) + pointers = [w.data_ptr() for invocation in workspaces for w in invocation] + assert len(pointers) == len(set(pointers)) + + +@pytest.mark.parametrize("ub,grouped_gemm", [(False, False), (True, False), (False, True)]) +def test_temporary_streams_release_workspaces(ub, grouped_gemm): + """Scratch is returned to the allocator when the invocation goes out of scope.""" + device = torch.cuda.current_device() + torch.cuda.synchronize() + gc.collect() + allocated = torch.cuda.memory_allocated() + for _ in range(8): + stream = torch.cuda.Stream() + with torch.cuda.stream(stream): + workspace = get_cublas_workspace(device, ub, grouped_gemm) + tensors = workspace if grouped_gemm else [workspace] + refs = [weakref.ref(tensor) for tensor in tensors] + for tensor in tensors: + tensor.fill_(1) + del tensor, tensors, workspace + stream.synchronize() + del stream + assert all(ref() is None for ref in refs) + assert torch.cuda.memory_allocated() == allocated + + +@pytest.mark.parametrize("ub,grouped_gemm", [(False, False), (True, False), (False, True)]) +def test_graph_workspaces_survive_replay_and_are_released(ub, grouped_gemm): + """Different captures on one stream must not borrow an earlier graph's scratch.""" + device = torch.cuda.current_device() + capture_stream = torch.cuda.Stream() + replay_streams = [torch.cuda.Stream(), torch.cuda.Stream()] + # Initialize the allocator/library state outside the graph, on a different stream. + workspace = get_cublas_workspace(device, ub, grouped_gemm) + count = len(workspace) if grouped_gemm else 1 + del workspace + outputs = [torch.empty(count, dtype=torch.uint8, device=device) for _ in range(2)] + torch.cuda.synchronize() + torch.cuda.empty_cache() + reserved = torch.cuda.memory_reserved() + allocated = torch.cuda.memory_allocated() + for _ in range(3): + graphs, pointers = [], [] + for value, output in enumerate(outputs, start=1): + graph = torch.cuda.CUDAGraph() + with torch.cuda.graph(graph, stream=capture_stream): + workspace = get_cublas_workspace(device, ub, grouped_gemm) + tensors = workspace if grouped_gemm else [workspace] + pointers.append({tensor.data_ptr() for tensor in tensors}) + refs = [weakref.ref(tensor) for tensor in tensors] + for index, tensor in enumerate(tensors): + tensor.fill_(value) + output[index : index + 1].copy_(tensor[:1]) + del tensor, tensors, workspace + assert all(ref() is None for ref in refs) + graphs.append(graph) + assert pointers[0].isdisjoint(pointers[1]) + for _ in range(8): + for graph, stream in zip(graphs, replay_streams): + with torch.cuda.stream(stream): + graph.replay() + torch.cuda.synchronize() + for value, output in enumerate(outputs, start=1): + torch.testing.assert_close(output, torch.full_like(output, value), rtol=0, atol=0) + for graph in graphs: + graph.reset() + del graph, graphs + gc.collect() + torch.cuda.empty_cache() + assert torch.cuda.memory_allocated() == allocated + assert torch.cuda.memory_reserved() == reserved @pytest.mark.parametrize("dtype", [torch.float32, torch.bfloat16]) -def test_concurrent_gemms_match_exact_reference(dtype): +@pytest.mark.parametrize("execution", ["eager", "graph"]) +def test_concurrent_gemms_match_exact_reference(dtype, execution): """Large reductions exercise cuBLAS algorithms that use scratch for partial sums.""" streams = [torch.cuda.Stream(), torch.cuda.Stream()] reduction, hidden, experts = 16384, 4096, 128 @@ -41,15 +121,122 @@ def test_concurrent_gemms_match_exact_reference(dtype): for stream in streams: stream.wait_event(ready) - outputs = [] + outputs, graphs = [], [] + if execution == "graph": + capture_stream = torch.cuda.Stream() + capture_stream.wait_event(ready) + for inp, grad in zip(inputs, gradients): + with torch.cuda.stream(capture_stream): + general_gemm(inp, grad, dtype, layout="NT", grad=True) + capture_stream.synchronize() + graph = torch.cuda.CUDAGraph() + with torch.cuda.graph(graph, stream=capture_stream): + out, *_ = general_gemm(inp, grad, dtype, layout="NT", grad=True) + graphs.append(graph) + outputs.append(out) for _ in range(8): - for stream, inp, grad in zip(streams, inputs, gradients): + for index, (stream, inp, grad) in enumerate(zip(streams, inputs, gradients)): with torch.cuda.stream(stream): - out, *_ = general_gemm(inp, grad, dtype, layout="NT", grad=True) - outputs.append(out) + if execution == "graph": + graphs[index].replay() + else: + 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) + + +@pytest.mark.parametrize("dtype", [torch.float32, torch.bfloat16]) +@pytest.mark.parametrize("execution", ["eager", "graph"]) +def test_grouped_gemms_match_exact_reference(dtype, execution): + """Internal GEMM streams must join the caller before its scratch is recycled.""" + reduction, hidden, experts, count = 8192, 512, 64, 4 + inputs = [torch.ones((reduction, hidden), dtype=dtype, device="cuda") for _ in range(count)] + gradients = [ + torch.full((reduction, experts), index + 1, dtype=dtype, device="cuda") + for index in range(count) + ] + outputs = [torch.empty((experts, hidden), dtype=dtype, device="cuda") for _ in range(count)] + stream = torch.cuda.Stream() + stream.wait_stream(torch.cuda.current_stream()) + + def run(): + general_grouped_gemm( + inputs, + gradients, + outputs, + [None] * count, + dtype, + layout="NT", + m_splits=[experts] * count, + grad=True, + ) + + with torch.cuda.stream(stream): + run() + stream.synchronize() + if execution == "graph": + graph = torch.cuda.CUDAGraph() + with torch.cuda.graph(graph, stream=stream): + run() + for _ in range(8): + with torch.cuda.stream(stream): + if execution == "graph": + graph.replay() + else: + run() + # Recycle same-size allocations immediately after the internal-stream join. + scratch = get_cublas_workspace(torch.cuda.current_device(), False, True) + for tensor in scratch: + tensor.fill_(0) + del tensor, scratch + stream.synchronize() + for index, output in enumerate(outputs): + expected = torch.full_like(output, reduction * (index + 1)) + torch.testing.assert_close(output, expected, rtol=0, atol=0) + + +@pytest.mark.skipif( + not hasattr(torch.library, "custom_op") + or not hasattr(torch.compiler, "cudagraph_mark_step_begin"), + reason="Custom operators and CUDA graph trees require newer PyTorch", +) +def test_compiled_gemm_graph_generations(monkeypatch): + """An opaque GEMM's scratch must remain valid when graph trees are rerecorded.""" + original = gemm.get_cublas_workspace + captured_workspaces = [] + + def allocate(*args, **kwargs): + workspace = original(*args, **kwargs) + if torch.cuda.is_current_stream_capturing(): + captured_workspaces.append(weakref.ref(workspace)) + return workspace + + monkeypatch.setattr(gemm, "get_cublas_workspace", allocate) + + @torch.library.custom_op("te_workspace_test::gemm", mutates_args=()) + def run(inp: torch.Tensor, grad: torch.Tensor) -> torch.Tensor: + return general_gemm(inp, grad, inp.dtype, layout="NT", grad=True)[0] + + @run.register_fake + def fake(inp, grad): + return inp.new_empty((grad.shape[1], inp.shape[1])) + + torch._dynamo.reset() + compiled = torch.compile(run, fullgraph=True, mode="reduce-overhead", dynamic=False) + for reduction, value in [(8192, 1), (16384, 2), (8192, 3)]: + inp = torch.full((reduction, 512), value, dtype=torch.bfloat16, device="cuda") + grad = torch.ones((reduction, 64), dtype=torch.bfloat16, device="cuda") + for _ in range(3): + torch.compiler.cudagraph_mark_step_begin() + out = compiled(inp, grad) + torch.testing.assert_close(out, torch.full_like(out, reduction * value), rtol=0, atol=0) + del out + assert len(captured_workspaces) >= 2, "Expected CUDA graph captures for different shapes" + assert all(ref() is None for ref in captured_workspaces) + del compiled + torch._dynamo.reset() diff --git a/tests/pytorch/test_torch_compile.py b/tests/pytorch/test_torch_compile.py index eae6f0a8a24..e0e28f9b882 100644 --- a/tests/pytorch/test_torch_compile.py +++ b/tests/pytorch/test_torch_compile.py @@ -1892,6 +1892,49 @@ def test_to_tensor_spec_quantized(factory, shape): # --------------------------------------------------------------------------- +@pytest.mark.skipif(not _opaque_available, reason="torch opaque object API not available") +def test_te_linear_cublas_workspace_graph_generations(monkeypatch): + """Compiled forward/backward scratch follows graph generations, not a Python cache.""" + import weakref + from transformer_engine.pytorch.cpp_extensions import gemm + + original = gemm.get_cublas_workspace + captured_workspaces = [] + + def allocate(*args, **kwargs): + workspace = original(*args, **kwargs) + if torch.cuda.is_current_stream_capturing(): + captured_workspaces.append(weakref.ref(workspace)) + return workspace + + monkeypatch.setattr(gemm, "get_cublas_workspace", allocate) + model = te.Linear(4096, 128, bias=False, params_dtype=torch.bfloat16, device="cuda") + with torch.no_grad(): + model.weight.fill_(1) + torch._dynamo.reset() + compiled = torch.compile(model, fullgraph=True, mode="reduce-overhead", dynamic=False) + with _assert_no_cudagraph_skips(True): + for batch, value in [(128, 1), (256, 2), (128, 3)]: + for _ in range(3): + torch.compiler.cudagraph_mark_step_begin() + model.zero_grad(set_to_none=True) + inp = torch.full( + (batch, 4096), value, dtype=torch.bfloat16, device="cuda", requires_grad=True + ) + out = compiled(inp) + out.sum().backward() + torch.testing.assert_close(out, torch.full_like(out, 4096 * value), rtol=0, atol=0) + torch.testing.assert_close(inp.grad, torch.full_like(inp, 128), rtol=0, atol=0) + torch.testing.assert_close( + model.weight.grad, torch.full_like(model.weight, batch * value), rtol=0, atol=0 + ) + del out, inp + assert len(captured_workspaces) >= 2, "Expected CUDA graph captures for different shapes" + assert all(ref() is None for ref in captured_workspaces) + del compiled + torch._dynamo.reset() + + @pytest.mark.skipif(not _opaque_available, reason="torch opaque object API not available") @pytest.mark.parametrize("compile_mode", _compile_modes) @pytest.mark.parametrize( diff --git a/transformer_engine/pytorch/cpp_extensions/gemm.py b/transformer_engine/pytorch/cpp_extensions/gemm.py index 56f23b872f5..e7a696cfa4b 100644 --- a/transformer_engine/pytorch/cpp_extensions/gemm.py +++ b/transformer_engine/pytorch/cpp_extensions/gemm.py @@ -49,16 +49,13 @@ def get_cublas_workspace_size_bytes() -> None: def get_cublas_workspace(device: int, ub: bool, grouped_gemm: bool) -> torch.Tensor: - """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) + """Allocate cuBLAS scratch for one invocation on the current CUDA 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.""" + The caching allocator safely reuses eager allocations on their stream. During + capture, the graph's private pool owns the allocation until the graph is + destroyed. Keeping the tensor in a Python cache would bypass that ownership + on later captures, including graphs replayed on different streams. + """ assert not (ub and grouped_gemm), "UB is unsupported for grouped GEMM." if ub: diff --git a/transformer_engine/pytorch/module/linear.py b/transformer_engine/pytorch/module/linear.py index 7f00e19644a..4e4c8d138d4 100644 --- a/transformer_engine/pytorch/module/linear.py +++ b/transformer_engine/pytorch/module/linear.py @@ -72,7 +72,6 @@ ) from ..cpp_extensions import ( general_gemm, - get_cublas_workspace, ) from ..constants import FP8BwdTensorIdx, FP8FwdTensorIdx, GemmParallelModes, dist_group_type from ..graph import is_graph_capturing @@ -2302,40 +2301,6 @@ def reset_parameters(self, defer_init=False): elif self.parallel_mode == "column": set_tensor_model_parallel_attributes(getattr(self, bias), True, 0, 1) - self._prealloc_cublas_workspace() - - def _apply(self, *args, **kwargs): - out = super()._apply(*args, **kwargs) - # Params may have just landed on CUDA (.to()/.cuda()/to_empty()). - self._prealloc_cublas_workspace() - return out - - def _prealloc_cublas_workspace(self) -> None: - """Allocate the process-global cuBLAS workspaces before the first GEMM: - under torch.compile it can run inside a cudagraph pool, and a workspace - first allocated there would not survive across graph generations.""" - # pylint: disable=import-outside-toplevel - from torch._guards import detect_fake_mode - from torch._subclasses.fake_tensor import is_fake - - weight = getattr(self, self.weight_names[0], None) - if weight is None or weight.device.type != "cuda": - return - if is_fake(weight) or detect_fake_mode(): - return - get_cublas_workspace(weight.device.index, False, False) - if self.ub_name is not None and any( - ( - self.ub_overlap_rs_fprop, - self.ub_overlap_ag_dgrad, - self.ub_overlap_ag_fprop, - self.ub_overlap_rs_dgrad, - self.ub_bulk_dgrad, - self.ub_bulk_wgrad, - ) - ): - get_cublas_workspace(weight.device.index, True, False) - def forward( self, inp: torch.Tensor, From ec5f5f9ce277bf0e92c9833f73597a123c9cb59c Mon Sep 17 00:00:00 2001 From: Matvey Saprykin Date: Mon, 5 Oct 2026 14:08:47 +0300 Subject: [PATCH 3/3] Test concurrent router dgrad workspace ownership 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 --- tests/pytorch/test_gemm_workspace.py | 55 +++++++++++++++++++++------- 1 file changed, 41 insertions(+), 14 deletions(-) diff --git a/tests/pytorch/test_gemm_workspace.py b/tests/pytorch/test_gemm_workspace.py index 013c126ea41..17bceddaa90 100644 --- a/tests/pytorch/test_gemm_workspace.py +++ b/tests/pytorch/test_gemm_workspace.py @@ -106,48 +106,75 @@ def test_graph_workspaces_survive_replay_and_are_released(ub, grouped_gemm): @pytest.mark.parametrize("dtype", [torch.float32, torch.bfloat16]) @pytest.mark.parametrize("execution", ["eager", "graph"]) -def test_concurrent_gemms_match_exact_reference(dtype, execution): - """Large reductions exercise cuBLAS algorithms that use scratch for partial sums.""" +@pytest.mark.parametrize("layout", ["NT", "NN"]) +def test_concurrent_gemms_match_exact_reference(dtype, execution, layout): + """Concurrent router wgrad and dgrad must own scratch in eager and graph execution.""" 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) - ] + hidden, experts = 4096, 128 + if layout == "NT": + # wgrad = grad_output.T @ input; reduce over tokens. + tokens = reduction = 16384 + input_shape, output_shape = (tokens, hidden), (experts, hidden) + else: + # dgrad = grad_output @ weight; match the router's NN [8192, 128, 4096] GEMM. + tokens, reduction = 8192, experts + input_shape, output_shape = (experts, hidden), (tokens, hidden) + inputs = [torch.full(input_shape, 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) + torch.full((tokens, experts), value, dtype=dtype, device="cuda") for value in (1.0, 3.0) ] + # Both operands and these analytic results are exactly representable in either dtype. + expected = [reduction, reduction * 6] ready = torch.cuda.Event() ready.record() for stream in streams: stream.wait_event(ready) - outputs, graphs = [], [] + outputs, graphs, checks = [], [], [] if execution == "graph": capture_stream = torch.cuda.Stream() capture_stream.wait_event(ready) for inp, grad in zip(inputs, gradients): with torch.cuda.stream(capture_stream): - general_gemm(inp, grad, dtype, layout="NT", grad=True) + general_gemm(inp, grad, dtype, layout=layout, grad=True) capture_stream.synchronize() + # Separate private pools, captured on one stream and replayed on two others. graph = torch.cuda.CUDAGraph() with torch.cuda.graph(graph, stream=capture_stream): - out, *_ = general_gemm(inp, grad, dtype, layout="NT", grad=True) + out, *_ = general_gemm(inp, grad, dtype, layout=layout, grad=True) + assert out.shape == output_shape + assert out.dtype == dtype graphs.append(graph) outputs.append(out) + del out for _ in range(8): for index, (stream, inp, grad) in enumerate(zip(streams, inputs, gradients)): with torch.cuda.stream(stream): if execution == "graph": graphs[index].replay() else: - out, *_ = general_gemm(inp, grad, dtype, layout="NT", grad=True) + out, *_ = general_gemm(inp, grad, dtype, layout=layout, grad=True) + assert out.shape == output_shape + assert out.dtype == dtype outputs.append(out) + del out + if execution == "eager": + # Submit both GEMMs before the checks. Keep only this pair of large NN outputs. + for index, (stream, out) in enumerate(zip(streams, outputs)): + with torch.cuda.stream(stream): + checks.append(torch.all(out == expected[index])) + del out + outputs.clear() 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) + checks.append(torch.all(out == expected[index])) + torch.testing.assert_close( + torch.stack(checks), + torch.ones(len(checks), dtype=torch.bool, device="cuda"), + rtol=0, + atol=0, + ) @pytest.mark.parametrize("dtype", [torch.float32, torch.bfloat16])