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()
70 changes: 63 additions & 7 deletions tests/pytorch/test_torch_compile.py
Original file line number Diff line number Diff line change
Expand Up @@ -1926,6 +1926,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 Expand Up @@ -2916,16 +2959,17 @@ def run(module, inp):
@pytest.mark.skipif(not torch.cuda.is_available(), reason="CUDA required")
def test_te_ops_linear_workspace_fake_mode(monkeypatch):
def unexpected_allocation(*args, **kwargs):
pytest.fail("Fake initialization must not populate the real workspace cache")
pytest.fail("Fake initialization must not allocate a real cuBLAS workspace")

monkeypatch.setattr(
"transformer_engine.pytorch.ops.basic.basic_linear.get_cublas_workspace",
"transformer_engine.pytorch.cpp_extensions.gemm.get_cublas_workspace",
unexpected_allocation,
)
with FakeTensorMode():
te.ops.BasicLinear(32, 64, device="cuda", dtype=torch.bfloat16)


@pytest.mark.skipif(not _opaque_available, reason="torch opaque object API not available")
@pytest.mark.skipif(not torch.cuda.is_available(), reason="CUDA required")
@pytest.mark.parametrize("training", [False, True])
@pytest.mark.parametrize("initial_device", ["cuda", "cpu", "meta"])
Expand All @@ -2935,23 +2979,33 @@ def test_te_ops_linear_bias_cold_cudagraphs(training, initial_device):
"""
import contextlib
import sys
import weakref
import torch
import transformer_engine.pytorch as te
from transformer_engine.pytorch.cpp_extensions.gemm import get_cublas_workspace
from transformer_engine.pytorch.cpp_extensions import gemm
from torch._dynamo.utils import counters

training, initial_device = sys.argv[1:]
training = training == "True"
assert get_cublas_workspace.cache_info().currsize == 0
workspaces, captured_workspaces = [], []
allocate_workspace = gemm.get_cublas_workspace

def allocate(*args, **kwargs):
workspace = allocate_workspace(*args, **kwargs)
ref = weakref.ref(workspace)
workspaces.append(ref)
if torch.cuda.is_current_stream_capturing():
captured_workspaces.append(ref)
return workspace

gemm.get_cublas_workspace = allocate
linear = te.ops.BasicLinear(32, 64, device=initial_device, dtype=torch.bfloat16)
if initial_device == "cpu":
assert get_cublas_workspace.cache_info().currsize == 0
linear.cuda()
elif initial_device == "meta":
assert get_cublas_workspace.cache_info().currsize == 0
linear.to_empty(device="cuda")
linear.reset_parameters()
assert get_cublas_workspace.cache_info().currsize > 0
assert not workspaces, "Module initialization must not allocate GEMM scratch"
model = te.ops.Sequential(linear, te.ops.Bias(64, dtype=torch.bfloat16))
x = torch.randn(32, 32, device="cuda", dtype=torch.bfloat16, requires_grad=training)
targets = (x, *model.parameters())
Expand All @@ -2977,6 +3031,8 @@ def forward(x):
del grads, expected_grads
del actual, expected
torch.cuda.synchronize()
assert captured_workspaces, "Expected workspace allocations during CUDA graph capture"
assert all(ref() is None for ref in workspaces), "Invocations must release workspace tensors"
assert not counters["inductor"]["cudagraph_skips"], counters["inductor"]
assert counters["inductor"]["cudagraph_recorded_non_static_inputs"] > 0, dict(
counters["inductor"]
Expand Down
Loading
Loading