From b1102589e1c5f5d7e71b310382000588d8374d7f Mon Sep 17 00:00:00 2001 From: alacour Date: Thu, 3 Sep 2026 19:20:40 -0700 Subject: [PATCH 1/6] =?UTF-8?q?edge=5Fframe=5Fkernel:=20ECENET=5FEF=5FWARP?= =?UTF-8?q?S=20=E2=80=94=20sweepable=20warps=20on=20every=20launch?= MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit The edge-frame family is the largest remaining kernel cost (~25 ms/step at 512 atoms: _ef_bwd_merged 2.33 ms/call, _pu_bwd_merged 1.98, fwd ~1.0), all running ~3.5x above their traffic bound. Unlike the RealSpace case the loads are already coalesced tiles, so the gap is some mix of padded ieee tl.dot arithmetic (9 valid of 16 in both dot dims), 36-byte row misalignment, and per-edge program overhead — not separable without measurement. First step: make warps-per-program env-tunable (default 4, the previous implicit value) on all ten launch sites so the cheap axis can be swept before any redesign. Co-Authored-By: Claude Fable 5 Claude-Session: https://claude.ai/code/session_019WyADHSVvKbQnYo1PAjGPK --- ecenet/edge_frame_kernel.py | 29 ++++++++++++++++++++++------- 1 file changed, 22 insertions(+), 7 deletions(-) diff --git a/ecenet/edge_frame_kernel.py b/ecenet/edge_frame_kernel.py index 6083ce3..979d7e1 100644 --- a/ecenet/edge_frame_kernel.py +++ b/ecenet/edge_frame_kernel.py @@ -40,6 +40,8 @@ repeat; ``cos_valid``/``sin_valid`` zero the |m| > l (and m=0 sin) slots. """ +import os + import torch from torch.profiler import record_function @@ -391,6 +393,13 @@ def _e2n_fwd_kernel(gc_ptr, gs_ptr, d_ptr, perm_ptr, aptr_ptr, _EF_TABLES: dict = {} +# Warps per program for every edge-frame kernel launch, overridable per GPU +# without a code change (read once at import): ECENET_EF_WARPS=2 python ... +# The kernels' tiles are small ((BLOCK_R, 16)-ish), so fewer warps than the +# Triton default of 4 may win; sweep alongside profile_step. +_EF_WARPS = int(os.environ.get('ECENET_EF_WARPS', 4)) + + def _next_pow2(x: int) -> int: """Smallest power of two ≥ x, floored at 16. tl.arange REQUIRES a power of two (a multiple of 16 like 48 or 96 compiles to "arange's range must @@ -492,7 +501,7 @@ def _e2n_forward_triton(g_cos, g_sin, edge_dst, D_block, n_atoms, n_base): g_cos.contiguous(), g_sin.contiguous(), D_block.contiguous(), perm, aptr, srcoff, okc, oks, Delta, n_base, S, P, - SP=_next_pow2(S), RB=_next_pow2(n_base)) + SP=_next_pow2(S), RB=_next_pow2(n_base), num_warps=_EF_WARPS) return Delta # Per-edge: pack+unrotate is exactly the dx backward kernel's math @@ -508,7 +517,8 @@ def _e2n_forward_triton(g_cos, g_sin, edge_dst, D_block, n_atoms, n_base): g_cos.contiguous(), g_sin.contiguous(), D_block.contiguous(), cos_col, cos_ok, sin_col, sin_ok, h_global, n_base, n_base, S, P, - SP=_next_pow2(S), P16=_next_pow2(P), BLOCK_R=block_r) + SP=_next_pow2(S), P16=_next_pow2(P), BLOCK_R=block_r, + num_warps=_EF_WARPS) Delta = torch.zeros(n_atoms, n_base, S, dtype=g_cos.dtype, device=g_cos.device) Delta.index_add_(0, edge_dst, h_global) @@ -539,7 +549,8 @@ def _ef_forward_triton(A_emb, edge_i, edge_j, D_block, n_ch, n_ang): cos_col, cos_ok, sin_col, sin_ok, A_cos, A_sin, C, R, S, P, - SP=_next_pow2(S), P16=_next_pow2(P), BLOCK_R=block_r) + SP=_next_pow2(S), P16=_next_pow2(P), BLOCK_R=block_r, + num_warps=_EF_WARPS) return A_cos, A_sin @@ -558,7 +569,8 @@ def _ef_backward_triton(dA_cos, dA_sin, A_emb, edge_i, edge_j, D_block, A_emb_c = A_emb.contiguous() args = (cos_col, cos_ok, sin_col, sin_ok) block_r = _ef_block_r(R) - kw = dict(SP=_next_pow2(S), P16=_next_pow2(P), BLOCK_R=block_r) + kw = dict(SP=_next_pow2(S), P16=_next_pow2(P), BLOCK_R=block_r, + num_warps=_EF_WARPS) def _scatter(dA_both): # scatter back to atoms (torch: well-optimized, visible to compile) @@ -838,7 +850,8 @@ def backward(ctx, dDelta): g_cos.contiguous(), g_sin.contiguous(), cos_col, cos_ok, sin_col, sin_ok, dD, n_base, n_base, n_sph, P, - SP=_next_pow2(n_sph), P16=_next_pow2(P), BLOCK_R=block_r) + SP=_next_pow2(n_sph), P16=_next_pow2(P), BLOCK_R=block_r, + num_warps=_EF_WARPS) return dg_cos, dg_sin, None, dD, None, None, None, None, None # Eager (double-differentiable) path. @@ -893,7 +906,8 @@ def forward(ctx, m_cos, m_sin, D_block, D_block.contiguous(), cos_col, cos_ok, sin_col, sin_ok, h_global, n_base, n_base, S, P, - SP=_next_pow2(S), P16=_next_pow2(P), BLOCK_R=block_r) + SP=_next_pow2(S), P16=_next_pow2(P), BLOCK_R=block_r, + num_warps=_EF_WARPS) else: h = _pack_grads(m_cos, m_sin, cos_flat_idx, sin_flat_idx, cos_valid, sin_valid, @@ -919,7 +933,8 @@ def backward(ctx, dh): P = (n_ch // n_base) * n_ang block_r = _ef_block_r(n_base) tabs = _ef_tables(S, n_ang, dh.device) - kw = dict(SP=_next_pow2(S), P16=_next_pow2(P), BLOCK_R=block_r) + kw = dict(SP=_next_pow2(S), P16=_next_pow2(P), BLOCK_R=block_r, + num_warps=_EF_WARPS) # Merged path: dh (the dominant read) loaded once → dm AND dD. if need_dm and need_dd and block_r >= n_base: From acbd1367a38b2c79b2f9b5431f0d3b8eeb224174 Mon Sep 17 00:00:00 2001 From: alacour Date: Thu, 3 Sep 2026 19:33:51 -0700 Subject: [PATCH 2/6] edge_frame_kernel: packed-D merged-backward variants (ECENET_EF_PACKD) MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit ncu on the merged backward kernels (A100, 512-atom box, 44k edges) settles the question the flat warps sweep left open: L1/TEX throughput 75-86% with DRAM at 9-18% and compute ~30% — the kernels are L1-bound, stalling on MIO short-scoreboard, while DRAM idles. The amplification is the per-element gathered D loads through the (l,m) column tables and the matching dD scatter-stores: hundreds of non-vectorizable L1 transactions per program. The packed variants spend the idle DRAM instead: the wrapper pre-packs D into dense (E, S, P)/(E, P, S) tensors once per call (plain torch indexing, ~14 MB at 44k edges) and unpacks the packed dD afterwards, so the kernels do only vectorized coalesced tile IO. Math and masks identical. Gated off by default behind ECENET_EF_PACKD=1 (module flag, monkeypatchable) pending an A/B on the A100; ncu's stall-fix estimate is ~33-37% on these kernels. test_triton_packed_d reruns the full test_triton_paths comparison set (both merged backwards, vs fp64 eager truth) with the flag on. Co-Authored-By: Claude Fable 5 Claude-Session: https://claude.ai/code/session_019WyADHSVvKbQnYo1PAjGPK --- ecenet/edge_frame_kernel.py | 151 ++++++++++++++++++++++++++++++++ tests/test_edge_frame_kernel.py | 20 +++++ 2 files changed, 171 insertions(+) diff --git a/ecenet/edge_frame_kernel.py b/ecenet/edge_frame_kernel.py index 979d7e1..5967d54 100644 --- a/ecenet/edge_frame_kernel.py +++ b/ecenet/edge_frame_kernel.py @@ -348,6 +348,109 @@ def _pu_bwd_merged_kernel(dh_ptr, mc_ptr, ms_ptr, d_ptr, tl.store(dd_base + sin_col[None, :], acc_s, mask=k_ok[:, None] & (sin_ok[None, :] > 0)) + # ── Packed-D variants (ECENET_EF_PACKD=1) ──────────────────────────────── + # ncu on the merged backward kernels (A100, 44k edges): L1/TEX throughput + # 75-86%, DRAM 9-18%, compute ~30% — L1-bound, with MIO short-scoreboard + # stalls. The scalar per-element gathers of D through the column tables + # (and the matching dD scatter-stores) are the amplification: hundreds of + # non-vectorizable L1 transactions per program while DRAM idles. These + # variants trade that for DRAM headroom: the wrapper pre-packs D into + # dense (E, S, P) / (E, P, S) tensors once per call (plain torch indexing) + # and unpacks dD afterwards, so the kernel does only vectorized coalesced + # tile IO. Math and masks identical to the table-gather kernels. + + @triton.jit + def _ef_bwd_merged_packed_kernel(dc_ptr, ds_ptr, a_ptr, ei_ptr, ej_ptr, + dct_ptr, dst_ptr, cosk_ptr, sink_ptr, + dab_ptr, ddc_ptr, dds_ptr, + C, R, S, P, + SP: tl.constexpr, P16: tl.constexpr, + BLOCK_R: tl.constexpr): + e = tl.program_id(0) + offs_r = tl.arange(0, BLOCK_R) + offs_k = tl.arange(0, SP) + offs_p = tl.arange(0, P16) + row_ok = offs_r < R + p_ok = offs_p < P + k_ok = offs_k < S + + cos_ok = tl.load(cosk_ptr + offs_p, mask=p_ok, other=0) + sin_ok = tl.load(sink_ptr + offs_p, mask=p_ok, other=0) + + g_off = e * (R * P) + offs_r[:, None] * P + offs_p[None, :] + dC = tl.load(dc_ptr + g_off, + mask=row_ok[:, None] & (cos_ok[None, :] > 0), other=0.0) + dS = tl.load(ds_ptr + g_off, + mask=row_ok[:, None] & (sin_ok[None, :] > 0), other=0.0) + + # dA_both = dC @ DcT + dS @ DsT — DxT pre-packed (E, P, S): plain tiles + t_off = e * (P * S) + offs_p[:, None] * S + offs_k[None, :] + t_mask = p_ok[:, None] & k_ok[None, :] + DcT = tl.load(dct_ptr + t_off, mask=t_mask, other=0.0) + DsT = tl.load(dst_ptr + t_off, mask=t_mask, other=0.0) + dA = tl.dot(dC, DcT, input_precision="ieee") \ + + tl.dot(dS, DsT, input_precision="ieee") + dab_off = e * (R * S) + offs_r[:, None] * S + offs_k[None, :] + tl.store(dab_ptr + dab_off, dA, mask=row_ok[:, None] & k_ok[None, :]) + + # dD = Aᵀ @ grads, stored packed (E, S, P); wrapper unpacks to columns + ei = tl.load(ei_ptr + e) + ej = tl.load(ej_ptr + e) + atom = tl.where(offs_r < C, ei, ej) + ch = tl.where(offs_r < C, offs_r, offs_r - C) + a_ptrs = a_ptr + (atom * C + ch)[:, None] * S + offs_k[None, :] + A = tl.load(a_ptrs, mask=row_ok[:, None] & k_ok[None, :], other=0.0) + acc_c = tl.dot(tl.trans(A), dC, input_precision="ieee") # (SP, P16) + acc_s = tl.dot(tl.trans(A), dS, input_precision="ieee") + o_off = e * (S * P) + offs_k[:, None] * P + offs_p[None, :] + o_mask = k_ok[:, None] & p_ok[None, :] + tl.store(ddc_ptr + o_off, acc_c, mask=o_mask) + tl.store(dds_ptr + o_off, acc_s, mask=o_mask) + + @triton.jit + def _pu_bwd_merged_packed_kernel(dh_ptr, mc_ptr, ms_ptr, + dcp_ptr, dsp_ptr, cosk_ptr, sink_ptr, + dmc_ptr, dms_ptr, ddc_ptr, dds_ptr, + R, S, P, + SP: tl.constexpr, P16: tl.constexpr, + BLOCK_R: tl.constexpr): + e = tl.program_id(0) + offs_r = tl.arange(0, BLOCK_R) + offs_k = tl.arange(0, SP) + offs_p = tl.arange(0, P16) + row_ok = offs_r < R + p_ok = offs_p < P + k_ok = offs_k < S + + cos_ok = tl.load(cosk_ptr + offs_p, mask=p_ok, other=0) + sin_ok = tl.load(sink_ptr + offs_p, mask=p_ok, other=0) + + dh_off = e * (R * S) + offs_r[:, None] * S + offs_k[None, :] + dh = tl.load(dh_ptr + dh_off, + mask=row_ok[:, None] & k_ok[None, :], other=0.0) + + # dm = dh @ D-cols — Dx pre-packed (E, S, P): plain tiles + d_off = e * (S * P) + offs_k[:, None] * P + offs_p[None, :] + d_mask = k_ok[:, None] & p_ok[None, :] + Dc = tl.load(dcp_ptr + d_off, mask=d_mask, other=0.0) + Ds = tl.load(dsp_ptr + d_off, mask=d_mask, other=0.0) + dmc = tl.dot(dh, Dc, input_precision="ieee") # (BLOCK_R, P16) + dms = tl.dot(dh, Ds, input_precision="ieee") + out_off = e * (R * P) + offs_r[:, None] * P + offs_p[None, :] + st_mask = row_ok[:, None] & p_ok[None, :] + tl.store(dmc_ptr + out_off, dmc, mask=st_mask) + tl.store(dms_ptr + out_off, dms, mask=st_mask) + + # dD = dhᵀ @ h, stored packed (E, S, P); wrapper unpacks to columns + mC = tl.load(mc_ptr + out_off, + mask=row_ok[:, None] & (cos_ok[None, :] > 0), other=0.0) + mS = tl.load(ms_ptr + out_off, + mask=row_ok[:, None] & (sin_ok[None, :] > 0), other=0.0) + acc_c = tl.dot(tl.trans(dh), mC, input_precision="ieee") # (SP, P16) + acc_s = tl.dot(tl.trans(dh), mS, input_precision="ieee") + tl.store(ddc_ptr + d_off, acc_c, mask=d_mask) + tl.store(dds_ptr + d_off, acc_s, mask=d_mask) + @triton.jit def _e2n_fwd_kernel(gc_ptr, gs_ptr, d_ptr, perm_ptr, aptr_ptr, srcoff_ptr, okc_ptr, oks_ptr, @@ -399,6 +502,31 @@ def _e2n_fwd_kernel(gc_ptr, gs_ptr, d_ptr, perm_ptr, aptr_ptr, # Triton default of 4 may win; sweep alongside profile_step. _EF_WARPS = int(os.environ.get('ECENET_EF_WARPS', 4)) +# Packed-D variants of the merged backward kernels (see the kernel comment +# block): trades the scalar table-gathered D loads / dD scatter-stores +# (L1-bound per ncu) for dense pre-packed tensors and vectorized tile IO. +# Off by default until benchmarked; module-level so tests can toggle it. +_EF_PACKD = os.environ.get('ECENET_EF_PACKD', '0') == '1' + + +def _pack_D(D_block, cos_col, cos_ok, sin_col, sin_ok): + """Dense per-edge packed D: Dx[e, k, p] = D[e, k, col(p)]·ok(p), (E, S, P).""" + Dc = D_block[:, :, cos_col.long()] * cos_ok.to(D_block.dtype) + Ds = D_block[:, :, sin_col.long()] * sin_ok.to(D_block.dtype) + return Dc.contiguous(), Ds.contiguous() + + +def _unpack_dD(ddc, dds, cos_col, cos_ok, sin_col, sin_ok, S): + """Inverse of _pack_D for the gradient: scatter packed (E, S, P) columns + back into (E, S, S). cos and sin column sets are disjoint; slots no p + touches stay 0 (matches the zeros-init of the table-scatter path).""" + dD = ddc.new_zeros(ddc.shape[0], S, S) + vc = cos_ok.bool() + vs = sin_ok.bool() + dD[:, :, cos_col[vc].long()] = ddc[:, :, vc] + dD[:, :, sin_col[vs].long()] = dds[:, :, vs] + return dD + def _next_pow2(x: int) -> int: """Smallest power of two ≥ x, floored at 16. tl.arange REQUIRES a power @@ -587,6 +715,18 @@ def _scatter(dA_both): # to cover the edge (block_r ≥ R; true for n_base/2C ≤ 128). if need_dx and need_dd and block_r >= R: dA_both = torch.empty(E, R, S, dtype=dA_cos.dtype, device=dA_cos.device) + if _EF_PACKD: + Dc, Ds = _pack_D(D_block, *args) + DcT = Dc.transpose(1, 2).contiguous() + DsT = Ds.transpose(1, 2).contiguous() + ddc = torch.empty(E, S, P, dtype=dA_cos.dtype, device=dA_cos.device) + dds = torch.empty_like(ddc) + _ef_bwd_merged_packed_kernel[(E,)]( + dA_cos, dA_sin, A_emb_c, edge_i.contiguous(), + edge_j.contiguous(), DcT, DsT, args[1], args[3], + dA_both, ddc, dds, + C, R, S, P, **kw) + return _scatter(dA_both), _unpack_dD(ddc, dds, *args, S) dD = torch.zeros_like(D_block) # columns no p touches stay 0 _ef_bwd_merged_kernel[(E,)]( dA_cos, dA_sin, A_emb_c, edge_i.contiguous(), edge_j.contiguous(), @@ -941,6 +1081,17 @@ def backward(ctx, dh): dm_cos = torch.empty(E, n_ch, n_ang, dtype=dh.dtype, device=dh.device) dm_sin = torch.empty_like(dm_cos) + if _EF_PACKD: + Dc, Ds = _pack_D(D_block, *tabs) + ddc = torch.empty(E, S, P, dtype=dh.dtype, device=dh.device) + dds = torch.empty_like(ddc) + _pu_bwd_merged_packed_kernel[(E,)]( + dh_c, m_cos.contiguous(), m_sin.contiguous(), + Dc, Ds, tabs[1], tabs[3], + dm_cos, dm_sin, ddc, dds, + n_base, S, P, **kw) + return (dm_cos, dm_sin, _unpack_dD(ddc, dds, *tabs, S), + None, None, None, None) dD = torch.zeros_like(D_block) _pu_bwd_merged_kernel[(E,)]( dh_c, m_cos.contiguous(), m_sin.contiguous(), diff --git a/tests/test_edge_frame_kernel.py b/tests/test_edge_frame_kernel.py index b9a8b5a..d632c8f 100644 --- a/tests/test_edge_frame_kernel.py +++ b/tests/test_edge_frame_kernel.py @@ -441,6 +441,25 @@ def test_triton_paths(): print("test_triton_paths[pack_unrotate]: OK") +def test_triton_packed_d(): + """CUDA-only: the ECENET_EF_PACKD merged-backward variants (dense packed-D + tile IO instead of per-element table gathers) pass the same fp64 + comparisons as the default kernels — test_triton_paths rerun with the + module flag on covers both EdgeFrameFused and PackUnrotateFused merged + backwards.""" + if not torch.cuda.is_available(): + print("test_triton_packed_d: SKIP (no CUDA)") + return + import ecenet.edge_frame_kernel as efk + old = efk._EF_PACKD + efk._EF_PACKD = True + try: + test_triton_paths() + finally: + efk._EF_PACKD = old + print("test_triton_packed_d: OK (packed-D merged backward)") + + def test_model_integration(): """Full ECENet (no MP): energy and autograd forces identical flag on/off.""" import ecenet as _ecenet @@ -524,6 +543,7 @@ def run(): test_pack_unrotate_matches_mp_ops() test_pack_unrotate_gradchecks() test_triton_paths() + test_triton_packed_d() test_model_integration() test_mp_integration() print("\nAll tests passed.") From 31c2ff3b71fd301c4856c492bc8607eecc030d60 Mon Sep 17 00:00:00 2001 From: alacour Date: Thu, 3 Sep 2026 19:54:48 -0700 Subject: [PATCH 3/6] edge_frame_kernel: amortize D packing once per step via the shared D_block MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit The packed kernels won at the kernel level (ncu: _ef_bwd_merged 2.94 -> 1.68 ms, _pu_bwd_merged 2.49 -> 1.95 ms) but the joint step stayed flat: per-call packing/transposing in the backward wrappers (~8 small torch ops x 7 backward calls/step) ate the entire ~5 ms win. The model builds ONE D_block per step and hands the same Python object to every fused edge-frame op, so _get_packed_D now caches the packed tensors as an attribute on that object — packing runs once per step. Function forwards fetch it and carry the tensors through save_for_backward (attributes do not survive the re-wrapping); the backward wrappers receive them instead of packing. A fresh step's fresh D_block starts clean, so no cross-step staleness. Still gated behind ECENET_EF_PACKD=1 pending the A/B. Co-Authored-By: Claude Fable 5 Claude-Session: https://claude.ai/code/session_019WyADHSVvKbQnYo1PAjGPK --- ecenet/edge_frame_kernel.py | 66 ++++++++++++++++++++++++++++++------- 1 file changed, 54 insertions(+), 12 deletions(-) diff --git a/ecenet/edge_frame_kernel.py b/ecenet/edge_frame_kernel.py index 5967d54..91aee5d 100644 --- a/ecenet/edge_frame_kernel.py +++ b/ecenet/edge_frame_kernel.py @@ -528,6 +528,30 @@ def _unpack_dD(ddc, dds, cos_col, cos_ok, sin_col, sin_ok, S): return dD +def _get_packed_D(D_block, n_ang): + """Once-per-step packed D — (Dc, Ds, DcT, DsT), each (E, S, P)/(E, P, S). + + Cached on the D_block PYTHON OBJECT: the model builds D_block once per + step and hands the same object to every fused edge-frame op (4x + EdgeFrameFused + 3x PackUnrotateFused at n_mp=4), so the packing cost is + paid once, not per backward — per-call packing was measured to eat the + packed kernels' entire ~5 ms/step win. A new step's fresh D_block starts + with no attribute, so there is no cross-step staleness. The attribute + does NOT survive save_for_backward's re-wrapping, so Functions must carry + the packed tensors through ctx themselves.""" + packed = getattr(D_block, '_ecenet_packed', None) + if packed is None or packed[0].shape[-1] != _ef_tables( + D_block.shape[-1], n_ang, D_block.device)[0].shape[0]: + with torch.no_grad(): + tabs = _ef_tables(D_block.shape[-1], n_ang, D_block.device) + Dc, Ds = _pack_D(D_block, *tabs) + packed = (Dc, Ds, + Dc.transpose(1, 2).contiguous(), + Ds.transpose(1, 2).contiguous()) + D_block._ecenet_packed = packed + return packed + + def _next_pow2(x: int) -> int: """Smallest power of two ≥ x, floored at 16. tl.arange REQUIRES a power of two (a multiple of 16 like 48 or 96 compiles to "arange's range must @@ -683,7 +707,7 @@ def _ef_forward_triton(A_emb, edge_i, edge_j, D_block, n_ch, n_ang): def _ef_backward_triton(dA_cos, dA_sin, A_emb, edge_i, edge_j, D_block, - n_ang, single, need_dx, need_dd): + n_ang, single, need_dx, need_dd, packedT=None): E = edge_i.shape[0] C = A_emb.shape[1] R = C if single else 2 * C @@ -715,10 +739,8 @@ def _scatter(dA_both): # to cover the edge (block_r ≥ R; true for n_base/2C ≤ 128). if need_dx and need_dd and block_r >= R: dA_both = torch.empty(E, R, S, dtype=dA_cos.dtype, device=dA_cos.device) - if _EF_PACKD: - Dc, Ds = _pack_D(D_block, *args) - DcT = Dc.transpose(1, 2).contiguous() - DsT = Ds.transpose(1, 2).contiguous() + if packedT is not None: + DcT, DsT = packedT ddc = torch.empty(E, S, P, dtype=dA_cos.dtype, device=dA_cos.device) dds = torch.empty_like(ddc) _ef_bwd_merged_packed_kernel[(E,)]( @@ -778,17 +800,28 @@ def forward(ctx, A_emb, edge_i, edge_j, D_block, .view(-1, n_ch, n_ang)) * cos_valid A_sin = (A_flat.index_select(1, sin_flat_idx) .view(-1, n_ch, n_ang)) * sin_valid + # Packed-D path: fetch the per-step packed D (amortized on the shared + # D_block object) in the FORWARD and carry it through ctx — attributes + # do not survive save_for_backward's re-wrapping. + extra = () + ctx.packed = _EF_PACKD and _ef_triton_ok(A_emb, edge_i.shape[0]) + if ctx.packed: + _, _, DcT, DsT = _get_packed_D(D_block, n_ang) + extra = (DcT, DsT) ctx.save_for_backward(A_emb, edge_i, edge_i if single else edge_j, D_block, - cos_flat_idx, sin_flat_idx, cos_valid, sin_valid) + cos_flat_idx, sin_flat_idx, cos_valid, sin_valid, + *extra) ctx.n_ang = n_ang ctx.single = single return A_cos, A_sin @staticmethod def backward(ctx, dA_cos, dA_sin): + saved = ctx.saved_tensors (A_emb, edge_i, edge_j, D_block, - cos_flat_idx, sin_flat_idx, cos_valid, sin_valid) = ctx.saved_tensors + cos_flat_idx, sin_flat_idx, cos_valid, sin_valid) = saved[:8] + packedT = saved[8:10] if ctx.packed else None E = dA_cos.shape[0] C = A_emb.shape[1] single = ctx.single @@ -803,7 +836,8 @@ def backward(ctx, dA_cos, dA_sin): dA_cos, dA_sin, A_emb, edge_i, edge_j, D_block, ctx.n_ang, single=single, need_dx=ctx.needs_input_grad[0], - need_dd=ctx.needs_input_grad[3]) + need_dd=ctx.needs_input_grad[3], + packedT=packedT) return dA_emb, None, None, dD, None, None, None, None # Adjoint of select: scatter masked grads into the rotated layout. @@ -1053,14 +1087,22 @@ def forward(ctx, m_cos, m_sin, D_block, cos_valid, sin_valid, n_base * S).view(E, n_base, S) h_global = torch.bmm(h, D_block.transpose(-1, -2)) + extra = () + ctx.packed = _EF_PACKD and _ef_triton_ok(m_cos, E) + if ctx.packed: + Dc, Ds, _, _ = _get_packed_D(D_block, n_ang) + extra = (Dc, Ds) ctx.save_for_backward(m_cos, m_sin, D_block, - cos_flat_idx, sin_flat_idx, cos_valid, sin_valid) + cos_flat_idx, sin_flat_idx, cos_valid, sin_valid, + *extra) return h_global @staticmethod def backward(ctx, dh): + saved = ctx.saved_tensors (m_cos, m_sin, D_block, - cos_flat_idx, sin_flat_idx, cos_valid, sin_valid) = ctx.saved_tensors + cos_flat_idx, sin_flat_idx, cos_valid, sin_valid) = saved[:7] + packed = saved[7:9] if ctx.packed else None E = m_cos.shape[0] n_ch, n_ang = cos_valid.shape S = D_block.shape[-1] @@ -1081,8 +1123,8 @@ def backward(ctx, dh): dm_cos = torch.empty(E, n_ch, n_ang, dtype=dh.dtype, device=dh.device) dm_sin = torch.empty_like(dm_cos) - if _EF_PACKD: - Dc, Ds = _pack_D(D_block, *tabs) + if packed is not None: + Dc, Ds = packed ddc = torch.empty(E, S, P, dtype=dh.dtype, device=dh.device) dds = torch.empty_like(ddc) _pu_bwd_merged_packed_kernel[(E,)]( From d06578017c7ed8c9b6e454bbb0eea0abec4082bd Mon Sep 17 00:00:00 2001 From: alacour Date: Thu, 3 Sep 2026 20:02:33 -0700 Subject: [PATCH 4/6] =?UTF-8?q?edge=5Fframe=5Fkernel:=20sync-free=20dD=20u?= =?UTF-8?q?npack=20=E2=80=94=20cached=20integer=20indices?= MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit torch.profiler with the packed path on showed the GPU win arriving (packed kernels 1.35/1.65 ms per call, CUDA total down) while wall time stayed flat, and named the thief: EdgeFrameFusedBackward at 10.6 ms CPU per call, with aten::index at 44% of CPU total. _unpack_dD's boolean-mask indexing (cos_col[vc], dD[:, :, col[vc]]) forces a host-device sync on every backward call — seven pipeline stalls per step. The masks are static, so the unpack now uses integer index tensors precomputed once per (S, n_ang, device); every op in the unpack is async. Co-Authored-By: Claude Fable 5 Claude-Session: https://claude.ai/code/session_019WyADHSVvKbQnYo1PAjGPK --- ecenet/edge_frame_kernel.py | 34 ++++++++++++++++++++++++++-------- 1 file changed, 26 insertions(+), 8 deletions(-) diff --git a/ecenet/edge_frame_kernel.py b/ecenet/edge_frame_kernel.py index 91aee5d..24791a7 100644 --- a/ecenet/edge_frame_kernel.py +++ b/ecenet/edge_frame_kernel.py @@ -516,15 +516,33 @@ def _pack_D(D_block, cos_col, cos_ok, sin_col, sin_ok): return Dc.contiguous(), Ds.contiguous() -def _unpack_dD(ddc, dds, cos_col, cos_ok, sin_col, sin_ok, S): +_UNPACK_IDX: dict = {} + + +def _unpack_indices(S, n_ang, device): + """Static integer index pairs for _unpack_dD, cached per (S, n_ang, + device). MUST be integer tensors: boolean-mask indexing at call time + forces a host-device sync per backward (measured ~10 ms of CPU stall per + EdgeFrameFusedBackward — it erased the packed kernels' entire win).""" + key = (S, n_ang, str(device)) + if key not in _UNPACK_IDX: + cos_col, cos_ok, sin_col, sin_ok = _ef_tables(S, n_ang, device) + vc = cos_ok.bool() + vs = sin_ok.bool() + _UNPACK_IDX[key] = (vc.nonzero().flatten(), cos_col[vc].long(), + vs.nonzero().flatten(), sin_col[vs].long()) + return _UNPACK_IDX[key] + + +def _unpack_dD(ddc, dds, S, n_ang): """Inverse of _pack_D for the gradient: scatter packed (E, S, P) columns back into (E, S, S). cos and sin column sets are disjoint; slots no p - touches stay 0 (matches the zeros-init of the table-scatter path).""" + touches stay 0 (matches the zeros-init of the table-scatter path). + Integer-index ops only — fully async, no host sync.""" + pc, cc, ps, sc = _unpack_indices(S, n_ang, ddc.device) dD = ddc.new_zeros(ddc.shape[0], S, S) - vc = cos_ok.bool() - vs = sin_ok.bool() - dD[:, :, cos_col[vc].long()] = ddc[:, :, vc] - dD[:, :, sin_col[vs].long()] = dds[:, :, vs] + dD[:, :, cc] = ddc[:, :, pc] + dD[:, :, sc] = dds[:, :, ps] return dD @@ -748,7 +766,7 @@ def _scatter(dA_both): edge_j.contiguous(), DcT, DsT, args[1], args[3], dA_both, ddc, dds, C, R, S, P, **kw) - return _scatter(dA_both), _unpack_dD(ddc, dds, *args, S) + return _scatter(dA_both), _unpack_dD(ddc, dds, S, n_ang) dD = torch.zeros_like(D_block) # columns no p touches stay 0 _ef_bwd_merged_kernel[(E,)]( dA_cos, dA_sin, A_emb_c, edge_i.contiguous(), edge_j.contiguous(), @@ -1132,7 +1150,7 @@ def backward(ctx, dh): Dc, Ds, tabs[1], tabs[3], dm_cos, dm_sin, ddc, dds, n_base, S, P, **kw) - return (dm_cos, dm_sin, _unpack_dD(ddc, dds, *tabs, S), + return (dm_cos, dm_sin, _unpack_dD(ddc, dds, S, n_ang), None, None, None, None) dD = torch.zeros_like(D_block) _pu_bwd_merged_kernel[(E,)]( From 8c61f76d8b321e95450adac1d56ad5cb1473225e Mon Sep 17 00:00:00 2001 From: alacour Date: Thu, 3 Sep 2026 20:06:50 -0700 Subject: [PATCH 5/6] edge_frame_kernel: packed-D forward kernels; flip ECENET_EF_PACKD default on With the sync-free unpack the packed backward finally reached the wall clock: joint force step 111.4 -> 107.8 ms on the A100 512-atom box (EdgeFrameFusedBackward CPU 10.6 ms -> 0.29 ms per call, aten::index off the CPU hot list, packed kernels at 1.35/1.65 ms in-app). This extends the same treatment to the forward side, which carries the identical table-gather pattern: _ef_fwd_packed_kernel (EdgeFrameFused forward, ~1.04 ms x 4/step) and _ef_bwd_dx_packed_kernel (PackUnrotate's forward contraction, ~1.17 ms x 3/step) read the pre-packed Dc/Ds / DcT/DsT tiles instead of gathering through the column tables; amortization is free since the packed tensors already live on the shared per-step D_block. Default flipped ON (ECENET_EF_PACKD=0 restores the table-gather kernels); test_triton_packed_d now runs the full fp64 comparison suite under both settings. Co-Authored-By: Claude Fable 5 Claude-Session: https://claude.ai/code/session_019WyADHSVvKbQnYo1PAjGPK --- ecenet/edge_frame_kernel.py | 121 ++++++++++++++++++++++++++++---- tests/test_edge_frame_kernel.py | 15 ++-- 2 files changed, 113 insertions(+), 23 deletions(-) diff --git a/ecenet/edge_frame_kernel.py b/ecenet/edge_frame_kernel.py index 24791a7..0b80432 100644 --- a/ecenet/edge_frame_kernel.py +++ b/ecenet/edge_frame_kernel.py @@ -451,6 +451,74 @@ def _pu_bwd_merged_packed_kernel(dh_ptr, mc_ptr, ms_ptr, tl.store(ddc_ptr + d_off, acc_c, mask=d_mask) tl.store(dds_ptr + d_off, acc_s, mask=d_mask) + @triton.jit + def _ef_fwd_packed_kernel(a_ptr, ei_ptr, ej_ptr, dcp_ptr, dsp_ptr, + outc_ptr, outs_ptr, + C, R, S, P, + SP: tl.constexpr, P16: tl.constexpr, + BLOCK_R: tl.constexpr): + # _ef_fwd_kernel with the column-gathered D loads replaced by tiles of + # the pre-packed (E, S, P) Dc/Ds (ok-zeros baked in by _pack_D). + e = tl.program_id(0) + offs_r = tl.program_id(1) * BLOCK_R + tl.arange(0, BLOCK_R) + offs_k = tl.arange(0, SP) + offs_p = tl.arange(0, P16) + row_ok = offs_r < R + + ei = tl.load(ei_ptr + e) + ej = tl.load(ej_ptr + e) + atom = tl.where(offs_r < C, ei, ej) + ch = tl.where(offs_r < C, offs_r, offs_r - C) + a_ptrs = a_ptr + (atom * C + ch)[:, None] * S + offs_k[None, :] + A = tl.load(a_ptrs, mask=row_ok[:, None] & (offs_k[None, :] < S), other=0.0) + + d_off = e * (S * P) + offs_k[:, None] * P + offs_p[None, :] + d_mask = (offs_k[:, None] < S) & (offs_p[None, :] < P) + Dc = tl.load(dcp_ptr + d_off, mask=d_mask, other=0.0) + Ds = tl.load(dsp_ptr + d_off, mask=d_mask, other=0.0) + + OC = tl.dot(A, Dc, input_precision="ieee") # (BLOCK_R, P16) + OS = tl.dot(A, Ds, input_precision="ieee") + out_off = e * (R * P) + offs_r[:, None] * P + offs_p[None, :] + st_mask = row_ok[:, None] & (offs_p[None, :] < P) + tl.store(outc_ptr + out_off, OC, mask=st_mask) + tl.store(outs_ptr + out_off, OS, mask=st_mask) + + @triton.jit + def _ef_bwd_dx_packed_kernel(dc_ptr, ds_ptr, dct_ptr, dst_ptr, + cosk_ptr, sink_ptr, dab_ptr, + C, R, S, P, + SP: tl.constexpr, P16: tl.constexpr, + BLOCK_R: tl.constexpr): + # _ef_bwd_dx_kernel on the pre-packed (E, P, S) transposed D. The + # ok-masks on the grad loads (the eager path's ·valid) stay. + e = tl.program_id(0) + offs_r = tl.program_id(1) * BLOCK_R + tl.arange(0, BLOCK_R) + offs_k = tl.arange(0, SP) + offs_p = tl.arange(0, P16) + row_ok = offs_r < R + p_ok = offs_p < P + k_ok = offs_k < S + + cos_ok = tl.load(cosk_ptr + offs_p, mask=p_ok, other=0) + sin_ok = tl.load(sink_ptr + offs_p, mask=p_ok, other=0) + + g_off = e * (R * P) + offs_r[:, None] * P + offs_p[None, :] + dC = tl.load(dc_ptr + g_off, + mask=row_ok[:, None] & (cos_ok[None, :] > 0), other=0.0) + dS = tl.load(ds_ptr + g_off, + mask=row_ok[:, None] & (sin_ok[None, :] > 0), other=0.0) + + t_off = e * (P * S) + offs_p[:, None] * S + offs_k[None, :] + t_mask = p_ok[:, None] & k_ok[None, :] + DcT = tl.load(dct_ptr + t_off, mask=t_mask, other=0.0) + DsT = tl.load(dst_ptr + t_off, mask=t_mask, other=0.0) + + dA = tl.dot(dC, DcT, input_precision="ieee") \ + + tl.dot(dS, DsT, input_precision="ieee") # (BLOCK_R, SP) + dab_off = e * (R * S) + offs_r[:, None] * S + offs_k[None, :] + tl.store(dab_ptr + dab_off, dA, mask=row_ok[:, None] & k_ok[None, :]) + @triton.jit def _e2n_fwd_kernel(gc_ptr, gs_ptr, d_ptr, perm_ptr, aptr_ptr, srcoff_ptr, okc_ptr, oks_ptr, @@ -502,11 +570,14 @@ def _e2n_fwd_kernel(gc_ptr, gs_ptr, d_ptr, perm_ptr, aptr_ptr, # Triton default of 4 may win; sweep alongside profile_step. _EF_WARPS = int(os.environ.get('ECENET_EF_WARPS', 4)) -# Packed-D variants of the merged backward kernels (see the kernel comment -# block): trades the scalar table-gathered D loads / dD scatter-stores -# (L1-bound per ncu) for dense pre-packed tensors and vectorized tile IO. -# Off by default until benchmarked; module-level so tests can toggle it. -_EF_PACKD = os.environ.get('ECENET_EF_PACKD', '0') == '1' +# Packed-D kernel variants (see the kernel comment block): trade the scalar +# table-gathered D loads / dD scatter-stores (L1-bound per ncu, 75-86% L1 +# with DRAM at 9-18%) for dense pre-packed tensors and vectorized tile IO, +# packed once per step on the shared D_block object. Measured on the A100 +# (512-atom box): merged backwards 2.33→1.35 / 1.98→1.65 ms per call, joint +# force step 111.4→107.8 ms. Default ON; ECENET_EF_PACKD=0 restores the +# table-gather kernels. Module-level so tests can toggle it. +_EF_PACKD = os.environ.get('ECENET_EF_PACKD', '1') == '1' def _pack_D(D_block, cos_col, cos_ok, sin_col, sin_ok): @@ -713,6 +784,15 @@ def _ef_forward_triton(A_emb, edge_i, edge_j, D_block, n_ch, n_ang): A_sin = torch.empty_like(A_cos) block_r = _ef_block_r(R) grid = (E, triton.cdiv(R, block_r)) + if _EF_PACKD: + Dc, Ds, _, _ = _get_packed_D(D_block, n_ang) + _ef_fwd_packed_kernel[grid]( + A_emb.contiguous(), edge_i.contiguous(), ej.contiguous(), + Dc, Ds, A_cos, A_sin, + C, R, S, P, + SP=_next_pow2(S), P16=_next_pow2(P), BLOCK_R=block_r, + num_warps=_EF_WARPS) + return A_cos, A_sin _ef_fwd_kernel[grid]( A_emb.contiguous(), edge_i.contiguous(), ej.contiguous(), D_block.contiguous(), @@ -1090,16 +1170,27 @@ def forward(ctx, m_cos, m_sin, D_block, device=m_cos.device) block_r = _ef_block_r(n_base) grid = (E, triton.cdiv(n_base, block_r)) - _ef_bwd_dx_kernel[grid]( - m_cos.contiguous(), m_sin.contiguous(), - # forward wants h @ Dᵀ where the dx kernel computes - # Σ_p g[·,p]·D[k, col(p)] — i.e. it contracts against D's - # COLUMNS, which is exactly the transpose we need. - D_block.contiguous(), - cos_col, cos_ok, sin_col, sin_ok, h_global, - n_base, n_base, S, P, - SP=_next_pow2(S), P16=_next_pow2(P), BLOCK_R=block_r, - num_warps=_EF_WARPS) + if _EF_PACKD: + # forward wants h @ Dᵀ; the packed dx kernel contracts + # against DcT/DsT (E, P, S), exactly that transpose. + _, _, DcT, DsT = _get_packed_D(D_block, n_ang) + _ef_bwd_dx_packed_kernel[grid]( + m_cos.contiguous(), m_sin.contiguous(), + DcT, DsT, cos_ok, sin_ok, h_global, + n_base, n_base, S, P, + SP=_next_pow2(S), P16=_next_pow2(P), BLOCK_R=block_r, + num_warps=_EF_WARPS) + else: + _ef_bwd_dx_kernel[grid]( + m_cos.contiguous(), m_sin.contiguous(), + # forward wants h @ Dᵀ where the dx kernel computes + # Σ_p g[·,p]·D[k, col(p)] — i.e. it contracts against + # D's COLUMNS, which is exactly the transpose we need. + D_block.contiguous(), + cos_col, cos_ok, sin_col, sin_ok, h_global, + n_base, n_base, S, P, + SP=_next_pow2(S), P16=_next_pow2(P), BLOCK_R=block_r, + num_warps=_EF_WARPS) else: h = _pack_grads(m_cos, m_sin, cos_flat_idx, sin_flat_idx, cos_valid, sin_valid, diff --git a/tests/test_edge_frame_kernel.py b/tests/test_edge_frame_kernel.py index d632c8f..c06da6e 100644 --- a/tests/test_edge_frame_kernel.py +++ b/tests/test_edge_frame_kernel.py @@ -442,22 +442,21 @@ def test_triton_paths(): def test_triton_packed_d(): - """CUDA-only: the ECENET_EF_PACKD merged-backward variants (dense packed-D - tile IO instead of per-element table gathers) pass the same fp64 - comparisons as the default kernels — test_triton_paths rerun with the - module flag on covers both EdgeFrameFused and PackUnrotateFused merged - backwards.""" + """CUDA-only: run test_triton_paths under BOTH _EF_PACKD settings, so the + packed-D kernels (the default) and the table-gather kernels (the + ECENET_EF_PACKD=0 fallback) both stay pinned to the fp64 reference.""" if not torch.cuda.is_available(): print("test_triton_packed_d: SKIP (no CUDA)") return import ecenet.edge_frame_kernel as efk old = efk._EF_PACKD - efk._EF_PACKD = True try: - test_triton_paths() + for packed in (True, False): + efk._EF_PACKD = packed + test_triton_paths() + print(f"test_triton_packed_d[packed={packed}]: OK") finally: efk._EF_PACKD = old - print("test_triton_packed_d: OK (packed-D merged backward)") def test_model_integration(): From 8bff38897f0b40643a067cf671388f4ee7b29ed3 Mon Sep 17 00:00:00 2001 From: alacour Date: Thu, 3 Sep 2026 20:24:31 -0700 Subject: [PATCH 6/6] edge_frame_kernel: refresh docstrings for the packed-D default Co-Authored-By: Claude Fable 5 Claude-Session: https://claude.ai/code/session_019WyADHSVvKbQnYo1PAjGPK --- ecenet/edge_frame_kernel.py | 11 +++++++---- 1 file changed, 7 insertions(+), 4 deletions(-) diff --git a/ecenet/edge_frame_kernel.py b/ecenet/edge_frame_kernel.py index 0b80432..f935bc6 100644 --- a/ecenet/edge_frame_kernel.py +++ b/ecenet/edge_frame_kernel.py @@ -873,7 +873,8 @@ def _scatter(dA_both): class EdgeFrameFused(torch.autograd.Function): """Fused gather → rotate → select with analytic, double-differentiable - backward. Saves (A_emb, D_block, indices) — NOT the (E, R, n_sph) + backward. Saves (A_emb, D_block, indices; plus the per-step packed D + under _EF_PACKD) — NOT the (E, R, n_sph) intermediates; the gathered rows are re-gathered in the backward. edge_j=None → single-source mode (MP steps 5-6): rows are A_emb[edge_i]'s @@ -1151,9 +1152,11 @@ class PackUnrotateFused(torch.autograd.Function): No gather, no scatter, no gate — the attention/message weighting and the node accumulation stay eager in the node frame exactly as the unfused path. The win is dropping the packed intermediate h from HBM (forward) and - from the bmm's saved set (backward). All three kernels are the existing - ones: forward = _ef_bwd_dx (its "grads" input is any per-edge (E, R, P) - tensor), backward dm = _ef_fwd with identity indices, dD = _ef_bwd_dd.""" + from the bmm's saved set (backward). The kernels are shared with the edge + frame: forward = _ef_bwd_dx (its "grads" input is any per-edge (E, R, P) + tensor) — or its packed-D variant under _EF_PACKD (default) — backward = + the merged _pu kernel (packed or table-gather), with dm = _ef_fwd with + identity indices / dD = _ef_bwd_dd as the non-merged fallbacks.""" @staticmethod def forward(ctx, m_cos, m_sin, D_block,