Files
kernbench2/docs/adr-proposed/ADR-0065-prog-flat-ops-composite-softmax-merge.md
T
ywkang 95cecccd8e gqa(adr): ADR-0064/0065 review fixes — composite size cap, type-aware logical_bytes, R recalibration
ADR-0064 Revision 2 review fixes:
- D7 NEW: composite size cap (MAX_COMPOSITE_LOGICAL_BYTES default 1024
  bytes). Oversized recipes deterministically segmented into N
  CompositeCmds; each segment incurs its own dispatch cost. Models real
  HW limits (descriptor queue entry, scheduler parser buffer, command
  SRAM) and prevents the model from rewarding pathologically-large
  fused composites.
- D2: type-aware extra-field byte counting (int/float=4, bool=1,
  tuple/list=1+4N, str=1) — replaces uniform 4 bytes per extra.
- D3: recalibrated defaults to FIXED=40 cycles, R=0.0625 cycles/byte
  (16 B/cycle — typical on-die descriptor queue width); anchor stays at
  ~43 ns for typical 1-OpSpec composite. Clarified anchor description:
  DMA stages do not appear in logical_bytes (auto-inserted by
  PE_SCHEDULER from operand.space per ADR-0065 D4).
- D4: removed clock_freq_ghz from pe_cost_model: override block;
  conversion uses the PE node's existing clock_freq_ghz attr. Added
  max_composite_logical_bytes knob.
- Context: emphasized command-count reduction (FIXED) as the primary
  signal; byte term as secondary refinement.
