From 5bbe293bea71adc5b65265c626d600a85d4e6f5d Mon Sep 17 00:00:00 2001 From: Yangwook Date: Wed, 10 Jun 2026 23:06:01 -0700 Subject: [PATCH] =?UTF-8?q?gqa(adr-0065):=20N3=20=E2=80=94=20GEMM-epilogue?= =?UTF-8?q?=20accumulate=20dataflow=20(O=20=3D=20O.corr=20+=20P.V)?= MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit Two fixes make the recipe's P.V GEMM fold into the rescaled accumulator: (1) op_log marks the composite gemm record accumulate=True when an 'add' epilogue targets the GEMM's own output; data_executor._execute_gemm then does O = O_existing + a@b instead of overwriting. (2) PE_SCHEDULER records the composite's numeric op at COMPLETION (in _retire_on_done) instead of at dispatch, so the GEMM's t_start sorts AFTER its prologue MATH ops -> it reads the computed P and the rescaled O. Verified end-to-end (DataExecutor): the full recipe produces O = O_old*corr + P@V matching a numpy flash-attention reference. Suite 816 pass / 3 pre-existing fail. Co-Authored-By: Claude Opus 4.8 (1M context) --- .../components/builtin/pe_scheduler.py | 26 ++++---- src/kernbench/sim_engine/data_executor.py | 16 ++++- src/kernbench/sim_engine/op_log.py | 8 +++ tests/test_recipe_data_mode.py | 63 +++++++++++++++++++ 4 files changed, 100 insertions(+), 13 deletions(-) diff --git a/src/kernbench/components/builtin/pe_scheduler.py b/src/kernbench/components/builtin/pe_scheduler.py index 6333808..d68b9c8 100644 --- a/src/kernbench/components/builtin/pe_scheduler.py +++ b/src/kernbench/components/builtin/pe_scheduler.py @@ -165,14 +165,9 @@ class PeSchedulerComponent(ComponentBase): from kernbench.components.builtin.pe_types import PipelineContext yield from self._hazard.admit(env, cmd, pe_txn.done) - env.process(self._retire_on_done(env, pe_txn.done, cmd.completion.id)) - - # Record the composite as a single numeric op for data-mode replay - # (ADR-0027): the tiles model latency; this zero-duration record - # carries the matmul (a@b → out) the DataExecutor replays. No-op - # without an op_logger (latency-only mode), so latency is unchanged. - self._on_process_start(env, cmd) - self._on_process_end(env, cmd) + # Retire from the hazard tracker AND record the composite's numeric + # op — both at completion (see _retire_on_done). + env.process(self._retire_on_done(env, pe_txn.done, cmd)) plan = self._generate_plan(cmd) @@ -192,11 +187,20 @@ class PeSchedulerComponent(ComponentBase): yield self._pending_feeds.put((plan, ctx)) def _retire_on_done( - self, env: simpy.Environment, done_event: Any, completion_id: str, + self, env: simpy.Environment, done_event: Any, cmd: Any, ) -> Generator: - """Retire a composite from the RW hazard tracker on completion.""" + """On composite completion: retire from the RW hazard tracker AND + record its numeric op (ADR-0027 / ADR-0065 N1+N3). + + Recording at completion (not dispatch) gives the composite's GEMM a + t_start AFTER its prologue MATH ops, so in data mode it reads their + results (P, the rescaled O) instead of stale values. No-op without an + op_logger → latency-only mode unchanged. + """ yield done_event - self._hazard.retire(completion_id) + self._hazard.retire(cmd.completion.id) + self._on_process_start(env, cmd) + self._on_process_end(env, cmd) def _feed_loop(self, env: simpy.Environment) -> Generator: """Single feeder process: FIFO command ordering (ADR-0014 D6). diff --git a/src/kernbench/sim_engine/data_executor.py b/src/kernbench/sim_engine/data_executor.py index f00e952..522a1cf 100644 --- a/src/kernbench/sim_engine/data_executor.py +++ b/src/kernbench/sim_engine/data_executor.py @@ -174,9 +174,21 @@ class DataExecutor: dtype_acc = p.get("dtype_acc", "f32") a_f = a.astype(_resolve_dtype(dtype_acc)) b_f = b.astype(_resolve_dtype(dtype_acc)) - result = np.matmul(a_f, b_f).astype(_resolve_dtype(dtype_out)) + result = np.matmul(a_f, b_f) - self.store.write(dst_space, p["dst_addr"], result) + # Accumulate into the existing output when requested (recipe O += + # P·V — the GEMM folds into the already-rescaled accumulator). The + # composite is recorded at completion so the rescale has run first. + if p.get("accumulate"): + try: + existing = self.store.read(dst_space, p["dst_addr"], + shape=result.shape, dtype=dtype_out) + result = existing.astype(_resolve_dtype(dtype_acc)) + result + except (KeyError, ValueError): + pass + + self.store.write(dst_space, p["dst_addr"], + result.astype(_resolve_dtype(dtype_out))) def _execute_math(self, op: OpRecord) -> None: """Execute math op: unary, binary, or reduction.""" diff --git a/src/kernbench/sim_engine/op_log.py b/src/kernbench/sim_engine/op_log.py index c2eea79..1b20bd9 100644 --- a/src/kernbench/sim_engine/op_log.py +++ b/src/kernbench/sim_engine/op_log.py @@ -262,10 +262,18 @@ def _extract_op_info(msg: Any) -> tuple[str, str, dict[str, Any]]: if gemm is not None: a = gemm.operands.get("a") b = gemm.operands.get("b") + # Accumulate (O += a@b) when an `add` epilogue targets the GEMM's + # own output (the recipe pattern: GEMM P·V folds into the already- + # rescaled O). Default overwrite for a plain GEMM composite. + accumulate = any( + o.kind == "add" and o.operands.get("other") is gemm.out + for o in msg.ops + ) params: dict[str, Any] = { "op": "gemm", "out_addr": gemm.out.addr, "out_nbytes": gemm.out.nbytes, + "accumulate": accumulate, } if a is not None and b is not None: params.update({ diff --git a/tests/test_recipe_data_mode.py b/tests/test_recipe_data_mode.py index 7dccf47..98a8a6e 100644 --- a/tests/test_recipe_data_mode.py +++ b/tests/test_recipe_data_mode.py @@ -93,3 +93,66 @@ def test_softmax_merge_computes_m_l_O_rescale(): assert np.allclose(m_got, m_new, atol=2e-2), f"m: {m_got.ravel()[:3]} vs {m_new.ravel()[:3]}" assert np.allclose(l_got, l_new, rtol=5e-2, atol=5e-2), "l mismatch" assert np.allclose(O_got, O_resc, rtol=5e-2, atol=5e-2), "O rescale mismatch" + + +def test_softmax_merge_full_accumulates_O(): + """The full recipe (prologue MATH + GEMM P·V + accumulate) in data mode: + O = O_old·corr + P·V (ADR-0065 N3). The GEMM record is recorded at + completion (sorts after the prologue) and accumulates into the rescaled O.""" + from kernbench.sim_engine.op_log import OpRecord, _extract_op_info + + rng = np.random.default_rng(1) + s = rng.standard_normal((G, TILE)).astype(np.float16) + m0 = rng.standard_normal((G, 1)).astype(np.float16) + l0 = np.abs(rng.standard_normal((G, 1))).astype(np.float16) + O0 = rng.standard_normal((G, D)).astype(np.float16) + Vd = rng.standard_normal((TILE, D)).astype(np.float16) + + sh = _tcm(0x1000, (G, TILE)) + mh = _tcm(0x2000, (G, 1)) + lh = _tcm(0x3000, (G, 1)) + Oh = _tcm(0x4000, (G, D)) + V = TensorHandle(id="V", addr=0x5000, shape=(TILE, D), dtype="f16", + nbytes=2 * TILE * D, space="hbm") + + tl = TLContext(pe_id=0, num_programs=1, scratch_base=0x100000, + scratch_size=1 << 20) + tl.composite( + prologue=[{"op": "softmax_merge", "s": sh, "m": mh, "l": lh, "O": Oh}], + op="gemm", b=V, out=Oh, + epilogue=[{"op": "add", "other": Oh}], + ) + cmd = [c for c in tl.commands if isinstance(c, CompositeCmd)][-1] + plan = generate_plan_from_ops(cmd.ops, tile_m=128, tile_k=128, tile_n=128, + bytes_per_element=2, + pe_prefix="sip0.cube0.pe0") + _kind, _name, gparams = _extract_op_info(cmd) + assert gparams.get("accumulate") is True + + store = MemoryStore() + store.write("tcm", sh.addr, s) + store.write("tcm", mh.addr, m0) + store.write("tcm", lh.addr, l0) + store.write("tcm", Oh.addr, O0) + store.write("hbm", V.addr, Vd) + + records = [ + OpRecord(t_start=float(i), t_end=float(i), component_id="pe_math", + op_kind="math", op_name=st.params["op"], params=dict(st.params)) + for i, st in enumerate(plan.prologue_stages) + ] + # GEMM recorded at completion → sorts after the prologue MATH ops. + records.append(OpRecord(t_start=100.0, t_end=100.0, component_id="pe_gemm", + op_kind="gemm", op_name="composite_gemm", + params=gparams)) + DataExecutor(records, store).run() + + sf = s.astype(np.float32) + m_new = np.maximum(m0.astype(np.float32), sf.max(-1, keepdims=True)) + corr = np.exp(m0.astype(np.float32) - m_new) + P = np.exp(sf - m_new) + O_expected = O0.astype(np.float32) * corr + P @ Vd.astype(np.float32) + + O_got = store.read("tcm", Oh.addr, shape=(G, D), dtype="f16").astype(np.float32) + assert np.allclose(O_got, O_expected, rtol=8e-2, atol=8e-2), \ + f"O accumulate mismatch: {O_got[0, :3]} vs {O_expected[0, :3]}"