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.
- atomic_ops-0.1.0/LICENSE +21 -0
- atomic_ops-0.1.0/MANIFEST.in +3 -0
- atomic_ops-0.1.0/PKG-INFO +296 -0
- atomic_ops-0.1.0/README.md +263 -0
- atomic_ops-0.1.0/atomic_ops/__init__.py +34 -0
- atomic_ops-0.1.0/atomic_ops/configs.py +136 -0
- atomic_ops-0.1.0/atomic_ops/fallback.py +22 -0
- atomic_ops-0.1.0/atomic_ops/gdn2_bwd.py +342 -0
- atomic_ops-0.1.0/atomic_ops/gdn2_fwd.py +395 -0
- atomic_ops-0.1.0/atomic_ops/gdn2_pipeline.py +152 -0
- atomic_ops-0.1.0/atomic_ops/py.typed +1 -0
- atomic_ops-0.1.0/atomic_ops/reference.py +159 -0
- atomic_ops-0.1.0/atomic_ops/utils.py +35 -0
- atomic_ops-0.1.0/atomic_ops.egg-info/PKG-INFO +296 -0
- atomic_ops-0.1.0/atomic_ops.egg-info/SOURCES.txt +24 -0
- atomic_ops-0.1.0/atomic_ops.egg-info/dependency_links.txt +1 -0
- atomic_ops-0.1.0/atomic_ops.egg-info/requires.txt +13 -0
- atomic_ops-0.1.0/atomic_ops.egg-info/top_level.txt +1 -0
- atomic_ops-0.1.0/pyproject.toml +90 -0
- atomic_ops-0.1.0/setup.cfg +4 -0
- atomic_ops-0.1.0/tests/test_clip_config_plumbing.py +240 -0
- atomic_ops-0.1.0/tests/test_configs.py +50 -0
- atomic_ops-0.1.0/tests/test_gdn2_full_math_correctness.py +309 -0
- atomic_ops-0.1.0/tests/test_imports.py +15 -0
- atomic_ops-0.1.0/tests/test_pallas.py +41 -0
- atomic_ops-0.1.0/tests/test_reference.py +20 -0
atomic_ops-0.1.0/LICENSE
ADDED
|
@@ -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,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
|
+
[](https://opensource.org/licenses/MIT)
|
|
39
|
+
[](https://www.python.org/downloads/)
|
|
40
|
+
[](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
|
+
[](https://opensource.org/licenses/MIT)
|
|
6
|
+
[](https://www.python.org/downloads/)
|
|
7
|
+
[](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
|
+
]
|