From a119cf41c85a16c4e6fed6b634d637fddb74762e Mon Sep 17 00:00:00 2001 From: amd-ruitang3 <145657428+amd-ruitang3@users.noreply.github.com> Date: Wed, 12 Aug 2026 18:01:00 +0800 Subject: [PATCH] Revert "Fix ASM split-K semaphore deadlock under CUDA graph capture (#4494)" This reverts commit c6ce60c7496b448b979d619a02d6b0394291f211. --- aiter/ops/gemm_op_a16w16.py | 15 ------ op_tests/test_gemm_a16w16_graph.py | 76 ------------------------------ 2 files changed, 91 deletions(-) delete mode 100644 op_tests/test_gemm_a16w16_graph.py diff --git a/aiter/ops/gemm_op_a16w16.py b/aiter/ops/gemm_op_a16w16.py index 8ea6729799..70d2a7083d 100644 --- a/aiter/ops/gemm_op_a16w16.py +++ b/aiter/ops/gemm_op_a16w16.py @@ -38,9 +38,6 @@ 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. @@ -56,19 +53,7 @@ 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) diff --git a/op_tests/test_gemm_a16w16_graph.py b/op_tests/test_gemm_a16w16_graph.py deleted file mode 100644 index eb1382d6ea..0000000000 --- a/op_tests/test_gemm_a16w16_graph.py +++ /dev/null @@ -1,76 +0,0 @@ -# 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()