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..17bceddaa90 --- /dev/null +++ b/tests/pytorch/test_gemm_workspace.py @@ -0,0 +1,269 @@ +# 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 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, + general_grouped_gemm, + get_cublas_workspace, +) + + +@pytest.mark.parametrize("ub,grouped_gemm", [(False, False), (True, False), (False, True)]) +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) + 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]) + 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]) +@pytest.mark.parametrize("execution", ["eager", "graph"]) +@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()] + 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((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, 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=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=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=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() + for index, out in enumerate(outputs): + 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]) +@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 6939847c8dc..e7a696cfa4b 100644 --- a/transformer_engine/pytorch/cpp_extensions/gemm.py +++ b/transformer_engine/pytorch/cpp_extensions/gemm.py @@ -48,9 +48,14 @@ 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.""" + """Allocate cuBLAS scratch for one invocation on the current CUDA stream. + + 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,