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


# A graph records launches, not the allocation that zeroed the counter. Give
# each recorded launch its own slot and record the zero-fill inside the graph.
# Public: a host that records more than this raises it before its first launch.
CAPTURE_SEMAPHORE_POOL_SLOTS = 4096
_capture_rings: dict[torch.device, Tensor] = {}
_capture_ring_next: dict[torch.device, int] = {}


def _prime_capture_pool(device: torch.device) -> Tensor:
"""Allocate the capture pool. Never called from inside a capture."""
ring = _capture_rings.get(device)
if ring is None:
ring = torch.zeros(
(CAPTURE_SEMAPHORE_POOL_SLOTS, *_SEMA_SHAPE),
dtype=torch.uint32,
device=device,
)
_capture_rings[device] = ring
return ring


def _get_captured_semaphore_workspace(device: torch.device) -> Tensor:
ring = _capture_rings.get(device)
if ring is None:
raise RuntimeError(
f"no split-K a16w16 semaphore pool on {device}: the pool is "
"allocated by an eager call that takes the splitK path (splitK "
"None or > 1); make one on this device before capturing a graph."
)
idx = _capture_ring_next.get(device, 0)
# Bound against the pool as allocated, not the constant: raising the
# constant after allocation must not let idx past the end of the tensor.
if idx >= ring.shape[0]:
raise RuntimeError(
f"split-K a16w16 capture semaphore pool exhausted on {device}: "
f"{ring.shape[0]} slots recorded. Raise "
"aiter.ops.gemm_op_a16w16.CAPTURE_SEMAPHORE_POOL_SLOTS before the "
"first split-K launch; the pool is sized when it is allocated."
)
_capture_ring_next[device] = idx + 1
sema = ring[idx]
sema.zero_() # a graph node, so every replay starts from zero
return sema


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

Expand All @@ -53,7 +98,30 @@ 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 capture this returns a slot from a per-device pool instead, with the
zero-fill recorded as a graph node so replay restores counter == 0. That
pool has to exist before capture starts, so the first eager splitK call on
a device allocates it: a fixed CAPTURE_SEMAPHORE_POOL_SLOTS * 4 KiB, paid
once per device even by a process that never captures a graph.
"""
# torch.device("cuda") means "the current device", which is not a stable
# key: resolve it now so the pool is keyed and allocated per GPU.
if device.index is None:
device = torch.device(device.type, torch.cuda.current_device())

# is_current_stream_capturing() answers about the current device; only pay
# the device switch (~1us, and this is a per-GEMM path) when it differs.
if device.index == torch.cuda.current_device():
capturing = torch.cuda.is_current_stream_capturing()
else:
with torch.cuda.device(device):
capturing = torch.cuda.is_current_stream_capturing()
if capturing:
return _get_captured_semaphore_workspace(device)

# Allocate here, never under capture: graph-pool memory cannot be freed.
_prime_capture_pool(device)
stream = torch.cuda.current_stream(device)
return _get_semaphore_workspace_keyed(device, stream.cuda_stream)

Expand Down
83 changes: 83 additions & 0 deletions op_tests/test_gemm_a16w16.py
Original file line number Diff line number Diff line change
Expand Up @@ -12,6 +12,7 @@
import aiter
from aiter import dtypes, hipb_create_extension, hipb_mm
from aiter.jit.utils.chip_info import get_gfx_runtime as get_gfx
from aiter.ops.gemm_op_a16w16 import get_semaphore_workspace
from aiter.ops.shuffle import shuffle_weight
from aiter.test_common import benchmark, checkAllclose, perftest
from aiter.tuned_gemm import tgemm, triton_gemm
Expand Down Expand Up @@ -428,6 +429,77 @@ def test_skinny_gemm():
return df


def check_graph(dtype, m, n, k, otype):
"""Capture the asm split-K GEMM in a HIP graph, replay it, compare against eager.

