atomic-ops 0.1.0__py3-none-any.whl

This diff represents the content of publicly available package versions that have been released to one of the supported registries. The information contained in this diff is provided for informational purposes only and reflects changes between package versions as they appear in their respective public registries.
atomic_ops/__init__.py ADDED
@@ -0,0 +1,34 @@
1
+ """
2
+ Atomic Ops — Fused Gated DeltaNet-2 kernels for TPU v5e (Pallas/JAX).
3
+ Ported from NVlabs DeltaNet Triton kernels.
4
+ """
5
+ from importlib.metadata import version as _version, PackageNotFoundError as _PkgNotFound
6
+
7
+ from .configs import KernelConfig, KAGGLE_SMALL, KAGGLE_MEDIUM, KAGGLE_LARGE, DEFAULT_CONFIG
8
+ from .utils import is_tpu_available, estimate_memory, get_recommended_config
9
+ from .fallback import gdn2_forward, gdn2_forward_trainable
10
+ from .gdn2_fwd import gdn2_pallas_forward
11
+ from .gdn2_pipeline import gdn2_pallas_forward_trainable
12
+ from .reference import gdn2_chunked_wy_reference, gdn2_token_serial_reference
13
+
14
+ try:
15
+ __version__ = _version("atomic_ops")
16
+ except _PkgNotFound:
17
+ __version__ = "0.0.0.dev0"
18
+
19
+ __all__ = [
20
+ "KernelConfig",
21
+ "KAGGLE_SMALL",
22
+ "KAGGLE_MEDIUM",
23
+ "KAGGLE_LARGE",
24
+ "DEFAULT_CONFIG",
25
+ "is_tpu_available",
26
+ "estimate_memory",
27
+ "get_recommended_config",
28
+ "gdn2_forward",
29
+ "gdn2_forward_trainable",
30
+ "gdn2_pallas_forward",
31
+ "gdn2_pallas_forward_trainable",
32
+ "gdn2_chunked_wy_reference",
33
+ "gdn2_token_serial_reference",
34
+ ]
atomic_ops/configs.py ADDED
@@ -0,0 +1,136 @@
1
+ """
2
+ Central configuration, numerical-safety helpers, and shape validation
3
+ shared across the forward/backward Pallas kernels.
4
+ """
5
+ from __future__ import annotations
6
+
7
+ import dataclasses as dc
8
+ import os
9
+
10
+ import jax
11
+ import jax.numpy as jnp
12
+
13
+ _HIGHEST = jax.lax.Precision.HIGHEST
14
+
15
+
16
+ @dc.dataclass(frozen=True)
17
+ class KernelConfig:
18
+ bt: int = 256
19
+ bc: int = 128
20
+ mb: int = 16
21
+ clip: float = 1e4
22
+ wy_eps: float = 0.0
23
+
24
+ @property
25
+ def n_sub(self) -> int:
26
+ return self.bt // self.bc
27
+
28
+ @property
29
+ def n_micro(self) -> int:
30
+ return self.bc // self.mb
31
+
32
+ def __post_init__(self):
33
+ if self.bt % self.bc != 0:
34
+ raise ValueError(f"bt={self.bt} must be divisible by bc={self.bc}")
35
+ if self.bt != 2 * self.bc:
36
+ raise ValueError(
37
+ f"bt={self.bt} must equal 2*bc (top-level WY-solve split "
38
+ f"supports only the 2-block case); got bc={self.bc}. "
39
+ f"Vary `mb` instead of `bc` to change solver granularity -- "
40
+ f"bc/mb do not affect numerical accuracy of the solve, only "
41
+ f"its speed (see grid_bt_bc_condition_diag.py Part 3)."
42
+ )
43
+ if self.bc % self.mb != 0:
44
+ raise ValueError(f"bc={self.bc} must be divisible by mb={self.mb}")
45
+ if not (0.0 <= self.wy_eps < 1.0):
46
+ raise ValueError(f"wy_eps={self.wy_eps} must be in [0, 1)")
47
+
48
+
49
+ KAGGLE_SMALL = KernelConfig(bt=128, bc=64, mb=16, clip=1e4, wy_eps=1e-3)
50
+ KAGGLE_MEDIUM = KernelConfig(bt=256, bc=128, mb=16, clip=1e4, wy_eps=1e-3)
51
+ KAGGLE_LARGE = KernelConfig(bt=256, bc=128, mb=16, clip=5e3, wy_eps=1e-3)
52
+ DEFAULT_CONFIG = KAGGLE_MEDIUM
53
+
54
+
55
+ def sanitize(x, config: KernelConfig = DEFAULT_CONFIG):
56
+ """Standard clip + nan_to_num defense used at every kernel boundary."""
57
+ c = config.clip
58
+ return jnp.nan_to_num(jnp.clip(x, -c, c), nan=0.0, posinf=c, neginf=-c)
59
+
60
+
61
+ def sanitize_h0(h0, config: KernelConfig = DEFAULT_CONFIG):
62
+ return sanitize(h0, config)
63
+
64
+
65
+ def clip_acc(x, config: KernelConfig = DEFAULT_CONFIG):
66
+ """Same as sanitize; separate name kept for call-site clarity in
67
+ read-modify-write accumulation loops (see gdn2_bwd.py, Kernel B4)."""
68
+ return sanitize(x, config)
69
+
70
+
71
+ def _reshape_to_chunks(t, bsz, n_chunks, H, D, bt):
72
+ t = t.reshape(bsz, n_chunks, bt, H, D)
73
+ return jnp.moveaxis(t, (1, 3), (2, 1))
74
+
75
+
76
+ def _reshape_from_chunks(t, bsz, n_chunks, bt, H, D):
77
+ t2 = jnp.moveaxis(t, (1, 2, 3), (3, 1, 2))
78
+ return t2.reshape(bsz, n_chunks * bt, H, D)
79
+
80
+
81
+ _GDN2_FWD_DIAG = os.environ.get("GDN2_FWD_DIAG", "0") == "1"
82
+ _LARGE_THRESHOLD = 1e6 # suspiciously large but still-finite trigger level
83
+
84
+
85
+ def _stage_diag(tag: str, x):
86
+ """No-op unless GDN2_FWD_DIAG=1. Diagnostic only -- never changes x."""
87
+ if not _GDN2_FWD_DIAG:
88
+ return x
89
+
90
+ finite_mask = jnp.isfinite(x)
91
+ all_finite = jnp.all(finite_mask)
92
+ n_nonfinite = jnp.sum(jnp.logical_not(finite_mask))
93
+ safe_x = jnp.where(finite_mask, x, 0.0)
94
+ max_abs = jnp.max(jnp.abs(safe_x))
95
+
96
+ def _report_nonfinite():
97
+ jax.debug.print(
98
+ "[GDN2-FWD-DIAG] non-finite at " + tag + ": n_nonfinite={n} max_abs_finite={m:.3e}",
99
+ n=n_nonfinite, m=max_abs,
100
+ )
101
+
102
+ def _report_large():
103
+ jax.debug.print(
104
+ "[GDN2-FWD-DIAG] suspiciously large (still finite) at " + tag + ": max_abs={m:.3e}",
105
+ m=max_abs,
106
+ )
107
+
108
+ jax.lax.cond(
109
+ jnp.logical_not(all_finite),
110
+ _report_nonfinite,
111
+ lambda: jax.lax.cond(max_abs > _LARGE_THRESHOLD, _report_large, lambda: None),
112
+ )
113
+ return x
114
+
115
+
116
+ def validate_inputs(q, k, v, w, b, g, scale, h0, config: KernelConfig):
117
+ """Shape/dtype sanity checks shared by forward and backward entry points."""
118
+ if q.ndim != 4:
119
+ raise ValueError(f"q must be (batch, seq_len, heads, d_head); got shape {q.shape}")
120
+ bsz, L, H, D = q.shape
121
+ if D != 128:
122
+ raise ValueError(f"Kernels assume d_head=128 (MXU tile); got D={D}.")
123
+ if L % config.bt != 0:
124
+ raise ValueError(f"seq_len={L} must be divisible by config.bt={config.bt}.")
125
+
126
+ for name, t in (("k", k), ("v", v), ("w", w), ("b", b), ("g", g)):
127
+ if t.shape != q.shape:
128
+ raise ValueError(f"{name}.shape={t.shape} must match q.shape={q.shape}")
129
+
130
+ if h0 is not None:
131
+ expected_h0 = (bsz, H, D, D)
132
+ if h0.shape != expected_h0:
133
+ raise ValueError(f"h0.shape={h0.shape} must be {expected_h0}")
134
+
135
+ n_chunks = L // config.bt
136
+ return bsz, L, H, D, n_chunks
atomic_ops/fallback.py ADDED
@@ -0,0 +1,22 @@
1
+ from __future__ import annotations
2
+ import jax.numpy as jnp
3
+ from .utils import is_tpu_available
4
+ from .gdn2_fwd import gdn2_pallas_forward as _pallas_fwd
5
+ from .gdn2_pipeline import gdn2_pallas_forward_trainable as _pallas_trainable
6
+ from .reference import gdn2_chunked_wy_reference
7
+
8
+
9
+ def gdn2_forward(q, k, v, w, b, g, scale, h0=None, config=None):
10
+ from .configs import DEFAULT_CONFIG
11
+ config = config or DEFAULT_CONFIG
12
+ if is_tpu_available() and q.shape[-1] == 128:
13
+ return _pallas_fwd(q, k, v, w, b, g, scale, h0=h0, config=config)
14
+ return gdn2_chunked_wy_reference(q, k, v, g, b, w, scale, chunk_size=config.bt, h0=h0)
15
+
16
+
17
+ def gdn2_forward_trainable(q, k, v, w, b, g, scale, h0=None, config=None):
18
+ from .configs import DEFAULT_CONFIG
19
+ config = config or DEFAULT_CONFIG
20
+ if is_tpu_available() and q.shape[-1] == 128:
21
+ return _pallas_trainable(q, k, v, w, b, g, scale, h0=h0, config=config)
22
+ return gdn2_chunked_wy_reference(q, k, v, g, b, w, scale, chunk_size=config.bt, h0=h0)
atomic_ops/gdn2_bwd.py ADDED
@@ -0,0 +1,342 @@
1
+ """
2
+ Backward kernels: B1 (state) -> B2 (dAqk/dv) -> B3 (WY/dqkg) -> B4 (intra) -> B5 (reverse cumsum).
3
+ """
4
+ from __future__ import annotations
5
+
6
+ import jax
7
+ import jax.numpy as jnp
8
+ from jax.experimental import pallas as pl
9
+ from jax.experimental.pallas import tpu as pltpu
10
+
11
+ from .configs import (
12
+ KernelConfig, DEFAULT_CONFIG, sanitize, clip_acc,
13
+ _reshape_to_chunks as _r2c, _reshape_from_chunks as _r2f,
14
+ )
15
+
16
+ _HIGHEST = jax.lax.Precision.HIGHEST
17
+
18
+
19
+ # ---------- B5 ----------
20
+ def reverse_cumsum_bwd(dgc, chunk_size: int, config: KernelConfig = DEFAULT_CONFIG):
21
+ C = chunk_size
22
+ idx = jnp.arange(C)
23
+ triu_ones = (idx[:, None] <= idx[None, :]).astype(jnp.float32)
24
+ dg_raw = jnp.einsum("ij,...jd->...id", triu_ones, dgc.astype(jnp.float32), precision=_HIGHEST)
25
+ return sanitize(dg_raw, config)
26
+
27
+
28
+ # ---------- B1 ----------
29
+ def gdn2_dhu_backward(do, dv_partial, w_pseudo, qg, kg, gc_last, scale, dht=None,
30
+ config: KernelConfig = DEFAULT_CONFIG):
31
+ bsz, H, n_chunks, BT, D = qg.shape
32
+ if dht is None:
33
+ dht = jnp.zeros((bsz, H, D, D), dtype=jnp.float32)
34
+ dht = sanitize(dht, config)
35
+
36
+ to_scan = tuple(jnp.moveaxis(x, 2, 0) for x in (do, dv_partial, w_pseudo, qg, kg, gc_last))
37
+
38
+ def step(dh_carry, inputs):
39
+ do_c, dvp_c, wp_c, qg_c, kg_c, gclast_c = inputs
40
+ decay_c = jnp.exp(gclast_c)[..., None]
41
+
42
+ dqh = scale * do_c
43
+ contrib_from_output = jnp.einsum("bhid,bhiv->bhdv", qg_c, dqh, precision=_HIGHEST)
44
+ contrib_from_state = dh_carry * decay_c
45
+
46
+ dv_write = jnp.einsum("bhid,bhdv->bhiv", kg_c, dh_carry, precision=_HIGHEST)
47
+ dv_new_c = dvp_c + dv_write
48
+ dv_new_c = sanitize(dv_new_c, config)
49
+
50
+ contrib_from_vnew = -jnp.einsum("bhjd,bhjv->bhdv", wp_c, dv_new_c, precision=_HIGHEST)
51
+
52
+ dh_pre_c = contrib_from_output + contrib_from_state + contrib_from_vnew
53
+ dh_pre_c = sanitize(dh_pre_c, config)
54
+ return dh_pre_c, (dh_pre_c, dv_new_c)
55
+
56
+ dh0, (dh_all_rev, dv_all_rev) = jax.lax.scan(step, dht, to_scan, reverse=True)
57
+ dh_all = jnp.moveaxis(dh_all_rev, 0, 2)
58
+ dv_all = jnp.moveaxis(dv_all_rev, 0, 2)
59
+ return dh_all, dh0, dv_all
60
+
61
+
62
+ # ---------- B2 ----------
63
+ def _kernel_b2_body(aqk_ref, vnew_ref, do_ref, daqk_ref, dvnew_ref, *, bt: int, config: KernelConfig):
64
+ Aqk = aqk_ref[0, 0, 0].astype(jnp.float32)
65
+ v_new = vnew_ref[0, 0, 0].astype(jnp.float32)
66
+ do = do_ref[0, 0, 0].astype(jnp.float32)
67
+
68
+ idx = jnp.arange(bt)
69
+ causal = (idx[:, None] >= idx[None, :]).astype(jnp.float32)
70
+
71
+ dAqk = jnp.dot(do, v_new.T, precision=_HIGHEST) * causal
72
+ dv_new = jnp.dot(Aqk.T, do, precision=_HIGHEST)
73
+
74
+ daqk_ref[0, 0, 0] = sanitize(dAqk, config)
75
+ dvnew_ref[0, 0, 0] = sanitize(dv_new, config)
76
+
77
+
78
+ def dav_backward_pallas(Aqk, v_new, do, config: KernelConfig = DEFAULT_CONFIG):
79
+ bsz, H, n_chunks, _BT, D = v_new.shape
80
+ grid = (bsz, H, n_chunks)
81
+ aqk_spec = pl.BlockSpec((1, 1, 1, config.bt, config.bt), lambda i, h, c: (i, h, c, 0, 0))
82
+ io_spec = pl.BlockSpec((1, 1, 1, config.bt, D), lambda i, h, c: (i, h, c, 0, 0))
83
+
84
+ dAqk, dv_new = pl.pallas_call(
85
+ lambda *refs: _kernel_b2_body(*refs, bt=config.bt, config=config),
86
+ grid=grid,
87
+ in_specs=[aqk_spec, io_spec, io_spec],
88
+ out_specs=[aqk_spec, io_spec],
89
+ out_shape=[
90
+ jax.ShapeDtypeStruct(Aqk.shape, jnp.float32),
91
+ jax.ShapeDtypeStruct(v_new.shape, jnp.float32),
92
+ ],
93
+ compiler_params=pltpu.CompilerParams(vmem_limit_bytes=64 * 1024 * 1024),
94
+ )(Aqk, v_new, do)
95
+ return dAqk, dv_new
96
+
97
+
98
+ # ---------- B3 ----------
99
+ def _kernel_b3_body(q_ref, k_ref, b_ref, w_ref, v_ref, gc_ref, a_ref, akk_ref,
100
+ hpre_ref, vnew_ref, do_ref, dv_ref, dhnext_ref,
101
+ dq_ref, dk_ref, db_ref, dw_ref, dvraw_ref, dgc_ref, dakk_ref,
102
+ *, scale: float, bt: int, wy_eps: float, config: KernelConfig):
103
+ q_c = q_ref[0, 0, 0].astype(jnp.float32)
104
+ k_c = k_ref[0, 0, 0].astype(jnp.float32)
105
+ b_c = b_ref[0, 0, 0].astype(jnp.float32)
106
+ w_c = w_ref[0, 0, 0].astype(jnp.float32)
107
+ v_c = v_ref[0, 0, 0].astype(jnp.float32)
108
+ gc = gc_ref[0, 0, 0].astype(jnp.float32)
109
+ A = a_ref[0, 0, 0].astype(jnp.float32)
110
+ h_pre = hpre_ref[0, 0, 0].astype(jnp.float32)
111
+ v_new = vnew_ref[0, 0, 0].astype(jnp.float32)
112
+ do = do_ref[0, 0, 0].astype(jnp.float32)
113
+ dv = dv_ref[0, 0, 0].astype(jnp.float32)
114
+ dh_next = dhnext_ref[0, 0, 0].astype(jnp.float32)
115
+
116
+ C = bt
117
+ gc_last = gc[C - 1]
118
+
119
+ kb_decayed = b_c * k_c * jnp.exp(gc)
120
+ kg = k_c * jnp.exp(gc_last[None, :] - gc)
121
+ qg = q_c * jnp.exp(gc)
122
+ wv = w_c * v_c
123
+
124
+ dqh_up = scale * do
125
+ dqg = jnp.dot(dqh_up, h_pre.T, precision=_HIGHEST)
126
+
127
+ dwh = -dv
128
+ dw_pseudo = jnp.dot(dwh, h_pre.T, precision=_HIGHEST)
129
+ du = dv
130
+
131
+ dkg = jnp.dot(v_new, dh_next.T, precision=_HIGHEST)
132
+
133
+ dA_from_w = jnp.dot(dw_pseudo, kb_decayed.T, precision=_HIGHEST)
134
+ dkb_decayed = jnp.dot(A.T, dw_pseudo, precision=_HIGHEST)
135
+
136
+ dA_from_u = jnp.dot(du, wv.T, precision=_HIGHEST)
137
+ dwv = jnp.dot(A.T, du, precision=_HIGHEST)
138
+
139
+ dA_total = dA_from_w + dA_from_u
140
+ dA_total = sanitize(dA_total, config)
141
+
142
+ idx = jnp.arange(C)
143
+ strict = (idx[:, None] > idx[None, :]).astype(jnp.float32)
144
+
145
+ tmp = jnp.dot(dA_total, A.T, precision=_HIGHEST)
146
+ tmp = sanitize(tmp, config)
147
+ dAkk_raw = -jnp.dot(A.T, tmp, precision=_HIGHEST)
148
+ dAkk_raw = dAkk_raw * (1.0 - wy_eps)
149
+ dAkk = dAkk_raw * strict
150
+
151
+ dk_from_kb = dkb_decayed * jnp.exp(gc) * b_c
152
+ db = dkb_decayed * jnp.exp(gc) * k_c
153
+ dgc_from_kb = dkb_decayed * kb_decayed
154
+
155
+ dx = dkg * kg
156
+ dk_from_kg = dkg * jnp.exp(gc_last[None, :] - gc)
157
+ dgc_from_kg = -dx
158
+ dgc_last_contrib = jnp.sum(dx, axis=0)
159
+
160
+ dq = dqg * jnp.exp(gc)
161
+ dgc_from_qg = dqg * qg
162
+
163
+ dw = dwv * v_c
164
+ dv_raw = dwv * w_c
165
+
166
+ dk = dk_from_kb + dk_from_kg
167
+ dgc = dgc_from_kb + dgc_from_qg + dgc_from_kg
168
+
169
+ decay_h_row = jnp.exp(gc_last)
170
+ dgc_last_from_decay = decay_h_row * jnp.sum(dh_next * h_pre, axis=-1)
171
+ dgc_last_total = dgc_last_contrib + dgc_last_from_decay
172
+
173
+ row_mask = (idx == (C - 1)).astype(jnp.float32)[:, None]
174
+ dgc = dgc + row_mask * dgc_last_total[None, :]
175
+
176
+ dq_ref[0, 0, 0] = sanitize(dq, config)
177
+ dk_ref[0, 0, 0] = sanitize(dk, config)
178
+ db_ref[0, 0, 0] = sanitize(db, config)
179
+ dw_ref[0, 0, 0] = sanitize(dw, config)
180
+ dvraw_ref[0, 0, 0] = sanitize(dv_raw, config)
181
+ dakk_ref[0, 0, 0] = sanitize(dAkk, config)
182
+ dgc_ref[0, 0, 0] = sanitize(dgc, config)
183
+
184
+
185
+ def wy_dqkg_backward_pallas(q, k, b, w, v, gc, A, Akk, h_pre_all, v_new_all,
186
+ do, dv, dh_next_all, scale, config: KernelConfig = DEFAULT_CONFIG):
187
+ bsz, H, n_chunks, _BT, D = q.shape
188
+ grid = (bsz, H, n_chunks)
189
+
190
+ io_spec = pl.BlockSpec((1, 1, 1, config.bt, D), lambda i, h, c: (i, h, c, 0, 0))
191
+ score_spec = pl.BlockSpec((1, 1, 1, config.bt, config.bt), lambda i, h, c: (i, h, c, 0, 0))
192
+ h_spec = pl.BlockSpec((1, 1, 1, D, D), lambda i, h, c: (i, h, c, 0, 0))
193
+
194
+ dq, dk, db, dw, dv_raw, dgc, dAkk = pl.pallas_call(
195
+ lambda *refs: _kernel_b3_body(*refs, scale=scale, bt=config.bt, wy_eps=config.wy_eps, config=config),
196
+ grid=grid,
197
+ in_specs=[io_spec, io_spec, io_spec, io_spec, io_spec, io_spec,
198
+ score_spec, score_spec, h_spec, io_spec, io_spec, io_spec, h_spec],
199
+ out_specs=[io_spec, io_spec, io_spec, io_spec, io_spec, io_spec, score_spec],
200
+ out_shape=[
201
+ jax.ShapeDtypeStruct((bsz, H, n_chunks, config.bt, D), jnp.float32),
202
+ jax.ShapeDtypeStruct((bsz, H, n_chunks, config.bt, D), jnp.float32),
203
+ jax.ShapeDtypeStruct((bsz, H, n_chunks, config.bt, D), jnp.float32),
204
+ jax.ShapeDtypeStruct((bsz, H, n_chunks, config.bt, D), jnp.float32),
205
+ jax.ShapeDtypeStruct((bsz, H, n_chunks, config.bt, D), jnp.float32),
206
+ jax.ShapeDtypeStruct((bsz, H, n_chunks, config.bt, D), jnp.float32),
207
+ jax.ShapeDtypeStruct((bsz, H, n_chunks, config.bt, config.bt), jnp.float32),
208
+ ],
209
+ compiler_params=pltpu.CompilerParams(vmem_limit_bytes=100 * 1024 * 1024),
210
+ )(q, k, b, w, v, gc, A, Akk, h_pre_all, v_new_all, do, dv, dh_next_all)
211
+
212
+ return dict(dq=dq, dk=dk, db=db, dw=dw, dv_raw=dv_raw, dgc=dgc, dAkk=dAkk)
213
+
214
+
215
+ # ---------- B4 ----------
216
+ def _dL_pair_sum(dM, edecay, R):
217
+ tmp = dM[:, :, None] * edecay
218
+ tmp = tmp * R[None, :, :]
219
+ return jnp.sum(tmp, axis=1)
220
+
221
+
222
+ def _dR_pair_sum(dM, edecay, L):
223
+ tmp = dM[:, :, None] * edecay
224
+ tmp = tmp * L[:, None, :]
225
+ return jnp.sum(tmp, axis=0)
226
+
227
+
228
+ def _dgc_pair_sum(dM, edecay, L, R, clipmask):
229
+ weight = dM[:, :, None] * L[:, None, :] * R[None, :, :] * edecay * clipmask
230
+ dgc_i = jnp.sum(weight, axis=1)
231
+ dgc_j = -jnp.sum(weight, axis=0)
232
+ return dgc_i, dgc_j
233
+
234
+
235
+ def _kernel_b4_body(q_ref, k_ref, b_ref, g_ref, daqk_ref, dakk_ref,
236
+ dq_ref, dk_ref, db_ref, dgc_ref, *, scale: float, bt: int, bc: int, n_sub: int,
237
+ config: KernelConfig):
238
+ q_full = q_ref[0, 0, 0].astype(jnp.float32)
239
+ k_full = k_ref[0, 0, 0].astype(jnp.float32)
240
+ b_full = b_ref[0, 0, 0].astype(jnp.float32)
241
+ g_raw = g_ref[0, 0, 0].astype(jnp.float32)
242
+ dAqk = daqk_ref[0, 0, 0].astype(jnp.float32)
243
+ dAkk = dakk_ref[0, 0, 0].astype(jnp.float32)
244
+
245
+ bt_idx = jnp.arange(bt)
246
+ tril_ones_bt = (bt_idx[:, None] >= bt_idx[None, :]).astype(jnp.float32)
247
+ gc = jnp.dot(tril_ones_bt, g_raw, precision=_HIGHEST)
248
+
249
+ bk_full = b_full * k_full
250
+
251
+ dq_ref[0, 0, 0] = jnp.zeros_like(q_full)
252
+ dk_ref[0, 0, 0] = jnp.zeros_like(k_full)
253
+ db_ref[0, 0, 0] = jnp.zeros_like(k_full)
254
+ dgc_ref[0, 0, 0] = jnp.zeros_like(g_raw)
255
+
256
+ for si in range(n_sub):
257
+ for sj in range(si + 1):
258
+ i0, i1 = si * bc, (si + 1) * bc
259
+ j0, j1 = sj * bc, (sj + 1) * bc
260
+
261
+ q_i = q_full[i0:i1]
262
+ k_i = k_full[i0:i1]
263
+ k_j = k_full[j0:j1]
264
+ b_i = b_full[i0:i1]
265
+ bk_i = bk_full[i0:i1]
266
+ gc_i = gc[i0:i1]
267
+ gc_j = gc[j0:j1]
268
+
269
+ dM_qk = dAqk[i0:i1, j0:j1]
270
+ dM_kk = dAkk[i0:i1, j0:j1]
271
+ if si == sj:
272
+ idx = jnp.arange(bc)
273
+ causal = (idx[:, None] >= idx[None, :]).astype(jnp.float32)
274
+ strict = (idx[:, None] > idx[None, :]).astype(jnp.float32)
275
+ dM_qk = dM_qk * causal
276
+ dM_kk = dM_kk * strict
277
+
278
+ decay_diff = gc_i[:, None, :] - gc_j[None, :, :]
279
+ clipmask = ((decay_diff >= -20.0) & (decay_diff <= 20.0)).astype(jnp.float32)
280
+ edecay = jnp.exp(jnp.clip(decay_diff, -20.0, 20.0))
281
+
282
+ L_qk = scale * q_i
283
+ R_qk = k_j
284
+ dL_qk = _dL_pair_sum(dM_qk, edecay, R_qk)
285
+ dR_qk = _dR_pair_sum(dM_qk, edecay, L_qk)
286
+ dgc_i_qk, dgc_j_qk = _dgc_pair_sum(dM_qk, edecay, L_qk, R_qk, clipmask)
287
+
288
+ L_kk = bk_i
289
+ R_kk = k_j
290
+ dL_kk = _dL_pair_sum(dM_kk, edecay, R_kk)
291
+ dR_kk = _dR_pair_sum(dM_kk, edecay, L_kk)
292
+ dgc_i_kk, dgc_j_kk = _dgc_pair_sum(dM_kk, edecay, L_kk, R_kk, clipmask)
293
+
294
+ dq_ref[0, 0, 0, i0:i1] = clip_acc(dq_ref[0, 0, 0, i0:i1] + dL_qk * scale, config)
295
+ db_ref[0, 0, 0, i0:i1] = clip_acc(db_ref[0, 0, 0, i0:i1] + dL_kk, config)
296
+ dk_ref[0, 0, 0, j0:j1] = clip_acc(dk_ref[0, 0, 0, j0:j1] + dR_qk + dR_kk, config)
297
+ dgc_ref[0, 0, 0, i0:i1] = clip_acc(dgc_ref[0, 0, 0, i0:i1] + dgc_i_qk + dgc_i_kk, config)
298
+ dgc_ref[0, 0, 0, j0:j1] = clip_acc(dgc_ref[0, 0, 0, j0:j1] + dgc_j_qk + dgc_j_kk, config)
299
+
300
+ dbk_final = db_ref[0, 0, 0]
301
+ dk_final = dk_ref[0, 0, 0] + dbk_final * b_full
302
+ db_final = dbk_final * k_full
303
+ dq_final = dq_ref[0, 0, 0]
304
+ dgc_final = dgc_ref[0, 0, 0]
305
+
306
+ dq_ref[0, 0, 0] = sanitize(dq_final, config)
307
+ dk_ref[0, 0, 0] = sanitize(dk_final, config)
308
+ db_ref[0, 0, 0] = sanitize(db_final, config)
309
+ dgc_ref[0, 0, 0] = sanitize(dgc_final, config)
310
+
311
+ def intra_backward_pallas(dAqk, dAkk, q, k, b, g, scale, config: KernelConfig = DEFAULT_CONFIG, interpret: bool = False):
312
+ bsz, L, H, D = q.shape
313
+ n_chunks = L // config.bt
314
+
315
+ def reshape_in(t):
316
+ return _r2c(t, bsz, n_chunks, H, D, config.bt)
317
+
318
+ q_r, k_r, b_r, g_r = map(reshape_in, (q, k, b, g))
319
+
320
+ grid = (bsz, H, n_chunks)
321
+ io_spec = pl.BlockSpec((1, 1, 1, config.bt, D), lambda i, h, c: (i, h, c, 0, 0))
322
+ score_spec = pl.BlockSpec((1, 1, 1, config.bt, config.bt), lambda i, h, c: (i, h, c, 0, 0))
323
+
324
+ dq, dk, db, dgc = pl.pallas_call(
325
+ lambda *refs: _kernel_b4_body(
326
+ *refs, scale=scale, bt=config.bt, bc=config.bc, n_sub=config.n_sub,
327
+ config=config,
328
+ ),
329
+ grid=grid,
330
+ in_specs=[io_spec, io_spec, io_spec, io_spec, score_spec, score_spec],
331
+ out_specs=[io_spec, io_spec, io_spec, io_spec],
332
+ out_shape=[
333
+ jax.ShapeDtypeStruct((bsz, H, n_chunks, config.bt, D), jnp.float32),
334
+ jax.ShapeDtypeStruct((bsz, H, n_chunks, config.bt, D), jnp.float32),
335
+ jax.ShapeDtypeStruct((bsz, H, n_chunks, config.bt, D), jnp.float32),
336
+ jax.ShapeDtypeStruct((bsz, H, n_chunks, config.bt, D), jnp.float32),
337
+ ],
338
+ compiler_params=pltpu.CompilerParams(vmem_limit_bytes=150 * 1024 * 1024),
339
+ interpret=interpret,
340
+ )(q_r, k_r, b_r, g_r, dAqk, dAkk)
341
+
342
+ return dq, dk, db, dgc