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.
- fast_attn_kernels-0.1.0/PKG-INFO +4 -0
- fast_attn_kernels-0.1.0/attns/__init__.py +1 -0
- fast_attn_kernels-0.1.0/attns/linear_attn.c +336 -0
- fast_attn_kernels-0.1.0/attns/linear_bind.py +414 -0
- fast_attn_kernels-0.1.0/fast_attn_kernels.egg-info/PKG-INFO +4 -0
- fast_attn_kernels-0.1.0/fast_attn_kernels.egg-info/SOURCES.txt +10 -0
- fast_attn_kernels-0.1.0/fast_attn_kernels.egg-info/dependency_links.txt +1 -0
- fast_attn_kernels-0.1.0/fast_attn_kernels.egg-info/requires.txt +1 -0
- fast_attn_kernels-0.1.0/fast_attn_kernels.egg-info/top_level.txt +1 -0
- fast_attn_kernels-0.1.0/pyproject.toml +8 -0
- fast_attn_kernels-0.1.0/setup.cfg +4 -0
- fast_attn_kernels-0.1.0/setup.py +16 -0
|
@@ -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,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 @@
|
|
|
1
|
+
|
|
@@ -0,0 +1 @@
|
|
|
1
|
+
torch
|
|
@@ -0,0 +1 @@
|
|
|
1
|
+
attns
|
|
@@ -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
|
+
)
|