Not part of the perf table: this is a pass/fail check that the split-K
semaphore survives capture/replay. The kernel needs its counter to be zero
when a launch starts; a graph records launches, not the memset that zeroed
the counter, so a replay can start dirty and the kernel spins forever.
"""
if dtype != dtypes.bf16 or otype not in (dtypes.bf16, dtypes.fp32):
return # the asm a16w16 path only takes bf16 in, bf16/fp32 out
if k % 64 or n % 64:
return

x = torch.randn(m, k, dtype=dtype, device="cuda")
weight = torch.randn(n, k, dtype=dtype, device="cuda")
wshuffle = shuffle_weight(weight, layout=(16, 16))
out = torch.empty(m, n, dtype=otype, device="cuda")

# Warm up outside capture: the first call loads the module and allocates the
# semaphore workspace, neither of which may happen inside a capture region.
aiter.gemm_a16w16_asm(x, wshuffle, out, bpreshuffle=wshuffle.is_shuffled)
torch.cuda.synchronize()
eager = out.clone()

# Leave the counter dirty on the stream the graph captures on -- the state
# production reaches, and the one an unfixed build cannot recover from.
stream = torch.cuda.Stream()
stream.wait_stream(torch.cuda.current_stream())
with torch.cuda.stream(stream):
get_semaphore_workspace(x.device).fill_(2)
stream.synchronize()

graph = torch.cuda.CUDAGraph()
with torch.cuda.graph(graph, stream=stream):
aiter.gemm_a16w16_asm(x, wshuffle, out, bpreshuffle=wshuffle.is_shuffled)
torch.cuda.current_stream().wait_stream(stream)

out.zero_()
graph.replay()
torch.cuda.synchronize()

# splitK reduces in whatever order the blocks finish, so the same inputs are
# not bit-identical run to run; compare with the file's usual tolerance.
err = checkAllclose(
eager,
out,
msg=f"graph dim: {(m, n, k)!s:<20} dtype: {dtype} otype: {otype}, replay vs eager: ",
catastrophic_check=True,
)
assert err == 0, f"graph replay diverged from eager at {(m, n, k)} {dtype} {otype}"

# The workspace a capture hands out must have its zero-fill recorded as a
# graph node, or replay N starts from whatever replay N-1 left. Checked on
# its own graph, with no kernel, so a regression asserts here in
# milliseconds instead of spinning in the GEMM above.
sink = torch.zeros(1, device=x.device)
probe = torch.cuda.CUDAGraph()
with torch.cuda.graph(probe, stream=stream):
sema = get_semaphore_workspace(x.device)
sink.add_(1) # a graph needs at least one node
torch.cuda.current_stream().wait_stream(stream)
torch.cuda.synchronize()
sema.view(torch.int32).fill_(2)
torch.cuda.synchronize()
probe.replay()
torch.cuda.synchronize()
assert (
int(sema.view(torch.int32).max().item()) == 0
), f"capture recorded no zero-fill for the splitK counter at {(m, n, k)}"


parser = argparse.ArgumentParser(
formatter_class=argparse.RawTextHelpFormatter,
description="config input of a16w16_gemm_test",
Expand Down Expand Up @@ -495,6 +567,13 @@ def test_skinny_gemm():
help="""Scale B.
e.g.: -sb 0.5""",
)
parser.add_argument(
"--graph",
action="store_true",
help="""Also run the HIP-graph capture/replay check for the asm splitK path
over the same sweep. Use shapes that select splitK (small m, large k).
e.g.: --graph -mnk 64,256,5120 32,512,8192 -d bf16 -o fp32""",
)
args = parser.parse_args()

df = []
Expand All @@ -503,6 +582,8 @@ def test_skinny_gemm():
for dtype in args.dtype:
for otype in args.otype:
for m, n, k in args.mnk:
if args.graph:
check_graph(dtype, m, n, k, otype)
ret = test_gemm(
dtype,
m,
Expand All @@ -521,3 +602,5 @@ def test_skinny_gemm():
df = pd.DataFrame(df)
df_md = df.to_markdown(index=False)
aiter.logger.info("gemm_a16w16 summary (markdown):\n%s", df_md)
if args.graph:
aiter.logger.info("all graph capture/replay checks passed")
Loading