Skip to content
Open
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
1 change: 1 addition & 0 deletions qa/L0_pytorch_unittest/test.sh
Original file line number Diff line number Diff line change
Expand Up @@ -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"
Expand Down
269 changes: 269 additions & 0 deletions tests/pytorch/test_gemm_workspace.py
Original file line number Diff line number Diff line change
@@ -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()
43 changes: 43 additions & 0 deletions tests/pytorch/test_torch_compile.py
Original file line number Diff line number Diff line change
Expand Up @@ -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(
Expand Down
9 changes: 7 additions & 2 deletions transformer_engine/pytorch/cpp_extensions/gemm.py
Original file line number Diff line number Diff line change
Expand Up @@ -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:
Expand Down
35 changes: 0 additions & 35 deletions transformer_engine/pytorch/module/linear.py
Original file line number Diff line number Diff line change
Expand Up @@ -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
Expand Down Expand Up @@ -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,
Expand Down
Loading