Files
kernbench2/src/kernbench/benches/_attention_mesh_mlo.py
T
mukesh d0f904c57e attention: docstring cleanup — clarify single/multi user per GQA study
Both mesh kernels and the milestone-gqa bench described single_user vs
multi_user in ways that were architecturally misleading. Source of truth
is the GQA Llama-70B sharding study at
llm_paper_review/notes/GQA_MHA_sharding/scripts/_gen_llama70b_1M_4cases.py:

  Single User (B=1) — TL prefill, BL decode:
    KV split @ PEs intra-cube. Each cube does its own 8-PE ring,
    independent of other cubes. Headline = 64 cubes serving one user.

  Multi User (B=8) — TR prefill, BR decode:
    KV split @ cubes inter-cube. 8 cubes/KV-group form the ring; inside
    each cube, "8 PEs each handle 1 different user → Batch on batch."

Updates the rank_axis docstring on both mesh kernels and the bench's
module docstring to make this accurate, including the v1 simplification
of the multi_user kernel (gating pe_id != 0 collapses B=8 → B=1; the
per-cube batch parallelism is deferred to sub-cycle 4c headline).

No behavior change; existing tests unchanged.

Co-Authored-By: Claude Opus 4.7 (1M context) <noreply@anthropic.com>
2026-06-04 12:37:25 -07:00

168 lines
7.1 KiB
Python
Raw Blame History

