atomic-ops 0.1.0__tar.gz

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.
@@ -0,0 +1,21 @@
1
+ MIT License
2
+
3
+ Copyright (c) 2026 Omirbay Akseleu
4
+
5
+ Permission is hereby granted, free of charge, to any person obtaining a copy
6
+ of this software and associated documentation files (the "Software"), to deal
7
+ in the Software without restriction, including without limitation the rights
8
+ to use, copy, modify, merge, publish, distribute, sublicense, and/or sell
9
+ copies of the Software, and to permit persons to whom the Software is
10
+ furnished to do so, subject to the following conditions:
11
+
12
+ The above copyright notice and this permission notice shall be included in all
13
+ copies or substantial portions of the Software.
14
+
15
+ THE SOFTWARE IS PROVIDED "AS IS", WITHOUT WARRANTY OF ANY KIND, EXPRESS OR
16
+ IMPLIED, INCLUDING BUT NOT LIMITED TO THE WARRANTIES OF MERCHANTABILITY,
17
+ FITNESS FOR A PARTICULAR PURPOSE AND NONINFRINGEMENT. IN NO EVENT SHALL THE
18
+ AUTHORS OR COPYRIGHT HOLDERS BE LIABLE FOR ANY CLAIM, DAMAGES OR OTHER
19
+ LIABILITY, WHETHER IN AN ACTION OF CONTRACT, TORT OR OTHERWISE, ARISING FROM,
20
+ OUT OF OR IN CONNECTION WITH THE SOFTWARE OR THE USE OR OTHER DEALINGS IN THE
21
+ SOFTWARE.
@@ -0,0 +1,3 @@
1
+ include LICENSE
2
+ include README.md
3
+ recursive-include atomic_ops py.typed
@@ -0,0 +1,296 @@
1
+ Metadata-Version: 2.4
2
+ Name: atomic_ops
3
+ Version: 0.1.0
4
+ Summary: Fused Gated DeltaNet-2 kernels for TPU v5e (JAX/Pallas)
5
+ Author-email: Omirbay Akseleu <omirbajakseleu5@gmail.com>
6
+ License: MIT
7
+ Project-URL: Homepage, https://github.com/Akseleu-J/atomic_ops
8
+ Project-URL: Repository, https://github.com/Akseleu-J/atomic_ops
9
+ Project-URL: Issues, https://github.com/Akseleu-J/atomic_ops/issues
10
+ Keywords: jax,pallas,tpu,gated-deltanet2,deltanet,linear-attention,gated-linear-attention
11
+ Classifier: Programming Language :: Python :: 3.10
12
+ Classifier: Programming Language :: Python :: 3.11
13
+ Classifier: Programming Language :: Python :: 3.12
14
+ Classifier: License :: OSI Approved :: MIT License
15
+ Classifier: Topic :: Scientific/Engineering :: Artificial Intelligence
16
+ Classifier: Operating System :: OS Independent
17
+ Classifier: Intended Audience :: Science/Research
18
+ Requires-Python: >=3.10
19
+ Description-Content-Type: text/markdown
20
+ License-File: LICENSE
21
+ Requires-Dist: jax<0.13,>=0.4.20
22
+ Requires-Dist: jaxlib<0.13,>=0.4.20
23
+ Requires-Dist: numpy>=1.24
24
+ Provides-Extra: dev
25
+ Requires-Dist: pytest>=7.0; extra == "dev"
26
+ Requires-Dist: pytest-cov; extra == "dev"
27
+ Requires-Dist: ruff>=0.6.0; extra == "dev"
28
+ Requires-Dist: pre-commit>=3.0; extra == "dev"
29
+ Provides-Extra: examples
30
+ Requires-Dist: flax<0.13,>=0.12.9; extra == "examples"
31
+ Requires-Dist: optax>=0.1.7; extra == "examples"
32
+ Dynamic: license-file
33
+
34
+ # atomic_ops
35
+
36
+ **Fused Gated DeltaNet-2 (GDN-2) kernels for TPU v5e, written in JAX/Pallas.**
37
+
38
+ [![License: MIT](https://img.shields.io/badge/License-MIT-yellow.svg)](https://opensource.org/licenses/MIT)
39
+ [![Python 3.10+](https://img.shields.io/badge/python-3.10%2B-blue.svg)](https://www.python.org/downloads/)
40
+ [![TPU v5e](https://img.shields.io/badge/TPU-v5e-orange.svg)](https://cloud.google.com/tpu/docs/v5e)
41
+
42
+ A from-scratch port of the [NVlabs Gated DeltaNet-2](https://github.com/NVlabs/GatedDeltaNet-2) Triton kernels
43
+ to `jax.experimental.pallas`, targeting **TPU v5e-8**. The backward pass is a single fused `custom_vjp`
44
+ that reuses forward residuals instead of recomputing them.
45
+
46
+ > **Headline numbers (measured, see [Benchmarks](#-benchmarks)):** on the full training shape
47
+ > (batch 8, seq 4096, 6 heads, d_head 128) the fused backward makes the training step
48
+ > **2.6× (FP32) / 3.4× (BF16) faster than the best pure-JAX WY baseline** and
49
+ > **27.2× (FP32) / 13.3× (BF16) faster than the widely used `associative_scan` baseline**.
50
+ > Across all measured shapes the gain vs `associative_scan` reaches **up to 38.8×** (FP32).
51
+ > The fused forward alone is currently **~1.6× slower** than the pure-JAX WY forward on TPU —
52
+ > training steps are backward-dominated, so end-to-end is still a clear win.
53
+ > A hybrid `JAX forward + Pallas backward` mode is planned (see [Limitations](#-limitations)).
54
+
55
+ ---
56
+
57
+ ## Features
58
+
59
+ - **Fully fused kernels** — forward (chunk scores → WY block solve → recompute → inter-chunk scan)
60
+ and backward (B1–B5 stage kernels), all in Pallas with TPU MXU tiling.
61
+ - **One-call training API** — `gdn2_forward_trainable` via `jax.custom_vjp`; no recomputation
62
+ of forward activations in the backward pass.
63
+ - **Automatic fallback** — on CPU/GPU or when `d_head != 128`, dispatches to a checkpointed
64
+ pure-JAX chunked-WY reference with identical damping semantics.
65
+ - **Numerical safety by default** — clipping + `nan_to_num` at every kernel boundary,
66
+ optional `wy_eps` Tikhonov-style damping of the WY solve, per-stage diagnostics
67
+ via the `GDN2_FWD_DIAG=1` environment flag.
68
+ - **Tunable configs** — `KernelConfig(bt, bc, mb, clip, wy_eps)` with presets tuned for
69
+ Kaggle TPU v5e (`KAGGLE_SMALL` / `KAGGLE_MEDIUM` / `KAGGLE_LARGE`) and an
70
+ `estimate_memory` / `get_recommended_config` helper.
71
+
72
+ ## Installation
73
+
74
+ ```bash
75
+ pip install atomic_ops
76
+
77
+ # or from source
78
+ pip install git+https://github.com/Akseleu-J/atomic_ops.git
79
+
80
+ # development
81
+ git clone https://github.com/Akseleu-J/atomic_ops.git
82
+ cd atomic_ops
83
+ pip install jax==0.11.1 jaxlib==0.11.1 libtpu==0.0.46 flax==0.12.9 optax==0.2.4
84
+ ```
85
+
86
+ Requires Python ≥ 3.10 and `jax>=0.4.20`. For TPU, install the matching `jaxlib`/`libtpu`
87
+ builds (see [JAX TPU installation](https://jax.readthedocs.io/en/latest/installation.html)).
88
+
89
+ ## Quick start
90
+
91
+ ```python
92
+ import jax.numpy as jnp
93
+ from atomic_ops import gdn2_forward_trainable
94
+
95
+ # (batch, seq_len, heads, d_head); d_head must be 128 on TPU
96
+ shape = (4, 2048, 6, 128)
97
+ q = jnp.ones(shape, dtype=jnp.float32)
98
+ k = jnp.ones(shape, dtype=jnp.float32)
99
+ v = jnp.ones(shape, dtype=jnp.float32)
100
+ w = jnp.ones(shape, dtype=jnp.float32)
101
+ b = jnp.ones(shape, dtype=jnp.float32)
102
+ g = jnp.ones(shape, dtype=jnp.float32) # log-decay gate, keep g <= 0
103
+ scale = shape[-1] ** -0.5
104
+
105
+ # Forward + backward (custom_vjp under the hood)
106
+ out, h_final = gdn2_forward_trainable(q, k, v, w, b, g, scale)
107
+
108
+ print(out.shape) # (4, 2048, 6, 128)
109
+ print(h_final.shape) # (4, 6, 128, 128)
110
+ ```
111
+
112
+ ### Runnable examples
113
+
114
+ Both examples work out of the box after `pip install atomic_ops`:
115
+
116
+ - [`examples/minimal_usage.py`](examples/minimal_usage.py) — 20-line minimal script:
117
+ build tensors, run `gdn2_forward_trainable`, print output shapes. Includes the
118
+ `is_tpu_available()` check so you can see which backend path was taken.
119
+ - [`examples/train_gdn2_layer.py`](examples/train_gdn2_layer.py) — a complete Flax +
120
+ Optax training step: a `GDN2Layer` module that projects `x` into `q,k,v,w,b,g`,
121
+ applies the gate via `-softplus(g)` (keeps `g <= 0`), auto-picks a config with
122
+ `get_recommended_config`, and runs one AdamW update step.
123
+
124
+ ```bash
125
+ python examples/minimal_usage.py # forward+backward, prints backend + shapes
126
+ python examples/train_gdn2_layer.py # full training step with Flax/Optax
127
+ ```
128
+
129
+ ## API
130
+
131
+ | Function | Description |
132
+ | --- | --- |
133
+ | `gdn2_forward_trainable(q, k, v, w, b, g, scale, h0=None, config=None)` | Forward **+** gradients. Dispatches to Pallas on TPU with `d_head=128`, otherwise to the pure-JAX reference. This is what you want for training. |
134
+ | `gdn2_forward(q, k, v, w, b, g, scale, h0=None, config=None)` | Forward only, same auto-dispatch. |
135
+ | `gdn2_pallas_forward_trainable(...)` | Forces the fused Pallas path (raises on non-TPU / `d_head != 128`). |
136
+ | `gdn2_pallas_forward(...)` / `gdn2_pallas_forward_with_residuals(...)` | Inference-oriented Pallas forward; the `_with_residuals` variant also returns the intermediate tensors the backward pass needs. |
137
+ | `gdn2_token_serial_reference(...)` | Ground-truth token-by-token scan. Slow, numerically exact — used as an independent check in tests. |
138
+ | `gdn2_chunked_wy_reference(...)` | Chunked-WY reference (the fallback path). |
139
+ | `KernelConfig(bt, bc, mb, clip, wy_eps)` | Kernel blocking / numerical-safety config. Constraints: `bt = 2*bc`, `bc % mb == 0`. |
140
+ | `KAGGLE_SMALL` / `KAGGLE_MEDIUM` / `KAGGLE_LARGE` | Presets for TPU v5e-8 (default: `KAGGLE_MEDIUM`). |
141
+ | `estimate_memory(...)` / `get_recommended_config(...)` | Rough per-chip HBM estimate and preset auto-selection. |
142
+ | `is_tpu_available()` | `True` if a TPU device is visible to JAX (used by the fallback dispatcher). |
143
+
144
+ All tensor arguments have shape `(batch, seq_len, heads, d_head)` with `seq_len % config.bt == 0`.
145
+ The initial recurrent state `h0` is optional, shape `(batch, heads, d_head, d_head)`, float32.
146
+
147
+ ## Benchmarks
148
+
149
+ Measured on **TPU v5e-8** with `jax==0.11.1`, `jaxlib==0.11.1`, `libtpu==0.0.46.1`.
150
+ Baselines: **OLD** = `jax.lax.associative_scan` GDN-2 (the formulation used in many research
151
+ codebases); **JAX_REF** = chunked WY recurrence in pure JAX. Mean steady-state over repeats, ms.
152
+
153
+ ### Full training step — batch 8, seq 4096, 6 heads, d_head 128
154
+
155
+ | Dtype | Stage | OLD (ms) | JAX_REF (ms) | **Pallas (ms)** | vs OLD | vs JAX_REF |
156
+ | --- | --- | --- | --- | --- | --- | --- |
157
+ | FP32 | fwd | 1067.24 | 63.24 | 102.08 | 10.5× | 0.62× |
158
+ | FP32 | bwd | 4767.24 | 458.99 | **175.33** | **27.2×** | **2.6×** |
159
+ | FP32 | fwd+bwd | 4762.70 | 460.03 | **175.20** | **27.2×** | **2.6×** |
160
+ | BF16 | fwd | 585.49 | 62.38 | 101.51 | 5.8× | 0.61× |
161
+ | BF16 | bwd | 2322.97 | 594.17 | **174.25** | **13.3×** | **3.4×** |
162
+ | BF16 | fwd+bwd | 2320.27 | 594.86 | **174.20** | **13.3×** | **3.4×** |
163
+
164
+ ### All measured shapes — full fwd+bwd cycle, speedup vs baselines
165
+
166
+ | Config (batch, seq_len) | Dtype | Pallas (ms) | vs OLD (scan) | vs JAX_REF |
167
+ | --- | --- | --- | --- | --- |
168
+ | small (B=1, L=1024) | FP32 | 5.92 | 24.3× | 3.7× |
169
+ | small (B=1, L=1024) | BF16 | 6.00 | 11.5× | 3.8× |
170
+ | medium (B=4, L=4096) | FP32 | 88.07 | 30.2× | 3.1× |
171
+ | medium (B=4, L=4096) | BF16 | 87.60 | 14.9× | 3.8× |
172
+ | **train shape (B=8, L=4096)** | FP32 | **175.20** | **27.2×** | **2.6×** |
173
+ | **train shape (B=8, L=4096)** | BF16 | **174.20** | **13.3×** | **3.4×** |
174
+ | KAGGLE_SMALL preset (B=4, L=2048) | FP32 | 30.86 | **38.8×** | 3.9× |
175
+ | KAGGLE_SMALL preset (B=4, L=2048) | BF16 | 30.66 | 18.7× | 3.9× |
176
+ | KAGGLE_LARGE preset (B=8, L=4096) | FP32 | 175.34 | 27.2× | 2.6× |
177
+ | KAGGLE_LARGE preset (B=8, L=4096) | BF16 | 174.12 | 13.3× | 3.4× |
178
+
179
+ Read it honestly: the fused **forward is ~1.6× slower than the pure-JAX WY forward** today;
180
+ the fused **backward is 2.6–3.9× faster** than JAX_REF, and since real training steps are
181
+ backward-dominated, the full cycle wins by **2.6–3.9×** over the fastest pure-JAX baseline and
182
+ by **11.5–38.8×** over the `associative_scan` baseline depending on shape and dtype.
183
+ Full per-stage tables for all configs and both dtypes:
184
+ [`benchmarks/raw/benchmark_speed_final_averaged.md`](benchmarks/raw/benchmark_speed_final_averaged.md).
185
+
186
+ ### Peak HBM — train shape (B=8, L=4096), forward+backward
187
+
188
+ | Dtype | OLD (MB) | JAX_REF (MB) | Pallas (MB) |
189
+ | --- | --- | --- | --- |
190
+ | FP32 | 1252.7 | 1250.0 | **1238.1** |
191
+ | BF16 | 915.5 | 915.5 | 915.5 |
192
+
193
+ Memory is on par with the pure-JAX reference (the backward reuses forward residuals instead of
194
+ recomputing them). Details:
195
+ [`benchmarks/raw/benchmark_memory_final_averaged.md`](benchmarks/raw/benchmark_memory_final_averaged.md).
196
+
197
+ ### Reproduce
198
+
199
+ ```bash
200
+ python benchmarks/run_speed_benchmark.py # writes JSON + markdown tables
201
+ python benchmarks/run_memory_benchmark.py # fork-isolated peak HBM measurement
202
+ ```
203
+
204
+ Every timing run is correctness-gated before measurement (Pallas output must match the reference
205
+ within tolerance, otherwise the row is rejected).
206
+
207
+ ## Correctness & testing
208
+
209
+ The test suite is deliberately layered — comparing implementations that share the same algebraic
210
+ derivation can hide derivation bugs, so the strongest checks are derivation-independent
211
+ (finite-difference gradients, token-serial scan). Full rationale:
212
+ [`docs/TESTING_STRATEGY.md`](docs/TESTING_STRATEGY.md).
213
+
214
+ - **CPU smoke tests** (seconds, no TPU needed, `interpret=True`):
215
+
216
+ ```bash
217
+ pytest tests/test_gdn2_full_math_correctness.py -v
218
+ ```
219
+
220
+ - **Full suite** (requires TPU): multi-seed sweeps, finite-difference gradient checks,
221
+ isolated B3–B5 backward-stage tests, BF16 dtype-contract checks, `wy_eps` damping coverage,
222
+ alternative `KAGGLE_SMALL` blocking:
223
+
224
+ ```bash
225
+ pytest tests/extended/test_gdn2_deep_correctness.py -v
226
+ ```
227
+
228
+ ## Limitations
229
+
230
+ - **TPU-only fused kernels.** The Pallas path assumes TPU MXU tiling and `d_head = 128`;
231
+ other backends/dtypes automatically fall back to the pure-JAX reference (slower, correct).
232
+ - **Fused forward is currently slower than the pure-JAX WY forward** (~0.6×). If your workload
233
+ is inference-only, use `gdn2_forward` / `gdn2_chunked_wy_reference` until the hybrid
234
+ `JAX forward + Pallas backward` mode lands (planned for v0.3.0).
235
+ - `seq_len` must be divisible by `config.bt` (256 by default, 128 for `KAGGLE_SMALL`).
236
+ - `KernelConfig.bt` must equal `2 * config.bc`; vary `mb` for solver granularity.
237
+ - The pairwise decay computation (Kernel A / B4) currently uses a VPU-bound broadcast-reduce
238
+ pattern rather than an MXU matmul. An MXU-factorized alternative has been sketched as a
239
+ post-beta hypothesis but is not implemented in this release; see
240
+ [`ROADMAP.md`](ROADMAP.md) and [`KNOWN_LIMITATIONS.md`](KNOWN_LIMITATIONS.md).
241
+
242
+ ## Debugging
243
+
244
+ Set `GDN2_FWD_DIAG=1` to get per-stage reports of non-finite or suspiciously large
245
+ (`>1e6`) activations at every kernel boundary. Diagnostic only — never changes values.
246
+
247
+ ```bash
248
+ GDN2_FWD_DIAG=1 python your_training_script.py
249
+ ```
250
+
251
+ ## Project layout
252
+
253
+ ```javascript
254
+ atomic_ops/
255
+ ├── atomic_ops/ # the package
256
+ │ ├── configs.py # KernelConfig, presets, sanitize/validate helpers
257
+ │ ├── gdn2_fwd.py # forward kernels: A (scores), B (WY solve), C (recompute), D (scan)
258
+ │ ├── gdn2_bwd.py # backward kernels B1–B5
259
+ │ ├── gdn2_pipeline.py # custom_vjp trainable wrapper
260
+ │ ├── reference.py # token-serial + chunked-WY pure-JAX references
261
+ │ ├── fallback.py # auto-dispatch (TPU+d_head=128 -> Pallas, else reference)
262
+ │ └── utils.py # is_tpu_available, estimate_memory, get_recommended_config
263
+ ├── benchmarks/ # speed & memory benchmarks + raw results
264
+ ├── tests/ # CPU smoke tests
265
+ │ └── extended/ # full TPU correctness suite
266
+ ├── examples/ # minimal_usage.py + Flax training step
267
+ ├── docs/TESTING_STRATEGY.md # why the tests are built this way
268
+ └── .github/workflows/ # CI (tests, lint), publish to PyPI
269
+ ```
270
+
271
+ ## Contributing
272
+
273
+ See [CONTRIBUTING.md](CONTRIBUTING.md). In short: `pytest tests/test_gdn2_full_math_correctness.py`
274
+ must pass, Ruff must be green, kernel changes require the full TPU suite referenced in the PR.
275
+
276
+ ## License
277
+
278
+ MIT — see [LICENSE](LICENSE). Kernels ported from the NVlabs Gated DeltaNet-2 Triton reference.
279
+
280
+ ## Citation
281
+
282
+ ```bibtex
283
+ @software{atomic_ops,
284
+ author = {Omirbay, Akseleu},
285
+ title = {atomic_ops: Fused Gated DeltaNet-2 kernels for TPU v5e in JAX/Pallas},
286
+ url = {https://github.com/Akseleu-J/atomic_ops},
287
+ license = {MIT},
288
+ year = {2026}
289
+ }
290
+ ```
291
+
292
+ ## Support
293
+
294
+ If this package is useful in your research, consider giving it a ⭐ — it helps other researchers find it.
295
+ Bug reports and questions go to [Issues](https://github.com/Akseleu-J/atomic_ops/issues).
296
+ # atomic-ops
@@ -0,0 +1,263 @@
1
+ # atomic_ops
2
+
3
+ **Fused Gated DeltaNet-2 (GDN-2) kernels for TPU v5e, written in JAX/Pallas.**
4
+
5
+ [![License: MIT](https://img.shields.io/badge/License-MIT-yellow.svg)](https://opensource.org/licenses/MIT)
6
+ [![Python 3.10+](https://img.shields.io/badge/python-3.10%2B-blue.svg)](https://www.python.org/downloads/)
7
+ [![TPU v5e](https://img.shields.io/badge/TPU-v5e-orange.svg)](https://cloud.google.com/tpu/docs/v5e)
8
+
9
+ A from-scratch port of the [NVlabs Gated DeltaNet-2](https://github.com/NVlabs/GatedDeltaNet-2) Triton kernels
10
+ to `jax.experimental.pallas`, targeting **TPU v5e-8**. The backward pass is a single fused `custom_vjp`
11
+ that reuses forward residuals instead of recomputing them.
12
+
13
+ > **Headline numbers (measured, see [Benchmarks](#-benchmarks)):** on the full training shape
14
+ > (batch 8, seq 4096, 6 heads, d_head 128) the fused backward makes the training step
15
+ > **2.6× (FP32) / 3.4× (BF16) faster than the best pure-JAX WY baseline** and
16
+ > **27.2× (FP32) / 13.3× (BF16) faster than the widely used `associative_scan` baseline**.
17
+ > Across all measured shapes the gain vs `associative_scan` reaches **up to 38.8×** (FP32).
18
+ > The fused forward alone is currently **~1.6× slower** than the pure-JAX WY forward on TPU —
19
+ > training steps are backward-dominated, so end-to-end is still a clear win.
20
+ > A hybrid `JAX forward + Pallas backward` mode is planned (see [Limitations](#-limitations)).
21
+
22
+ ---
23
+
24
+ ## Features
25
+
26
+ - **Fully fused kernels** — forward (chunk scores → WY block solve → recompute → inter-chunk scan)
27
+ and backward (B1–B5 stage kernels), all in Pallas with TPU MXU tiling.
28
+ - **One-call training API** — `gdn2_forward_trainable` via `jax.custom_vjp`; no recomputation
29
+ of forward activations in the backward pass.
30
+ - **Automatic fallback** — on CPU/GPU or when `d_head != 128`, dispatches to a checkpointed
31
+ pure-JAX chunked-WY reference with identical damping semantics.
32
+ - **Numerical safety by default** — clipping + `nan_to_num` at every kernel boundary,
33
+ optional `wy_eps` Tikhonov-style damping of the WY solve, per-stage diagnostics
34
+ via the `GDN2_FWD_DIAG=1` environment flag.
35
+ - **Tunable configs** — `KernelConfig(bt, bc, mb, clip, wy_eps)` with presets tuned for
36
+ Kaggle TPU v5e (`KAGGLE_SMALL` / `KAGGLE_MEDIUM` / `KAGGLE_LARGE`) and an
37
+ `estimate_memory` / `get_recommended_config` helper.
38
+
39
+ ## Installation
40
+
41
+ ```bash
42
+ pip install atomic_ops
43
+
44
+ # or from source
45
+ pip install git+https://github.com/Akseleu-J/atomic_ops.git
46
+
47
+ # development
48
+ git clone https://github.com/Akseleu-J/atomic_ops.git
49
+ cd atomic_ops
50
+ pip install jax==0.11.1 jaxlib==0.11.1 libtpu==0.0.46 flax==0.12.9 optax==0.2.4
51
+ ```
52
+
53
+ Requires Python ≥ 3.10 and `jax>=0.4.20`. For TPU, install the matching `jaxlib`/`libtpu`
54
+ builds (see [JAX TPU installation](https://jax.readthedocs.io/en/latest/installation.html)).
55
+
56
+ ## Quick start
57
+
58
+ ```python
59
+ import jax.numpy as jnp
60
+ from atomic_ops import gdn2_forward_trainable
61
+
62
+ # (batch, seq_len, heads, d_head); d_head must be 128 on TPU
63
+ shape = (4, 2048, 6, 128)
64
+ q = jnp.ones(shape, dtype=jnp.float32)
65
+ k = jnp.ones(shape, dtype=jnp.float32)
66
+ v = jnp.ones(shape, dtype=jnp.float32)
67
+ w = jnp.ones(shape, dtype=jnp.float32)
68
+ b = jnp.ones(shape, dtype=jnp.float32)
69
+ g = jnp.ones(shape, dtype=jnp.float32) # log-decay gate, keep g <= 0
70
+ scale = shape[-1] ** -0.5
71
+
72
+ # Forward + backward (custom_vjp under the hood)
73
+ out, h_final = gdn2_forward_trainable(q, k, v, w, b, g, scale)
74
+
75
+ print(out.shape) # (4, 2048, 6, 128)
76
+ print(h_final.shape) # (4, 6, 128, 128)
77
+ ```
78
+
79
+ ### Runnable examples
80
+
81
+ Both examples work out of the box after `pip install atomic_ops`:
82
+
83
+ - [`examples/minimal_usage.py`](examples/minimal_usage.py) — 20-line minimal script:
84
+ build tensors, run `gdn2_forward_trainable`, print output shapes. Includes the
85
+ `is_tpu_available()` check so you can see which backend path was taken.
86
+ - [`examples/train_gdn2_layer.py`](examples/train_gdn2_layer.py) — a complete Flax +
87
+ Optax training step: a `GDN2Layer` module that projects `x` into `q,k,v,w,b,g`,
88
+ applies the gate via `-softplus(g)` (keeps `g <= 0`), auto-picks a config with
89
+ `get_recommended_config`, and runs one AdamW update step.
90
+
91
+ ```bash
92
+ python examples/minimal_usage.py # forward+backward, prints backend + shapes
93
+ python examples/train_gdn2_layer.py # full training step with Flax/Optax
94
+ ```
95
+
96
+ ## API
97
+
98
+ | Function | Description |
99
+ | --- | --- |
100
+ | `gdn2_forward_trainable(q, k, v, w, b, g, scale, h0=None, config=None)` | Forward **+** gradients. Dispatches to Pallas on TPU with `d_head=128`, otherwise to the pure-JAX reference. This is what you want for training. |
101
+ | `gdn2_forward(q, k, v, w, b, g, scale, h0=None, config=None)` | Forward only, same auto-dispatch. |
102
+ | `gdn2_pallas_forward_trainable(...)` | Forces the fused Pallas path (raises on non-TPU / `d_head != 128`). |
103
+ | `gdn2_pallas_forward(...)` / `gdn2_pallas_forward_with_residuals(...)` | Inference-oriented Pallas forward; the `_with_residuals` variant also returns the intermediate tensors the backward pass needs. |
104
+ | `gdn2_token_serial_reference(...)` | Ground-truth token-by-token scan. Slow, numerically exact — used as an independent check in tests. |
105
+ | `gdn2_chunked_wy_reference(...)` | Chunked-WY reference (the fallback path). |
106
+ | `KernelConfig(bt, bc, mb, clip, wy_eps)` | Kernel blocking / numerical-safety config. Constraints: `bt = 2*bc`, `bc % mb == 0`. |
107
+ | `KAGGLE_SMALL` / `KAGGLE_MEDIUM` / `KAGGLE_LARGE` | Presets for TPU v5e-8 (default: `KAGGLE_MEDIUM`). |
108
+ | `estimate_memory(...)` / `get_recommended_config(...)` | Rough per-chip HBM estimate and preset auto-selection. |
109
+ | `is_tpu_available()` | `True` if a TPU device is visible to JAX (used by the fallback dispatcher). |
110
+
111
+ All tensor arguments have shape `(batch, seq_len, heads, d_head)` with `seq_len % config.bt == 0`.
112
+ The initial recurrent state `h0` is optional, shape `(batch, heads, d_head, d_head)`, float32.
113
+
114
+ ## Benchmarks
115
+
116
+ Measured on **TPU v5e-8** with `jax==0.11.1`, `jaxlib==0.11.1`, `libtpu==0.0.46.1`.
117
+ Baselines: **OLD** = `jax.lax.associative_scan` GDN-2 (the formulation used in many research
118
+ codebases); **JAX_REF** = chunked WY recurrence in pure JAX. Mean steady-state over repeats, ms.
119
+
120
+ ### Full training step — batch 8, seq 4096, 6 heads, d_head 128
121
+
122
+ | Dtype | Stage | OLD (ms) | JAX_REF (ms) | **Pallas (ms)** | vs OLD | vs JAX_REF |
123
+ | --- | --- | --- | --- | --- | --- | --- |
124
+ | FP32 | fwd | 1067.24 | 63.24 | 102.08 | 10.5× | 0.62× |
125
+ | FP32 | bwd | 4767.24 | 458.99 | **175.33** | **27.2×** | **2.6×** |
126
+ | FP32 | fwd+bwd | 4762.70 | 460.03 | **175.20** | **27.2×** | **2.6×** |
127
+ | BF16 | fwd | 585.49 | 62.38 | 101.51 | 5.8× | 0.61× |
128
+ | BF16 | bwd | 2322.97 | 594.17 | **174.25** | **13.3×** | **3.4×** |
129
+ | BF16 | fwd+bwd | 2320.27 | 594.86 | **174.20** | **13.3×** | **3.4×** |
130
+
131
+ ### All measured shapes — full fwd+bwd cycle, speedup vs baselines
132
+
133
+ | Config (batch, seq_len) | Dtype | Pallas (ms) | vs OLD (scan) | vs JAX_REF |
134
+ | --- | --- | --- | --- | --- |
135
+ | small (B=1, L=1024) | FP32 | 5.92 | 24.3× | 3.7× |
136
+ | small (B=1, L=1024) | BF16 | 6.00 | 11.5× | 3.8× |
137
+ | medium (B=4, L=4096) | FP32 | 88.07 | 30.2× | 3.1× |
138
+ | medium (B=4, L=4096) | BF16 | 87.60 | 14.9× | 3.8× |
139
+ | **train shape (B=8, L=4096)** | FP32 | **175.20** | **27.2×** | **2.6×** |
140
+ | **train shape (B=8, L=4096)** | BF16 | **174.20** | **13.3×** | **3.4×** |
141
+ | KAGGLE_SMALL preset (B=4, L=2048) | FP32 | 30.86 | **38.8×** | 3.9× |
142
+ | KAGGLE_SMALL preset (B=4, L=2048) | BF16 | 30.66 | 18.7× | 3.9× |
143
+ | KAGGLE_LARGE preset (B=8, L=4096) | FP32 | 175.34 | 27.2× | 2.6× |
144
+ | KAGGLE_LARGE preset (B=8, L=4096) | BF16 | 174.12 | 13.3× | 3.4× |
145
+
146
+ Read it honestly: the fused **forward is ~1.6× slower than the pure-JAX WY forward** today;
147
+ the fused **backward is 2.6–3.9× faster** than JAX_REF, and since real training steps are
148
+ backward-dominated, the full cycle wins by **2.6–3.9×** over the fastest pure-JAX baseline and
149
+ by **11.5–38.8×** over the `associative_scan` baseline depending on shape and dtype.
150
+ Full per-stage tables for all configs and both dtypes:
151
+ [`benchmarks/raw/benchmark_speed_final_averaged.md`](benchmarks/raw/benchmark_speed_final_averaged.md).
152
+
153
+ ### Peak HBM — train shape (B=8, L=4096), forward+backward
154
+
155
+ | Dtype | OLD (MB) | JAX_REF (MB) | Pallas (MB) |
156
+ | --- | --- | --- | --- |
157
+ | FP32 | 1252.7 | 1250.0 | **1238.1** |
158
+ | BF16 | 915.5 | 915.5 | 915.5 |
159
+
160
+ Memory is on par with the pure-JAX reference (the backward reuses forward residuals instead of
161
+ recomputing them). Details:
162
+ [`benchmarks/raw/benchmark_memory_final_averaged.md`](benchmarks/raw/benchmark_memory_final_averaged.md).
163
+
164
+ ### Reproduce
165
+
166
+ ```bash
167
+ python benchmarks/run_speed_benchmark.py # writes JSON + markdown tables
168
+ python benchmarks/run_memory_benchmark.py # fork-isolated peak HBM measurement
169
+ ```
170
+
171
+ Every timing run is correctness-gated before measurement (Pallas output must match the reference
172
+ within tolerance, otherwise the row is rejected).
173
+
174
+ ## Correctness & testing
175
+
176
+ The test suite is deliberately layered — comparing implementations that share the same algebraic
177
+ derivation can hide derivation bugs, so the strongest checks are derivation-independent
178
+ (finite-difference gradients, token-serial scan). Full rationale:
179
+ [`docs/TESTING_STRATEGY.md`](docs/TESTING_STRATEGY.md).
180
+
181
+ - **CPU smoke tests** (seconds, no TPU needed, `interpret=True`):
182
+
183
+ ```bash
184
+ pytest tests/test_gdn2_full_math_correctness.py -v
185
+ ```
186
+
187
+ - **Full suite** (requires TPU): multi-seed sweeps, finite-difference gradient checks,
188
+ isolated B3–B5 backward-stage tests, BF16 dtype-contract checks, `wy_eps` damping coverage,
189
+ alternative `KAGGLE_SMALL` blocking:
190
+
191
+ ```bash
192
+ pytest tests/extended/test_gdn2_deep_correctness.py -v
193
+ ```
194
+
195
+ ## Limitations
196
+
197
+ - **TPU-only fused kernels.** The Pallas path assumes TPU MXU tiling and `d_head = 128`;
198
+ other backends/dtypes automatically fall back to the pure-JAX reference (slower, correct).
199
+ - **Fused forward is currently slower than the pure-JAX WY forward** (~0.6×). If your workload
200
+ is inference-only, use `gdn2_forward` / `gdn2_chunked_wy_reference` until the hybrid
201
+ `JAX forward + Pallas backward` mode lands (planned for v0.3.0).
202
+ - `seq_len` must be divisible by `config.bt` (256 by default, 128 for `KAGGLE_SMALL`).
203
+ - `KernelConfig.bt` must equal `2 * config.bc`; vary `mb` for solver granularity.
204
+ - The pairwise decay computation (Kernel A / B4) currently uses a VPU-bound broadcast-reduce
205
+ pattern rather than an MXU matmul. An MXU-factorized alternative has been sketched as a
206
+ post-beta hypothesis but is not implemented in this release; see
207
+ [`ROADMAP.md`](ROADMAP.md) and [`KNOWN_LIMITATIONS.md`](KNOWN_LIMITATIONS.md).
208
+
209
+ ## Debugging
210
+
211
+ Set `GDN2_FWD_DIAG=1` to get per-stage reports of non-finite or suspiciously large
212
+ (`>1e6`) activations at every kernel boundary. Diagnostic only — never changes values.
213
+
214
+ ```bash
215
+ GDN2_FWD_DIAG=1 python your_training_script.py
216
+ ```
217
+
218
+ ## Project layout
219
+
220
+ ```javascript
221
+ atomic_ops/
222
+ ├── atomic_ops/ # the package
223
+ │ ├── configs.py # KernelConfig, presets, sanitize/validate helpers
224
+ │ ├── gdn2_fwd.py # forward kernels: A (scores), B (WY solve), C (recompute), D (scan)
225
+ │ ├── gdn2_bwd.py # backward kernels B1–B5
226
+ │ ├── gdn2_pipeline.py # custom_vjp trainable wrapper
227
+ │ ├── reference.py # token-serial + chunked-WY pure-JAX references
228
+ │ ├── fallback.py # auto-dispatch (TPU+d_head=128 -> Pallas, else reference)
229
+ │ └── utils.py # is_tpu_available, estimate_memory, get_recommended_config
230
+ ├── benchmarks/ # speed & memory benchmarks + raw results
231
+ ├── tests/ # CPU smoke tests
232
+ │ └── extended/ # full TPU correctness suite
233
+ ├── examples/ # minimal_usage.py + Flax training step
234
+ ├── docs/TESTING_STRATEGY.md # why the tests are built this way
235
+ └── .github/workflows/ # CI (tests, lint), publish to PyPI
236
+ ```
237
+
238
+ ## Contributing
239
+
240
+ See [CONTRIBUTING.md](CONTRIBUTING.md). In short: `pytest tests/test_gdn2_full_math_correctness.py`
241
+ must pass, Ruff must be green, kernel changes require the full TPU suite referenced in the PR.
242
+
243
+ ## License
244
+
245
+ MIT — see [LICENSE](LICENSE). Kernels ported from the NVlabs Gated DeltaNet-2 Triton reference.
246
+
247
+ ## Citation
248
+
249
+ ```bibtex
250
+ @software{atomic_ops,
251
+ author = {Omirbay, Akseleu},
252
+ title = {atomic_ops: Fused Gated DeltaNet-2 kernels for TPU v5e in JAX/Pallas},
253
+ url = {https://github.com/Akseleu-J/atomic_ops},
254
+ license = {MIT},
255
+ year = {2026}
256
+ }
257
+ ```
258
+
259
+ ## Support
260
+
261
+ If this package is useful in your research, consider giving it a ⭐ — it helps other researchers find it.
262
+ Bug reports and questions go to [Issues](https://github.com/Akseleu-J/atomic_ops/issues).
263
+ # atomic-ops
@@ -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
+ ]