fast-attn-kernels 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,4 @@
1
+ Metadata-Version: 2.4
2
+ Name: fast-attn-kernels
3
+ Version: 0.1.0
4
+ Requires-Dist: torch
@@ -0,0 +1 @@
1
+ from .linear_bind import LinearAttention, linear_attention
@@ -0,0 +1,336 @@
1
+ #include <stdlib.h>
2
+ #include <string.h>
3
+ #include <math.h>
4
+
5
+ #ifdef _OPENMP
6
+ #include <omp.h>
7
+ #endif
8
+
9
+ static inline float phi(float x) { return x > 0.0f ? x + 1.0f : expf(x); }
10
+
11
+ static inline float dphi(float x, float phix) { return x > 0.0f ? 1.0f : phix; }
12
+
13
+ /* ============================================================
14
+ Forward pass: chunked causal linear attention
15
+ ============================================================ */
16
+ void linear_attention_forward(
17
+ const float* q, const float* k, const float* v,
18
+ float* out, float* kv_states, float* z_states,
19
+ long B, long H, long N, long D, long E,
20
+ long chunk_size, long num_chunks)
21
+ {
22
+ long BH = B * H;
23
+ long qk_bh_stride = N * D;
24
+ long v_bh_stride = N * E;
25
+ long kv_bh_stride = num_chunks * D * E;
26
+ long z_bh_stride = num_chunks * D;
27
+
28
+ #pragma omp parallel for schedule(dynamic)
29
+ for (long bh = 0; bh < BH; ++bh) {
30
+ const float* Qb = q + bh * qk_bh_stride;
31
+ const float* Kb = k + bh * qk_bh_stride;
32
+ const float* Vb = v + bh * v_bh_stride;
33
+ float* Ob = out + bh * v_bh_stride;
34
+ float* KVb = kv_states + bh * kv_bh_stride;
35
+ float* Zb = z_states + bh * z_bh_stride;
36
+
37
+ float* S = calloc((size_t)(D * E), sizeof(float));
38
+ float* zstate = calloc((size_t)D, sizeof(float));
39
+ float* qphi = malloc((size_t)(chunk_size * D) * sizeof(float));
40
+ float* kphi = malloc((size_t)(chunk_size * D) * sizeof(float));
41
+ float* L = malloc((size_t)(chunk_size * chunk_size) * sizeof(float));
42
+ float* out_row = malloc((size_t)E * sizeof(float));
43
+
44
+ for (long c = 0; c < num_chunks; ++c) {
45
+ long c_start = c * chunk_size;
46
+ long rem = N - c_start;
47
+ long C = chunk_size < rem ? chunk_size : rem;
48
+
49
+ memcpy(KVb + c * D * E, S, (size_t)(D * E) * sizeof(float));
50
+ memcpy(Zb + c * D, zstate, (size_t)D * sizeof(float));
51
+
52
+ for (long i = 0; i < C; ++i) {
53
+ const float* qi = Qb + (c_start + i) * D;
54
+ const float* ki = Kb + (c_start + i) * D;
55
+ for (long d = 0; d < D; ++d) {
56
+ qphi[i * D + d] = phi(qi[d]);
57
+ kphi[i * D + d] = phi(ki[d]);
58
+ }
59
+ }
60
+
61
+ for (long i = 0; i < C; ++i) {
62
+ const float* qi = &qphi[i * D];
63
+ for (long j = 0; j <= i; ++j) {
64
+ const float* kj = &kphi[j * D];
65
+ float acc = 0.0f;
66
+ for (long d = 0; d < D; ++d) acc += qi[d] * kj[d];
67
+ L[i * chunk_size + j] = acc;
68
+ }
69
+ }
70
+
71
+ for (long i = 0; i < C; ++i) {
72
+ const float* qi = &qphi[i * D];
73
+ for (long e = 0; e < E; ++e) out_row[e] = 0.0f;
74
+
75
+ float z_intra = 0.0f;
76
+ for (long j = 0; j <= i; ++j) {
77
+ float lij = L[i * chunk_size + j];
78
+ z_intra += lij;
79
+ const float* vj = Vb + (c_start + j) * E;
80
+ for (long e = 0; e < E; ++e) out_row[e] += lij * vj[e];
81
+ }
82
+
83
+ float z_inter = 0.0f;
84
+ for (long d = 0; d < D; ++d) z_inter += qi[d] * zstate[d];
85
+ for (long e = 0; e < E; ++e) {
86
+ float acc = 0.0f;
87
+ for (long d = 0; d < D; ++d) acc += qi[d] * S[d * E + e];
88
+ out_row[e] += acc;
89
+ }
90
+
91
+ float z = z_intra + z_inter;
92
+ if (z < 1e-6f) z = 1e-6f;
93
+
94
+ float* orow = Ob + (c_start + i) * E;
95
+ for (long e = 0; e < E; ++e) orow[e] = out_row[e] / z;
96
+ }
97
+
98
+ for (long i = 0; i < C; ++i) {
99
+ const float* ki = &kphi[i * D];
100
+ const float* vi = Vb + (c_start + i) * E;
101
+ for (long d = 0; d < D; ++d) {
102
+ zstate[d] += ki[d];
103
+ float kd = ki[d];
104
+ float* Srow = &S[d * E];
105
+ for (long e = 0; e < E; ++e) Srow[e] += kd * vi[e];
106
+ }
107
+ }
108
+ }
109
+
110
+ free(S);
111
+ free(zstate);
112
+ free(qphi);
113
+ free(kphi);
114
+ free(L);
115
+ free(out_row);
116
+ }
117
+ }
118
+
119
+ /* ============================================================
120
+ Backward pass: recompute-based, chunked, reverse order
121
+ ============================================================ */
122
+ void linear_attention_backward(
123
+ const float* q, const float* k, const float* v, const float* dout,
124
+ const float* kv_states, const float* z_states,
125
+ float* dq, float* dk, float* dv,
126
+ long B, long H, long N, long D, long E,
127
+ long chunk_size, long num_chunks)
128
+ {
129
+ long BH = B * H;
130
+ long qk_bh_stride = N * D;
131
+ long v_bh_stride = N * E;
132
+ long kv_bh_stride = num_chunks * D * E;
133
+ long z_bh_stride = num_chunks * D;
134
+
135
+ #pragma omp parallel for schedule(dynamic)
136
+ for (long bh = 0; bh < BH; ++bh) {
137
+ const float* Qb = q + bh * qk_bh_stride;
138
+ const float* Kb = k + bh * qk_bh_stride;
139
+ const float* Vb = v + bh * v_bh_stride;
140
+ const float* DOb = dout + bh * v_bh_stride;
141
+ const float* KVb = kv_states + bh * kv_bh_stride;
142
+ const float* Zb = z_states + bh * z_bh_stride;
143
+ float* DQb = dq + bh * qk_bh_stride;
144
+ float* DKb = dk + bh * qk_bh_stride;
145
+ float* DVb = dv + bh * v_bh_stride;
146
+
147
+ float* dS = calloc((size_t)(D * E), sizeof(float));
148
+ float* dzstate = calloc((size_t)D, sizeof(float));
149
+ float* qraw = malloc((size_t)(chunk_size * D) * sizeof(float));
150
+ float* kraw = malloc((size_t)(chunk_size * D) * sizeof(float));
151
+ float* qphi = malloc((size_t)(chunk_size * D) * sizeof(float));
152
+ float* kphi = malloc((size_t)(chunk_size * D) * sizeof(float));
153
+ float* L = malloc((size_t)(chunk_size * chunk_size) * sizeof(float));
154
+ float* dL = malloc((size_t)(chunk_size * chunk_size) * sizeof(float));
155
+ float* dnum = malloc((size_t)(chunk_size * E) * sizeof(float));
156
+ float* dden = malloc((size_t)chunk_size * sizeof(float));
157
+ float* dq_total = malloc((size_t)(chunk_size * D) * sizeof(float));
158
+ float* dk_total = malloc((size_t)(chunk_size * D) * sizeof(float));
159
+ float* dv_total = malloc((size_t)(chunk_size * E) * sizeof(float));
160
+ float* dS_local = malloc((size_t)(D * E) * sizeof(float));
161
+ float* dz_local = malloc((size_t)D * sizeof(float));
162
+ float* num_row = malloc((size_t)E * sizeof(float));
163
+
164
+ for (long c = num_chunks - 1; c >= 0; --c) {
165
+ long c_start = c * chunk_size;
166
+ long rem = N - c_start;
167
+ long C = chunk_size < rem ? chunk_size : rem;
168
+
169
+ const float* Sb = KVb + c * D * E;
170
+ const float* zb = Zb + c * D;
171
+
172
+ /* load + apply feature map */
173
+ for (long i = 0; i < C; ++i) {
174
+ const float* qi = Qb + (c_start + i) * D;
175
+ const float* ki = Kb + (c_start + i) * D;
176
+ for (long d = 0; d < D; ++d) {
177
+ float qr = qi[d];
178
+ float kr = ki[d];
179
+ qraw[i * D + d] = qr;
180
+ kraw[i * D + d] = kr;
181
+ qphi[i * D + d] = phi(qr);
182
+ kphi[i * D + d] = phi(kr);
183
+ }
184
+ }
185
+
186
+ /* recompute L = tril(q @ k^T) */
187
+ for (long i = 0; i < C; ++i) {
188
+ for (long j = 0; j <= i; ++j) {
189
+ float acc = 0.0f;
190
+ for (long d = 0; d < D; ++d) acc += qphi[i * D + d] * kphi[j * D + d];
191
+ L[i * chunk_size + j] = acc;
192
+ }
193
+ }
194
+
195
+ /* recompute forward output -> derive dnum, dden */
196
+ for (long i = 0; i < C; ++i) {
197
+ const float* qi = &qphi[i * D];
198
+ for (long e = 0; e < E; ++e) num_row[e] = 0.0f;
199
+
200
+ float z_intra = 0.0f;
201
+ for (long j = 0; j <= i; ++j) {
202
+ float lij = L[i * chunk_size + j];
203
+ z_intra += lij;
204
+ const float* vj = Vb + (c_start + j) * E;
205
+ for (long e = 0; e < E; ++e) num_row[e] += lij * vj[e];
206
+ }
207
+
208
+ float z_inter = 0.0f;
209
+ for (long d = 0; d < D; ++d) z_inter += qi[d] * zb[d];
210
+ for (long e = 0; e < E; ++e) {
211
+ float acc = 0.0f;
212
+ for (long d = 0; d < D; ++d) acc += qi[d] * Sb[d * E + e];
213
+ num_row[e] += acc;
214
+ }
215
+
216
+ float den = z_intra + z_inter;
217
+ if (den < 1e-6f) den = 1e-6f;
218
+
219
+ const float* dorow = DOb + (c_start + i) * E;
220
+ float dot_do_out = 0.0f;
221
+ for (long e = 0; e < E; ++e) {
222
+ float out_e = num_row[e] / den;
223
+ dnum[i * E + e] = dorow[e] / den;
224
+ dot_do_out += dorow[e] * out_e;
225
+ }
226
+ dden[i] = -dot_do_out / den;
227
+ }
228
+
229
+ memset(dq_total, 0, (size_t)(C * D) * sizeof(float));
230
+ memset(dk_total, 0, (size_t)(C * D) * sizeof(float));
231
+ memset(dv_total, 0, (size_t)(C * E) * sizeof(float));
232
+ memset(dS_local, 0, (size_t)(D * E) * sizeof(float));
233
+ memset(dz_local, 0, (size_t)D * sizeof(float));
234
+
235
+ /* grads from the inter-chunk term (q @ S_before, q . z_before) */
236
+ for (long i = 0; i < C; ++i) {
237
+ const float* qi = &qphi[i * D];
238
+
239
+ for (long d = 0; d < D; ++d) {
240
+ float acc = 0.0f;
241
+ for (long e = 0; e < E; ++e) acc += dnum[i * E + e] * Sb[d * E + e];
242
+ acc += dden[i] * zb[d];
243
+ dq_total[i * D + d] += acc;
244
+ }
245
+
246
+ for (long d = 0; d < D; ++d) {
247
+ float qd = qi[d];
248
+ for (long e = 0; e < E; ++e) dS_local[d * E + e] += qd * dnum[i * E + e];
249
+ }
250
+
251
+ for (long d = 0; d < D; ++d) dz_local[d] += qi[d] * dden[i];
252
+ }
253
+
254
+ /* grads from the intra-chunk term (causal local attention) */
255
+ for (long i = 0; i < C; ++i) {
256
+ for (long j = 0; j <= i; ++j) {
257
+ float dl = 0.0f;
258
+ const float* vj = Vb + (c_start + j) * E;
259
+ for (long e = 0; e < E; ++e) dl += dnum[i * E + e] * vj[e];
260
+ dl += dden[i];
261
+ dL[i * chunk_size + j] = dl;
262
+ }
263
+ }
264
+
265
+ for (long i = 0; i < C; ++i) {
266
+ const float* qi = &qphi[i * D];
267
+ for (long j = 0; j <= i; ++j) {
268
+ float dl = dL[i * chunk_size + j];
269
+ const float* kj = &kphi[j * D];
270
+ for (long d = 0; d < D; ++d) {
271
+ dq_total[i * D + d] += dl * kj[d];
272
+ dk_total[j * D + d] += dl * qi[d];
273
+ }
274
+ }
275
+ }
276
+
277
+ for (long i = 0; i < C; ++i) {
278
+ for (long j = 0; j <= i; ++j) {
279
+ float lij = L[i * chunk_size + j];
280
+ for (long e = 0; e < E; ++e) dv_total[j * E + e] += lij * dnum[i * E + e];
281
+ }
282
+ }
283
+
284
+ /* grads from the recurrent state update: S += k^T v, zstate += sum(k) */
285
+ for (long i = 0; i < C; ++i) {
286
+ const float* vi = Vb + (c_start + i) * E;
287
+ const float* ki = &kphi[i * D];
288
+
289
+ for (long d = 0; d < D; ++d) {
290
+ float acc = 0.0f;
291
+ for (long e = 0; e < E; ++e) acc += vi[e] * dS[d * E + e];
292
+ dk_total[i * D + d] += acc + dzstate[d];
293
+ }
294
+
295
+ for (long e = 0; e < E; ++e) {
296
+ float acc = 0.0f;
297
+ for (long d = 0; d < D; ++d) acc += ki[d] * dS[d * E + e];
298
+ dv_total[i * E + e] += acc;
299
+ }
300
+ }
301
+
302
+ /* propagate accumulated state-gradient to the previous chunk */
303
+ for (long idx = 0; idx < D * E; ++idx) dS[idx] += dS_local[idx];
304
+ for (long d = 0; d < D; ++d) dzstate[d] += dz_local[d];
305
+
306
+ /* apply feature-map derivative, write results */
307
+ for (long i = 0; i < C; ++i) {
308
+ float* dqrow = DQb + (c_start + i) * D;
309
+ float* dkrow = DKb + (c_start + i) * D;
310
+ for (long d = 0; d < D; ++d) {
311
+ dqrow[d] = dq_total[i * D + d] * dphi(qraw[i * D + d], qphi[i * D + d]);
312
+ dkrow[d] = dk_total[i * D + d] * dphi(kraw[i * D + d], kphi[i * D + d]);
313
+ }
314
+ float* dvrow = DVb + (c_start + i) * E;
315
+ for (long e = 0; e < E; ++e) dvrow[e] = dv_total[i * E + e];
316
+ }
317
+ }
318
+
319
+ free(dS);
320
+ free(dzstate);
321
+ free(qraw);
322
+ free(kraw);
323
+ free(qphi);
324
+ free(kphi);
325
+ free(L);
326
+ free(dL);
327
+ free(dnum);
328
+ free(dden);
329
+ free(dq_total);
330
+ free(dk_total);
331
+ free(dv_total);
332
+ free(dS_local);
333
+ free(dz_local);
334
+ free(num_row);
335
+ }
336
+ }
@@ -0,0 +1,414 @@
1
+ import os
2
+ import sys
3
+ import ctypes
4
+ import platform
5
+ import torch
6
+
7
+ _THIS_DIR = os.path.dirname(os.path.abspath(__file__))
8
+
9
+
10
+ def _find_lib():
11
+ # possible names depending on OS
12
+ if platform.system() == "Windows": names = ["liblinear_attn.dll", "linear_attn.dll"]
13
+ elif platform.system() == "Darwin": names = ["liblinear_attn.dylib"]
14
+ else: names = ["liblinear_attn.so"]
15
+
16
+ # possible locations (depends on CMake generator: single-config vs multi-config)
17
+ search_dirs = [
18
+ os.path.join(_THIS_DIR, "..", "build"),
19
+ os.path.join(_THIS_DIR, "..", "build", "Debug"),
20
+ os.path.join(_THIS_DIR, "..", "build", "Release"),
21
+ os.path.join(_THIS_DIR, "build"),
22
+ _THIS_DIR,
23
+ ]
24
+
25
+ for d in search_dirs:
26
+ for n in names:
27
+ p = os.path.join(d, n)
28
+ if os.path.isfile(p): return p, d
29
+
30
+ raise FileNotFoundError(f"linear_attn shared library nahi mili. Search kiya: {search_dirs} me {names}")
31
+
32
+
33
+ _LIB_PATH, _LIB_DIR = _find_lib()
34
+
35
+ # Windows pe MinGW-built DLL ki runtime dependencies (libgomp, libwinpthread)
36
+ # resolve karne ke liye us folder ko DLL search path me add karo
37
+ if platform.system() == "Windows":
38
+ if hasattr(os, "add_dll_directory"): os.add_dll_directory(_LIB_DIR)
39
+ os.environ["PATH"] = _LIB_DIR + os.pathsep + os.environ.get("PATH", "")
40
+
41
+ _lib = ctypes.CDLL(_LIB_PATH)
42
+
43
+ c_long = ctypes.c_long
44
+ c_float_p = ctypes.POINTER(ctypes.c_float)
45
+
46
+ _lib.linear_attention_forward.argtypes = [
47
+ c_float_p, c_float_p, c_float_p,
48
+ c_float_p, c_float_p, c_float_p,
49
+ c_long, c_long, c_long, c_long, c_long, c_long, c_long,
50
+ ]
51
+ _lib.linear_attention_forward.restype = None
52
+
53
+ _lib.linear_attention_backward.argtypes = [
54
+ c_float_p, c_float_p, c_float_p, c_float_p,
55
+ c_float_p, c_float_p,
56
+ c_float_p, c_float_p, c_float_p,
57
+ c_long, c_long, c_long, c_long, c_long, c_long, c_long,
58
+ ]
59
+ _lib.linear_attention_backward.restype = None
60
+
61
+
62
+ def _ptr(t: torch.Tensor):
63
+ return ctypes.cast(ctypes.c_void_p(t.data_ptr()), c_float_p)
64
+
65
+
66
+ class _LinearAttnCFn(torch.autograd.Function):
67
+ @staticmethod
68
+ def forward(ctx, q, k, v, chunk_size):
69
+ assert not q.is_cuda, "yeh C backend sirf CPU tensors ke liye hai"
70
+ q = q.contiguous().float()
71
+ k = k.contiguous().float()
72
+ v = v.contiguous().float()
73
+
74
+ B, H, N, D = q.shape
75
+ E = v.shape[-1]
76
+ num_chunks = (N + chunk_size - 1) // chunk_size
77
+
78
+ out = torch.empty((B, H, N, E), dtype=torch.float32)
79
+ kv_states = torch.zeros((B, H, num_chunks, D, E), dtype=torch.float32)
80
+ z_states = torch.zeros((B, H, num_chunks, D), dtype=torch.float32)
81
+
82
+ _lib.linear_attention_forward(
83
+ _ptr(q), _ptr(k), _ptr(v),
84
+ _ptr(out), _ptr(kv_states), _ptr(z_states),
85
+ c_long(B), c_long(H), c_long(N), c_long(D), c_long(E),
86
+ c_long(chunk_size), c_long(num_chunks),
87
+ )
88
+
89
+ ctx.save_for_backward(q, k, v, kv_states, z_states)
90
+ ctx.chunk_size = chunk_size
91
+ return out
92
+
93
+ @staticmethod
94
+ def backward(ctx, dout):
95
+ q, k, v, kv_states, z_states = ctx.saved_tensors
96
+ dout = dout.contiguous().float()
97
+
98
+ B, H, N, D = q.shape
99
+ E = v.shape[-1]
100
+ num_chunks = kv_states.shape[2]
101
+
102
+ dq = torch.zeros_like(q)
103
+ dk = torch.zeros_like(k)
104
+ dv = torch.zeros_like(v)
105
+
106
+ _lib.linear_attention_backward(
107
+ _ptr(q), _ptr(k), _ptr(v), _ptr(dout),
108
+ _ptr(kv_states), _ptr(z_states),
109
+ _ptr(dq), _ptr(dk), _ptr(dv),
110
+ c_long(B), c_long(H), c_long(N), c_long(D), c_long(E),
111
+ c_long(ctx.chunk_size), c_long(num_chunks),
112
+ )
113
+
114
+ return dq, dk, dv, None
115
+
116
+
117
+ import torch.nn as nn
118
+
119
+ try:
120
+ import triton
121
+ import triton.language as tl
122
+ _HAS_TRITON = True
123
+
124
+ # ============================================================
125
+ # Forward kernel: chunked causal linear attention
126
+ # ============================================================
127
+ @triton.jit
128
+ def _la_fwd_kernel(
129
+ Q, K, V, Out, KV_STATES, Z_STATES,
130
+ stride_qb, stride_qh, stride_qm, stride_qd,
131
+ stride_kb, stride_kh, stride_kn, stride_kd,
132
+ stride_vb, stride_vh, stride_vn, stride_ve,
133
+ stride_ob, stride_oh, stride_om, stride_oe,
134
+ stride_kvb, stride_kvh, stride_kvc, stride_kvd, stride_kve,
135
+ stride_zb, stride_zh, stride_zc, stride_zd,
136
+ seq_len, head_dim, v_dim, num_chunks,
137
+ BLOCK_C: tl.constexpr, BLOCK_D: tl.constexpr, BLOCK_E: tl.constexpr,
138
+ ):
139
+ pid_bh = tl.program_id(0)
140
+ Q += pid_bh * stride_qh; K += pid_bh * stride_kh
141
+ V += pid_bh * stride_vh; Out += pid_bh * stride_oh
142
+ KV_STATES += pid_bh * stride_kvh
143
+ Z_STATES += pid_bh * stride_zh
144
+
145
+ offs_d = tl.arange(0, BLOCK_D)
146
+ offs_e = tl.arange(0, BLOCK_E)
147
+ offs_c = tl.arange(0, BLOCK_C)
148
+
149
+ S = tl.zeros((BLOCK_D, BLOCK_E), dtype=tl.float32)
150
+ z_state = tl.zeros((BLOCK_D,), dtype=tl.float32)
151
+
152
+ for c in range(0, num_chunks):
153
+ c_start = c * BLOCK_C
154
+ offs_m = c_start + offs_c
155
+ row_mask = offs_m < seq_len
156
+
157
+ q_ptrs = Q + offs_m[:, None] * stride_qm + offs_d[None, :] * stride_qd
158
+ k_ptrs = K + offs_m[:, None] * stride_kn + offs_d[None, :] * stride_kd
159
+ v_ptrs = V + offs_m[:, None] * stride_vn + offs_e[None, :] * stride_ve
160
+ qk_mask = row_mask[:, None] & (offs_d[None, :] < head_dim)
161
+ v_mask = row_mask[:, None] & (offs_e[None, :] < v_dim)
162
+
163
+ q = tl.load(q_ptrs, mask=qk_mask, other=0.0).to(tl.float32)
164
+ k = tl.load(k_ptrs, mask=qk_mask, other=0.0).to(tl.float32)
165
+ v = tl.load(v_ptrs, mask=v_mask, other=0.0).to(tl.float32)
166
+
167
+ q = tl.where(q > 0, q + 1.0, tl.exp(q))
168
+ k = tl.where(k > 0, k + 1.0, tl.exp(k))
169
+
170
+ # store the state as it was BEFORE this chunk (needed for backward)
171
+ kv_ptrs = KV_STATES + c * stride_kvc + offs_d[:, None] * stride_kvd + offs_e[None, :] * stride_kve
172
+ z_ptrs = Z_STATES + c * stride_zc + offs_d * stride_zd
173
+ tl.store(kv_ptrs, S.to(tl.float32))
174
+ tl.store(z_ptrs, z_state.to(tl.float32))
175
+
176
+ local_scores = tl.dot(q, tl.trans(k))
177
+ causal_mask = offs_c[:, None] >= offs_c[None, :]
178
+ local_scores = tl.where(causal_mask, local_scores, 0.0)
179
+
180
+ out_intra = tl.dot(local_scores.to(tl.float32), v)
181
+ z_intra = tl.sum(local_scores, axis=1)
182
+
183
+ out_inter = tl.dot(q, S)
184
+ z_inter = tl.sum(q * z_state[None, :], axis=1)
185
+
186
+ out = out_intra + out_inter
187
+ z = z_intra + z_inter
188
+ z = tl.maximum(z, 1e-6)
189
+ out = out / z[:, None]
190
+
191
+ out_ptrs = Out + offs_m[:, None] * stride_om + offs_e[None, :] * stride_oe
192
+ tl.store(out_ptrs, out.to(Out.dtype.element_ty), mask=v_mask)
193
+
194
+ S += tl.dot(tl.trans(k), v)
195
+ z_state += tl.sum(k, axis=0)
196
+
197
+
198
+ # ============================================================
199
+ # Backward kernel: recompute-based (no N x N matrices stored)
200
+ # ============================================================
201
+ @triton.jit
202
+ def _la_bwd_kernel(
203
+ Q, K, V, DOut, DQ, DK, DV, KV_STATES, Z_STATES,
204
+ stride_qb, stride_qh, stride_qm, stride_qd,
205
+ stride_kb, stride_kh, stride_kn, stride_kd,
206
+ stride_vb, stride_vh, stride_vn, stride_ve,
207
+ stride_ob, stride_oh, stride_om, stride_oe,
208
+ stride_kvb, stride_kvh, stride_kvc, stride_kvd, stride_kve,
209
+ stride_zb, stride_zh, stride_zc, stride_zd,
210
+ seq_len, head_dim, v_dim, num_chunks,
211
+ BLOCK_C: tl.constexpr, BLOCK_D: tl.constexpr, BLOCK_E: tl.constexpr,
212
+ ):
213
+ pid_bh = tl.program_id(0)
214
+ Q += pid_bh * stride_qh; K += pid_bh * stride_kh
215
+ V += pid_bh * stride_vh; DOut += pid_bh * stride_oh
216
+ DQ += pid_bh * stride_qh; DK += pid_bh * stride_kh; DV += pid_bh * stride_vh
217
+ KV_STATES += pid_bh * stride_kvh
218
+ Z_STATES += pid_bh * stride_zh
219
+
220
+ offs_d = tl.arange(0, BLOCK_D)
221
+ offs_e = tl.arange(0, BLOCK_E)
222
+ offs_c = tl.arange(0, BLOCK_C)
223
+
224
+ # running gradient-state accumulators (for chunks processed so far, going backward)
225
+ dS = tl.zeros((BLOCK_D, BLOCK_E), dtype=tl.float32)
226
+ dz_state = tl.zeros((BLOCK_D,), dtype=tl.float32)
227
+
228
+ for c in range(num_chunks - 1, -1, -1):
229
+ c_start = c * BLOCK_C
230
+ offs_m = c_start + offs_c
231
+ row_mask = offs_m < seq_len
232
+ qk_mask = row_mask[:, None] & (offs_d[None, :] < head_dim)
233
+ v_mask = row_mask[:, None] & (offs_e[None, :] < v_dim)
234
+
235
+ q_ptrs = Q + offs_m[:, None] * stride_qm + offs_d[None, :] * stride_qd
236
+ k_ptrs = K + offs_m[:, None] * stride_kn + offs_d[None, :] * stride_kd
237
+ v_ptrs = V + offs_m[:, None] * stride_vn + offs_e[None, :] * stride_ve
238
+ do_ptrs = DOut + offs_m[:, None] * stride_om + offs_e[None, :] * stride_oe
239
+
240
+ q_raw = tl.load(q_ptrs, mask=qk_mask, other=0.0).to(tl.float32)
241
+ k_raw = tl.load(k_ptrs, mask=qk_mask, other=0.0).to(tl.float32)
242
+ v = tl.load(v_ptrs, mask=v_mask, other=0.0).to(tl.float32)
243
+ do = tl.load(do_ptrs, mask=v_mask, other=0.0).to(tl.float32)
244
+
245
+ q = tl.where(q_raw > 0, q_raw + 1.0, tl.exp(q_raw))
246
+ k = tl.where(k_raw > 0, k_raw + 1.0, tl.exp(k_raw))
247
+ dphi_q = tl.where(q_raw > 0, 1.0, q) # d(elu+1)/dx = q itself when x<=0 (=exp(x)), else 1
248
+ dphi_k = tl.where(k_raw > 0, 1.0, k)
249
+
250
+ kv_ptrs = KV_STATES + c * stride_kvc + offs_d[:, None] * stride_kvd + offs_e[None, :] * stride_kve
251
+ z_ptrs = Z_STATES + c * stride_zc + offs_d * stride_zd
252
+ S_before = tl.load(kv_ptrs) # state BEFORE this chunk
253
+ z_before = tl.load(z_ptrs)
254
+
255
+ # recompute forward quantities needed
256
+ local_scores = tl.dot(q, tl.trans(k))
257
+ causal_mask = offs_c[:, None] >= offs_c[None, :]
258
+ local_scores = tl.where(causal_mask, local_scores, 0.0)
259
+
260
+ out_intra = tl.dot(local_scores, v)
261
+ z_intra = tl.sum(local_scores, axis=1)
262
+ out_inter = tl.dot(q, S_before)
263
+ z_inter = tl.sum(q * z_before[None, :], axis=1)
264
+ z = tl.maximum(z_intra + z_inter, 1e-6)
265
+
266
+ out = (out_intra + out_inter) / z[:, None]
267
+
268
+ # dOut/dz and dOut/dnumerator
269
+ d_num = do / z[:, None] # (C, E)
270
+ d_z = -tl.sum(do * out, axis=1) / z # (C,)
271
+
272
+ # ---- grads through inter-chunk term ----
273
+ # out_inter = q @ S_before ; z_inter = sum(q * z_before)
274
+ dq_inter = tl.dot(d_num, tl.trans(S_before)) + d_z[:, None] * z_before[None, :]
275
+ dS_local = tl.dot(tl.trans(q), d_num) # contribution to dS_before from this chunk's numerator
276
+ dz_state_local = tl.sum(q * d_z[:, None], axis=0)
277
+
278
+ # ---- grads through intra-chunk term ----
279
+ # local_scores = causal_mask * (q @ k^T)
280
+ d_local_scores = tl.dot(d_num, tl.trans(v)) + d_z[:, None]
281
+ d_local_scores = tl.where(causal_mask, d_local_scores, 0.0)
282
+
283
+ dq_intra = tl.dot(d_local_scores, k)
284
+ dk_intra = tl.dot(tl.trans(d_local_scores), q)
285
+ dv_intra = tl.dot(tl.trans(local_scores), d_num)
286
+
287
+ dq_total = dq_intra + dq_inter
288
+ dk_total = dk_intra
289
+ dv_total = dv_intra
290
+
291
+ # ---- grads from state update S += k^T v, z_state += sum(k) ----
292
+ # accumulated dS (from later chunks, passed backward) also flows into k, v of THIS chunk's update
293
+ dk_from_state = tl.dot(dS, tl.trans(v)) if False else tl.dot(v, tl.trans(dS)) # placeholder unused
294
+ # correct: S_next = S_before + k^T @ v -> dK += dS_next @ v^T-like term, dV += k @ dS_next
295
+ dk_total += tl.dot(v, tl.trans(dS)) # (C,D) via (C,E)@(E,D)
296
+ dv_total += tl.dot(k, dS) # (C,E) via (C,D)@(D,E)
297
+ dk_total += dz_state[None, :] * 1.0 # from z_state += sum(k): broadcast grad
298
+ dk_total = dk_total # dz_state contributes uniformly across rows: handled below properly
299
+
300
+ # z_state_next = z_before_next... actually z_state accumulates sum(k) per chunk:
301
+ # gradient of z_state w.r.t. k is all-ones over rows, scaled by dz_state
302
+ dk_total += tl.broadcast_to(dz_state[None, :], (BLOCK_C, BLOCK_D))
303
+
304
+ # propagate accumulated grads to previous (earlier) chunk's state
305
+ dS += dS_local
306
+ dz_state += dz_state_local
307
+
308
+ # apply feature-map derivative
309
+ dq_final = dq_total * dphi_q
310
+ dk_final = dk_total * dphi_k
311
+
312
+ dq_ptrs = DQ + offs_m[:, None] * stride_qm + offs_d[None, :] * stride_qd
313
+ dk_ptrs = DK + offs_m[:, None] * stride_kn + offs_d[None, :] * stride_kd
314
+ dv_ptrs = DV + offs_m[:, None] * stride_vn + offs_e[None, :] * stride_ve
315
+ tl.store(dq_ptrs, dq_final.to(DQ.dtype.element_ty), mask=qk_mask)
316
+ tl.store(dk_ptrs, dk_final.to(DK.dtype.element_ty), mask=qk_mask)
317
+ tl.store(dv_ptrs, dv_total.to(DV.dtype.element_ty), mask=v_mask)
318
+
319
+
320
+ # ============================================================
321
+ # Autograd wrapper
322
+ # ============================================================
323
+ class _LinearAttnFn(torch.autograd.Function):
324
+ @staticmethod
325
+ def forward(ctx, q, k, v, chunk_size):
326
+ B, H, N, D = q.shape
327
+ E = v.shape[-1]
328
+ BLOCK_C = chunk_size
329
+ BLOCK_D = triton.next_power_of_2(D)
330
+ BLOCK_E = triton.next_power_of_2(E)
331
+ num_chunks = triton.cdiv(N, BLOCK_C)
332
+
333
+ out = torch.empty((B, H, N, E), device=q.device, dtype=q.dtype)
334
+ kv_states = torch.empty((B, H, num_chunks, D, E), device=q.device, dtype=torch.float32)
335
+ z_states = torch.empty((B, H, num_chunks, D), device=q.device, dtype=torch.float32)
336
+
337
+ grid = (B * H,)
338
+ _la_fwd_kernel[grid](
339
+ q, k, v, out, kv_states, z_states,
340
+ q.stride(0), q.stride(1), q.stride(2), q.stride(3),
341
+ k.stride(0), k.stride(1), k.stride(2), k.stride(3),
342
+ v.stride(0), v.stride(1), v.stride(2), v.stride(3),
343
+ out.stride(0), out.stride(1), out.stride(2), out.stride(3),
344
+ kv_states.stride(0), kv_states.stride(1), kv_states.stride(2), kv_states.stride(3), kv_states.stride(4),
345
+ z_states.stride(0), z_states.stride(1), z_states.stride(2), z_states.stride(3),
346
+ N, D, E, num_chunks,
347
+ BLOCK_C=BLOCK_C, BLOCK_D=BLOCK_D, BLOCK_E=BLOCK_E,
348
+ )
349
+ ctx.save_for_backward(q, k, v, kv_states, z_states)
350
+ ctx.BLOCK_C, ctx.BLOCK_D, ctx.BLOCK_E = BLOCK_C, BLOCK_D, BLOCK_E
351
+ ctx.num_chunks, ctx.N, ctx.D, ctx.E = num_chunks, N, D, E
352
+ return out
353
+
354
+ @staticmethod
355
+ def backward(ctx, dout):
356
+ q, k, v, kv_states, z_states = ctx.saved_tensors
357
+ B, H, N, D = q.shape
358
+ E = v.shape[-1]
359
+ dq = torch.empty_like(q)
360
+ dk = torch.empty_like(k)
361
+ dv = torch.empty_like(v)
362
+
363
+ grid = (B * H,)
364
+ _la_bwd_kernel[grid](
365
+ q, k, v, dout.contiguous(), dq, dk, dv, kv_states, z_states,
366
+ q.stride(0), q.stride(1), q.stride(2), q.stride(3),
367
+ k.stride(0), k.stride(1), k.stride(2), k.stride(3),
368
+ v.stride(0), v.stride(1), v.stride(2), v.stride(3),
369
+ dout.stride(0), dout.stride(1), dout.stride(2), dout.stride(3),
370
+ kv_states.stride(0), kv_states.stride(1), kv_states.stride(2), kv_states.stride(3), kv_states.stride(4),
371
+ z_states.stride(0), z_states.stride(1), z_states.stride(2), z_states.stride(3),
372
+ N, D, E, ctx.num_chunks,
373
+ BLOCK_C=ctx.BLOCK_C, BLOCK_D=ctx.BLOCK_D, BLOCK_E=ctx.BLOCK_E,
374
+ )
375
+ return dq, dk, dv, None
376
+
377
+ except ImportError: _HAS_TRITON = False
378
+
379
+
380
+ def linear_attention(q, k, v, chunk_size=64):
381
+ if q.is_cuda:
382
+ if not _HAS_TRITON: raise RuntimeError("Triton not installed, GPU path unavailable")
383
+ return _LinearAttnFn.apply(q, k, v, chunk_size)
384
+ else: return _LinearAttnCFn.apply(q, k, v, chunk_size)
385
+
386
+
387
+ # ============================================================
388
+ # nn.Module: drop-in linear attention layer
389
+ # ============================================================
390
+ class LinearAttention(nn.Module):
391
+ def __init__(self, dim, num_heads, chunk_size=64, qkv_bias=False):
392
+ super().__init__()
393
+ assert dim % num_heads == 0
394
+ self.num_heads = num_heads
395
+ self.head_dim = dim // num_heads
396
+ self.chunk_size = chunk_size
397
+
398
+ self.q_proj = nn.Linear(dim, dim, bias=qkv_bias)
399
+ self.k_proj = nn.Linear(dim, dim, bias=qkv_bias)
400
+ self.v_proj = nn.Linear(dim, dim, bias=qkv_bias)
401
+ self.out_proj = nn.Linear(dim, dim, bias=qkv_bias)
402
+
403
+ def forward(self, x):
404
+ B, N, C = x.shape
405
+ H, Dh = self.num_heads, self.head_dim
406
+
407
+ q = self.q_proj(x).view(B, N, H, Dh).transpose(1, 2).contiguous()
408
+ k = self.k_proj(x).view(B, N, H, Dh).transpose(1, 2).contiguous()
409
+ v = self.v_proj(x).view(B, N, H, Dh).transpose(1, 2).contiguous()
410
+
411
+ out = linear_attention(q, k, v, chunk_size=self.chunk_size) # (B, H, N, Dh)
412
+
413
+ out = out.transpose(1, 2).contiguous().view(B, N, C)
414
+ return self.out_proj(out)
@@ -0,0 +1,4 @@
1
+ Metadata-Version: 2.4
2
+ Name: fast-attn-kernels
3
+ Version: 0.1.0
4
+ Requires-Dist: torch
@@ -0,0 +1,10 @@
1
+ pyproject.toml
2
+ setup.py
3
+ attns/__init__.py
4
+ attns/linear_attn.c
5
+ attns/linear_bind.py
6
+ fast_attn_kernels.egg-info/PKG-INFO
7
+ fast_attn_kernels.egg-info/SOURCES.txt
8
+ fast_attn_kernels.egg-info/dependency_links.txt
9
+ fast_attn_kernels.egg-info/requires.txt
10
+ fast_attn_kernels.egg-info/top_level.txt
@@ -0,0 +1,8 @@
1
+ [build-system]
2
+ requires = ["setuptools", "wheel"]
3
+ build-backend = "setuptools.build_meta"
4
+
5
+ [project]
6
+ name = "fast-attn-kernels"
7
+ version = "0.1.0"
8
+ dependencies = ["torch"]
@@ -0,0 +1,4 @@
1
+ [egg_info]
2
+ tag_build =
3
+ tag_date = 0
4
+
@@ -0,0 +1,16 @@
1
+ from setuptools import setup, find_packages, Extension
2
+
3
+ setup(
4
+ name="fast-attn-kernels",
5
+ version="0.1.0",
6
+ packages=find_packages(),
7
+ ext_modules=[
8
+ Extension(
9
+ "attns.liblinear_attn",
10
+ sources=["attns/linear_attn.c"],
11
+ extra_compile_args=["-O3", "-fopenmp"],
12
+ extra_link_args=["-fopenmp"],
13
+ )
14
+ ],
15
+ install_requires=["torch"],
16
+ )