Proposed prerequisites surfaced while evaluating the GQA fused-attention ADR against the actual kernbench tl/sim_engine implementation: - ADR-0061 tl.broadcast: data-faithful GQA head reuse (fixes the MemoryStore nbytes check that forces h_q==h_kv==1 today). - ADR-0062 tl.load_async: non-blocking HBM tile load for KV prefetch (KV-load-bound decode/long-context overlap). - ADR-0063 tl.scratch_scope: per-tile scratch recycling (removes the 1 MiB bump-allocator ceiling that caps context at S=16). Co-Authored-By: Claude Opus 4.8 (1M context) <noreply@anthropic.com>
6.0 KiB
ADR-0062: tl.load_async — non-blocking HBM tile load for KV prefetch
Status
Proposed
Supporting ADR for ADR-0060 (AHBM GQA Fused Attention). Decode and long-context attention are KV-load-bound; the dominant lever is overlapping the next K/V tile's DMA with the current tile's compute. Today the only async primitive is for IPCQ comms, not for HBM loads.
Context
What overlap requires
FlashAttention streams K/V tiles: while the MATH/GEMM engine works on
tile j (Q·Kⱼᵀ → softmax → P·V), the DMA engine should already be
pulling tile j+1 (and j+2 at prefetch depth 2). With that overlap a
KV-load-bound kernel runs at roughly max(compute, dma) per tile instead
of compute + dma.
What exists
tl.load(ptr, shape, dtype)is blocking: it emitsDmaReadCmdand the greenlet kernel suspends until PE_DMA signals completion (src/kernbench/triton_emu/tl_context.py:177-203; greenlet drive inkernel_runner.py). No twotl.loads can be in flight from one kernel.- An async pattern does exist, but only for IPCQ:
tl.recv_async(dir, ...) -> RecvFuture+tl.wait(future)(tl_context.py:543-560,tl_context.py:660-693). It proves the machinery — a non-blocking command that returns a future, resolved later bytl.wait— works in the greenlet model. - The DMA engine already models a read channel as a SimPy resource
separate from the write channel
(
src/kernbench/components/builtin/pe_dma.py), so concurrent in-flight reads are representable at the component layer; only the kernel-facing API serialises them.
There is no tl.load_async. Prefetch/double-buffering cannot be
expressed.
Decision
Add a non-blocking HBM load that mirrors the existing recv_async/wait
contract.
D1. tl surface
def load_async(self, ptr: int, shape: tuple[int, ...],
dtype: str = "f16") -> LoadFuture:
"""Issue a DMA read and return immediately. Resolve with tl.wait(fut)
-> TensorHandle. Multiple loads may be in flight; ordering of
resolution is by tl.wait calls, not issue order."""
tl.wait is extended to accept a LoadFuture (it already dispatches on
CompletionHandle vs RecvFuture — add a third arm), returning the
loaded TensorHandle.
D2. Command
Reuse the existing DmaReadCmd with a blocking=False flag (mirroring
IpcqRecvCmd.blocking), or a thin DmaReadAsyncCmd sibling — whichever
keeps PE_DMA's handler simplest. The op_log entry is unchanged
(memory/dma_read); asynchrony is a scheduling property, not a new op
kind, so existing op_log consumers and the dma_read_count metric
(milestone bench) keep working.
D3. Latency / overlap semantics
load_asynccharges the issue (descriptor push) only; the kernel proceeds.- The DMA transfer occupies the read channel for its modelled duration in parallel with whatever compute the kernel issues next.
tl.wait(fut)blocks only if the transfer has not finished; if it has, it returns immediately (same fast-path asrecv_async/wait).- Determinism is preserved: completion is a scheduled event on the modelled DMA channel/link (SPEC §0.1, R8). No magic — the overlap is real modelled concurrency, not a hand-waved latency subtraction.
D4. Double-buffer usage (the intended pattern)
# Prime: issue first two tile loads.
f0 = tl.load_async(K_base + 0*tile_bytes, (TILE, d))
f1 = tl.load_async(K_base + 1*tile_bytes, (TILE, d))
for j in range(n_tiles):
Kj = tl.wait(f_cur) # resolves the j-th load
if j + 2 < n_tiles: # keep depth-2 pipeline full
f_next2 = tl.load_async(K_base + (j+2)*tile_bytes, (TILE, d))
s = tl.dot(q, tl.trans(Kj)) # compute overlaps f_{j+1}'s DMA
...
(V tiles double-buffer the same way; ADR-0060 §3 schedules the V load so it overlaps the Q·Kᵀ + softmax of the same tile.)
Alternatives
A1. Use a tiled tl.composite instead
tl.composite already pipelines tiles inside PE_SCHEDULER (DMA of tile
j+1 overlaps compute of tile j) and recycles per-tile scratch — so
the composite path gets prefetch "for free." Rejected as the sole
mechanism because today's composite is a single GEMM head + math
epilogues (pe_commands.py:144-162); it cannot express the
two-GEMM-with-running-(m,l,O) flash inner loop. Building a bespoke
"flash composite" kind is a much larger change than load_async, which
is a small, general primitive that also helps non-composite kernels.
load_async and a future flash-composite are not mutually exclusive
(ADR-0060 §8).
A2. Rely on recv_async only (no HBM async)
Only works for the Ring-Attention path where KV arrives over IPCQ (comms, already overlappable). The decode and no-SP paths load KV from local HBM, which has no async path today. Insufficient.
Consequences
Positive
- Enables the headline KV-load-bound optimisation for decode and long-context, in the proven greenlet model.
- Symmetric with
recv_async/wait— low conceptual surface area. - General: any bandwidth-bound kernel can prefetch.
Negative
- Kernels must manage futures and buffer lifetimes explicitly (the double-buffer dance). Mitigated by keeping the pattern in a shared helper used by all four ADR-0060 cases.
- Interacts with scratch lifetime (ADR-0063): in-flight buffers must not
be recycled before their
tl.wait. The recycling scope must exclude live prefetch buffers.
Test Requirements
- Overlap is real: a kernel that issues
load_asyncthen a GEMM of comparable duration completes in ≈max(load, gemm), notload+gemm(assert end-to-end latency strictly below the serial sum). - Correctness:
tl.wait(load_async(...))returns the same bytes as blockingtl.load(...)for the same address (Phase 2). - Two in flight: two outstanding
load_asyncs to distinct addresses both resolve to correct, independent tensors. - op_log compatibility: each
load_asyncstill logs exactly onememory/dma_read, sodma_read_countmetrics are unchanged vs the blocking version.