Files
kernbench2/docs/adr-proposed/ADR-0062-prog-tl-async-load.md
T
ywkang f4f55b2c1c gqa(adr): add supporting feature ADRs 0061-0063 for GQA fused attention
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>
2026-06-03 17:38:55 -07:00

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 emits DmaReadCmd and the greenlet kernel suspends until PE_DMA signals completion (src/kernbench/triton_emu/tl_context.py:177-203; greenlet drive in kernel_runner.py). No two tl.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 by tl.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_async charges 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 as recv_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

  1. Overlap is real: a kernel that issues load_async then a GEMM of comparable duration completes in ≈max(load, gemm), not load+gemm (assert end-to-end latency strictly below the serial sum).
  2. Correctness: tl.wait(load_async(...)) returns the same bytes as blocking tl.load(...) for the same address (Phase 2).
  3. Two in flight: two outstanding load_asyncs to distinct addresses both resolve to correct, independent tensors.
  4. op_log compatibility: each load_async still logs exactly one memory/dma_read, so dma_read_count metrics are unchanged vs the blocking version.