- Open review: added large-composite scheduler-cost stress test.
- Test req: added composite-size-cap (#8) and R-sensitivity sweep (#9).

ADR-0065 + DDD-0065 follow-on updates:
- opt2 vs opt3 dispatch ratio updated 2.4× → ≈4.0× under new defaults
  (FIXED-dominated, reflecting the corrected framing).
- Test req #9: decode opt2 composite fits within 1024-byte cap; no
  segmentation needed for the GQA workload.
- DDD §6: TLContext lowering checks logical_bytes against cap (step 8).
- DDD §11: performance model recomputed with new defaults + sensitivity
  table across R ∈ {0.25, 0.0625, 0.03125} confirming opt2 < opt3 holds.
- DDD §9 P6 gate: ratio band 2.4×±10% → 4.0×±15%; sensitivity sweep added.

Co-Authored-By: Claude Opus 4.7 (1M context) <noreply@anthropic.com>
2026-06-10 16:21:58 -07:00

17 KiB
Raw Blame History

ADR-0065: Flat-ops CompositeCmd + first stateful recipe (softmax_merge)

Status

Proposed (verification gated on ADR-0064 Revision 2 land)

Supporting ADR for ADR-0060 (AHBM GQA Fused Attention). Implements the §5.6 / §8 item 4 carve-out exactly: opt2 of decode = (#1 existing GEMM composite for Q·Kᵀ) + (#2 a single composite carrying softmax + P·V + online-softmax merge). The "ex_composite" / "flash_pv_merge" name referenced in ADR-0060 §5.6 is retired — this ADR fixes the canonical shape using the existing tl.composite entry with two structural additions: (a) a flat ops tuple, and (b) a first stateful recipe softmax_merge that lowers to a sequence of MATH micro-ops.

Context

What ADR-0060 left to this ADR

ADR-0060 §5.6 shipped opt3 (software pipelining) as the recommended now, and explicitly carved out opt2 — fewer per-tile dispatches, K-before-V DMA priority, measurable only after ADR-0064's cost model — as the follow-up.

§8 item 4 sized the carve-out: "if revisited, split it into two composites — #1 = Q·Kᵀ (existing composite + scale), #2 = softmax + P·V + online-softmax accumulator merge". This ADR implements #2.

Why composite needs a structural change

Today's CompositeCmd has the shape (op, a, b, out_addr, ops=(head, *epi)) with implicit conventions:

  • op selects the head engine (gemm | math)
  • ops[0] is the head OpSpec, ops[1:] are epilogue OpSpecs that run after the head per-output-tile or per-K-tile (scope-determined)

For decode opt2 we need MATH ops to run before the head GEMM (the online-softmax m, l, O update + emit P for the GEMM input). The current shape has no slot for that.

Two failed paths considered earlier:

  • Embedding softmax_merge as an epilogue of #1 (Q·Kᵀ) — wrong: it needs to read Sj (= #1's output) and writes a fresh P consumed by the next GEMM (#2's P·V); circular.
  • A single mixed-engine recipe flash_pv_merge absorbing MATH + GEMM + MATH — violates the engine boundary, complicates PE_SCHEDULER.

The clean shape is: drop the head/epilogue distinction from the command format. Make ops a flat ordered tuple; let each op's scope + position determine its tile-loop placement. The "prologue" concept stays in the user-facing tl.composite(...) API as ergonomics — but the HW command never sees it.

Why a recipe (not a generic op list)

Decode opt2's #2 MATH chain is 8 steps with a fixed structure (online softmax merge). Letting the kernel author wire those 8 steps would be both verbose and a hazard surface (correctness depends on step order). A recipe — a single named op (softmax_merge) that TLContext expands to the 8 steps with addresses pre-filled — keeps the kernel concise and the structure under one source of truth. The recipe table lives in TLContext (the compiler analog); PE_SCHEDULER never reads it. See boundary in D5.

Decision

D1. CompositeCmd is a flat ordered op list

@dataclass(frozen=True)
class CompositeCmd:
    completion: CompletionHandle
    ops: tuple[OpSpec, ...]                       # ordered MATH/GEMM ops
    rw_handles: tuple[TensorHandle, ...] = ()     # cross-composite hazard
    data_op: bool = True

    @property
    def logical_bytes(self) -> int: ...           # ADR-0064 D2

The previous (op, a, b, out_addr, out_nbytes, math_op) fields are absorbed into ops. The first OpSpec in ops (if kind == "gemm") is the head GEMM; preceding OpSpecs are pre-loop ops, following OpSpecs are post-loop / per-tile epilogue ops depending on scope.

D2. OpSpec carries named operands

@dataclass(frozen=True)
class OpSpec:
    kind: str                                     # "gemm" | "rmax" | ...
    scope: Scope                                  # KERNEL | K_TILE | OUTPUT_TILE
    operands: dict[str, TensorHandle]             # named — see D5
    extra: dict[str, Any]                         # scalars, axes, m/k/n, etc.
    out: TensorHandle | None = None               # explicit write-back handle

    @property
    def logical_bytes(self) -> int: ...

Named operands let PE_SCHEDULER route inputs to engine ports (e.g., GEMM takes operands["a"], operands["b"]) without positional convention.

D3. Position + scope determines tile-loop placement

PE_SCHEDULER scans cmd.ops for the GEMM op (≤1 per CompositeCmd, see D6). Calling its index g:

OpSpec position scope Placement in plan
0 .. g-1 KERNEL pre-loop stages (run once before tile loop)
g KERNEL head GEMM (drives tile loop with extra["m"], extra["k"], extra["n"])
g+1 .. K_TILE per K-tile epilogue (inside accumulation loop)
g+1 .. OUTPUT_TILE per (m, n) tile epilogue (after K accumulation)
g+1 .. KERNEL post-loop stages (run once after tile loop)

A composite with no GEMM op (e.g., MATH-only composite) treats all ops as KERNEL-scope sequential.

D4. PE_SCHEDULER auto-inserts DMAs from operand space

PE_SCHEDULER scans the GEMM's operands and out. For each:

  • space == "hbm" → emit DMA_READ Stage (or DMA_WRITE for out) before the FETCH/GEMM Stages
  • space == "tcm" → no DMA Stage, operand consumed in place

This is already partially modelled today via a_pinned/b_pinned; D4 makes the rule uniform and visible: the kernel never emits explicit DMA cmds for composite operands; PE_SCHEDULER infers them from handles. tl.ref(addr, shape) returns a handle with space="hbm"; tl.zeros / tl.full / tl.load outputs return space="tcm".

Non-GEMM ops (MATH chain in prologue / epilogue) have TCM-resident operands (kernel allocates via tl.zeros / tl.full); no DMA inserted for them.

D5. RECIPE_DESCRIPTORS — TLContext-internal, NOT in pe_commands.py

Recipes live in a new TLContext-adjacent module (src/kernbench/triton_emu/tl_recipes.py). PE_SCHEDULER does not import it. The first recipe:

@dataclass(frozen=True)
class RecipeDescriptor:
    # Operand contract: name → "R" (read-only) | "RW" (read-write)
    operands: dict[str, str]
    # Type / shape rule for the implicit primary output (auto-bound to
    # head's first matrix operand). None means recipe emits no value.
    primary_out: PrimaryOutSpec | None
    # Single-shot prologue (full reduction) vs tile-aligned execution.
    tile_alignment: Literal["single_shot", "tile_aligned"]
    # Internal scratch budget (bytes), function of operand dims.
    internal_scratch_bytes_fn: Callable[..., int]
    # Engine micro-op sequence: TLContext expands this into flat OpSpecs.
    engine_seq: tuple[EngineOp, ...]


RECIPE_DESCRIPTORS["softmax_merge"] = RecipeDescriptor(
    operands={"s": "R", "m": "RW", "l": "RW", "O": "RW"},
    primary_out=PrimaryOutSpec(
        from_shape="s", from_dtype="s", transform="identity",
    ),
    tile_alignment="single_shot",
    internal_scratch_bytes_fn=lambda G, TILE, d, bpe: bpe * (
        G + G + G + G * TILE + G  # m_loc + m_new + corr + P + l_loc
    ),
    engine_seq=(
        EngineOp("MATH", "rmax",       src="s",         dst="m_loc",  reduce_axis=-1),
        EngineOp("MATH", "max_elem",   src_a="m",       src_b="m_loc", dst="m_new"),
        EngineOp("MATH", "exp_diff",   src_a="m",       src_b="m_new", dst="corr"),
        EngineOp("MATH", "exp_diff",   src_a="s",       src_b="m_new", dst="P",     bcast_axis=0),
        EngineOp("MATH", "rsum",       src="P",         dst="l_loc",   reduce_axis=-1),
        EngineOp("MATH", "fma",        src_a="l",       src_b="corr",  src_c="l_loc", dst="l"),
        EngineOp("MATH", "mul_bcast",  src_a="O",       src_b="corr",  dst="O",     bcast_axis=1),
        EngineOp("MATH", "copy",       src="m_new",     dst="m"),
    ),
)

TLContext, on tl.composite(prologue=[{"op": "softmax_merge", ...}], op="gemm", ...):

  1. Validates operand R/RW against RECIPE_DESCRIPTORS["softmax_merge"].
  2. Allocates scratch (m_loc, m_new, corr, P, l_loc) via the scratch_scope helper (ADR-0063).
  3. Derives the primary output P's shape from s.shape (identity).
  4. Expands engine_seq into 8 flat MATH OpSpecs with all addresses / sizes filled. Each OpSpec gets scope=KERNEL.
  5. Auto-binds head GEMM's operands["a"] = P_handle.
  6. Emits CompositeCmd(ops=(8 MATH + 1 GEMM + epilogue), rw_handles= (m, l, O)).

RECIPE_DESCRIPTORS appears nowhere in the HW path. PE_SCHEDULER sees only the flat ops list.

D6. Invariants

  1. GEMM count. 0 ≤ count(op.kind == "gemm" in cmd.ops) ≤ 1. A composite with 0 GEMMs is MATH-only.
  2. softmax_merge has no V operand. RECIPE_DESCRIPTORS["softmax_merge" ].operands does not include any HBM-resident handle → no DMA inserted during the prologue → V DMA can only happen during the GEMM head → natural K-before-V DMA priority across #1 / #2 of decode opt2.
  3. Strict FIFO across composites. PE_SCHEDULER tracks in-flight rw_handles. A new composite whose rw_handles (or any of its op inputs) intersect with an in-flight composite's rw_handles waits for all earlier composites to complete — no out-of-order reorder.
  4. No backwards-incompatibility for legacy callers. The user-facing API tl.composite(op="gemm", a, b, epilogue=[...]) is preserved. TLContext internally lowers it to a flat-ops CompositeCmd.
  5. PE_SCHEDULER recipe-free. PE_SCHEDULER does not import RECIPE_DESCRIPTORS, does not know that "softmax_merge" exists, and has no recipe-specific code branches.

D7. Boundary summary (compiler vs scheduler vs engine)

HOST                                       DEVICE
TLContext (compiler analog)          PE_SCHEDULER (scheduler/dispatcher)
  - validate against RECIPE_DESC.       - scan flat ops list
  - derive primary_out shape            - identify GEMM (≤1)
  - allocate scratch                    - position-based stage placement
  - expand recipe → flat OpSpecs        - auto-insert DMA stages from space
  - compute rw_handles                  - strict-FIFO RW hazard tracking
  - emit CompositeCmd                   - emit Stage list
  imports: RECIPE_DESCRIPTORS            imports: nothing recipe-related

                ENGINES (PE_MATH / PE_GEMM / PE_DMA)
                  - read Stage.params (op_kind, addresses, n_elements)
                  - existing latency models unchanged
                  imports: nothing recipe-related

Alternatives

A1. Three composites per tile (no prologue concept)

Split #2 further into softmax_merge + gemm(P·V) + add(O) as three separate CompositeCmds. Pros: no flat-ops restructuring (add reuses existing epilogue path). Cons: 3 dispatches per tile instead of 2 — diverges from ADR-0060 §8 item 4 sizing ("split into two"); loses ~32 ns × N_tiles of fixed-cost savings. Rejected.

A2. Mixed-engine recipe (flash_pv_merge)

One recipe absorbing MATH + GEMM + MATH (ADR-0060 §5.6 original phrasing). Pros: tightest dispatch (1 composite). Cons: PE_SCHEDULER must know engine-boundary recipes, violates the "engines see only Stage.op_kind" boundary, recipe table becomes ISA-like and pollutes HW interface. Rejected.

A3. Generic op-list DAG (tl.ex_composite([...]))

User-defined sequence of any ops with named slots and dataflow analysis. Pros: maximally flexible. Cons: a single user (softmax_merge) does not justify a generic DAG language; PE_SCHEDULER would need dataflow analysis; reintroduces ADR-0060 §8 item 4's "absorb the inner loop" trap. Rejected (Simplicity First).

A4. RW-aware scheduler reorder (not strict FIFO)

PE_SCHEDULER could allow a later CompositeCmd to overtake an in-flight one when their rw_handles don't intersect (e.g., next-tile #1 could overlap with this-tile #2 since #1 doesn't touch m/l/O). Pros: more overlap. Cons: more scheduler state, gains are bounded (next #1 waits for K DMA anyway). Deferred — strict FIFO ships first; RW-aware reorder may be a future follow-up ADR.

A5. Embedding softmax_merge as an epilogue of #1 (Q·Kᵀ)

Wrong: softmax_merge's output P is the input of #2's P·V GEMM, which runs after #1's epilogue. Embedding it would force #2 to read its own input from #1's epilogue scratch — circular tile-loop dependency. Rejected (incorrect).

Consequences

Positive

  • ADR-0060 §5.6 / §8 item 4 carve-out implemented exactly: opt2 = 2 composites per tile, K-before-V DMA priority natural.
  • Composite contract becomes uniform: flat ops list with scope-based placement. No structural distinction between "head" and "epilogue" in the HW cmd.
  • DMAs auto-inferred from operand space; kernel surface simpler.
  • HW interface (CompositeCmd shape) stays small even as recipes grow — recipe expansion happens host-side.
  • New fused patterns (linear+gelu+norm, rmsnorm+linear, …) can be added by inserting one RECIPE_DESCRIPTORS entry; no PE_SCHEDULER changes.

Negative

  • CompositeCmd's struct shape changes — a refactor for all current callers. Semantics preserved (D6.4); existing goldens unchanged once ADR-0064 Revision 2 lands.
  • OpSpec.operands changes from positional tuple[Any, ...] to dict[str, TensorHandle] — touches the existing epilogue lowering.
  • New TLContext code: recipe expansion + scratch-slot allocation + primary_out auto-bind. ~80120 LOC.
  • PE_SCHEDULER's _generate_plan gets a position-based scan + DMA auto-insertion + strict-FIFO RW tracker. ~60100 LOC.

Open review items

  1. Recipe shape derivation rule beyond identity. Future recipes may need non-identity transforms (e.g., transposed primary_out). The PrimaryOutSpec.transform field is forward-compat; add new transforms as needed.
  2. Recipe scratch allocator integration with ADR-0063's scratch_scope. Initial wiring: TLContext allocates inside the current scratch_scope context if active, else in a kernel-scope persistent slot.
  3. tl_recipes.py location. New module under triton_emu/. Keeps recipe knowledge co-located with the compiler analog (TLContext).
  4. GEMM out=TensorHandle. New form alongside legacy out_ptr: int. Both accepted; handle form preferred for new code (enables strict-FIFO RW tracking by handle identity).

Test Requirements

  1. CompositeCmd flat-ops refactor — meaning preservation. Every existing bench's op_log is byte-equal pre/post refactor (legacy API tl.composite(op="gemm", a, b, epilogue=[...]) lowers to the same Stage sequence).
  2. softmax_merge recipe lowering. TLContext call tl.composite( prologue=[{"op": "softmax_merge", "s": Sj, "m": m, "l": l, "O": O}], op="gemm", b=V_ref, out=O, epilogue=[{"op": "add", "other": O}]) produces a CompositeCmd with exactly 10 ops in the expected order (8 MATH + 1 GEMM + 1 MATH(add)) and rw_handles == (m, l, O).
  3. K-before-V DMA priority invariant. op_log of decode opt2 shows no V-related DMA during #2's MATH prologue (#2's V DMA begins only when the GEMM head starts).
  4. Strict-FIFO RW serialization. Two consecutive composites sharing any rw_handle complete in dispatch order; their stages do not interleave.
  5. DMA auto-insertion from space. A GEMM composite with b=tl.ref(...) (space=hbm) emits a DMA_READ Stage; with b from tl.zeros(...) (space=tcm) it does not. Output handle with space=hbm emits DMA_WRITE; space=tcm does not.
  6. opt2 numeric parity with opt3. In data mode, opt2's final (m, l, O) matches opt3's within fp tolerance.
  7. opt2 dispatch ratio (after ADR-0064 Rev2). opt3 vs opt2 per-tile PE_CPU dispatch cycles ratio ≈ 4.0× at default calibration (FIXED=40 cycles, R=0.0625 cycles/byte). Ratio is FIXED-dominated, reflecting command-count reduction as the primary signal.
  8. GEMM-count invariant. A composite with two GEMM OpSpecs raises a validation error at TLContext emit.
  9. Within composite size cap (ADR-0064 D7). Decode opt2's #2 composite (10 ops, ~310 logical bytes) fits comfortably under the default MAX_COMPOSITE_LOGICAL_BYTES=1024 — no segmentation. The recipe lowering test asserts cmd.logical_bytes < 1024.

Dependencies

  • ADR-0060 §5.6, §8 item 4 — the carve-out this ADR implements.
  • ADR-0064 Revision 2 — dispatch cost model; lands first, makes opt2's fewer-issues win measurable.
  • ADR-0063 scratch_scope — recipe intermediates allocated inside; m, l, O allocated outside as persistent arena (per DDD-0060 §6.2).
  • ADR-0046 tl_context contract — extended by prologue=[...] kwarg, out=TensorHandle, RECIPE_DESCRIPTORS lookup.
  • ADR-0042 tile plan generators — extended (not redesigned) to process the flat ops list and auto-insert DMAs.
  • ADR-0014 PE pipeline — boundary preserved: scheduler stays recipe-free, engines stay op_kind-opaque.

Migration

Land order:

  1. ADR-0064 Revision 2 (separate PR): structural dispatch cost + logical_bytes property + topology config override + one-time goldens regeneration.
  2. ADR-0065 Phase 1 (this ADR, PR a): CompositeCmd flat-ops restructure + legacy-API lowering. Refactor only; goldens unchanged.
  3. ADR-0065 Phase 2 (this ADR, PR b): tl_recipes.py + softmax_merge recipe + PE_SCHEDULER position-scan + DMA auto-insertion + strict-FIFO RW tracker.
  4. ADR-0065 Phase 3 (this ADR, PR c): _gqa_decode_long.py opt2 variant + numeric parity test + dispatch-ratio measurement against opt3 baseline.

Detailed implementation plan: see DDD-0065.