Skip to content
Merged
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
15 changes: 15 additions & 0 deletions aiter/ops/gemm_op_a16w16.py
Original file line number Diff line number Diff line change
Expand Up @@ -38,6 +38,9 @@ def _get_semaphore_workspace_keyed(device: torch.device, stream_id: int) -> Tens
return torch.zeros(_SEMA_SHAPE, dtype=torch.uint32, device=device)


_captured_semaphore_keepalive: list[Tensor] = []


def get_semaphore_workspace(device: torch.device) -> Tensor:
"""Return a per-(device, stream) zero-initialized semaphore workspace.

Expand All @@ -53,7 +56,19 @@ def get_semaphore_workspace(device: torch.device) -> Tensor:
Workspace size is small (~4 KB) and stream count per process is typically
< 8, so the LRU cap of 64 leaves plenty of headroom before any in-flight
workspace risks being evicted.

Under CUDA graph capture this returns a fresh workspace per launch instead
of the cached one: a captured graph bakes in the pointer and replays on a
stream other than the capture stream, so the cached counter can be left
non-zero and the reduction never fires. Allocating under capture also
records the zero-fill as a graph node, re-establishing the counter==0 entry
invariant on every replay. It is retained for the process lifetime because
aiter cannot observe when a graph dies.
"""
if torch.cuda.is_current_stream_capturing():
w = torch.zeros(_SEMA_SHAPE, dtype=torch.uint32, device=device)
_captured_semaphore_keepalive.append(w)
return w
stream = torch.cuda.current_stream(device)
return _get_semaphore_workspace_keyed(device, stream.cuda_stream)

Expand Down
76 changes: 76 additions & 0 deletions op_tests/test_gemm_a16w16_graph.py
Original file line number Diff line number Diff line change
@@ -0,0 +1,76 @@
# SPDX-License-Identifier: MIT
# Copyright (C) 2024-2026, Advanced Micro Devices, Inc. All rights reserved.

"""Regression test for the ASM split-K semaphore under graph replay."""

import os
import subprocess
import sys

import pytest
import torch

from aiter.jit.utils.chip_info import get_gfx_runtime

_CHILD_ENV = "AITER_ASM_SPLITK_GRAPH_CHILD"
_KERNEL = "_ZN5aiter39bf16gemm_fp32bf16_tn_64x64_splitk_cleanE"


def _run_child() -> None:
from aiter.ops.gemm_op_a16w16 import (
gemm_a16w16_asm,
get_semaphore_workspace,
)

torch.manual_seed(0)
device = torch.device("cuda:0")
x = torch.randn(64, 5120, dtype=torch.bfloat16, device=device)
weight = torch.randn(256, 5120, dtype=torch.bfloat16, device=device)

eager = torch.empty(64, 256, dtype=torch.float32, device=device)
gemm_a16w16_asm(x, weight, eager, splitK=13, kernelName=_KERNEL)
torch.cuda.synchronize()

capture_stream = torch.cuda.Stream(device=device)
with torch.cuda.stream(capture_stream):
stale = get_semaphore_workspace(device)
stale.fill_(2)
capture_stream.synchronize()

first = torch.empty_like(eager)
second = torch.empty_like(eager)
graph = torch.cuda.CUDAGraph()
with torch.cuda.graph(graph, stream=capture_stream):
gemm_a16w16_asm(x, weight, first, splitK=13, kernelName=_KERNEL)
gemm_a16w16_asm(x, weight, second, splitK=13, kernelName=_KERNEL)

for _ in range(4):
graph.replay()
torch.cuda.synchronize()
torch.testing.assert_close(first, eager, rtol=1e-2, atol=1e-2)
torch.testing.assert_close(second, eager, rtol=1e-2, atol=1e-2)


@pytest.mark.skipif(not torch.cuda.is_available(), reason="requires a ROCm GPU")
def test_asm_splitk_graph_replay_uses_fresh_semaphores() -> None:
if get_gfx_runtime() != "gfx950":
pytest.skip("the production regression was observed on gfx950")

env = os.environ.copy()
env[_CHILD_ENV] = "1"
proc = subprocess.run(
[sys.executable, __file__],
capture_output=True,
env=env,
text=True,
timeout=120,
check=False,
)
assert proc.returncode == 0, (
"captured ASM split-K GEMM hung or returned an incorrect result "
f"(exit {proc.returncode})\nstdout:\n{proc.stdout}\nstderr:\n{proc.stderr}"
)


if __name__ == "__main__" and os.environ.get(_CHILD_ENV) == "1":
_run_child()
Loading