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 +34 -0
- atomic_ops/configs.py +136 -0
- atomic_ops/fallback.py +22 -0
- atomic_ops/gdn2_bwd.py +342 -0
- atomic_ops/gdn2_fwd.py +395 -0
- atomic_ops/gdn2_pipeline.py +152 -0
- atomic_ops/py.typed +1 -0
- atomic_ops/reference.py +159 -0
- atomic_ops/utils.py +35 -0
- atomic_ops-0.1.0.dist-info/METADATA +296 -0
- atomic_ops-0.1.0.dist-info/RECORD +14 -0
- atomic_ops-0.1.0.dist-info/WHEEL +5 -0
- atomic_ops-0.1.0.dist-info/licenses/LICENSE +21 -0
- atomic_ops-0.1.0.dist-info/top_level.txt +1 -0
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
|