This file contains ambiguous Unicode characters
This file contains Unicode characters that might be confused with other characters. If you think that this is intentional, you can safely ignore this warning. Use the Escape button to reveal them.
"""Mesh-native bidirectional AllReduce-mlo attention — decode (ADR-0059 Proposed).
Every rank holds the full Q (replicated, small at ``S_q=1``) and 1/n_ranks
of KV (sequence-sharded). Each rank computes its partial attention
against own KV in ONE shot, then runs a bidirectional fan-out of the
``(m, , o)`` triplet: the triplet originating at rank i reaches rank j at
step ``|i - j|``. Every rank merges every other rank's triplet exactly
once over ``n_ranks - 1`` steps, ending with the final answer replicated
on every rank.
Supersedes ADR-0056's closed-ring ``_attention_allreduce_mlo.py``. Both
modules stay on disk during the transition; this one runs on the
hardware's actual open-mesh wiring (no closed-ring SFR install required).
Imported by ``milestone_gqa_llama70b`` (after the bench's Phase 2 switches
its imports) and invoked through ``torch.launch(...)`` — not through
``dist.all_reduce(...)``. See ADR-0056 Context for why this kernel is not
backend-dispatched via ADR-0050's algorithm-module contract.
"""
from __future__ import annotations
from kernbench.common.pe_commands import TensorHandle
def _view(handle: TensorHandle, new_shape: tuple[int, ...]) -> TensorHandle:
"""Reshape — metadata only, no command emitted (cf. ``tl.trans``)."""
return TensorHandle(
id=handle.id,
addr=handle.addr,
shape=new_shape,
dtype=handle.dtype,
nbytes=handle.nbytes,
data=handle.data,
space=handle.space,
pinned=handle.pinned,
)
def attention_mesh_mlo_kernel(
q_ptr: int,
k_ptr: int,
v_ptr: int,
o_ptr: int,
S_q: int,
S_kv_per_rank: int,
h_q: int,
h_kv: int,
d_head: int,
n_ranks: int,
rank_axis: int = 0,
*,
tl,
) -> None:
"""Mesh-native bidirectional AllReduce-mlo — see module docstring.
``rank_axis`` selects which program-id dimension carries the ring rank,
matching the GQA Llama-70B sharding study's TL/BL vs TR/BR distinction
(`llm_paper_review/notes/GQA_MHA_sharding/scripts/_gen_llama70b_1M_4cases.py`):
0 — single_user_* panels (TL/BL): rank == tl.program_id(axis=0) (PE
id in cube). KV is split @ PEs **intra-cube**; ring runs over
the 8 PEs of one cube (NOC-only). At Llama-70B headline scale
this kernel launches once per cube; 64 such cubes run in
parallel for one user (1 Q-head per cube × 8 cubes per KV-group
× 8 KV-groups). The PE-level ring inside each cube is
independent of the others.
1 — multi_user_* panels (TR/BR): rank == tl.program_id(axis=1)
(cube id). KV is split @ cubes **inter-cube**; ring runs over
the cubes of one KV-group. The kernel gates ``pe_id != 0`` to
return early — a v1 simplification: at headline scale (B=8) the
study's "Batch on batch" pattern would have all 8 PEs each handle
one user's batch element instead of staying silent. Validation
shipped with B=1 to focus on the inter-cube ring's correctness;
adding the per-cube batch dimension is sub-cycle 4c headline work.
"""
# For multi_user (rank_axis=1) only PE 0 in each cube runs the ring.
if rank_axis != 0 and tl.program_id(axis=0) != 0:
return
rank = tl.program_id(axis=rank_axis)
has_E = rank < n_ranks - 1
has_W = rank > 0
# Q is replicated on every rank — loaded once.
Q = tl.load(q_ptr, shape=(S_q, h_q * d_head), dtype="f16")
# Local KV chunk. KV is sequence-sharded and stays put on this rank for
# the entire fan-out — distinguishing decode from prefill (ADR-0059 D3)
# where KV circulates.
K = tl.load(k_ptr, shape=(S_kv_per_rank, h_kv, d_head), dtype="f16")
V = tl.load(v_ptr, shape=(S_kv_per_rank, h_kv, d_head), dtype="f16")
# ── One-shot local partial attention ──────────────────────────
K_2d_T = _view(K, (h_q * d_head, S_kv_per_rank))
V_2d = _view(V, (S_kv_per_rank, h_q * d_head))
scores = tl.dot(Q, K_2d_T)
m = tl.max(scores, axis=-1)
P = tl.softmax(scores, axis=-1)
scores_centered = scores - m
exp_scores = tl.exp(scores_centered)
ell = tl.sum(exp_scores, axis=-1)
o = tl.dot(P, V_2d)
# Seed bidirectional waves with own triplet (step-1 send).
to_send_east_m: TensorHandle | None = m
to_send_east_ell: TensorHandle | None = ell
to_send_east_o: TensorHandle | None = o
to_send_west_m: TensorHandle | None = m
to_send_west_ell: TensorHandle | None = ell
to_send_west_o: TensorHandle | None = o
# Bidirectional fan-out of (m, , o) triplets — n_ranks - 1 steps.
for step in range(1, n_ranks):
# Send eastbound triplet (own at step 1; forwarded at later steps).
if has_E and to_send_east_m is not None:
tl.send(dir="E", src=to_send_east_m)
tl.send(dir="E", src=to_send_east_ell)
tl.send(dir="E", src=to_send_east_o)
# Send westbound triplet.
if has_W and to_send_west_m is not None:
tl.send(dir="W", src=to_send_west_m)
tl.send(dir="W", src=to_send_west_ell)
tl.send(dir="W", src=to_send_west_o)
# Receive eastbound triplet from W (originated at rank - step).
m_from_W: TensorHandle | None = None
ell_from_W: TensorHandle | None = None
o_from_W: TensorHandle | None = None
if has_W and (rank - step) >= 0:
m_from_W = tl.recv(dir="W", shape=m.shape, dtype="f16")
ell_from_W = tl.recv(dir="W", shape=ell.shape, dtype="f16")
o_from_W = tl.recv(dir="W", shape=o.shape, dtype="f16")
m_combined = tl.maximum(m, m_from_W)
scale_old = tl.exp(m - m_combined)
scale_new = tl.exp(m_from_W - m_combined)
ell = ell * scale_old + ell_from_W * scale_new
o = o * scale_old + o_from_W * scale_new
m = m_combined
# Receive westbound triplet from E (originated at rank + step).
m_from_E: TensorHandle | None = None
ell_from_E: TensorHandle | None = None
o_from_E: TensorHandle | None = None
if has_E and (rank + step) < n_ranks:
m_from_E = tl.recv(dir="E", shape=m.shape, dtype="f16")
ell_from_E = tl.recv(dir="E", shape=ell.shape, dtype="f16")
o_from_E = tl.recv(dir="E", shape=o.shape, dtype="f16")
m_combined = tl.maximum(m, m_from_E)
scale_old = tl.exp(m - m_combined)
scale_new = tl.exp(m_from_E - m_combined)
ell = ell * scale_old + ell_from_E * scale_new
o = o * scale_old + o_from_E * scale_new
m = m_combined
# Forward the original received triplet (not the merged running state)
# so neighbors get the original wave. ``None`` propagates if nothing
# arrived this step.
to_send_east_m = m_from_W
to_send_east_ell = ell_from_W
to_send_east_o = o_from_W
to_send_west_m = m_from_E
to_send_west_ell = ell_from_E
to_send_west_o = o_from_E
# Final normalize: O := o / .
O_final = o / ell
tl.store(o_ptr, O_final)