1use crate::GpuCtx;
11use crate::encoder_weights::{Act, MaskKind};
12use crate::forward::ShaderModuleTuned as _;
13use crate::forward::{make_bg, pipeline, uni};
14use anyhow::Result;
15
16pub(crate) fn act_code(act: Option<Act>) -> u32 {
18 match act {
19 None => 0,
20 Some(Act::GeluErf) => 1,
21 Some(Act::GeluTanh) => 2,
22 Some(Act::Silu) => 3,
23 Some(Act::Tanh) => 4,
24 Some(Act::Relu) => 5,
25 }
26}
27
28const ACT_FNS: &str = r#"
31fn erf_as(x: f32) -> f32 {
32 let s = sign(x);
33 let a = abs(x);
34 let t = 1.0 / (1.0 + 0.3275911 * a);
35 let y = 1.0 - (((((1.061405429 * t - 1.453152027) * t) + 1.421413741) * t - 0.284496736) * t
36 + 0.254829592) * t * exp(-a * a);
37 return s * y;
38}
39fn apply_act(v: f32, code: u32) -> f32 {
40 switch code {
41 case 1u: { return 0.5 * v * (1.0 + erf_as(v * 0.70710678)); }
42 // tanh-GELU. The argument is CLAMPED, exactly as the decoder's kernels already do
43 // (forward.rs / lib.rs) — WGSL's `tanh` overflows to NaN on some backends for |arg| >~ 88,
44 // because a naive (e^2x - 1)/(e^2x + 1) expansion divides inf by inf. `arg` reaches 105 at a
45 // pre-activation of only 13.8, which the text encoders never hit and a ViT's outlier tokens
46 // hit on the first block. tanh is ±1 to well inside f32 epsilon by |arg| = 20, so this is
47 // EXACT, not an approximation.
48 case 2u: {
49 let a = clamp(0.7978845608 * (v + 0.044715 * v * v * v), -20.0, 20.0);
50 return 0.5 * v * (1.0 + tanh(a));
51 }
52 case 3u: { return v / (1.0 + exp(-v)); }
53 case 4u: { return tanh(v); }
54 case 5u: { return max(v, 0.0); }
55 default: { return v; }
56 }
57}
58"#;
59
60pub(crate) fn enc_gemm_src(f16: bool) -> String {
63 let (enable, wty) = if f16 {
64 ("enable f16;", "f16")
65 } else {
66 ("", "f32")
67 };
68 format!(
69 r#"{enable}
70struct Meta {{ m: u32, n: u32, k: u32, flags: u32 }}
71@group(0) @binding(0) var<storage, read> x: array<f32>;
72@group(0) @binding(1) var<storage, read> w: array<{wty}>;
73@group(0) @binding(2) var<storage, read> bias: array<f32>;
74@group(0) @binding(3) var<storage, read_write> y: array<f32>;
75@group(0) @binding(4) var<uniform> mt: Meta;
76var<workgroup> xt: array<f32, 256>;
77var<workgroup> wt: array<f32, 256>;
78{ACT_FNS}
79@compute @workgroup_size(16, 16)
80fn main(@builtin(workgroup_id) wg: vec3<u32>, @builtin(local_invocation_id) li: vec3<u32>) {{
81 let row = wg.y * 16u + li.y;
82 let col = wg.x * 16u + li.x;
83 var acc = 0.0;
84 let ktiles = (mt.k + 15u) / 16u;
85 for (var kb = 0u; kb < ktiles; kb++) {{
86 let kx = kb * 16u + li.x;
87 var xv = 0.0;
88 if (row < mt.m && kx < mt.k) {{ xv = x[row * mt.k + kx]; }}
89 xt[li.y * 16u + li.x] = xv;
90 let kw = kb * 16u + li.y;
91 var wv = 0.0;
92 if (col < mt.n && kw < mt.k) {{ wv = f32(w[col * mt.k + kw]); }}
93 wt[li.y * 16u + li.x] = wv;
94 workgroupBarrier();
95 for (var kk = 0u; kk < 16u; kk++) {{
96 acc += xt[li.y * 16u + kk] * wt[kk * 16u + li.x];
97 }}
98 workgroupBarrier();
99 }}
100 if (row < mt.m && col < mt.n) {{
101 if ((mt.flags & 1u) != 0u) {{ acc += bias[col]; }}
102 y[row * mt.n + col] = apply_act(acc, mt.flags >> 8u);
103 }}
104}}
105"#
106 )
107}
108
109pub fn enc_gemm2_src(f16: bool, bm: usize, bn: usize) -> String {
141 let (enable, wty) = if f16 {
142 ("enable f16;", "f16")
143 } else {
144 ("", "f32")
145 };
146 const BK: usize = 16;
147 let tm = bm / 16; let cpt = bn / 16; let x_iters = bm * BK / 256; let w_iters = bn / 16; let bn_ = bn;
152 debug_assert!(tm >= 1 && bm.is_multiple_of(16) && x_iters * 256 == bm * BK);
153 debug_assert!(
154 matches!(cpt, 2 | 4 | 8),
155 "cols/thread must be 2, 4 or 8, got {cpt}"
156 );
157
158 let (inner, acc_decl, epi) = if cpt <= 4 {
163 let mut inner = String::new();
164 for kk in 0..BK {
165 let comps = (0..cpt)
166 .map(|j| format!("ws[{kk}u * {bn_}u + col0 + {j}u]"))
167 .collect::<Vec<_>>()
168 .join(", ");
169 inner.push_str(&format!(" {{ let b = vec{cpt}<f32>({comps});\n"));
170 for i in 0..tm {
171 inner.push_str(&format!(
172 " acc{i} += xs[(row0 + {i}u) * {BK}u + {kk}u] * b;\n"
173 ));
174 }
175 inner.push_str(" }\n");
176 }
177 let acc_decl = (0..tm)
178 .map(|i| format!(" var acc{i} = vec{cpt}<f32>(0.0);"))
179 .collect::<Vec<_>>()
180 .join("\n");
181 let mut epi = String::new();
182 for i in 0..tm {
183 epi.push_str(&format!(
184 " {{ let r = mrow + {i}u; if (r < mt.m) {{\n let acc = acc{i};\n"
185 ));
186 for j in 0..cpt {
187 epi.push_str(&format!(
188 " {{ let c = ncol + {j}u; if (c < mt.n) {{ var v = acc[{j}u]; \
189 if ((mt.flags & 1u) != 0u) {{ v += bias[c]; }} \
190 y[r * mt.n + c] = apply_act(v, mt.flags >> 8u); }} }}\n"
191 ));
192 }
193 epi.push_str(" } }\n");
194 }
195 (inner, acc_decl, epi)
196 } else {
197 const NV: usize = 2; let mut inner = String::new();
201 for kk in 0..BK {
202 for b in 0..NV {
203 let comps = (0..4)
204 .map(|j| format!("ws[{kk}u * {bn_}u + col0 + {}u]", b * 4 + j))
205 .collect::<Vec<_>>()
206 .join(", ");
207 inner.push_str(&format!(" let b{b}_{kk} = vec4<f32>({comps});\n"));
208 }
209 for i in 0..tm {
210 inner.push_str(&format!(
212 " let a{i}_{kk} = xs[(row0 + {i}u) * {BK}u + {kk}u];\n"
213 ));
214 for b in 0..NV {
215 inner.push_str(&format!(" acc{i}_{b} += a{i}_{kk} * b{b}_{kk};\n"));
216 }
217 }
218 }
219 let acc_decl = (0..tm)
220 .flat_map(|i| (0..NV).map(move |b| format!(" var acc{i}_{b} = vec4<f32>(0.0);")))
221 .collect::<Vec<_>>()
222 .join("\n");
223 let mut epi = String::new();
224 for i in 0..tm {
225 epi.push_str(&format!(" {{ let r = mrow + {i}u; if (r < mt.m) {{\n"));
226 for b in 0..NV {
227 for j in 0..4 {
228 let c = b * 4 + j;
229 epi.push_str(&format!(
230 " {{ let c = ncol + {c}u; if (c < mt.n) {{ var v = acc{i}_{b}[{j}u]; \
231 if ((mt.flags & 1u) != 0u) {{ v += bias[c]; }} \
232 y[r * mt.n + c] = apply_act(v, mt.flags >> 8u); }} }}\n"
233 ));
234 }
235 }
236 epi.push_str(" } }\n");
237 }
238 (inner, acc_decl, epi)
239 };
240
241 format!(
242 r#"{enable}
243struct Meta {{ m: u32, n: u32, k: u32, flags: u32 }}
244@group(0) @binding(0) var<storage, read> x: array<f32>;
245@group(0) @binding(1) var<storage, read> w: array<{wty}>;
246@group(0) @binding(2) var<storage, read> bias: array<f32>;
247@group(0) @binding(3) var<storage, read_write> y: array<f32>;
248@group(0) @binding(4) var<uniform> mt: Meta;
249var<workgroup> xs: array<f32, {xs_len}>; // [BM][BK]
250var<workgroup> ws: array<f32, {ws_len}>; // [BK][BN], transposed at stage
251{ACT_FNS}
252@compute @workgroup_size(16, 16)
253fn main(@builtin(workgroup_id) wg: vec3<u32>, @builtin(local_invocation_id) li: vec3<u32>) {{
254 let t = li.y * 16u + li.x; // 0..255
255 let mbase = wg.y * {bm}u; // first row of this tile
256 let nbase = wg.x * {bn_}u; // first col of this tile
257 let row0 = li.y * {tm}u; // thread's rows, tile-local
258 let col0 = li.x * {cpt}u; // thread's cols, tile-local
259 let mrow = mbase + row0;
260 let ncol = nbase + col0;
261{acc_decl}
262
263 let ktiles = (mt.k + {BK}u - 1u) / {BK}u;
264 for (var kb = 0u; kb < ktiles; kb++) {{
265 let koff = kb * {BK}u;
266 // Stage x[BM][BK]: consecutive threads read consecutive k (coalesced; x is [m,k]).
267 for (var i = 0u; i < {x_iters}u; i++) {{
268 let r = i * 16u + t / {BK}u;
269 let c = t % {BK}u;
270 let gr = mbase + r;
271 let gc = koff + c;
272 var v = 0.0;
273 if (gr < mt.m && gc < mt.k) {{ v = x[gr * mt.k + gc]; }}
274 xs[r * {BK}u + c] = v;
275 }}
276 // Stage w[BN][BK] TRANSPOSED into ws[BK][BN]: consecutive threads read consecutive k of
277 // one output column (w is [n,k]), so the global reads coalesce and the transpose happens
278 // in shared, where it is free.
279 for (var i = 0u; i < {w_iters}u; i++) {{
280 let c = i * 16u + t / {BK}u;
281 let kk = t % {BK}u;
282 let gc = nbase + c;
283 let gk = koff + kk;
284 var v = 0.0;
285 if (gc < mt.n && gk < mt.k) {{ v = f32(w[gc * mt.k + gk]); }}
286 ws[kk * {bn_}u + c] = v;
287 }}
288 workgroupBarrier();
289{inner}
290 workgroupBarrier();
291 }}
292{epi}}}
293"#,
294 xs_len = bm * BK,
295 ws_len = BK * bn,
296 )
297}
298
299pub fn enc_gemm3_src(f16: bool, bm: usize, bn: usize, bk: usize) -> String {
318 let (enable, wty) = if f16 {
319 ("enable f16;", "f16")
320 } else {
321 ("", "f32")
322 };
323 let tm = bm / 16;
324 let cpt = bn / 16;
325 let nv = cpt / 4; let kq = bk / 4; assert!(cpt.is_multiple_of(4) && bk.is_multiple_of(4) && bm.is_multiple_of(16));
328 let bnq = bn / 4; let xs_len = bm * kq; let ws_len = bk * bnq; let x_rounds = xs_len.div_ceil(256);
332 let w_rounds = ws_len.div_ceil(256);
333
334 let acc_decl = (0..tm)
335 .flat_map(|i| (0..nv).map(move |b| format!(" var acc{i}_{b} = vec4<f32>(0.0);")))
336 .collect::<Vec<_>>()
337 .join("\n");
338
339 let mut inner = String::new();
342 for q in 0..kq {
343 for i in 0..tm {
344 inner.push_str(&format!(
345 " let a{q}_{i} = xs[(row0 + {i}u) * {kq}u + {q}u];\n"
346 ));
347 }
348 for kk in 0..4 {
349 for b in 0..nv {
350 inner.push_str(&format!(
351 " let w{q}_{kk}_{b} = ws[({}u) * {bnq}u + colq + {b}u];\n",
352 q * 4 + kk
353 ));
354 }
355 for i in 0..tm {
356 for b in 0..nv {
357 inner.push_str(&format!(
358 " acc{i}_{b} += a{q}_{i}[{kk}u] * w{q}_{kk}_{b};\n"
359 ));
360 }
361 }
362 }
363 }
364
365 let mut epi = String::new();
366 for i in 0..tm {
367 epi.push_str(&format!(" {{ let r = mrow + {i}u; if (r < mt.m) {{\n"));
368 for b in 0..nv {
369 for j in 0..4 {
370 let c = b * 4 + j;
371 epi.push_str(&format!(
372 " {{ let c = ncol + {c}u; if (c < mt.n) {{ var v = acc{i}_{b}[{j}u]; \
373 if ((mt.flags & 1u) != 0u) {{ v += bias[c]; }} \
374 y[r * mt.n + c] = apply_act(v, mt.flags >> 8u); }} }}\n"
375 ));
376 }
377 }
378 epi.push_str(" } }\n");
379 }
380
381 format!(
382 r#"{enable}
383struct Meta {{ m: u32, n: u32, k: u32, flags: u32 }}
384@group(0) @binding(0) var<storage, read> x: array<f32>; // [m, k]
385@group(0) @binding(1) var<storage, read> w: array<{wty}>; // [k, n] -- TRANSPOSED at upload
386@group(0) @binding(2) var<storage, read> bias: array<f32>;
387@group(0) @binding(3) var<storage, read_write> y: array<f32>;
388@group(0) @binding(4) var<uniform> mt: Meta;
389var<workgroup> xs: array<vec4<f32>, {xs_len}>; // [bm][BK/4]
390var<workgroup> ws: array<vec4<f32>, {ws_len}>; // [BK][bn/4]
391{ACT_FNS}
392@compute @workgroup_size(16, 16)
393fn main(@builtin(workgroup_id) wg: vec3<u32>, @builtin(local_invocation_id) li: vec3<u32>) {{
394 let t = li.y * 16u + li.x;
395 let mbase = wg.y * {bm}u;
396 let nbase = wg.x * {bn}u;
397 let row0 = li.y * {tm}u;
398 let colq = li.x * {nv}u; // this thread's first COLUMN-QUAD in the tile
399 let mrow = mbase + row0;
400 let ncol = nbase + colq * 4u;
401{acc_decl}
402
403 let ktiles = (mt.k + {bk}u - 1u) / {bk}u;
404 for (var kb = 0u; kb < ktiles; kb++) {{
405 let koff = kb * {bk}u;
406 // Stage x as vec4 along k (x is [m,k]: contiguous, and consecutive threads take consecutive
407 // k-quads of a row, so the warp's reads are contiguous too).
408 for (var i = 0u; i < {x_rounds}u; i++) {{
409 let idx = i * 256u + t;
410 if (idx < {xs_len}u) {{
411 let r = idx / {kq}u;
412 let q = idx % {kq}u;
413 let gr = mbase + r;
414 let gk = koff + q * 4u;
415 var v = vec4<f32>(0.0);
416 if (gr < mt.m) {{
417 if (gk + 0u < mt.k) {{ v.x = x[gr * mt.k + gk + 0u]; }}
418 if (gk + 1u < mt.k) {{ v.y = x[gr * mt.k + gk + 1u]; }}
419 if (gk + 2u < mt.k) {{ v.z = x[gr * mt.k + gk + 2u]; }}
420 if (gk + 3u < mt.k) {{ v.w = x[gr * mt.k + gk + 3u]; }}
421 }}
422 xs[r * {kq}u + q] = v;
423 }}
424 }}
425 // Stage wT as vec4 along n. THIS is what the transpose bought: at a fixed k the columns are
426 // contiguous, so one thread loads four of them in one go and consecutive threads load the
427 // next four.
428 for (var i = 0u; i < {w_rounds}u; i++) {{
429 let idx = i * 256u + t;
430 if (idx < {ws_len}u) {{
431 let kk = idx / {bnq}u;
432 let cq = idx % {bnq}u;
433 let gk = koff + kk;
434 let gc = nbase + cq * 4u;
435 var v = vec4<f32>(0.0);
436 if (gk < mt.k) {{
437 if (gc + 0u < mt.n) {{ v.x = f32(w[gk * mt.n + gc + 0u]); }}
438 if (gc + 1u < mt.n) {{ v.y = f32(w[gk * mt.n + gc + 1u]); }}
439 if (gc + 2u < mt.n) {{ v.z = f32(w[gk * mt.n + gc + 2u]); }}
440 if (gc + 3u < mt.n) {{ v.w = f32(w[gk * mt.n + gc + 3u]); }}
441 }}
442 ws[kk * {bnq}u + cq] = v;
443 }}
444 }}
445 workgroupBarrier();
446{inner}
447 workgroupBarrier();
448 }}
449{epi}}}
450"#
451 )
452}
453
454pub fn enc_gemm3_f16a_src(bm: usize, bn: usize, bk: usize) -> String {
469 let tm = bm / 16;
470 let cpt = bn / 16;
471 let nv = cpt / 4;
472 let kq = bk / 4;
473 assert!(cpt.is_multiple_of(4) && bk.is_multiple_of(4) && bm.is_multiple_of(16));
474 let bnq = bn / 4;
475 let xs_len = bm * kq;
476 let ws_len = bk * bnq;
477 let x_rounds = xs_len.div_ceil(256);
478 let w_rounds = ws_len.div_ceil(256);
479
480 let acc_decl = (0..tm)
481 .flat_map(|i| (0..nv).map(move |b| format!(" var acc{i}_{b} = vec4<f32>(0.0);")))
482 .collect::<Vec<_>>()
483 .join("\n");
484 let hacc_decl = (0..tm)
486 .flat_map(|i| (0..nv).map(move |b| format!(" var h{i}_{b} = vec4<f16>(0.0);")))
487 .collect::<Vec<_>>()
488 .join("\n");
489 let fold = (0..tm)
490 .flat_map(|i| (0..nv).map(move |b| format!(" acc{i}_{b} += vec4<f32>(h{i}_{b});")))
491 .collect::<Vec<_>>()
492 .join("\n");
493
494 let mut inner = String::new();
495 for q in 0..kq {
496 for i in 0..tm {
497 inner.push_str(&format!(
498 " let a{q}_{i} = xs[(row0 + {i}u) * {kq}u + {q}u];\n"
499 ));
500 }
501 for kk in 0..4 {
502 for b in 0..nv {
503 inner.push_str(&format!(
504 " let w{q}_{kk}_{b} = ws[({}u) * {bnq}u + colq + {b}u];\n",
505 q * 4 + kk
506 ));
507 }
508 for i in 0..tm {
509 for b in 0..nv {
510 inner.push_str(&format!(
511 " h{i}_{b} = fma(vec4<f16>(a{q}_{i}[{kk}u]), w{q}_{kk}_{b}, h{i}_{b});\n"
512 ));
513 }
514 }
515 }
516 }
517
518 let mut epi = String::new();
519 for i in 0..tm {
520 epi.push_str(&format!(" {{ let r = mrow + {i}u; if (r < mt.m) {{\n"));
521 for b in 0..nv {
522 for j in 0..4 {
523 let c = b * 4 + j;
524 epi.push_str(&format!(
525 " {{ let c = ncol + {c}u; if (c < mt.n) {{ var v = acc{i}_{b}[{j}u]; \
526 if ((mt.flags & 1u) != 0u) {{ v += bias[c]; }} \
527 y[r * mt.n + c] = apply_act(v, mt.flags >> 8u); }} }}\n"
528 ));
529 }
530 }
531 epi.push_str(" } }\n");
532 }
533
534 format!(
535 r#"enable f16;
536struct Meta {{ m: u32, n: u32, k: u32, flags: u32 }}
537@group(0) @binding(0) var<storage, read> x: array<f32>; // [m, k]
538@group(0) @binding(1) var<storage, read> w: array<vec4<f16>>; // [k, n/4] -- TRANSPOSED f16
539@group(0) @binding(2) var<storage, read> bias: array<f32>;
540@group(0) @binding(3) var<storage, read_write> y: array<f32>;
541@group(0) @binding(4) var<uniform> mt: Meta;
542var<workgroup> xs: array<vec4<f16>, {xs_len}>; // [bm][BK/4], converted once at stage
543var<workgroup> ws: array<vec4<f16>, {ws_len}>; // [BK][bn/4]
544{ACT_FNS}
545@compute @workgroup_size(16, 16)
546fn main(@builtin(workgroup_id) wg: vec3<u32>, @builtin(local_invocation_id) li: vec3<u32>) {{
547 let t = li.y * 16u + li.x;
548 let mbase = wg.y * {bm}u;
549 let nbase = wg.x * {bn}u;
550 let row0 = li.y * {tm}u;
551 let colq = li.x * {nv}u;
552 let mrow = mbase + row0;
553 let ncol = nbase + colq * 4u;
554{acc_decl}
555
556 let ktiles = (mt.k + {bk}u - 1u) / {bk}u;
557 for (var kb = 0u; kb < ktiles; kb++) {{
558 let koff = kb * {bk}u;
559 for (var i = 0u; i < {x_rounds}u; i++) {{
560 let idx = i * 256u + t;
561 if (idx < {xs_len}u) {{
562 let r = idx / {kq}u;
563 let q = idx % {kq}u;
564 let gr = mbase + r;
565 let gk = koff + q * 4u;
566 var v = vec4<f32>(0.0);
567 if (gr < mt.m) {{
568 if (gk + 0u < mt.k) {{ v.x = x[gr * mt.k + gk + 0u]; }}
569 if (gk + 1u < mt.k) {{ v.y = x[gr * mt.k + gk + 1u]; }}
570 if (gk + 2u < mt.k) {{ v.z = x[gr * mt.k + gk + 2u]; }}
571 if (gk + 3u < mt.k) {{ v.w = x[gr * mt.k + gk + 3u]; }}
572 }}
573 xs[idx] = vec4<f16>(v);
574 }}
575 }}
576 for (var i = 0u; i < {w_rounds}u; i++) {{
577 let idx = i * 256u + t;
578 if (idx < {ws_len}u) {{
579 let kk = idx / {bnq}u;
580 let cq = idx % {bnq}u;
581 let gk = koff + kk;
582 var v = vec4<f16>(0.0);
583 if (gk < mt.k && (nbase + cq * 4u) < mt.n) {{ v = w[(gk * mt.n) / 4u + nbase / 4u + cq]; }}
584 ws[kk * {bnq}u + cq] = v;
585 }}
586 }}
587 workgroupBarrier();
588{hacc_decl}
589{inner}
590{fold}
591 workgroupBarrier();
592 }}
593{epi}}}
594"#
595 )
596}
597
598pub fn enc_gemm4_f16_src(bm: usize, bn: usize, bk: usize) -> String {
644 enc_gemm4_typed_src(bm, bn, bk, true)
645}
646
647pub fn enc_gemm4_src(bm: usize, bn: usize, bk: usize) -> String {
648 enc_gemm4_typed_src(bm, bn, bk, false)
649}
650
651pub fn enc_gemm4_f16w_src(bm: usize, bn: usize, bk: usize) -> String {
654 let src = enc_gemm4_typed_src(bm, bn, bk, true)
655 .replace(
656 "@group(0) @binding(1) var<storage, read> w: array<f32>;",
657 "@group(0) @binding(1) var<storage, read> w: array<f16>;",
658 )
659 .replace(
661 "{ v = w[gk * mt.n + gc]; }",
662 "{ v = f32(w[gk * mt.n + gc]); }",
663 );
664 assert!(
665 src.contains("array<f16>;") && src.contains("f32(w[gk"),
666 "v4 w-binding/stage drifted"
667 );
668 src
669}
670
671fn enc_gemm4_typed_src(bm: usize, bn: usize, bk: usize, ab_f16: bool) -> String {
672 assert!(bm.is_multiple_of(8) && bn.is_multiple_of(8) && bk.is_multiple_of(8));
673 const SG_ROWS: usize = 2; const SG_COLS: usize = 4;
675 let nsg = SG_ROWS * SG_COLS;
676 let threads = nsg * 32;
677 let rows_per_sg = bm / SG_ROWS; let cols_per_sg = bn / SG_COLS; let ai = rows_per_sg / 8; let aj = cols_per_sg / 8; let ksteps = bk / 8;
682
683 let abty = if ab_f16 { "f16" } else { "f32" };
686 let f16_enable = if ab_f16 { "enable f16;\n" } else { "" };
687 let acc_decl = (0..ai)
688 .flat_map(|i| {
689 (0..aj).map(move |j| {
690 format!(
691 " var acc{i}_{j} = coopLoadT<coop_mat8x8<f32, C>>(&bias8[nbase + qc * {cols_per_sg}u + {}u], mt.n);",
692 j * 8
693 )
694 })
695 })
696 .collect::<Vec<_>>()
697 .join("\n");
698
699 let mut inner = String::new();
700 for kk in 0..ksteps {
701 for i in 0..ai {
702 inner.push_str(&format!(
703 " let a{kk}_{i} = coopLoadT<coop_mat8x8<{abty}, A>>(&xs[(qr * {rows_per_sg}u + {}u) * {bk}u + {}u], {bk}u);\n",
704 i * 8, kk * 8
705 ));
706 }
707 for j in 0..aj {
708 inner.push_str(&format!(
709 " let b{kk}_{j} = coopLoadT<coop_mat8x8<{abty}, B>>(&ws[{}u * {bn}u + qc * {cols_per_sg}u + {}u], {bn}u);\n",
710 kk * 8, j * 8
711 ));
712 }
713 for i in 0..ai {
714 for j in 0..aj {
715 inner.push_str(&format!(
716 " acc{i}_{j} = coopMultiplyAdd(a{kk}_{i}, b{kk}_{j}, acc{i}_{j});\n"
717 ));
718 }
719 }
720 }
721
722 let mut store = String::new();
725 for i in 0..ai {
726 for j in 0..aj {
727 store.push_str(&format!(
728 " {{ let r = mbase + qr * {rows_per_sg}u + {}u; let c = nbase + qc * {cols_per_sg}u + {}u;\n \
729 if (r < mt.m && c < mt.n) {{ coopStoreT(acc{i}_{j}, &y[r * mt.n + c], mt.n); }} }}\n",
730 i * 8, j * 8
731 ));
732 }
733 }
734
735 let xs_len = bm * bk;
736 let ws_len = bk * bn;
737 let x_rounds = xs_len.div_ceil(threads);
738 let w_rounds = ws_len.div_ceil(threads);
739
740 format!(
741 r#"{f16_enable}enable wgpu_cooperative_matrix;
742struct Meta {{ m: u32, n: u32, k: u32, flags: u32 }}
743@group(0) @binding(0) var<storage, read> x: array<f32>; // [m, k]
744@group(0) @binding(1) var<storage, read> w: array<f32>; // [k, n] -- TRANSPOSED at upload
745@group(0) @binding(2) var<storage, read> bias8: array<f32>; // [8, n] -- bias broadcast to 8 rows
746@group(0) @binding(3) var<storage, read_write> y: array<f32>; // [m, n]
747@group(0) @binding(4) var<uniform> mt: Meta;
748var<workgroup> xs: array<{abty}, {xs_len}>; // [bm][bk]
749var<workgroup> ws: array<{abty}, {ws_len}>; // [bk][bn]
750@compute @workgroup_size({threads})
751fn main(
752 @builtin(workgroup_id) wg: vec3<u32>,
753 @builtin(local_invocation_index) t: u32,
754 @builtin(subgroup_id) sg: u32,
755) {{
756 let mbase = wg.y * {bm}u;
757 let nbase = wg.x * {bn}u;
758 let qr = sg / {SG_COLS}u;
759 let qc = sg % {SG_COLS}u;
760{acc_decl}
761
762 let ktiles = (mt.k + {bk}u - 1u) / {bk}u;
763 for (var kb = 0u; kb < ktiles; kb++) {{
764 let koff = kb * {bk}u;
765 for (var i = 0u; i < {x_rounds}u; i++) {{
766 let idx = i * {threads}u + t;
767 if (idx < {xs_len}u) {{
768 let r = idx / {bk}u;
769 let c = idx % {bk}u;
770 let gr = mbase + r;
771 let gc = koff + c;
772 var v = 0.0;
773 if (gr < mt.m && gc < mt.k) {{ v = x[gr * mt.k + gc]; }}
774 xs[idx] = {abty}(v);
775 }}
776 }}
777 for (var i = 0u; i < {w_rounds}u; i++) {{
778 let idx = i * {threads}u + t;
779 if (idx < {ws_len}u) {{
780 let kk = idx / {bn}u;
781 let c = idx % {bn}u;
782 let gk = koff + kk;
783 let gc = nbase + c;
784 var v = 0.0;
785 if (gk < mt.k && gc < mt.n) {{ v = w[gk * mt.n + gc]; }}
786 ws[idx] = {abty}(v);
787 }}
788 }}
789 workgroupBarrier();
790{inner}
791 workgroupBarrier();
792 }}
793
794{store}}}
795"#
796 )
797}
798
799pub fn enc_gemm3_sk_src(f16: bool, bm: usize, bn: usize, bk: usize) -> String {
808 gemm3_sk_patch(enc_gemm3_src(f16, bm, bn, bk), bk)
809}
810
811pub fn enc_gemm3_sk_f16a_src(bm: usize, bn: usize, bk: usize) -> String {
815 gemm3_sk_patch(enc_gemm3_f16a_src(bm, bn, bk), bk)
816}
817
818fn gemm3_sk_patch(base: String, bk: usize) -> String {
821 let frags = [
824 "fn main(@builtin(workgroup_id) wg: vec3<u32>, @builtin(local_invocation_id) li: vec3<u32>) {",
825 " let ktiles = (mt.k + ",
826 "@group(0) @binding(3) var<storage, read_write> y: array<f32>;",
827 ];
828 for f in frags {
829 assert!(base.contains(f), "v3 source drifted: {f}");
830 }
831 let src = base
832 .replace(
833 "fn main(@builtin(workgroup_id) wg: vec3<u32>, @builtin(local_invocation_id) li: vec3<u32>) {",
834 "fn main(@builtin(workgroup_id) wg: vec3<u32>, @builtin(local_invocation_id) li: vec3<u32>,\n @builtin(num_workgroups) nwg: vec3<u32>) {",
835 )
836 .replace(
837 "@group(0) @binding(3) var<storage, read_write> y: array<f32>;",
838 "@group(0) @binding(3) var<storage, read_write> y: array<f32>; // [nz, m, n] partials",
839 )
840 .replace(
843 " let t = li.y * 16u + li.x;",
844 " let t = li.y * 16u + li.x;\n if (mt.flags == 0xFFFFFFFFu) { y[0] = bias[0]; } // keep bias bound (never true)",
845 );
846 let (patched, klooped) = {
848 let needle = format!(
849 " let ktiles = (mt.k + {bk}u - 1u) / {bk}u;\n for (var kb = 0u; kb < ktiles; kb++) {{\n let koff = kb * {bk}u;"
850 );
851 let replacement = format!(
852 " let ktiles = (mt.k + {bk}u - 1u) / {bk}u;\n let per = (ktiles + nwg.z - 1u) / nwg.z;\n let kb0 = wg.z * per;\n let kb1 = min(kb0 + per, ktiles);\n for (var kb = kb0; kb < kb1; kb++) {{\n let koff = kb * {bk}u;"
853 );
854 let ok = src.contains(&needle);
855 (src.replace(&needle, &replacement), ok)
856 };
857 assert!(klooped, "v3 k-loop shape drifted");
858 let mut out = patched;
860 assert!(
861 out.contains("var v = acc") || out.contains("var v = f32(acc"),
862 "v3/f16a epilogue drifted"
863 );
864 out = rewrite_v3_epilogue_for_sk(&out);
867 out
868}
869
870fn rewrite_v3_epilogue_for_sk(src: &str) -> String {
874 let mut out = String::with_capacity(src.len());
875 for line in src.lines() {
876 let t = line.trim_start();
877 if let Some(rest) = t.strip_prefix("{ let c = ncol + ") {
878 let j = rest.split("u;").next().expect("column offset");
880 let acc = rest
881 .split("var v = ")
882 .nth(1)
883 .and_then(|s| s.split(';').next())
884 .expect("accumulator expr");
885 let indent = &line[..line.len() - t.len()];
886 out.push_str(&format!(
887 "{indent}{{ let c = ncol + {j}u; if (c < mt.n) {{ y[wg.z * mt.m * mt.n + r * mt.n + c] = {acc}; }} }}\n"
888 ));
889 } else {
890 out.push_str(line);
891 out.push('\n');
892 }
893 }
894 out
895}
896
897fn enc_gemm3_sk_reduce_src() -> String {
901 format!(
902 r#"
903struct Meta {{ m: u32, n: u32, flags: u32, sk: u32 }}
904@group(0) @binding(0) var<storage, read> part: array<f32>;
905@group(0) @binding(1) var<storage, read> bias: array<f32>;
906@group(0) @binding(2) var<storage, read_write> y: array<f32>;
907@group(0) @binding(3) var<uniform> mt: Meta;
908{ACT_FNS}
909@compute @workgroup_size(256)
910fn main(@builtin(global_invocation_id) gid: vec3<u32>) {{
911 let i = gid.x;
912 let total = mt.m * mt.n;
913 if (i >= total) {{ return; }}
914 var acc = 0.0;
915 for (var c = 0u; c < mt.sk; c++) {{ acc += part[c * total + i]; }}
916 if ((mt.flags & 1u) != 0u) {{ acc += bias[i % mt.n]; }}
917 y[i] = apply_act(acc, mt.flags >> 8u);
918}}
919"#
920 )
921}
922
923pub(crate) const GEMM2_TILES: [(usize, usize); 9] = [
937 (128, 128),
944 (128, 64),
945 (64, 128),
946 (64, 64),
947 (32, 64),
948 (64, 32),
949 (32, 32),
950 (16, 64),
951 (16, 32),
952];
953
954const GEMM2_MIN_WGS: usize = 96;
966
967pub fn gemm2_tile(m: usize, n: usize) -> (usize, usize) {
975 if let Ok(v) = std::env::var("OSFKB_ENC_TILE")
979 && let Some((a, b)) = v.split_once('x')
980 && let (Ok(a), Ok(b)) = (a.parse(), b.parse())
981 && GEMM2_TILES.contains(&(a, b))
982 {
983 return (a, b);
984 }
985 let wgs = |(bm, bn): (usize, usize)| n.div_ceil(bn) * m.div_ceil(bm);
986 GEMM2_TILES
987 .into_iter()
988 .find(|&t| wgs(t) >= GEMM2_MIN_WGS)
989 .unwrap_or_else(|| {
992 GEMM2_TILES
993 .into_iter()
994 .rev()
995 .max_by_key(|&t| wgs(t))
996 .expect("non-empty tile table")
997 })
998}
999
1000pub(crate) fn gemm2_tier(tile: (usize, usize)) -> usize {
1002 GEMM2_TILES
1003 .iter()
1004 .position(|&t| t == tile)
1005 .expect("tile came from GEMM2_TILES")
1006}
1007pub(crate) fn gemm2_enabled() -> bool {
1010 !matches!(
1011 std::env::var("OSFKB_ENC_GEMM2").ok().as_deref(),
1012 Some("0") | Some("off")
1013 )
1014}
1015
1016pub(crate) const GEMM3_TILES: [(usize, usize, usize); 5] = [
1021 (128, 128, 16),
1022 (64, 128, 16),
1023 (64, 64, 16),
1024 (32, 64, 8),
1025 (16, 64, 8),
1026];
1027
1028fn gemm3_min_wgs() -> usize {
1036 static V: std::sync::OnceLock<usize> = std::sync::OnceLock::new();
1037 *V.get_or_init(|| {
1038 std::env::var("OSFKB_ENC_G3_MINWGS")
1039 .ok()
1040 .and_then(|v| v.parse().ok())
1041 .unwrap_or(64)
1042 })
1043}
1044
1045pub(crate) fn gemm3_grid_ok(m: usize, tile: (usize, usize, usize), wgs: usize) -> bool {
1060 wgs >= gemm3_min_wgs() || (wgs >= 48 && tile.0 >= 64 && m / tile.0 >= 3)
1061}
1062
1063pub(crate) fn gemm3_tile(m: usize, n: usize) -> (usize, usize, usize) {
1064 let wgs = |(bm, bn, _): (usize, usize, usize)| n.div_ceil(bn) * m.div_ceil(bm);
1065 let mut pick = GEMM3_TILES
1066 .into_iter()
1067 .find(|&t| gemm3_grid_ok(m, t, wgs(t)))
1068 .unwrap_or_else(|| {
1069 GEMM3_TILES
1070 .into_iter()
1071 .rev()
1072 .max_by_key(|&t| wgs(t))
1073 .expect("non-empty tile table")
1074 });
1075 if std::env::var("OSFKB_ENC_G3_DEMOTE").ok().as_deref() == Some("0") {
1085 return pick;
1086 }
1087 let mut idx = gemm3_tier(pick);
1088 while pick.0 >= 64 && wgs(pick) < 80 && idx + 1 < GEMM3_TILES.len() {
1089 let next = GEMM3_TILES[idx + 1];
1090 if wgs(next) >= 128 {
1091 pick = next;
1092 idx += 1;
1093 } else {
1094 break;
1095 }
1096 }
1097 pick
1098}
1099
1100pub(crate) fn gemm3_tier(tile: (usize, usize, usize)) -> usize {
1102 GEMM3_TILES
1103 .iter()
1104 .position(|&t| t == tile)
1105 .expect("tile came from GEMM3_TILES")
1106}
1107
1108pub(crate) fn gemm3_eligible(n: usize) -> bool {
1120 (n >= 1152 || n <= 512)
1121 && gemm2_enabled() && !matches!(
1123 std::env::var("OSFKB_ENC_GEMM3").ok().as_deref(),
1124 Some("0") | Some("off")
1125 )
1126}
1127
1128pub(crate) fn gemm3_sk_band(n: usize) -> bool {
1134 n > 512
1135 && n < 1152
1136 && n.is_multiple_of(64)
1137 && gemm2_enabled()
1138 && !matches!(
1139 std::env::var("OSFKB_ENC_GEMM3").ok().as_deref(),
1140 Some("0") | Some("off")
1141 )
1142 && std::env::var("OSFKB_ENC_GEMM3_SK").ok().as_deref() != Some("0")
1143}
1144
1145pub(crate) fn gemm3_smalln_sk(n: usize) -> bool {
1153 n <= 512
1154 && n.is_multiple_of(64)
1155 && std::env::var("OSFKB_ENC_SK_SMALLN").ok().as_deref() != Some("0")
1156 && std::env::var("OSFKB_ENC_GEMM3_SK").ok().as_deref() != Some("0")
1157}
1158
1159pub(crate) fn f16a_enabled() -> bool {
1164 std::env::var("OSFKB_ENC_F16A").ok().as_deref() == Some("1")
1165}
1166
1167pub(crate) fn gemm3_sk_chunks(k: usize) -> u32 {
1170 if k >= 2048 { 4 } else { 2 }
1171}
1172
1173pub(crate) fn gemm3_sk_plan(m: usize, n: usize, k: usize) -> (bool, u32) {
1183 if std::env::var("OSFKB_ENC_SK_DYNZ").ok().as_deref() == Some("0") {
1187 return (false, gemm3_sk_chunks(k));
1188 }
1189 let cols = n.div_ceil(GEMM3_SK_TILE.1);
1190 let zpick = |base: usize| {
1191 let mut z = 8u32;
1192 while z > 2 && base * z as usize > 160 {
1193 z >>= 1;
1194 }
1195 z.max(gemm3_sk_chunks(k))
1196 };
1197 let base32 = cols * m.div_ceil(GEMM3_SK_TILE32.0);
1198 let z32 = zpick(base32);
1199 if base32 * z32 as usize >= 96 && std::env::var("OSFKB_ENC_SK32").ok().as_deref() != Some("0") {
1200 return (true, z32);
1201 }
1202 (false, zpick(cols * m.div_ceil(GEMM3_SK_TILE.0)))
1203}
1204
1205pub(crate) const SK_PART_F32: usize = 8 * 81 * 1024;
1212
1213pub(crate) const GEMM3_SK_TILE: (usize, usize, usize) = (16, 64, 8);
1215
1216pub(crate) const GEMM3_SK_TILE32: (usize, usize, usize) = (32, 64, 8);
1220
1221const ENC_LAYERNORM: &str = r#"
1225struct Meta { h: u32, flags: u32, eps: f32, pad: u32 }
1226@group(0) @binding(0) var<storage, read> x: array<f32>;
1227@group(0) @binding(1) var<storage, read> res: array<f32>;
1228@group(0) @binding(2) var<storage, read> w: array<f32>;
1229@group(0) @binding(3) var<storage, read> b: array<f32>;
1230@group(0) @binding(4) var<storage, read_write> out: array<f32>;
1231@group(0) @binding(5) var<uniform> mt: Meta;
1232var<workgroup> sh: array<f32, 256>;
1233fn val(base: u32, i: u32) -> f32 {
1234 var v = x[base + i];
1235 if ((mt.flags & 1u) != 0u) { v += res[base + i]; }
1236 return v;
1237}
1238@compute @workgroup_size(256)
1239fn main(@builtin(workgroup_id) wg: vec3<u32>, @builtin(local_invocation_index) t: u32) {
1240 let base = wg.x * mt.h;
1241 var sum = 0.0;
1242 for (var i = t; i < mt.h; i += 256u) { sum += val(base, i); }
1243 sh[t] = sum;
1244 workgroupBarrier();
1245 for (var s = 128u; s > 0u; s >>= 1u) {
1246 if (t < s) { sh[t] += sh[t + s]; }
1247 workgroupBarrier();
1248 }
1249 // flags bit 2: RMS mode — no mean-centering (variance below becomes mean-of-squares).
1250 var mean = sh[0] / f32(mt.h);
1251 if ((mt.flags & 4u) != 0u) { mean = 0.0; }
1252 workgroupBarrier();
1253 var sq = 0.0;
1254 for (var i = t; i < mt.h; i += 256u) {
1255 let d = val(base, i) - mean;
1256 sq += d * d;
1257 }
1258 sh[t] = sq;
1259 workgroupBarrier();
1260 for (var s = 128u; s > 0u; s >>= 1u) {
1261 if (t < s) { sh[t] += sh[t + s]; }
1262 workgroupBarrier();
1263 }
1264 let inv = 1.0 / sqrt(sh[0] / f32(mt.h) + mt.eps);
1265 for (var i = t; i < mt.h; i += 256u) {
1266 var o = (val(base, i) - mean) * inv * w[i];
1267 if ((mt.flags & 2u) != 0u) { o += b[i]; }
1268 out[base + i] = o;
1269 }
1270}
1271"#;
1272
1273const ENC_ATTN: &str = r#"
1279struct Meta { nrows: u32, n_heads: u32, hd: u32, mode: u32, window: u32, n_kv_heads: u32, packed: u32, p2: u32 }
1280@group(0) @binding(0) var<storage, read> q: array<f32>;
1281@group(0) @binding(1) var<storage, read> k: array<f32>;
1282@group(0) @binding(2) var<storage, read> v: array<f32>;
1283@group(0) @binding(3) var<storage, read_write> out: array<f32>;
1284@group(0) @binding(4) var<storage, read> seq_starts: array<u32>;
1285@group(0) @binding(5) var<storage, read> seq_of: array<u32>;
1286@group(0) @binding(6) var<storage, read> valid: array<u32>;
1287@group(0) @binding(7) var<uniform> mt: Meta;
1288var<workgroup> qsh: array<f32, 128>;
1289var<workgroup> psh: array<f32, 64>;
1290var<workgroup> red: array<f32, 64>;
1291@compute @workgroup_size(64)
1292fn main(@builtin(workgroup_id) wg: vec3<u32>, @builtin(local_invocation_index) t: u32) {
1293 let i = wg.x;
1294 let head = wg.y;
1295 let h = mt.n_heads * mt.hd;
1296 // GQA: q spans n_heads, k/v span n_kv_heads; each kv head serves n_heads/n_kv_heads q heads.
1297 let kvh = mt.n_kv_heads * mt.hd;
1298 let kv_head = head / (mt.n_heads / mt.n_kv_heads);
1299 // mt.packed == 1: q/k/v are three VIEWS of one [T, q|k|v] buffer (the fused-QKV GEMM's
1300 // output, bound to all three slots) — row stride h+2·kvh, k at offset h, v at h+kvh.
1301 // mt.packed == 0: the historical three separate [T, ·] buffers.
1302 var qstride = h;
1303 var kstride = kvh;
1304 var koff = 0u;
1305 var voff = 0u;
1306 if (mt.packed == 1u) {
1307 qstride = h + 2u * kvh;
1308 kstride = h + 2u * kvh;
1309 koff = h;
1310 voff = h + kvh;
1311 }
1312 let qbase = i * qstride + head * mt.hd;
1313 let obase = i * h + head * mt.hd;
1314 for (var d = t; d < mt.hd; d += 64u) { qsh[d] = q[qbase + d]; }
1315 workgroupBarrier();
1316
1317 let s = seq_of[i];
1318 let a = seq_starts[s];
1319 let b = seq_starts[s + 1u];
1320 var jstart = a;
1321 var jend = b;
1322 if (mt.mode == 1u) { jend = min(jend, i + 1u); }
1323 if (mt.window > 0u) {
1324 let half_w = mt.window / 2u;
1325 if (i - a > half_w) { jstart = max(jstart, i - half_w); }
1326 jend = min(jend, i + half_w + 1u);
1327 }
1328
1329 let scale = 1.0 / sqrt(f32(mt.hd));
1330 var m_run = -3.0e38;
1331 var l_run = 0.0;
1332 // Per-thread output dims: d = t, t + 64 (hd ≤ 128).
1333 var acc0 = 0.0;
1334 var acc1 = 0.0;
1335 let ntiles = (jend - jstart + 63u) / 64u;
1336 for (var tile = 0u; tile < ntiles; tile++) {
1337 let j = jstart + tile * 64u + t;
1338 var score = -3.0e38;
1339 // valid[j] == 0 (ColBERT expansion pads): never an attention key.
1340 if (j < jend && valid[j] != 0u) {
1341 var dot = 0.0;
1342 let kbase = j * kstride + koff + kv_head * mt.hd;
1343 for (var d = 0u; d < mt.hd; d++) { dot += qsh[d] * k[kbase + d]; }
1344 score = dot * scale;
1345 }
1346 psh[t] = score;
1347 red[t] = score;
1348 workgroupBarrier();
1349 for (var r = 32u; r > 0u; r >>= 1u) {
1350 if (t < r) { red[t] = max(red[t], red[t + r]); }
1351 workgroupBarrier();
1352 }
1353 let tile_max = red[0];
1354 workgroupBarrier();
1355 let new_m = max(m_run, tile_max);
1356 let rescale = exp(m_run - new_m);
1357 // exp of masked lanes is exp(-inf) = 0 — they contribute nothing.
1358 var p = 0.0;
1359 if (psh[t] > -3.0e37) { p = exp(psh[t] - new_m); }
1360 psh[t] = p;
1361 red[t] = p;
1362 workgroupBarrier();
1363 for (var r = 32u; r > 0u; r >>= 1u) {
1364 if (t < r) { red[t] += red[t + r]; }
1365 workgroupBarrier();
1366 }
1367 l_run = l_run * rescale + red[0];
1368 workgroupBarrier();
1369 acc0 *= rescale;
1370 acc1 *= rescale;
1371 let tile_len = min(64u, jend - (jstart + tile * 64u));
1372 let d0 = t;
1373 let d1 = t + 64u;
1374 for (var jj = 0u; jj < tile_len; jj++) {
1375 let vbase = (jstart + tile * 64u + jj) * kstride + voff + kv_head * mt.hd;
1376 let p_j = psh[jj];
1377 if (d0 < mt.hd) { acc0 += p_j * v[vbase + d0]; }
1378 if (d1 < mt.hd) { acc1 += p_j * v[vbase + d1]; }
1379 }
1380 m_run = new_m;
1381 workgroupBarrier();
1382 }
1383 var inv = 0.0;
1384 if (l_run > 0.0) { inv = 1.0 / l_run; }
1385 if (t < mt.hd) { out[obase + t] = acc0 * inv; }
1386 if (t + 64u < mt.hd) { out[obase + t + 64u] = acc1 * inv; }
1387}
1388"#;
1389
1390const ENC_ATTN_V4: &str = r#"
1398struct Meta { nrows: u32, n_heads: u32, hd: u32, mode: u32, window: u32, n_kv_heads: u32, packed: u32, p2: u32 }
1399@group(0) @binding(0) var<storage, read> q: array<vec4<f32>>;
1400@group(0) @binding(1) var<storage, read> k: array<vec4<f32>>;
1401@group(0) @binding(2) var<storage, read> v: array<vec4<f32>>;
1402@group(0) @binding(3) var<storage, read_write> out: array<vec4<f32>>;
1403@group(0) @binding(4) var<storage, read> seq_starts: array<u32>;
1404@group(0) @binding(5) var<storage, read> seq_of: array<u32>;
1405@group(0) @binding(6) var<storage, read> valid: array<u32>;
1406@group(0) @binding(7) var<uniform> mt: Meta;
1407var<workgroup> qsh: array<vec4<f32>, 32>;
1408var<workgroup> psh: array<f32, 64>;
1409var<workgroup> red: array<f32, 64>;
1410@compute @workgroup_size(64)
1411fn main(@builtin(workgroup_id) wg: vec3<u32>, @builtin(local_invocation_index) t: u32) {
1412 let i = wg.x;
1413 let head = wg.y;
1414 let h = mt.n_heads * mt.hd;
1415 let kvh = mt.n_kv_heads * mt.hd;
1416 let kv_head = head / (mt.n_heads / mt.n_kv_heads);
1417 let hd4 = mt.hd / 4u;
1418 var qstride = h;
1419 var kstride = kvh;
1420 var koff = 0u;
1421 var voff = 0u;
1422 if (mt.packed == 1u) {
1423 qstride = h + 2u * kvh;
1424 kstride = h + 2u * kvh;
1425 koff = h;
1426 voff = h + kvh;
1427 }
1428 let qbase4 = (i * qstride + head * mt.hd) / 4u;
1429 let obase4 = (i * h + head * mt.hd) / 4u;
1430 for (var d = t; d < hd4; d += 64u) { qsh[d] = q[qbase4 + d]; }
1431 workgroupBarrier();
1432
1433 let s = seq_of[i];
1434 let a = seq_starts[s];
1435 let b = seq_starts[s + 1u];
1436 var jstart = a;
1437 var jend = b;
1438 if (mt.mode == 1u) { jend = min(jend, i + 1u); }
1439 if (mt.window > 0u) {
1440 let half_w = mt.window / 2u;
1441 if (i - a > half_w) { jstart = max(jstart, i - half_w); }
1442 jend = min(jend, i + half_w + 1u);
1443 }
1444
1445 let scale = 1.0 / sqrt(f32(mt.hd));
1446 var m_run = -3.0e38;
1447 var l_run = 0.0;
1448 // Per-thread output: dims [4t, 4t+4) — one vec4 (hd ≤ 128 ⇒ hd4 ≤ 32 active threads).
1449 var acc = vec4<f32>(0.0);
1450 let ntiles = (jend - jstart + 63u) / 64u;
1451 for (var tile = 0u; tile < ntiles; tile++) {
1452 let j = jstart + tile * 64u + t;
1453 var score = -3.0e38;
1454 if (j < jend && valid[j] != 0u) {
1455 var dot4 = 0.0;
1456 let kbase4 = (j * kstride + koff + kv_head * mt.hd) / 4u;
1457 for (var d = 0u; d < hd4; d++) { dot4 += dot(qsh[d], k[kbase4 + d]); }
1458 score = dot4 * scale;
1459 }
1460 psh[t] = score;
1461 red[t] = score;
1462 workgroupBarrier();
1463 for (var r = 32u; r > 0u; r >>= 1u) {
1464 if (t < r) { red[t] = max(red[t], red[t + r]); }
1465 workgroupBarrier();
1466 }
1467 let tile_max = red[0];
1468 workgroupBarrier();
1469 let new_m = max(m_run, tile_max);
1470 let rescale = exp(m_run - new_m);
1471 var p = 0.0;
1472 if (psh[t] > -3.0e37) { p = exp(psh[t] - new_m); }
1473 psh[t] = p;
1474 red[t] = p;
1475 workgroupBarrier();
1476 for (var r = 32u; r > 0u; r >>= 1u) {
1477 if (t < r) { red[t] += red[t + r]; }
1478 workgroupBarrier();
1479 }
1480 l_run = l_run * rescale + red[0];
1481 workgroupBarrier();
1482 acc *= rescale;
1483 let tile_len = min(64u, jend - (jstart + tile * 64u));
1484 for (var jj = 0u; jj < tile_len; jj++) {
1485 let vbase4 = ((jstart + tile * 64u + jj) * kstride + voff + kv_head * mt.hd) / 4u;
1486 let p_j = psh[jj];
1487 if (t < hd4) { acc += p_j * v[vbase4 + t]; }
1488 }
1489 m_run = new_m;
1490 workgroupBarrier();
1491 }
1492 var inv = 0.0;
1493 if (l_run > 0.0) { inv = 1.0 / l_run; }
1494 if (t < hd4) { out[obase4 + t] = acc * inv; }
1495}
1496"#;
1497
1498const ENC_ATTN_RB4: &str = r#"
1508struct Meta { nrows: u32, n_heads: u32, hd: u32, mode: u32, window: u32, n_kv_heads: u32, packed: u32, p2: u32 }
1509@group(0) @binding(0) var<storage, read> q: array<vec4<f32>>;
1510@group(0) @binding(1) var<storage, read> k: array<vec4<f32>>;
1511@group(0) @binding(2) var<storage, read> v: array<vec4<f32>>;
1512@group(0) @binding(3) var<storage, read_write> out: array<vec4<f32>>;
1513@group(0) @binding(4) var<storage, read> seq_starts: array<u32>;
1514@group(0) @binding(5) var<storage, read> seq_of: array<u32>;
1515@group(0) @binding(6) var<storage, read> valid: array<u32>;
1516@group(0) @binding(7) var<uniform> mt: Meta;
1517const RB: u32 = 4u;
1518var<workgroup> qsh: array<vec4<f32>, 64>; // [RB][hd/4 <= 16]
1519var<workgroup> ssh: array<f32, 256>; // [RB][64] scores -> probabilities
1520var<workgroup> red: array<f32, 64>; // [RB][16] reduction ladder
1521var<workgroup> ja_r: array<u32, 4>;
1522var<workgroup> jb_r: array<u32, 4>;
1523var<workgroup> m_run: array<f32, 4>;
1524var<workgroup> l_run: array<f32, 4>;
1525var<workgroup> m_new: array<f32, 4>;
1526var<workgroup> resc: array<f32, 4>;
1527@compute @workgroup_size(64)
1528fn main(@builtin(workgroup_id) wg: vec3<u32>, @builtin(local_invocation_index) t: u32) {
1529 let head = wg.y;
1530 let r0 = wg.x * RB;
1531 let h = mt.n_heads * mt.hd;
1532 let kvh = mt.n_kv_heads * mt.hd;
1533 let kv_head = head / (mt.n_heads / mt.n_kv_heads);
1534 let hd4 = mt.hd / 4u;
1535 var qstride = h;
1536 var kstride = kvh;
1537 var koff = 0u;
1538 var voff = 0u;
1539 if (mt.packed == 1u) {
1540 qstride = h + 2u * kvh;
1541 kstride = h + 2u * kvh;
1542 koff = h;
1543 voff = h + kvh;
1544 }
1545 // Per-row bounds (ragged: rows may sit in different sequences; window/causal are per row).
1546 if (t < RB) {
1547 let i = r0 + t;
1548 var ja = 0u;
1549 var jb = 0u;
1550 if (i < mt.nrows) {
1551 let s = seq_of[i];
1552 ja = seq_starts[s];
1553 jb = seq_starts[s + 1u];
1554 if (mt.mode == 1u) { jb = min(jb, i + 1u); }
1555 if (mt.window > 0u) {
1556 let half_w = mt.window / 2u;
1557 if (i - ja > half_w) { ja = max(ja, i - half_w); }
1558 jb = min(jb, i + half_w + 1u);
1559 }
1560 }
1561 ja_r[t] = ja;
1562 jb_r[t] = jb;
1563 m_run[t] = -3.0e38;
1564 l_run[t] = 0.0;
1565 }
1566 // Stage the RB query rows (cooperative: 64 threads over RB·hd4 <= 64 vec4s).
1567 if (t < RB * hd4) {
1568 let r = t / hd4;
1569 let d = t % hd4;
1570 let i = r0 + r;
1571 var qv = vec4<f32>(0.0);
1572 if (i < mt.nrows) { qv = q[(i * qstride + head * mt.hd) / 4u + d]; }
1573 qsh[r * hd4 + d] = qv;
1574 }
1575 workgroupBarrier();
1576 let jlo = min(min(ja_r[0], ja_r[1]), min(ja_r[2], ja_r[3]));
1577 let jhi = max(max(jb_r[0], jb_r[1]), max(jb_r[2], jb_r[3]));
1578
1579 let scale = 1.0 / sqrt(f32(mt.hd));
1580 // Thread (rr, dd) owns output dims [4·dd, 4·dd+4) of row r0+rr.
1581 let rr = t / 16u;
1582 let dd = t % 16u;
1583 var acc = vec4<f32>(0.0);
1584 let span = jhi - jlo;
1585 let ntiles = (span + 63u) / 64u;
1586 for (var tile = 0u; tile < ntiles; tile++) {
1587 let j = jlo + tile * 64u + t;
1588 // Score key j against ALL RB rows: k[j] is read ONCE and dotted four times.
1589 var d0 = -3.0e38;
1590 var d1 = -3.0e38;
1591 var d2 = -3.0e38;
1592 var d3 = -3.0e38;
1593 if (j < jhi && valid[j] != 0u) {
1594 var s0 = 0.0;
1595 var s1 = 0.0;
1596 var s2 = 0.0;
1597 var s3 = 0.0;
1598 let kbase4 = (j * kstride + koff + kv_head * mt.hd) / 4u;
1599 for (var d = 0u; d < hd4; d++) {
1600 let kv = k[kbase4 + d];
1601 s0 += dot(kv, qsh[d]);
1602 s1 += dot(kv, qsh[hd4 + d]);
1603 s2 += dot(kv, qsh[2u * hd4 + d]);
1604 s3 += dot(kv, qsh[3u * hd4 + d]);
1605 }
1606 if (j >= ja_r[0] && j < jb_r[0]) { d0 = s0 * scale; }
1607 if (j >= ja_r[1] && j < jb_r[1]) { d1 = s1 * scale; }
1608 if (j >= ja_r[2] && j < jb_r[2]) { d2 = s2 * scale; }
1609 if (j >= ja_r[3] && j < jb_r[3]) { d3 = s3 * scale; }
1610 }
1611 ssh[t] = d0;
1612 ssh[64u + t] = d1;
1613 ssh[128u + t] = d2;
1614 ssh[192u + t] = d3;
1615 workgroupBarrier();
1616 // Row max, all four rows at once: lane L of row R reduces keys L, L+16, L+32, L+48.
1617 let lane = dd;
1618 var mx = ssh[rr * 64u + lane];
1619 mx = max(mx, ssh[rr * 64u + lane + 16u]);
1620 mx = max(mx, ssh[rr * 64u + lane + 32u]);
1621 mx = max(mx, ssh[rr * 64u + lane + 48u]);
1622 red[t] = mx;
1623 workgroupBarrier();
1624 for (var r = 8u; r > 0u; r >>= 1u) {
1625 if (lane < r) { red[t] = max(red[t], red[t + r]); }
1626 workgroupBarrier();
1627 }
1628 if (t < RB) {
1629 let nm = max(m_run[t], red[t * 16u]);
1630 m_new[t] = nm;
1631 resc[t] = exp(m_run[t] - nm);
1632 }
1633 workgroupBarrier();
1634 // exp + row sums, same ladder.
1635 var ps = 0.0;
1636 for (var rr2 = 0u; rr2 < RB; rr2++) {
1637 let sc = ssh[rr2 * 64u + t];
1638 var p = 0.0;
1639 if (sc > -3.0e37) { p = exp(sc - m_new[rr2]); }
1640 ssh[rr2 * 64u + t] = p;
1641 }
1642 workgroupBarrier();
1643 ps = ssh[rr * 64u + lane] + ssh[rr * 64u + lane + 16u]
1644 + ssh[rr * 64u + lane + 32u] + ssh[rr * 64u + lane + 48u];
1645 red[t] = ps;
1646 workgroupBarrier();
1647 for (var r = 8u; r > 0u; r >>= 1u) {
1648 if (lane < r) { red[t] += red[t + r]; }
1649 workgroupBarrier();
1650 }
1651 if (t < RB) { l_run[t] = l_run[t] * resc[t] + red[t * 16u]; }
1652 // P·V: thread (rr, dd) accumulates its vec4 of row rr.
1653 acc *= resc[rr];
1654 let tile_len = min(64u, jhi - (jlo + tile * 64u));
1655 if (dd < hd4) {
1656 for (var jj = 0u; jj < tile_len; jj++) {
1657 let p_j = ssh[rr * 64u + jj];
1658 let vbase4 = ((jlo + tile * 64u + jj) * kstride + voff + kv_head * mt.hd) / 4u;
1659 acc += p_j * v[vbase4 + dd];
1660 }
1661 }
1662 if (t < RB) { m_run[t] = m_new[t]; }
1663 workgroupBarrier();
1664 }
1665 let i = r0 + rr;
1666 if (i < mt.nrows && dd < hd4) {
1667 var inv = 0.0;
1668 if (l_run[rr] > 0.0) { inv = 1.0 / l_run[rr]; }
1669 out[(i * h + head * mt.hd) / 4u + dd] = acc * inv;
1670 }
1671}
1672"#;
1673
1674const ENC_ATTN_DISENT: &str = r#"
1682struct Meta { nrows: u32, n_heads: u32, hd: u32, span: u32, max_rel: u32, sf: u32, c2p: u32, p2c: u32 }
1683@group(0) @binding(0) var<storage, read> q: array<f32>;
1684@group(0) @binding(1) var<storage, read> k: array<f32>;
1685@group(0) @binding(2) var<storage, read> v: array<f32>;
1686@group(0) @binding(3) var<storage, read_write> out: array<f32>;
1687@group(0) @binding(4) var<storage, read> seq_starts: array<u32>;
1688@group(0) @binding(5) var<storage, read> seq_of: array<u32>;
1689@group(0) @binding(6) var<storage, read> valid: array<u32>;
1690@group(0) @binding(7) var<storage, read> pos_k: array<f32>;
1691@group(0) @binding(8) var<storage, read> pos_q: array<f32>;
1692@group(0) @binding(9) var<uniform> mt: Meta;
1693var<workgroup> qsh: array<f32, 128>;
1694var<workgroup> psh: array<f32, 64>;
1695var<workgroup> red: array<f32, 64>;
1696// HF make_log_bucket_position: identity within ±span/2, logarithmic (ceil) beyond. Odd in rel,
1697// which is why one index m serves both the c2p and p2c gathers.
1698fn bucket(rel: i32, span: u32, max_rel: u32) -> i32 {
1699 let mid = i32(span / 2u);
1700 let a = abs(rel);
1701 if (a <= mid) { return rel; }
1702 let mid_f = f32(mid);
1703 let lp = ceil(log(f32(a) / mid_f) / log((f32(max_rel) - 1.0) / mid_f) * (mid_f - 1.0)) + mid_f;
1704 let sgn = select(1.0, -1.0, rel < 0);
1705 return i32(sgn * lp);
1706}
1707@compute @workgroup_size(64)
1708fn main(@builtin(workgroup_id) wg: vec3<u32>, @builtin(local_invocation_index) t: u32) {
1709 let i = wg.x;
1710 let head = wg.y;
1711 let h = mt.n_heads * mt.hd;
1712 let p2c_on = (mt.p2c & 1u) != 0u;
1713 var qkstride = h;
1714 var koff = 0u;
1715 var voff = 0u;
1716 if ((mt.p2c & 2u) != 0u) {
1717 qkstride = 3u * h;
1718 koff = h;
1719 voff = 2u * h;
1720 }
1721 let qbase = i * qkstride + head * mt.hd;
1722 let obase = i * h + head * mt.hd;
1723 for (var d = t; d < mt.hd; d += 64u) { qsh[d] = q[qbase + d]; }
1724 workgroupBarrier();
1725
1726 let s = seq_of[i];
1727 let a = seq_starts[s];
1728 let b = seq_starts[s + 1u];
1729 let qi_local = i32(i - a);
1730 let two_span = i32(2u * mt.span);
1731
1732 let scale = 1.0 / sqrt(f32(mt.hd) * f32(mt.sf));
1733 var m_run = -3.0e38;
1734 var l_run = 0.0;
1735 var acc0 = 0.0;
1736 var acc1 = 0.0;
1737 let ntiles = (b - a + 63u) / 64u;
1738 for (var tile = 0u; tile < ntiles; tile++) {
1739 let j = a + tile * 64u + t;
1740 var score = -3.0e38;
1741 if (j < b && valid[j] != 0u) {
1742 let kbase = j * qkstride + koff + head * mt.hd;
1743 // Relative-table row for (i, j): shared by the c2p and p2c terms.
1744 let m = u32(clamp(bucket(qi_local - i32(j - a), mt.span, mt.max_rel) + i32(mt.span),
1745 0, two_span - 1));
1746 let rbase = m * h + head * mt.hd;
1747 var dot = 0.0;
1748 for (var d = 0u; d < mt.hd; d++) {
1749 let kd = k[kbase + d];
1750 dot += qsh[d] * kd; // content → content
1751 if (mt.c2p != 0u) { dot += qsh[d] * pos_k[rbase + d]; } // content q → position k
1752 if (p2c_on) { dot += kd * pos_q[rbase + d]; } // position q -> content k
1753 }
1754 score = dot * scale;
1755 }
1756 psh[t] = score;
1757 red[t] = score;
1758 workgroupBarrier();
1759 for (var r = 32u; r > 0u; r >>= 1u) {
1760 if (t < r) { red[t] = max(red[t], red[t + r]); }
1761 workgroupBarrier();
1762 }
1763 let tile_max = red[0];
1764 workgroupBarrier();
1765 let new_m = max(m_run, tile_max);
1766 let rescale = exp(m_run - new_m);
1767 var p = 0.0;
1768 if (psh[t] > -3.0e37) { p = exp(psh[t] - new_m); }
1769 psh[t] = p;
1770 red[t] = p;
1771 workgroupBarrier();
1772 for (var r = 32u; r > 0u; r >>= 1u) {
1773 if (t < r) { red[t] += red[t + r]; }
1774 workgroupBarrier();
1775 }
1776 l_run = l_run * rescale + red[0];
1777 workgroupBarrier();
1778 acc0 *= rescale;
1779 acc1 *= rescale;
1780 let tile_len = min(64u, b - (a + tile * 64u));
1781 let d0 = t;
1782 let d1 = t + 64u;
1783 for (var jj = 0u; jj < tile_len; jj++) {
1784 let vbase = (a + tile * 64u + jj) * qkstride + voff + head * mt.hd;
1785 let p_j = psh[jj];
1786 if (d0 < mt.hd) { acc0 += p_j * v[vbase + d0]; }
1787 if (d1 < mt.hd) { acc1 += p_j * v[vbase + d1]; }
1788 }
1789 m_run = new_m;
1790 workgroupBarrier();
1791 }
1792 var inv = 0.0;
1793 if (l_run > 0.0) { inv = 1.0 / l_run; }
1794 if (t < mt.hd) { out[obase + t] = acc0 * inv; }
1795 if (t + 64u < mt.hd) { out[obase + t + 64u] = acc1 * inv; }
1796}
1797"#;
1798
1799const ENC_ATTN_DISENT_RB4: &str = r#"
1807struct Meta { nrows: u32, n_heads: u32, hd: u32, span: u32, max_rel: u32, sf: u32, c2p: u32, p2c: u32 }
1808@group(0) @binding(0) var<storage, read> q: array<vec4<f32>>;
1809@group(0) @binding(1) var<storage, read> k: array<vec4<f32>>;
1810@group(0) @binding(2) var<storage, read> v: array<vec4<f32>>;
1811@group(0) @binding(3) var<storage, read_write> out: array<vec4<f32>>;
1812@group(0) @binding(4) var<storage, read> seq_starts: array<u32>;
1813@group(0) @binding(5) var<storage, read> seq_of: array<u32>;
1814@group(0) @binding(6) var<storage, read> valid: array<u32>;
1815@group(0) @binding(7) var<storage, read> pos_k: array<vec4<f32>>;
1816@group(0) @binding(8) var<storage, read> pos_q: array<vec4<f32>>;
1817@group(0) @binding(9) var<uniform> mt: Meta;
1818const RB: u32 = 4u;
1819var<workgroup> qsh: array<vec4<f32>, 64>; // [RB][hd/4 <= 16]
1820var<workgroup> ssh: array<f32, 256>; // [RB][64]
1821var<workgroup> red: array<f32, 64>;
1822var<workgroup> ja_r: array<u32, 4>;
1823var<workgroup> jb_r: array<u32, 4>;
1824var<workgroup> qi_r: array<i32, 4>;
1825var<workgroup> m_run: array<f32, 4>;
1826var<workgroup> l_run: array<f32, 4>;
1827var<workgroup> m_new: array<f32, 4>;
1828var<workgroup> resc: array<f32, 4>;
1829fn bucket(rel: i32, span: u32, max_rel: u32) -> i32 {
1830 let mid = i32(span / 2u);
1831 let a = abs(rel);
1832 if (a <= mid) { return rel; }
1833 let mid_f = f32(mid);
1834 let lp = ceil(log(f32(a) / mid_f) / log((f32(max_rel) - 1.0) / mid_f) * (mid_f - 1.0)) + mid_f;
1835 let sgn = select(1.0, -1.0, rel < 0);
1836 return i32(sgn * lp);
1837}
1838@compute @workgroup_size(64)
1839fn main(@builtin(workgroup_id) wg: vec3<u32>, @builtin(local_invocation_index) t: u32) {
1840 let head = wg.y;
1841 let r0 = wg.x * RB;
1842 let h = mt.n_heads * mt.hd;
1843 let hd4 = mt.hd / 4u;
1844 // mt.p2c: bit 0 = the p2c score term; bit 1 = PACKED q|k|v (the fused-QKV GEMM's output
1845 // bound to all three slots — same trick as ENC_ATTN's packed mode; no GQA here).
1846 let p2c_on = (mt.p2c & 1u) != 0u;
1847 var qkstride = h;
1848 var koff = 0u;
1849 var voff = 0u;
1850 if ((mt.p2c & 2u) != 0u) {
1851 qkstride = 3u * h;
1852 koff = h;
1853 voff = 2u * h;
1854 }
1855 if (t < RB) {
1856 let i = r0 + t;
1857 var ja = 0u;
1858 var jb = 0u;
1859 var qi = 0;
1860 if (i < mt.nrows) {
1861 let s = seq_of[i];
1862 ja = seq_starts[s];
1863 jb = seq_starts[s + 1u];
1864 qi = i32(i - ja);
1865 }
1866 ja_r[t] = ja;
1867 jb_r[t] = jb;
1868 qi_r[t] = qi;
1869 m_run[t] = -3.0e38;
1870 l_run[t] = 0.0;
1871 }
1872 if (t < RB * hd4) {
1873 let r = t / hd4;
1874 let d = t % hd4;
1875 let i = r0 + r;
1876 var qv = vec4<f32>(0.0);
1877 if (i < mt.nrows) { qv = q[(i * qkstride + head * mt.hd) / 4u + d]; }
1878 qsh[r * hd4 + d] = qv;
1879 }
1880 workgroupBarrier();
1881 let jlo = min(min(ja_r[0], ja_r[1]), min(ja_r[2], ja_r[3]));
1882 let jhi = max(max(jb_r[0], jb_r[1]), max(jb_r[2], jb_r[3]));
1883
1884 let scale = 1.0 / sqrt(f32(mt.hd) * f32(mt.sf));
1885 let two_span = i32(2u * mt.span);
1886 let rr = t / 16u;
1887 let dd = t % 16u;
1888 var acc = vec4<f32>(0.0);
1889 let span_keys = jhi - jlo;
1890 let ntiles = (span_keys + 63u) / 64u;
1891 for (var tile = 0u; tile < ntiles; tile++) {
1892 let j = jlo + tile * 64u + t;
1893 var sc0 = -3.0e38;
1894 var sc1 = -3.0e38;
1895 var sc2 = -3.0e38;
1896 var sc3 = -3.0e38;
1897 if (j < jhi && valid[j] != 0u) {
1898 let kbase4 = (j * qkstride + koff + head * mt.hd) / 4u;
1899 // Content dots: k[j] read ONCE, dotted with all four staged q rows.
1900 var c0 = 0.0;
1901 var c1 = 0.0;
1902 var c2 = 0.0;
1903 var c3 = 0.0;
1904 for (var d = 0u; d < hd4; d++) {
1905 let kv = k[kbase4 + d];
1906 c0 += dot(kv, qsh[d]);
1907 c1 += dot(kv, qsh[hd4 + d]);
1908 c2 += dot(kv, qsh[2u * hd4 + d]);
1909 c3 += dot(kv, qsh[3u * hd4 + d]);
1910 }
1911 // Positional terms per row (the LUT row m depends on i−j — not shareable).
1912 for (var r = 0u; r < RB; r++) {
1913 if (j < ja_r[r] || j >= jb_r[r]) { continue; }
1914 let m = u32(clamp(
1915 bucket(qi_r[r] - i32(j - ja_r[r]), mt.span, mt.max_rel) + i32(mt.span),
1916 0, two_span - 1));
1917 let rbase4 = (m * h + head * mt.hd) / 4u;
1918 var pterm = 0.0;
1919 for (var d = 0u; d < hd4; d++) {
1920 if (mt.c2p != 0u) { pterm += dot(qsh[r * hd4 + d], pos_k[rbase4 + d]); }
1921 if (p2c_on) { pterm += dot(k[kbase4 + d], pos_q[rbase4 + d]); }
1922 }
1923 var content = c0;
1924 if (r == 1u) { content = c1; }
1925 if (r == 2u) { content = c2; }
1926 if (r == 3u) { content = c3; }
1927 let sc = (content + pterm) * scale;
1928 if (r == 0u) { sc0 = sc; }
1929 if (r == 1u) { sc1 = sc; }
1930 if (r == 2u) { sc2 = sc; }
1931 if (r == 3u) { sc3 = sc; }
1932 }
1933 }
1934 ssh[t] = sc0;
1935 ssh[64u + t] = sc1;
1936 ssh[128u + t] = sc2;
1937 ssh[192u + t] = sc3;
1938 workgroupBarrier();
1939 let lane = dd;
1940 var mx = ssh[rr * 64u + lane];
1941 mx = max(mx, ssh[rr * 64u + lane + 16u]);
1942 mx = max(mx, ssh[rr * 64u + lane + 32u]);
1943 mx = max(mx, ssh[rr * 64u + lane + 48u]);
1944 red[t] = mx;
1945 workgroupBarrier();
1946 for (var r = 8u; r > 0u; r >>= 1u) {
1947 if (lane < r) { red[t] = max(red[t], red[t + r]); }
1948 workgroupBarrier();
1949 }
1950 if (t < RB) {
1951 let nm = max(m_run[t], red[t * 16u]);
1952 m_new[t] = nm;
1953 resc[t] = exp(m_run[t] - nm);
1954 }
1955 workgroupBarrier();
1956 for (var rr2 = 0u; rr2 < RB; rr2++) {
1957 let sc = ssh[rr2 * 64u + t];
1958 var p = 0.0;
1959 if (sc > -3.0e37) { p = exp(sc - m_new[rr2]); }
1960 ssh[rr2 * 64u + t] = p;
1961 }
1962 workgroupBarrier();
1963 let ps = ssh[rr * 64u + lane] + ssh[rr * 64u + lane + 16u]
1964 + ssh[rr * 64u + lane + 32u] + ssh[rr * 64u + lane + 48u];
1965 red[t] = ps;
1966 workgroupBarrier();
1967 for (var r = 8u; r > 0u; r >>= 1u) {
1968 if (lane < r) { red[t] += red[t + r]; }
1969 workgroupBarrier();
1970 }
1971 if (t < RB) { l_run[t] = l_run[t] * resc[t] + red[t * 16u]; }
1972 acc *= resc[rr];
1973 let tile_len = min(64u, jhi - (jlo + tile * 64u));
1974 if (dd < hd4) {
1975 for (var jj = 0u; jj < tile_len; jj++) {
1976 let p_j = ssh[rr * 64u + jj];
1977 let vbase4 = ((jlo + tile * 64u + jj) * qkstride + voff + head * mt.hd) / 4u;
1978 acc += p_j * v[vbase4 + dd];
1979 }
1980 }
1981 if (t < RB) { m_run[t] = m_new[t]; }
1982 workgroupBarrier();
1983 }
1984 let i = r0 + rr;
1985 if (i < mt.nrows && dd < hd4) {
1986 var inv = 0.0;
1987 if (l_run[rr] > 0.0) { inv = 1.0 / l_run[rr]; }
1988 out[(i * h + head * mt.hd) / 4u + dd] = acc * inv;
1989 }
1990}
1991"#;
1992
1993const ENC_MEAN_POOL: &str = r#"
1995struct Meta { h: u32, pad0: u32, pad1: u32, pad2: u32 }
1996@group(0) @binding(0) var<storage, read> hidden: array<f32>;
1997@group(0) @binding(1) var<storage, read_write> out: array<f32>;
1998@group(0) @binding(2) var<storage, read> seq_starts: array<u32>;
1999@group(0) @binding(3) var<uniform> mt: Meta;
2000@compute @workgroup_size(256)
2001fn main(@builtin(workgroup_id) wg: vec3<u32>, @builtin(local_invocation_index) t: u32) {
2002 let s = wg.x;
2003 let a = seq_starts[s];
2004 let b = seq_starts[s + 1u];
2005 let inv = 1.0 / f32(max(b - a, 1u));
2006 for (var i = t; i < mt.h; i += 256u) {
2007 var sum = 0.0;
2008 for (var r = a; r < b; r++) { sum += hidden[r * mt.h + i]; }
2009 out[s * mt.h + i] = sum * inv;
2010 }
2011}
2012"#;
2013
2014const ENC_L2NORM: &str = r#"
2017struct Meta { dim: u32, pad0: u32, pad1: u32, pad2: u32 }
2018@group(0) @binding(0) var<storage, read_write> buf: array<f32>;
2019@group(0) @binding(1) var<uniform> mt: Meta;
2020var<workgroup> sh: array<f32, 256>;
2021@compute @workgroup_size(256)
2022fn main(@builtin(workgroup_id) wg: vec3<u32>, @builtin(local_invocation_index) t: u32) {
2023 let base = wg.x * mt.dim;
2024 var sq = 0.0;
2025 for (var i = t; i < mt.dim; i += 256u) {
2026 let v = buf[base + i];
2027 sq += v * v;
2028 }
2029 sh[t] = sq;
2030 workgroupBarrier();
2031 for (var s = 128u; s > 0u; s >>= 1u) {
2032 if (t < s) { sh[t] += sh[t + s]; }
2033 workgroupBarrier();
2034 }
2035 let n = sqrt(sh[0]);
2036 if (n > 0.0) {
2037 let inv = 1.0 / n;
2038 for (var i = t; i < mt.dim; i += 256u) { buf[base + i] *= inv; }
2039 }
2040}
2041"#;
2042
2043const ROPE_ENC: &str = r#"
2047struct Meta { nrows: u32, nh: u32, hd: u32, theta: f32 }
2048@group(0) @binding(0) var<storage, read_write> x: array<f32>;
2049@group(0) @binding(1) var<storage, read> seq_starts: array<u32>;
2050@group(0) @binding(2) var<storage, read> seq_of: array<u32>;
2051@group(0) @binding(3) var<uniform> mt: Meta;
2052@compute @workgroup_size(64)
2053fn main(@builtin(workgroup_id) wg: vec3<u32>, @builtin(local_invocation_index) t: u32) {
2054 let i = wg.x;
2055 let p = f32(i - seq_starts[seq_of[i]]);
2056 let half = mt.hd / 2u;
2057 let pairs = mt.nh * half;
2058 let h = mt.nh * mt.hd;
2059 for (var pj = t; pj < pairs; pj += 64u) {
2060 let head = pj / half;
2061 let j = pj % half;
2062 let freq = pow(mt.theta, -2.0 * f32(j) / f32(mt.hd));
2063 let ang = p * freq;
2064 let c = cos(ang);
2065 let s = sin(ang);
2066 let base = i * h + head * mt.hd;
2067 let a = x[base + j];
2068 let b = x[base + j + half];
2069 x[base + j] = a * c - b * s;
2070 x[base + j + half] = a * s + b * c;
2071 }
2072}
2073"#;
2074
2075const GLU_SPLIT: &str = const_format_glu();
2078
2079const fn const_format_glu() -> &'static str {
2080 r#"
2083struct Meta { i_width: u32, act: u32, p0: u32, p1: u32 }
2084@group(0) @binding(0) var<storage, read> mid: array<f32>;
2085@group(0) @binding(1) var<storage, read_write> out: array<f32>;
2086@group(0) @binding(2) var<uniform> mt: Meta;
2087fn erf_as(x: f32) -> f32 {
2088 let s = sign(x);
2089 let a = abs(x);
2090 let t = 1.0 / (1.0 + 0.3275911 * a);
2091 let y = 1.0 - (((((1.061405429 * t - 1.453152027) * t) + 1.421413741) * t - 0.284496736) * t
2092 + 0.254829592) * t * exp(-a * a);
2093 return s * y;
2094}
2095fn apply_act(v: f32, code: u32) -> f32 {
2096 switch code {
2097 case 1u: { return 0.5 * v * (1.0 + erf_as(v * 0.70710678)); }
2098 // tanh-GELU. The argument is CLAMPED, exactly as the decoder's kernels already do
2099 // (forward.rs / lib.rs) — WGSL's `tanh` overflows to NaN on some backends for |arg| >~ 88,
2100 // because a naive (e^2x - 1)/(e^2x + 1) expansion divides inf by inf. `arg` reaches 105 at a
2101 // pre-activation of only 13.8, which the text encoders never hit and a ViT's outlier tokens
2102 // hit on the first block. tanh is ±1 to well inside f32 epsilon by |arg| = 20, so this is
2103 // EXACT, not an approximation.
2104 case 2u: {
2105 let a = clamp(0.7978845608 * (v + 0.044715 * v * v * v), -20.0, 20.0);
2106 return 0.5 * v * (1.0 + tanh(a));
2107 }
2108 case 3u: { return v / (1.0 + exp(-v)); }
2109 case 4u: { return tanh(v); }
2110 case 5u: { return max(v, 0.0); }
2111 default: { return v; }
2112 }
2113}
2114@compute @workgroup_size(256)
2115fn main(@builtin(workgroup_id) wg: vec3<u32>, @builtin(local_invocation_index) t: u32) {
2116 let row = wg.x;
2117 for (var j = t; j < mt.i_width; j += 256u) {
2118 let a = mid[row * 2u * mt.i_width + j];
2119 let g = mid[row * 2u * mt.i_width + mt.i_width + j];
2120 out[row * mt.i_width + j] = apply_act(a, mt.act) * g;
2121 }
2122}
2123"#
2124}
2125
2126const ENC_ADD: &str = r#"
2129struct Meta { n: u32, p0: u32, p1: u32, p2: u32 }
2130@group(0) @binding(0) var<storage, read_write> dst: array<f32>;
2131@group(0) @binding(1) var<storage, read> src: array<f32>;
2132@group(0) @binding(2) var<uniform> mt: Meta;
2133@compute @workgroup_size(256)
2134fn main(@builtin(global_invocation_id) gid: vec3<u32>) {
2135 if (gid.x < mt.n) { dst[gid.x] += src[gid.x]; }
2136}
2137"#;
2138
2139const ENC_QK_NORM: &str = r#"
2142struct Meta { nrows: u32, n_heads: u32, hd: u32, eps: f32 }
2143@group(0) @binding(0) var<storage, read_write> x: array<f32>;
2144@group(0) @binding(1) var<storage, read> w: array<f32>;
2145@group(0) @binding(2) var<uniform> mt: Meta;
2146@compute @workgroup_size(64)
2147fn main(@builtin(workgroup_id) wg: vec3<u32>, @builtin(local_invocation_index) t: u32) {
2148 let i = wg.x;
2149 let h = mt.n_heads * mt.hd;
2150 for (var head = t; head < mt.n_heads; head += 64u) {
2151 let base = i * h + head * mt.hd;
2152 var ms = 0.0;
2153 for (var d = 0u; d < mt.hd; d++) { let v = x[base + d]; ms += v * v; }
2154 let inv = 1.0 / sqrt(ms / f32(mt.hd) + mt.eps);
2155 for (var d = 0u; d < mt.hd; d++) { x[base + d] = x[base + d] * inv * w[d]; }
2156 }
2157}
2158"#;
2159
2160const ENC_LAST_POOL: &str = r#"
2162struct Meta { h: u32, p0: u32, p1: u32, p2: u32 }
2163@group(0) @binding(0) var<storage, read> hidden: array<f32>;
2164@group(0) @binding(1) var<storage, read_write> out: array<f32>;
2165@group(0) @binding(2) var<storage, read> seq_starts: array<u32>;
2166@group(0) @binding(3) var<uniform> mt: Meta;
2167@compute @workgroup_size(256)
2168fn main(@builtin(workgroup_id) wg: vec3<u32>, @builtin(local_invocation_index) t: u32) {
2169 let s = wg.x;
2170 let last = seq_starts[s + 1u] - 1u;
2171 for (var i = t; i < mt.h; i += 256u) {
2172 out[s * mt.h + i] = hidden[last * mt.h + i];
2173 }
2174}
2175"#;
2176
2177const ENC_CLS_POOL: &str = r#"
2180struct Meta { h: u32, p0: u32, p1: u32, p2: u32 }
2181@group(0) @binding(0) var<storage, read> hidden: array<f32>;
2182@group(0) @binding(1) var<storage, read_write> out: array<f32>;
2183@group(0) @binding(2) var<storage, read> seq_starts: array<u32>;
2184@group(0) @binding(3) var<uniform> mt: Meta;
2185@compute @workgroup_size(256)
2186fn main(@builtin(workgroup_id) wg: vec3<u32>, @builtin(local_invocation_index) t: u32) {
2187 let s = wg.x;
2188 let first = seq_starts[s];
2189 for (var i = t; i < mt.h; i += 256u) {
2190 out[s * mt.h + i] = hidden[first * mt.h + i];
2191 }
2192}
2193"#;
2194
2195const ENC_COPY: &str = r#"
2199struct Meta { n: u32, p0: u32, p1: u32, p2: u32 }
2200@group(0) @binding(0) var<storage, read_write> dst: array<f32>;
2201@group(0) @binding(1) var<storage, read> src: array<f32>;
2202@group(0) @binding(2) var<uniform> mt: Meta;
2203@compute @workgroup_size(256)
2204fn main(@builtin(global_invocation_id) gid: vec3<u32>) {
2205 if (gid.x < mt.n) { dst[gid.x] = src[gid.x]; }
2206}
2207"#;
2208
2209const ENC_CONV: &str = r#"
2213struct Meta { h: u32, l: u32, p0: u32, p1: u32 }
2214@group(0) @binding(0) var<storage, read> bcx: array<f32>;
2215@group(0) @binding(1) var<storage, read> conv_w: array<f32>;
2216@group(0) @binding(2) var<storage, read_write> y: array<f32>;
2217@group(0) @binding(3) var<storage, read> seq_starts: array<u32>;
2218@group(0) @binding(4) var<storage, read> seq_of: array<u32>;
2219@group(0) @binding(5) var<storage, read> valid: array<u32>;
2220@group(0) @binding(6) var<uniform> mt: Meta;
2221@compute @workgroup_size(256)
2222fn main(@builtin(workgroup_id) wg: vec3<u32>, @builtin(local_invocation_index) t: u32) {
2223 let i = wg.x;
2224 let a = seq_starts[seq_of[i]];
2225 let w3 = 3u * mt.h;
2226 // HF zeroes the conv-block input at padded positions: invalid row ⇒ zero output.
2227 let own_valid = valid[i] != 0u;
2228 for (var c = t; c < mt.h; c += 256u) {
2229 if (!own_valid) {
2230 y[i * mt.h + c] = 0.0;
2231 continue;
2232 }
2233 let cc = bcx[i * w3 + mt.h + c];
2234 var acc = 0.0;
2235 for (var k = 0u; k < mt.l; k++) {
2236 let back = mt.l - 1u - k;
2237 if (i >= a + back) {
2238 let j = i - back;
2239 if (valid[j] != 0u) {
2240 let bx = bcx[j * w3 + c] * bcx[j * w3 + 2u * mt.h + c];
2241 acc += conv_w[c * mt.l + k] * bx;
2242 }
2243 }
2244 }
2245 y[i * mt.h + c] = cc * acc;
2246 }
2247}
2248"#;
2249
2250pub struct EncKernels {
2253 gemm_f32: wgpu::ComputePipeline,
2254 gemm_f16: Option<wgpu::ComputePipeline>,
2255 gemm2_f32: Vec<wgpu::ComputePipeline>,
2258 gemm2_f16: Option<Vec<wgpu::ComputePipeline>>,
2259 gemm3_f32: Vec<wgpu::ComputePipeline>,
2262 gemm3_f16: Option<Vec<wgpu::ComputePipeline>>,
2263 gemm3_sk_f32: wgpu::ComputePipeline,
2266 gemm3_sk_f16: Option<wgpu::ComputePipeline>,
2267 gemm3_sk32_f32: wgpu::ComputePipeline,
2269 gemm3_sk32_f16: Option<wgpu::ComputePipeline>,
2270 gemm3_sk_reduce: wgpu::ComputePipeline,
2271 gemm3_f16a: Option<Vec<wgpu::ComputePipeline>>,
2273 gemm3_sk_f16a: Option<wgpu::ComputePipeline>,
2275 gemm3_sk32_f16a: Option<wgpu::ComputePipeline>,
2276 gemm4_f16w: Option<wgpu::ComputePipeline>,
2279 layernorm: wgpu::ComputePipeline,
2280 attn: wgpu::ComputePipeline,
2281 attn4: wgpu::ComputePipeline,
2284 attn_rb4: wgpu::ComputePipeline,
2286 attn_disent: wgpu::ComputePipeline,
2287 attn_disent_rb4: wgpu::ComputePipeline,
2289 mean_pool: wgpu::ComputePipeline,
2290 l2norm: wgpu::ComputePipeline,
2291 rope: wgpu::ComputePipeline,
2292 glu: wgpu::ComputePipeline,
2293 add: wgpu::ComputePipeline,
2294 qk_norm: wgpu::ComputePipeline,
2295 last_pool: wgpu::ComputePipeline,
2296 cls_pool: wgpu::ComputePipeline,
2297 copy: wgpu::ComputePipeline,
2298 conv: wgpu::ComputePipeline,
2299}
2300
2301pub fn gelu_tanh_on_gpu(ctx: &GpuCtx, xs: &[f32]) -> anyhow::Result<Vec<f32>> {
2305 let src = format!(
2306 "{ACT_FNS}\n\
2307 @group(0) @binding(0) var<storage, read_write> v: array<f32>;\n\
2308 @compute @workgroup_size(64)\n\
2309 fn main(@builtin(global_invocation_id) g: vec3<u32>) {{\n\
2310 \x20 if (g.x < {}u) {{ v[g.x] = apply_act(v[g.x], 2u); }}\n\
2311 }}\n",
2312 xs.len()
2313 );
2314 let module = ctx
2315 .device
2316 .shader_module_tuned(wgpu::ShaderModuleDescriptor {
2317 label: Some("gelu_probe"),
2318 source: wgpu::ShaderSource::Wgsl(src.into()),
2319 });
2320 let pl = ctx
2321 .device
2322 .create_compute_pipeline(&wgpu::ComputePipelineDescriptor {
2323 label: Some("gelu_probe"),
2324 layout: None,
2325 module: &module,
2326 entry_point: Some("main"),
2327 compilation_options: Default::default(),
2328 cache: None,
2329 });
2330 let buf = ctx.storage(xs);
2331 let bg = make_bg1_pub(ctx, &pl, &buf);
2332 dispatch(ctx, &pl, &bg, (xs.len() as u32).div_ceil(64), 1);
2333 ctx.read(&buf, xs.len())
2334}
2335
2336fn make_bg1_pub(ctx: &GpuCtx, pl: &wgpu::ComputePipeline, buf: &wgpu::Buffer) -> wgpu::BindGroup {
2337 ctx.device.create_bind_group(&wgpu::BindGroupDescriptor {
2338 label: None,
2339 layout: &pl.get_bind_group_layout(0),
2340 entries: &[wgpu::BindGroupEntry {
2341 binding: 0,
2342 resource: buf.as_entire_binding(),
2343 }],
2344 })
2345}
2346
2347impl EncKernels {
2348 pub fn new(ctx: &GpuCtx) -> Self {
2350 Self {
2351 gemm_f32: pipeline(ctx, "enc_gemm_f32", &enc_gemm_src(false)),
2352 gemm_f16: ctx
2353 .f16
2354 .then(|| pipeline(ctx, "enc_gemm_f16", &enc_gemm_src(true))),
2355 gemm2_f32: GEMM2_TILES
2356 .iter()
2357 .map(|&(bm, bn)| {
2358 pipeline(
2359 ctx,
2360 &format!("enc_gemm2_{bm}x{bn}_f32"),
2361 &enc_gemm2_src(false, bm, bn),
2362 )
2363 })
2364 .collect(),
2365 gemm2_f16: ctx.f16.then(|| {
2366 GEMM2_TILES
2367 .iter()
2368 .map(|&(bm, bn)| {
2369 pipeline(
2370 ctx,
2371 &format!("enc_gemm2_{bm}x{bn}_f16"),
2372 &enc_gemm2_src(true, bm, bn),
2373 )
2374 })
2375 .collect()
2376 }),
2377 gemm3_f32: GEMM3_TILES
2378 .iter()
2379 .map(|&(bm, bn, bk)| {
2380 pipeline(
2381 ctx,
2382 &format!("enc_gemm3_{bm}x{bn}x{bk}_f32"),
2383 &enc_gemm3_src(false, bm, bn, bk),
2384 )
2385 })
2386 .collect(),
2387 gemm3_f16: ctx.f16.then(|| {
2388 GEMM3_TILES
2389 .iter()
2390 .map(|&(bm, bn, bk)| {
2391 pipeline(
2392 ctx,
2393 &format!("enc_gemm3_{bm}x{bn}x{bk}_f16"),
2394 &enc_gemm3_src(true, bm, bn, bk),
2395 )
2396 })
2397 .collect()
2398 }),
2399 gemm3_sk_f32: pipeline(
2400 ctx,
2401 "enc_gemm3_sk_f32",
2402 &enc_gemm3_sk_src(false, GEMM3_SK_TILE.0, GEMM3_SK_TILE.1, GEMM3_SK_TILE.2),
2403 ),
2404 gemm3_sk_f16: ctx.f16.then(|| {
2405 pipeline(
2406 ctx,
2407 "enc_gemm3_sk_f16",
2408 &enc_gemm3_sk_src(true, GEMM3_SK_TILE.0, GEMM3_SK_TILE.1, GEMM3_SK_TILE.2),
2409 )
2410 }),
2411 gemm3_sk32_f32: pipeline(
2412 ctx,
2413 "enc_gemm3_sk32_f32",
2414 &enc_gemm3_sk_src(
2415 false,
2416 GEMM3_SK_TILE32.0,
2417 GEMM3_SK_TILE32.1,
2418 GEMM3_SK_TILE32.2,
2419 ),
2420 ),
2421 gemm3_sk32_f16: ctx.f16.then(|| {
2422 pipeline(
2423 ctx,
2424 "enc_gemm3_sk32_f16",
2425 &enc_gemm3_sk_src(
2426 true,
2427 GEMM3_SK_TILE32.0,
2428 GEMM3_SK_TILE32.1,
2429 GEMM3_SK_TILE32.2,
2430 ),
2431 )
2432 }),
2433 gemm3_sk_reduce: pipeline(ctx, "enc_gemm3_sk_reduce", &enc_gemm3_sk_reduce_src()),
2434 gemm3_f16a: ctx.f16.then(|| {
2435 GEMM3_TILES
2436 .iter()
2437 .map(|&(bm, bn, bk)| {
2438 pipeline(
2439 ctx,
2440 &format!("enc_gemm3_f16a_{bm}x{bn}x{bk}"),
2441 &enc_gemm3_f16a_src(bm, bn, bk),
2442 )
2443 })
2444 .collect()
2445 }),
2446 gemm3_sk_f16a: ctx.f16.then(|| {
2447 pipeline(
2448 ctx,
2449 "enc_gemm3_sk_f16a",
2450 &enc_gemm3_sk_f16a_src(GEMM3_SK_TILE.0, GEMM3_SK_TILE.1, GEMM3_SK_TILE.2),
2451 )
2452 }),
2453 gemm3_sk32_f16a: ctx.f16.then(|| {
2454 pipeline(
2455 ctx,
2456 "enc_gemm3_sk32_f16a",
2457 &enc_gemm3_sk_f16a_src(GEMM3_SK_TILE32.0, GEMM3_SK_TILE32.1, GEMM3_SK_TILE32.2),
2458 )
2459 }),
2460 gemm4_f16w: (ctx.f16 && ctx.coop_matrix)
2461 .then(|| pipeline(ctx, "enc_gemm4_f16w", &enc_gemm4_f16w_src(128, 128, 8))),
2462 layernorm: pipeline(ctx, "enc_layernorm", ENC_LAYERNORM),
2463 attn: pipeline(ctx, "enc_attn", ENC_ATTN),
2464 attn4: pipeline(ctx, "enc_attn_v4", ENC_ATTN_V4),
2465 attn_rb4: pipeline(ctx, "enc_attn_rb4", ENC_ATTN_RB4),
2466 attn_disent: pipeline(ctx, "enc_attn_disent", ENC_ATTN_DISENT),
2467 attn_disent_rb4: pipeline(ctx, "enc_attn_disent_rb4", ENC_ATTN_DISENT_RB4),
2468 mean_pool: pipeline(ctx, "enc_mean_pool", ENC_MEAN_POOL),
2469 l2norm: pipeline(ctx, "enc_l2norm", ENC_L2NORM),
2470 rope: pipeline(ctx, "enc_rope", ROPE_ENC),
2471 glu: pipeline(ctx, "enc_glu", GLU_SPLIT),
2472 add: pipeline(ctx, "enc_add", ENC_ADD),
2473 qk_norm: pipeline(ctx, "enc_qk_norm", ENC_QK_NORM),
2474 last_pool: pipeline(ctx, "enc_last_pool", ENC_LAST_POOL),
2475 cls_pool: pipeline(ctx, "enc_cls_pool", ENC_CLS_POOL),
2476 copy: pipeline(ctx, "enc_copy", ENC_COPY),
2477 conv: pipeline(ctx, "enc_conv", ENC_CONV),
2478 }
2479 }
2480
2481 pub fn rope_for_tests(
2483 &self,
2484 ctx: &GpuCtx,
2485 x: &[f32],
2486 seq_starts: &[u32],
2487 n_heads: usize,
2488 hd: usize,
2489 theta: f32,
2490 ) -> Result<Vec<f32>> {
2491 let h = n_heads * hd;
2492 let nrows = x.len() / h;
2493 let xb = ctx.storage(x);
2494 let sb = ctx.storage_bytes(bytemuck::cast_slice(seq_starts));
2495 let mb = ctx.storage_bytes(bytemuck::cast_slice(&row_to_seq(seq_starts, nrows)));
2496 let meta = uni(
2497 ctx,
2498 bytemuck::cast_slice(&[nrows as u32, n_heads as u32, hd as u32, theta.to_bits()]),
2499 );
2500 let bg = make_bg(ctx, &self.rope, &[&xb, &sb, &mb], &meta);
2501 dispatch(ctx, &self.rope, &bg, nrows as u32, 1);
2502 ctx.read(&xb, x.len())
2503 }
2504
2505 pub fn glu_for_tests(
2507 &self,
2508 ctx: &GpuCtx,
2509 mid: &[f32],
2510 i_width: usize,
2511 act: Act,
2512 ) -> Result<Vec<f32>> {
2513 let rows = mid.len() / (2 * i_width);
2514 let mb = ctx.storage(mid);
2515 let ob = ctx.storage(&vec![0f32; rows * i_width]);
2516 let meta = uni(
2517 ctx,
2518 bytemuck::cast_slice(&[i_width as u32, act_code(Some(act)), 0, 0]),
2519 );
2520 let bg = make_bg(ctx, &self.glu, &[&mb, &ob], &meta);
2521 dispatch(ctx, &self.glu, &bg, rows as u32, 1);
2522 ctx.read(&ob, rows * i_width)
2523 }
2524
2525 pub(crate) fn ln_pl(&self) -> &wgpu::ComputePipeline {
2532 &self.layernorm
2533 }
2534 pub(crate) fn attn_pl(&self) -> &wgpu::ComputePipeline {
2535 &self.attn
2536 }
2537 pub(crate) fn add_pl(&self) -> &wgpu::ComputePipeline {
2538 &self.add
2539 }
2540
2541 pub(crate) fn gemm2_pipeline_tile(
2542 &self,
2543 f16: bool,
2544 tile: (usize, usize),
2545 ) -> Result<&wgpu::ComputePipeline> {
2546 let set = if f16 {
2547 self.gemm2_f16
2548 .as_ref()
2549 .ok_or_else(|| anyhow::anyhow!("adapter has no SHADER_F16"))?
2550 } else {
2551 &self.gemm2_f32
2552 };
2553 Ok(&set[gemm2_tier(tile)])
2554 }
2555
2556 pub(crate) fn gemm3_f16a_pipeline_tile(
2558 &self,
2559 tile: (usize, usize, usize),
2560 ) -> Option<&wgpu::ComputePipeline> {
2561 self.gemm3_f16a.as_ref().map(|set| &set[gemm3_tier(tile)])
2562 }
2563
2564 pub(crate) fn gemm3_sk32_pipeline(&self, f16: bool) -> Result<&wgpu::ComputePipeline> {
2566 if f16 {
2567 self.gemm3_sk32_f16
2568 .as_ref()
2569 .ok_or_else(|| anyhow::anyhow!("adapter has no SHADER_F16"))
2570 } else {
2571 Ok(&self.gemm3_sk32_f32)
2572 }
2573 }
2574
2575 pub(crate) fn gemm3_sk_pipeline(&self, f16: bool) -> Result<&wgpu::ComputePipeline> {
2577 if f16 {
2578 self.gemm3_sk_f16
2579 .as_ref()
2580 .ok_or_else(|| anyhow::anyhow!("adapter has no SHADER_F16"))
2581 } else {
2582 Ok(&self.gemm3_sk_f32)
2583 }
2584 }
2585
2586 pub(crate) fn gemm3_pipeline_tile(
2588 &self,
2589 f16: bool,
2590 tile: (usize, usize, usize),
2591 ) -> Result<&wgpu::ComputePipeline> {
2592 let set = if f16 {
2593 self.gemm3_f16
2594 .as_ref()
2595 .ok_or_else(|| anyhow::anyhow!("adapter has no SHADER_F16"))?
2596 } else {
2597 &self.gemm3_f32
2598 };
2599 Ok(&set[gemm3_tier(tile)])
2600 }
2601
2602 pub(crate) fn gemm_pipeline(&self, f16: bool) -> Result<&wgpu::ComputePipeline> {
2604 if f16 {
2605 self.gemm_f16
2606 .as_ref()
2607 .ok_or_else(|| anyhow::anyhow!("adapter has no SHADER_F16"))
2608 } else {
2609 Ok(&self.gemm_f32)
2610 }
2611 }
2612
2613 #[allow(clippy::too_many_arguments)]
2617 pub fn gemm_for_tests(
2618 &self,
2619 ctx: &GpuCtx,
2620 x: &[f32],
2621 w: &[f32],
2622 bias: Option<&[f32]>,
2623 m: usize,
2624 n: usize,
2625 k: usize,
2626 act: Option<Act>,
2627 w_f16: bool,
2628 ) -> Result<Vec<f32>> {
2629 let (pl, bm, bn) = if gemm2_enabled() {
2632 let t = gemm2_tile(m, n);
2633 (self.gemm2_pipeline_tile(w_f16, t)?, t.0, t.1)
2634 } else {
2635 (self.gemm_pipeline(w_f16)?, 16, 16)
2636 };
2637 let xb = ctx.storage(x);
2638 let wb = if w_f16 {
2639 let h: Vec<u8> = w
2640 .iter()
2641 .flat_map(|&v| half::f16::from_f32(v).to_le_bytes())
2642 .collect();
2643 ctx.storage_bytes(&h)
2644 } else {
2645 ctx.storage(w)
2646 };
2647 let bb = ctx.storage(bias.unwrap_or(&[0.0]));
2648 let yb = ctx.storage(&vec![0f32; m * n]);
2649 let flags = u32::from(bias.is_some()) | (act_code(act) << 8);
2650 let meta = uni(
2651 ctx,
2652 bytemuck::cast_slice(&[m as u32, n as u32, k as u32, flags]),
2653 );
2654 let bg = make_bg(ctx, pl, &[&xb, &wb, &bb, &yb], &meta);
2655 dispatch(
2656 ctx,
2657 pl,
2658 &bg,
2659 (n as u32).div_ceil(bn as u32),
2660 (m as u32).div_ceil(bm as u32),
2661 );
2662 ctx.read(&yb, m * n)
2663 }
2664
2665 #[allow(clippy::too_many_arguments)]
2675 pub fn gemm_resident(
2676 &self,
2677 ctx: &GpuCtx,
2678 x: &[f32],
2679 wb: &wgpu::Buffer,
2680 bb: &wgpu::Buffer,
2681 m: usize,
2682 n: usize,
2683 k: usize,
2684 act: Option<Act>,
2685 ) -> Result<Vec<f32>> {
2686 let (pl, bm, bn) = if gemm2_enabled() {
2688 let t = gemm2_tile(m, n);
2689 (self.gemm2_pipeline_tile(false, t)?, t.0, t.1)
2690 } else {
2691 (self.gemm_pipeline(false)?, 16, 16)
2692 };
2693 let xb = ctx.storage(x);
2694 let yb = ctx.storage(&vec![0f32; m * n]);
2695 let flags = 1u32 | (act_code(act) << 8); let meta = uni(
2697 ctx,
2698 bytemuck::cast_slice(&[m as u32, n as u32, k as u32, flags]),
2699 );
2700 let bg = make_bg(ctx, pl, &[&xb, wb, bb, &yb], &meta);
2701 dispatch(
2702 ctx,
2703 pl,
2704 &bg,
2705 (n as u32).div_ceil(bn as u32),
2706 (m as u32).div_ceil(bm as u32),
2707 );
2708 ctx.read(&yb, m * n)
2709 }
2710
2711 #[allow(clippy::too_many_arguments)]
2720 pub fn gemm_resident_into(
2721 &self,
2722 ctx: &GpuCtx,
2723 xb: &wgpu::Buffer,
2724 wb: &wgpu::Buffer,
2725 bb: &wgpu::Buffer,
2726 yb: &wgpu::Buffer,
2727 m: usize,
2728 n: usize,
2729 k: usize,
2730 act: Option<Act>,
2731 ) -> Result<()> {
2732 let (pl, bm, bn) = if gemm2_enabled() {
2733 let t = gemm2_tile(m, n);
2734 (self.gemm2_pipeline_tile(false, t)?, t.0, t.1)
2735 } else {
2736 (self.gemm_pipeline(false)?, 16, 16)
2737 };
2738 let flags = 1u32 | (act_code(act) << 8); let meta = uni(
2740 ctx,
2741 bytemuck::cast_slice(&[m as u32, n as u32, k as u32, flags]),
2742 );
2743 let bg = make_bg(ctx, pl, &[xb, wb, bb, yb], &meta);
2744 dispatch(
2745 ctx,
2746 pl,
2747 &bg,
2748 (n as u32).div_ceil(bn as u32),
2749 (m as u32).div_ceil(bm as u32),
2750 );
2751 Ok(())
2752 }
2753
2754 #[allow(clippy::too_many_arguments)]
2760 pub fn gemm_bench(
2761 &self,
2762 ctx: &GpuCtx,
2763 m: usize,
2764 n: usize,
2765 k: usize,
2766 w_f16: bool,
2767 v2: bool,
2768 reps: usize,
2769 ) -> Result<f64> {
2770 let (pl, bm, bn) = if v2 {
2771 let t = gemm2_tile(m, n);
2772 (self.gemm2_pipeline_tile(w_f16, t)?, t.0, t.1)
2773 } else {
2774 (self.gemm_pipeline(w_f16)?, 16, 16)
2775 };
2776 let xb = ctx.storage(&vec![0.5f32; m * k]);
2777 let wb = if w_f16 {
2778 ctx.storage_bytes(&vec![0u8; n * k * 2])
2779 } else {
2780 ctx.storage(&vec![0.25f32; n * k])
2781 };
2782 let bb = ctx.storage(&[0.0f32]);
2783 let yb = ctx.storage(&vec![0f32; m * n]);
2784 let meta = uni(
2785 ctx,
2786 bytemuck::cast_slice(&[m as u32, n as u32, k as u32, 0u32]),
2787 );
2788 let bg = make_bg(ctx, pl, &[&xb, &wb, &bb, &yb], &meta);
2789 let (gx, gy) = (
2790 (n as u32).div_ceil(bn as u32),
2791 (m as u32).div_ceil(bm as u32),
2792 );
2793
2794 let run = |reps: usize| -> f64 {
2795 let mut enc = ctx
2796 .device
2797 .create_command_encoder(&wgpu::CommandEncoderDescriptor { label: None });
2798 {
2799 let mut pass = enc.begin_compute_pass(&wgpu::ComputePassDescriptor::default());
2800 pass.set_pipeline(pl);
2801 pass.set_bind_group(0, &bg, &[]);
2802 for _ in 0..reps {
2803 pass.dispatch_workgroups(gx, gy, 1);
2804 }
2805 }
2806 let t0 = std::time::Instant::now();
2807 ctx.queue.submit([enc.finish()]);
2808 let _ = ctx.device.poll(wgpu::PollType::wait_indefinitely());
2809 t0.elapsed().as_secs_f64()
2810 };
2811 run(2); let mut best = f64::MAX;
2813 for _ in 0..5 {
2814 best = best.min(run(reps));
2815 }
2816 let flop = 2.0 * m as f64 * n as f64 * k as f64 * reps as f64;
2817 Ok(flop / best / 1e9)
2818 }
2819
2820 #[allow(clippy::too_many_arguments)]
2826 pub fn gemm3_bench(
2827 &self,
2828 ctx: &GpuCtx,
2829 m: usize,
2830 n: usize,
2831 k: usize,
2832 bm: usize,
2833 bn: usize,
2834 bk: usize,
2835 reps: usize,
2836 ) -> Result<f64> {
2837 let pl = pipeline(
2838 ctx,
2839 &format!("enc_gemm3_{bm}x{bn}x{bk}_lab"),
2840 &enc_gemm3_src(false, bm, bn, bk),
2841 );
2842 let xb = ctx.storage(&vec![0.5f32; m * k]);
2843 let wb = ctx.storage(&vec![0.25f32; k * n]); let bb = ctx.storage(&[0.0f32]);
2845 let yb = ctx.storage(&vec![0f32; m * n]);
2846 let meta = uni(
2847 ctx,
2848 bytemuck::cast_slice(&[m as u32, n as u32, k as u32, 0u32]),
2849 );
2850 let bg = make_bg(ctx, &pl, &[&xb, &wb, &bb, &yb], &meta);
2851 let (gx, gy) = (
2852 (n as u32).div_ceil(bn as u32),
2853 (m as u32).div_ceil(bm as u32),
2854 );
2855 let run = |reps: usize| -> f64 {
2856 let mut enc = ctx
2857 .device
2858 .create_command_encoder(&wgpu::CommandEncoderDescriptor { label: None });
2859 {
2860 let mut pass = enc.begin_compute_pass(&wgpu::ComputePassDescriptor::default());
2861 pass.set_pipeline(&pl);
2862 pass.set_bind_group(0, &bg, &[]);
2863 for _ in 0..reps {
2864 pass.dispatch_workgroups(gx, gy, 1);
2865 }
2866 }
2867 let t0 = std::time::Instant::now();
2868 ctx.queue.submit([enc.finish()]);
2869 let _ = ctx.device.poll(wgpu::PollType::wait_indefinitely());
2870 t0.elapsed().as_secs_f64()
2871 };
2872 run(2);
2873 let mut best = f64::MAX;
2874 for _ in 0..5 {
2875 best = best.min(run(reps));
2876 }
2877 let flop = 2.0 * m as f64 * n as f64 * k as f64 * reps as f64;
2878 Ok(flop / best / 1e9)
2879 }
2880
2881 #[allow(clippy::too_many_arguments)]
2886 pub fn gemm3_sk_bench(
2887 &self,
2888 ctx: &GpuCtx,
2889 m: usize,
2890 n: usize,
2891 k: usize,
2892 bm: usize,
2893 bn: usize,
2894 bk: usize,
2895 chunks: u32,
2896 reps: usize,
2897 ) -> Result<f64> {
2898 let pl = pipeline(
2899 ctx,
2900 &format!("enc_gemm3_sk_{bm}x{bn}x{bk}_lab"),
2901 &enc_gemm3_sk_src(false, bm, bn, bk),
2902 );
2903 let rpl = pipeline(ctx, "enc_gemm3_sk_reduce_lab", &enc_gemm3_sk_reduce_src());
2904 let xb = ctx.storage(&vec![0.5f32; m * k]);
2905 let wb = ctx.storage(&vec![0.25f32; k * n]); let bb = ctx.storage(&[0.0f32]);
2907 let part = ctx.storage(&vec![0f32; chunks as usize * m * n]);
2908 let yb = ctx.storage(&vec![0f32; m * n]);
2909 let meta = uni(
2910 ctx,
2911 bytemuck::cast_slice(&[m as u32, n as u32, k as u32, 0u32]),
2912 );
2913 let meta_r = uni(
2914 ctx,
2915 bytemuck::cast_slice(&[m as u32, n as u32, 0u32, chunks]),
2916 );
2917 let bg = make_bg(ctx, &pl, &[&xb, &wb, &bb, &part], &meta);
2918 let rbg = make_bg(ctx, &rpl, &[&part, &bb, &yb], &meta_r);
2919 let (gx, gy) = (
2920 (n as u32).div_ceil(bn as u32),
2921 (m as u32).div_ceil(bm as u32),
2922 );
2923 let rgx = ((m * n) as u32).div_ceil(256);
2924 let run = |reps: usize| -> f64 {
2925 let mut enc = ctx
2926 .device
2927 .create_command_encoder(&wgpu::CommandEncoderDescriptor { label: None });
2928 {
2929 let mut pass = enc.begin_compute_pass(&wgpu::ComputePassDescriptor::default());
2930 for _ in 0..reps {
2931 pass.set_pipeline(&pl);
2932 pass.set_bind_group(0, &bg, &[]);
2933 pass.dispatch_workgroups(gx, gy, chunks);
2934 pass.set_pipeline(&rpl);
2935 pass.set_bind_group(0, &rbg, &[]);
2936 pass.dispatch_workgroups(rgx, 1, 1);
2937 }
2938 }
2939 let t0 = std::time::Instant::now();
2940 ctx.queue.submit([enc.finish()]);
2941 let _ = ctx.device.poll(wgpu::PollType::wait_indefinitely());
2942 t0.elapsed().as_secs_f64()
2943 };
2944 run(2);
2945 let mut best = f64::MAX;
2946 for _ in 0..5 {
2947 best = best.min(run(reps));
2948 }
2949 let flop = 2.0 * m as f64 * n as f64 * k as f64 * reps as f64;
2950 Ok(flop / best / 1e9)
2951 }
2952
2953 #[allow(clippy::too_many_arguments)]
2955 pub fn layernorm_for_tests(
2956 &self,
2957 ctx: &GpuCtx,
2958 x: &[f32],
2959 res: Option<&[f32]>,
2960 w: &[f32],
2961 b: Option<&[f32]>,
2962 t: usize,
2963 h: usize,
2964 eps: f32,
2965 ) -> Result<Vec<f32>> {
2966 let xb = ctx.storage(x);
2967 let rb = ctx.storage(res.unwrap_or(&[0.0]));
2968 let wb = ctx.storage(w);
2969 let bb = ctx.storage(b.unwrap_or(&[0.0]));
2970 let ob = ctx.storage(&vec![0f32; t * h]);
2971 let flags = u32::from(res.is_some()) | (u32::from(b.is_some()) << 1);
2972 let meta = uni(
2973 ctx,
2974 bytemuck::cast_slice(&[h as u32, flags, eps.to_bits(), 0]),
2975 );
2976 let bg = make_bg(ctx, &self.layernorm, &[&xb, &rb, &wb, &bb, &ob], &meta);
2977 dispatch(ctx, &self.layernorm, &bg, t as u32, 1);
2978 ctx.read(&ob, t * h)
2979 }
2980
2981 #[allow(clippy::too_many_arguments)]
2983 #[allow(clippy::too_many_arguments)]
2984 pub fn attention_for_tests(
2985 &self,
2986 ctx: &GpuCtx,
2987 q: &[f32],
2988 k: &[f32],
2989 v: &[f32],
2990 seq_starts: &[u32],
2991 n_heads: usize,
2992 n_kv_heads: usize,
2993 hd: usize,
2994 mask: MaskKind,
2995 window: u32,
2996 ) -> Result<Vec<f32>> {
2997 anyhow::ensure!(hd <= 128, "encoder attention supports head_dim ≤ 128");
2998 let h = n_heads * hd;
2999 let nrows = q.len() / h;
3000 let seq_of = row_to_seq(seq_starts, nrows);
3001 let qb = ctx.storage(q);
3002 let kb = ctx.storage(k);
3003 let vb = ctx.storage(v);
3004 let ob = ctx.storage(&vec![0f32; nrows * h]);
3005 let sb = ctx.storage_bytes(bytemuck::cast_slice(seq_starts));
3006 let mb = ctx.storage_bytes(bytemuck::cast_slice(&seq_of));
3007 let mode = match mask {
3008 MaskKind::Bidirectional => 0u32,
3009 MaskKind::Causal => 1u32,
3010 };
3011 let meta = uni(
3012 ctx,
3013 bytemuck::cast_slice(&[
3014 nrows as u32,
3015 n_heads as u32,
3016 hd as u32,
3017 mode,
3018 window,
3019 n_kv_heads as u32,
3020 0,
3021 0,
3022 ]),
3023 );
3024 let valid = ctx.storage_bytes(bytemuck::cast_slice(&vec![1u32; nrows]));
3025 let bg = make_bg(
3026 ctx,
3027 &self.attn,
3028 &[&qb, &kb, &vb, &ob, &sb, &mb, &valid],
3029 &meta,
3030 );
3031 dispatch(ctx, &self.attn, &bg, nrows as u32, n_heads as u32);
3032 ctx.read(&ob, nrows * h)
3033 }
3034
3035 pub fn mean_pool_l2_for_tests(
3037 &self,
3038 ctx: &GpuCtx,
3039 hidden: &[f32],
3040 h: usize,
3041 seq_starts: &[u32],
3042 ) -> Result<Vec<Vec<f32>>> {
3043 let n_seqs = seq_starts.len() - 1;
3044 let hb = ctx.storage(hidden);
3045 let ob = ctx.storage(&vec![0f32; n_seqs * h]);
3046 let sb = ctx.storage_bytes(bytemuck::cast_slice(seq_starts));
3047 let meta = uni(ctx, bytemuck::cast_slice(&[h as u32, 0, 0, 0]));
3048 let bg = make_bg(ctx, &self.mean_pool, &[&hb, &ob, &sb], &meta);
3049 dispatch(ctx, &self.mean_pool, &bg, n_seqs as u32, 1);
3050 let meta2 = uni(ctx, bytemuck::cast_slice(&[h as u32, 0, 0, 0]));
3051 let bg2 = make_bg(ctx, &self.l2norm, &[&ob], &meta2);
3052 dispatch(ctx, &self.l2norm, &bg2, n_seqs as u32, 1);
3053 let flat = ctx.read(&ob, n_seqs * h)?;
3054 Ok(flat.chunks_exact(h).map(<[f32]>::to_vec).collect())
3055 }
3056}
3057
3058fn cpu_layer_norm_rows(x: &mut [f32], h: usize, w: &[f32], b: &[f32], eps: f32) {
3064 for row in x.chunks_exact_mut(h) {
3065 let mean = row.iter().sum::<f32>() / h as f32;
3066 let var = row.iter().map(|v| (v - mean) * (v - mean)).sum::<f32>() / h as f32;
3067 let inv = 1.0 / (var + eps).sqrt();
3068 for (j, v) in row.iter_mut().enumerate() {
3069 *v = (*v - mean) * inv * w[j] + b[j];
3070 }
3071 }
3072}
3073
3074fn cpu_linear(x: &[f32], w: &[f32], bias: Option<&[f32]>, n: usize, k: usize) -> Vec<f32> {
3077 let rows = x.len() / k;
3078 let mut out = vec![0f32; rows * n];
3079 for (o, xr) in out.chunks_exact_mut(n).zip(x.chunks_exact(k)) {
3081 for (nn, o_n) in o.iter_mut().enumerate() {
3082 let wr = &w[nn * k..(nn + 1) * k];
3083 let mut acc = bias.map_or(0.0, |b| b[nn]);
3084 for (xk, wk) in xr.iter().zip(wr) {
3085 acc += xk * wk;
3086 }
3087 *o_n = acc;
3088 }
3089 }
3090 out
3091}
3092
3093pub(crate) fn dispatch(
3094 ctx: &GpuCtx,
3095 pl: &wgpu::ComputePipeline,
3096 bg: &wgpu::BindGroup,
3097 gx: u32,
3098 gy: u32,
3099) {
3100 let mut enc = ctx
3101 .device
3102 .create_command_encoder(&wgpu::CommandEncoderDescriptor { label: None });
3103 {
3104 let mut pass = enc.begin_compute_pass(&wgpu::ComputePassDescriptor::default());
3105 pass.set_pipeline(pl);
3106 pass.set_bind_group(0, bg, &[]);
3107 pass.dispatch_workgroups(gx, gy, 1);
3108 }
3109 ctx.queue.submit([enc.finish()]);
3110}
3111
3112pub(crate) fn row_to_seq(seq_starts: &[u32], t: usize) -> Vec<u32> {
3114 let mut map = vec![0u32; t];
3115 for (s, w) in seq_starts.windows(2).enumerate() {
3116 for r in w[0]..w[1] {
3117 map[r as usize] = s as u32;
3118 }
3119 }
3120 map
3121}
3122
3123use crate::encoder_weights::{EncArch, EncBatch, EncoderConfig, MlpKind, NormKind, PosKind};
3128use crate::pooling::{EmbedOut, Pooling};
3129use crate::weights::{LazySt, f32_to_f16_bytes};
3130use anyhow::Context;
3131use std::path::Path;
3132
3133enum EncMatBuf {
3135 F16(wgpu::Buffer),
3136 F32(wgpu::Buffer),
3137}
3138
3139struct EncGpuLinear {
3141 w: EncMatBuf,
3142 b: Option<wgpu::Buffer>,
3143 n: u32,
3144 k: u32,
3145 v3: bool,
3148 sk: bool,
3151 f16a: bool,
3154 b_host: Option<Vec<f32>>,
3157}
3158
3159struct CeHeadGpu {
3171 bg_pooler: GemmBg,
3173 bg_classifier: GemmBg,
3175 metas: Vec<wgpu::Buffer>,
3177 out: wgpu::Buffer,
3179 n_labels: usize,
3180 fused: Option<(wgpu::ComputePipeline, wgpu::BindGroup)>,
3183 _mid: wgpu::Buffer,
3185 _weights: Vec<EncGpuLinear>,
3186}
3187
3188const ENC_CE_FUSED: &str = r#"
3194struct Meta { h: u32, n_labels: u32, p0: u32, p1: u32 }
3195@group(0) @binding(0) var<storage, read> hidden: array<f32>; // [T, h]
3196@group(0) @binding(1) var<storage, read> seq_starts: array<u32>;
3197@group(0) @binding(2) var<storage, read> pw: array<f32>; // pooler [h, h] (HF [n,k])
3198@group(0) @binding(3) var<storage, read> pb: array<f32>; // pooler bias [h]
3199@group(0) @binding(4) var<storage, read> cw: array<f32>; // classifier [n_labels, h]
3200@group(0) @binding(5) var<storage, read> cb: array<f32>; // classifier bias [n_labels]
3201@group(0) @binding(6) var<storage, read_write> out: array<f32>; // [b, n_labels]
3202@group(0) @binding(7) var<uniform> mt: Meta;
3203var<workgroup> cls: array<f32, 2048>;
3204var<workgroup> pooled: array<f32, 2048>;
3205var<workgroup> red: array<f32, 256>;
3206@compute @workgroup_size(256)
3207fn main(@builtin(workgroup_id) wg: vec3<u32>, @builtin(local_invocation_index) t: u32) {
3208 let s = wg.x;
3209 let h = mt.h;
3210 let base = seq_starts[s] * h; // the sequence's FIRST token = CLS
3211 for (var j = t; j < h; j += 256u) { cls[j] = hidden[base + j]; }
3212 workgroupBarrier();
3213 // pooled[j] = tanh(pw[j]·cls + pb[j]) — each thread owns rows j, j+256, …
3214 for (var j = t; j < h; j += 256u) {
3215 var acc = 0.0;
3216 let rb = j * h;
3217 for (var d = 0u; d < h; d++) { acc += pw[rb + d] * cls[d]; }
3218 pooled[j] = tanh(clamp(acc + pb[j], -20.0, 20.0));
3219 }
3220 workgroupBarrier();
3221 // classifier: one ladder-reduced dot per label (n_labels is 1-2 in practice).
3222 for (var l = 0u; l < mt.n_labels; l++) {
3223 var part = 0.0;
3224 let rb = l * h;
3225 for (var d = t; d < h; d += 256u) { part += cw[rb + d] * pooled[d]; }
3226 red[t] = part;
3227 workgroupBarrier();
3228 for (var r = 128u; r > 0u; r >>= 1u) {
3229 if (t < r) { red[t] += red[t + r]; }
3230 workgroupBarrier();
3231 }
3232 if (t == 0u) { out[s * mt.n_labels + l] = red[0] + cb[l]; }
3233 workgroupBarrier();
3234 }
3235}
3236"#;
3237
3238pub(crate) fn v2_bgs(
3240 ctx: &GpuCtx,
3241 kernels: &EncKernels,
3242 f16: bool,
3243 bufs: &[&wgpu::Buffer],
3244 meta: &wgpu::Buffer,
3245) -> Result<Vec<wgpu::BindGroup>> {
3246 GEMM2_TILES
3247 .into_iter()
3248 .map(|t| {
3249 Ok(make_bg(
3250 ctx,
3251 kernels.gemm2_pipeline_tile(f16, t)?,
3252 bufs,
3253 meta,
3254 ))
3255 })
3256 .collect()
3257}
3258
3259pub(crate) fn v3_bgs(
3262 ctx: &GpuCtx,
3263 kernels: &EncKernels,
3264 f16: bool,
3265 bufs: &[&wgpu::Buffer],
3266 meta: &wgpu::Buffer,
3267) -> Result<Vec<wgpu::BindGroup>> {
3268 GEMM3_TILES
3269 .into_iter()
3270 .map(|t| {
3271 Ok(make_bg(
3272 ctx,
3273 kernels.gemm3_pipeline_tile(f16, t)?,
3274 bufs,
3275 meta,
3276 ))
3277 })
3278 .collect()
3279}
3280
3281#[allow(clippy::too_many_arguments)]
3287fn head_gemm(
3288 ctx: &GpuCtx,
3289 kernels: &EncKernels,
3290 use_f16: bool,
3291 w: &[f32],
3292 bias: &[f32],
3293 n: u32,
3294 k: u32,
3295 x: &wgpu::Buffer,
3296 y: &wgpu::Buffer,
3297 act: Option<Act>,
3298) -> Result<(GemmBg, wgpu::Buffer, EncGpuLinear)> {
3299 anyhow::ensure!(
3300 w.len() == (n as usize) * (k as usize),
3301 "head weight shape {} != {n}×{k}",
3302 w.len()
3303 );
3304 let mat = if use_f16 {
3305 EncMatBuf::F16(ctx.storage_bytes(&f32_to_f16_bytes(w)))
3306 } else {
3307 EncMatBuf::F32(ctx.storage(w))
3308 };
3309 let bbuf = ctx.storage(bias);
3310 let flags = 1u32 | (act_code(act) << 8);
3312 let meta = uni(ctx, bytemuck::cast_slice(&[0u32, n, k, flags]));
3313 let wbuf = match &mat {
3314 EncMatBuf::F16(b) | EncMatBuf::F32(b) => b,
3315 };
3316 let bufs = [x, wbuf, &bbuf, y];
3317 let bg = if gemm2_enabled() {
3318 GemmBg::V2(v2_bgs(ctx, kernels, use_f16, &bufs, &meta)?)
3320 } else {
3321 GemmBg::V1(make_bg(ctx, kernels.gemm_pipeline(use_f16)?, &bufs, &meta))
3322 };
3323 let keep = EncGpuLinear {
3324 w: mat,
3325 b: Some(bbuf),
3326 n,
3327 k,
3328 v3: false, sk: false,
3330 f16a: false,
3331 b_host: None,
3332 };
3333 Ok((bg, meta, keep))
3334}
3335
3336struct EncGpuNorm {
3338 w: wgpu::Buffer,
3339 b: Option<wgpu::Buffer>,
3340}
3341
3342enum GemmBg {
3348 V1(wgpu::BindGroup),
3349 V2(Vec<wgpu::BindGroup>),
3353 V3(Vec<wgpu::BindGroup>),
3357 V3Sk {
3362 plain: Vec<wgpu::BindGroup>,
3363 part: wgpu::BindGroup,
3364 part32: wgpu::BindGroup,
3366 reduce: wgpu::BindGroup,
3367 k: u32,
3368 f16a: bool,
3370 },
3371 F16A {
3376 tiles: Vec<wgpu::BindGroup>,
3377 coop: Option<wgpu::BindGroup>,
3378 },
3379}
3380
3381type StepDisp<'s> = (
3383 &'s wgpu::ComputePipeline,
3384 &'s wgpu::BindGroup,
3385 u32,
3386 u32,
3387 u32,
3388 &'static str,
3389);
3390
3391enum EStep {
3392 Gemm { bg: GemmBg, n: u32 },
3394 Ln { bg: wgpu::BindGroup },
3396 Rope { bg: wgpu::BindGroup },
3398 QkNorm { bg: wgpu::BindGroup },
3400 Glu { bg: wgpu::BindGroup },
3402 Attn { bg: wgpu::BindGroup },
3404 Attn4 { bg: wgpu::BindGroup },
3406 AttnRb { bg: wgpu::BindGroup },
3408 AttnDisent { bg: wgpu::BindGroup },
3411 AttnDisentRb { bg: wgpu::BindGroup },
3413 Add { bg: wgpu::BindGroup },
3415 Copy { bg: wgpu::BindGroup },
3417 Conv { bg: wgpu::BindGroup },
3419 L2Rows { bg: wgpu::BindGroup },
3421}
3422
3423pub struct EncoderGpu {
3431 cfg: EncoderConfig,
3432 kernels: EncKernels,
3433 use_f16: bool,
3434 max_tokens: usize,
3435 word: Vec<f32>,
3437 pos_table: Option<Vec<f32>>,
3438 ttype: Option<Vec<f32>>,
3439 emb_in: wgpu::Buffer,
3441 hidden_buf: wgpu::Buffer,
3444 _pool_src_seq_buf: wgpu::Buffer,
3447 pooled: wgpu::Buffer,
3448 seq_starts_buf: wgpu::Buffer,
3449 seq_of_buf: wgpu::Buffer,
3450 valid_buf: wgpu::Buffer,
3451 steps: Vec<EStep>,
3453 bg_pool: Option<wgpu::BindGroup>,
3455 bg_l2: Option<wgpu::BindGroup>,
3456 ptb: Option<wgpu::Buffer>,
3458 metas_tokens: Vec<wgpu::Buffer>,
3460 metas_sk: Vec<(wgpu::Buffer, u32, u32)>,
3463 metas_elems: Vec<wgpu::Buffer>,
3465 ce_head: Option<CeHeadGpu>,
3467 staging: wgpu::Buffer,
3474 staged_read: bool,
3475 _weights: Vec<EncGpuLinear>,
3477 _norms: Vec<EncGpuNorm>,
3478 _scratch: Vec<wgpu::Buffer>,
3479}
3480
3481pub const MAX_VISION_BATCH: usize = 32;
3485
3486impl EncoderGpu {
3487 pub fn kernels(&self) -> &EncKernels {
3491 &self.kernels
3492 }
3493
3494 pub fn load(ctx: &GpuCtx, dir: &Path, max_tokens: usize) -> Result<Self> {
3513 let want_f16 = matches!(std::env::var("OSFKB_ENC_F16").ok().as_deref(), Some("1"));
3514 Self::load_with(ctx, dir, max_tokens, want_f16 && ctx.f16)
3515 }
3516
3517 pub fn load_f32(ctx: &GpuCtx, dir: &Path, max_tokens: usize) -> Result<Self> {
3520 Self::load_with(ctx, dir, max_tokens, false)
3521 }
3522
3523 pub fn load_siglip_vision(ctx: &GpuCtx, dir: &Path) -> Result<Self> {
3529 let want_f16 = matches!(std::env::var("OSFKB_ENC_F16").ok().as_deref(), Some("1"));
3530 let (_, spec) = crate::encoder_weights::siglip_configs_from_dir(dir)?;
3531 let n = (spec.image_size / spec.patch_size).pow(2);
3532 Self::load_with_cfg(ctx, dir, spec.config, MAX_VISION_BATCH * n, want_f16 && ctx.f16)
3535 }
3536
3537 pub fn load_siglip_vision_f32(ctx: &GpuCtx, dir: &Path) -> Result<Self> {
3540 let (_, spec) = crate::encoder_weights::siglip_configs_from_dir(dir)?;
3541 let n = (spec.image_size / spec.patch_size).pow(2);
3542 Self::load_with_cfg(ctx, dir, spec.config, MAX_VISION_BATCH * n, false)
3543 }
3544
3545 fn load_with(ctx: &GpuCtx, dir: &Path, max_tokens: usize, use_f16: bool) -> Result<Self> {
3546 let cfg = crate::encoder_weights::encoder_config_from_dir(dir)?;
3547 Self::load_with_cfg(ctx, dir, cfg, max_tokens, use_f16)
3548 }
3549
3550 fn load_with_cfg(
3555 ctx: &GpuCtx,
3556 dir: &Path,
3557 cfg: EncoderConfig,
3558 max_tokens: usize,
3559 use_f16: bool,
3560 ) -> Result<Self> {
3561 anyhow::ensure!(
3562 cfg.head_dim <= 128,
3563 "encoder attention supports head_dim ≤ 128"
3564 );
3565 let st = LazySt::open(dir)?;
3566 let kernels = EncKernels::new(ctx);
3567 if use_f16 {
3568 kernels.gemm_pipeline(true)?; }
3570 let mut b = PlanBuilder::new(ctx, &kernels, &cfg, max_tokens, use_f16);
3571 let (word, pos_table, ttype) = match cfg.arch {
3572 EncArch::Bert | EncArch::XlmRoberta => b.build_bert_family(&st)?,
3573 EncArch::DebertaV2 => b.build_deberta(&st)?,
3574 EncArch::ModernBert => b.build_modernbert(&st)?,
3575 EncArch::Qwen3Embed => b.build_qwen3_embed(&st)?,
3576 EncArch::Lfm2Colbert => b.build_lfm2_colbert(&st, dir)?,
3577 EncArch::NomicBert => b.build_nomic(&st)?,
3578 EncArch::SiglipVision => b.build_siglip_vision(&st)?,
3579 other => anyhow::bail!(
3580 "GPU encoder tensor table for {other:?} pending verification against a real \
3581 checkpoint"
3582 ),
3583 };
3584 let PlanBuilder {
3585 emb_in,
3586 cur,
3587 pool_src_seq_buf,
3588 pooled,
3589 ptb,
3590 valid_buf,
3591 seq_starts_buf,
3592 seq_of_buf,
3593 steps,
3594 bg_pool,
3595 bg_l2,
3596 metas_tokens,
3597 metas_sk,
3598 metas_elems,
3599 weights,
3600 norms,
3601 scratch,
3602 ..
3603 } = b;
3604 if !matches!(cfg.pooling, Pooling::PerToken { .. } | Pooling::MapHead) {
3607 anyhow::ensure!(
3608 bg_pool.is_some() && bg_l2.is_some(),
3609 "plan built no pooling bind groups"
3610 );
3611 }
3612
3613 let ce_head = if let Pooling::CrossEncoder { n_labels } = cfg.pooling {
3617 let h = cfg.hidden;
3618 let prefix = ["", "bert.", "roberta."]
3622 .into_iter()
3623 .find(|p| st.has(&format!("{p}embeddings.word_embeddings.weight")))
3624 .context("cross-encoder: word embeddings not under '', 'bert.' or 'roberta.'")?;
3625 let cap = max_tokens;
3628 let mid = ctx.storage(&vec![0f32; cap * h]);
3629 let out = ctx.storage(&vec![0f32; cap * n_labels]);
3630 let (pool_w, pool_b, cls_w, cls_b) = if st.has(&format!("{prefix}pooler.dense.weight"))
3636 {
3637 (
3638 format!("{prefix}pooler.dense.weight"),
3639 format!("{prefix}pooler.dense.bias"),
3640 "classifier.weight".to_string(),
3641 "classifier.bias".to_string(),
3642 )
3643 } else {
3644 (
3645 "classifier.dense.weight".to_string(),
3646 "classifier.dense.bias".to_string(),
3647 "classifier.out_proj.weight".to_string(),
3648 "classifier.out_proj.bias".to_string(),
3649 )
3650 };
3651 let (bg_pooler, m1, w1) = head_gemm(
3652 ctx,
3653 &kernels,
3654 use_f16,
3655 &st.tensor_f32(&pool_w)?,
3656 &st.tensor_f32(&pool_b)?,
3657 h as u32,
3658 h as u32,
3659 &pooled,
3660 &mid,
3661 Some(Act::Tanh),
3662 )?;
3663 let (bg_classifier, m2, w2) = head_gemm(
3664 ctx,
3665 &kernels,
3666 use_f16,
3667 &st.tensor_f32(&cls_w)?,
3668 &st.tensor_f32(&cls_b)?,
3669 n_labels as u32,
3670 h as u32,
3671 &mid,
3672 &out,
3673 None,
3674 )?;
3675 let fused = if !use_f16
3678 && h <= 2048
3679 && std::env::var("OSFKB_ENC_CE_FUSED").ok().as_deref() != Some("0")
3680 {
3681 let (EncMatBuf::F32(pwb) | EncMatBuf::F16(pwb)) = &w1.w;
3682 let (EncMatBuf::F32(cwb) | EncMatBuf::F16(cwb)) = &w2.w;
3683 let pl = pipeline(ctx, "enc_ce_fused", ENC_CE_FUSED);
3684 let meta = uni(
3685 ctx,
3686 bytemuck::cast_slice(&[h as u32, n_labels as u32, 0, 0]),
3687 );
3688 let bg = make_bg(
3689 ctx,
3690 &pl,
3691 &[
3692 &cur,
3693 &seq_starts_buf,
3694 pwb,
3695 w1.b.as_ref().expect("pooler bias"),
3696 cwb,
3697 w2.b.as_ref().expect("classifier bias"),
3698 &out,
3699 ],
3700 &meta,
3701 );
3702 Some((pl, bg))
3704 } else {
3705 None
3706 };
3707 Some(CeHeadGpu {
3708 bg_pooler,
3709 bg_classifier,
3710 metas: vec![m1, m2],
3711 out,
3712 n_labels,
3713 fused,
3714 _mid: mid,
3715 _weights: vec![w1, w2],
3716 })
3717 } else {
3718 None
3719 };
3720
3721 let stage_floats = max_tokens
3724 * cfg.hidden.max(match cfg.pooling {
3725 Pooling::PerToken { dim } => dim,
3726 _ => 0,
3727 });
3728 let staging = ctx.device.create_buffer(&wgpu::BufferDescriptor {
3729 label: Some("enc_readback_staging"),
3730 size: (stage_floats * 4) as u64,
3731 usage: wgpu::BufferUsages::MAP_READ | wgpu::BufferUsages::COPY_DST,
3732 mapped_at_creation: false,
3733 });
3734 let staged_read = std::env::var("OSFKB_ENC_STAGED_READ").ok().as_deref() == Some("1");
3735
3736 Ok(Self {
3737 cfg,
3738 kernels,
3739 use_f16,
3740 max_tokens,
3741 word,
3742 pos_table,
3743 ttype,
3744 emb_in,
3745 hidden_buf: cur,
3746 _pool_src_seq_buf: pool_src_seq_buf,
3747 pooled,
3748 seq_starts_buf,
3749 seq_of_buf,
3750 valid_buf,
3751 steps,
3752 bg_pool,
3753 bg_l2,
3754 ptb,
3755 metas_tokens,
3756 metas_sk,
3757 metas_elems,
3758 ce_head,
3759 staging,
3760 staged_read,
3761 _weights: weights,
3762 _norms: norms,
3763 _scratch: scratch,
3764 })
3765 }
3766
3767 pub fn config(&self) -> &EncoderConfig {
3769 &self.cfg
3770 }
3771
3772 pub fn precision(&self) -> &'static str {
3774 if self.use_f16 { "f16" } else { "f32" }
3775 }
3776
3777 pub fn forward_hidden(&mut self, ctx: &GpuCtx, batch: &EncBatch) -> Result<Vec<f32>> {
3785 if std::env::var("OSFKB_ENC_TS").ok().as_deref() == Some("1") {
3787 static ONCE: std::sync::Once = std::sync::Once::new();
3788 let mut rows = None;
3789 ONCE.call_once(|| rows = Some(self.profile_encode(ctx, batch)));
3790 if let Some(Ok(rows)) = rows {
3791 let total: f64 = rows.iter().map(|r| r.1).sum();
3792 eprintln!(" [enc ts] plan kernels {total:8.1} µs (head/readback excluded)");
3793 for (name, us, cnt) in rows {
3794 eprintln!(
3795 " [enc ts] {name:16} {us:8.1} µs ({:4.1}%) ×{cnt}",
3796 us / total * 100.0
3797 );
3798 }
3799 }
3800 }
3801 let t = self.dispatch(ctx, batch)?;
3802 ctx.read(&self.hidden_buf, t * self.cfg.hidden)
3803 }
3804
3805 pub fn run_layers_gpu(&mut self, ctx: &GpuCtx, rows: &[f32], n: usize) -> Result<Vec<f32>> {
3812 self.run_layers_gpu_batched(ctx, rows, n, 1)
3813 }
3814
3815 pub fn run_layers_gpu_batched(
3823 &mut self,
3824 ctx: &GpuCtx,
3825 rows: &[f32],
3826 seq_len: usize,
3827 batch: usize,
3828 ) -> Result<Vec<f32>> {
3829 let h = self.cfg.hidden;
3830 let t = batch * seq_len;
3831 anyhow::ensure!(
3832 rows.len() == t * h,
3833 "rows {} != batch {batch} × seq_len {seq_len} × hidden {h}",
3834 rows.len()
3835 );
3836 anyhow::ensure!(
3837 t <= self.max_tokens,
3838 "batch·seq_len {t} exceeds max_tokens {}",
3839 self.max_tokens
3840 );
3841 ctx.queue
3842 .write_buffer(&self.emb_in, 0, bytemuck::cast_slice(rows));
3843 let seq_starts: Vec<u32> = (0..=batch).map(|i| (i * seq_len) as u32).collect();
3845 let seq_of = row_to_seq(&seq_starts, t);
3846 let valid = vec![1u32; t];
3847 ctx.queue
3848 .write_buffer(&self.seq_starts_buf, 0, bytemuck::cast_slice(&seq_starts));
3849 ctx.queue
3850 .write_buffer(&self.seq_of_buf, 0, bytemuck::cast_slice(&seq_of));
3851 ctx.queue
3852 .write_buffer(&self.valid_buf, 0, bytemuck::cast_slice(&valid));
3853 for m in &self.metas_tokens {
3856 ctx.queue
3857 .write_buffer(m, 0, bytemuck::cast_slice(&[t as u32]));
3858 }
3859 for (m, nn, k) in &self.metas_sk {
3860 let z = gemm3_sk_plan(t, *nn as usize, *k as usize).1;
3861 ctx.queue.write_buffer(m, 12, bytemuck::cast_slice(&[z]));
3862 }
3863 for m in &self.metas_elems {
3864 ctx.queue
3865 .write_buffer(m, 0, bytemuck::cast_slice(&[(t * h) as u32]));
3866 }
3867 let mut cmd = ctx
3868 .device
3869 .create_command_encoder(&wgpu::CommandEncoderDescriptor { label: None });
3870 {
3871 let mut pass = cmd.begin_compute_pass(&wgpu::ComputePassDescriptor::default());
3872 for step in &self.steps {
3873 let (d, d2) = self.resolve_step(step, t)?;
3874 for (pl, bg, gx, gy, gz, _) in std::iter::once(d).chain(d2) {
3875 pass.set_pipeline(pl);
3876 pass.set_bind_group(0, bg, &[]);
3877 pass.dispatch_workgroups(gx, gy, gz);
3878 }
3879 }
3880 }
3881 ctx.queue.submit([cmd.finish()]);
3882 ctx.read(&self.hidden_buf, t * h)
3883 }
3884
3885 pub fn encode(&mut self, ctx: &GpuCtx, batch: &EncBatch) -> Result<EmbedOut> {
3886 let _trace = std::env::var("OSFKB_ENC_GPU_TRACE").is_ok();
3892 if std::env::var("OSFKB_ENC_TS").ok().as_deref() == Some("1") {
3895 static ONCE: std::sync::Once = std::sync::Once::new();
3896 let mut rows = None;
3897 ONCE.call_once(|| rows = Some(self.profile_encode(ctx, batch)));
3898 if let Some(rows) = rows {
3899 match rows {
3900 Ok(rows) => {
3901 let total: f64 = rows.iter().map(|r| r.1).sum();
3902 eprintln!(
3903 " [enc ts] plan kernels {total:8.1} µs (pool/head/readback excluded)"
3904 );
3905 for (name, us, cnt) in rows {
3906 eprintln!(
3907 " [enc ts] {name:16} {us:8.1} µs ({:4.1}%) ×{cnt}",
3908 us / total * 100.0
3909 );
3910 }
3911 }
3912 Err(e) => eprintln!(" [enc ts] profile failed: {e}"),
3913 }
3914 }
3915 }
3916 let _t0 = std::time::Instant::now();
3917 let t = self.dispatch(ctx, batch)?;
3918 if _trace {
3919 eprintln!(
3920 " [enc gpu] dispatch (gather+upload+GPU) {:?}",
3921 _t0.elapsed()
3922 );
3923 }
3924 let _t1 = std::time::Instant::now();
3925 let (cfg, h, b_seqs) = (&self.cfg, self.cfg.hidden, batch.n_seqs());
3926
3927 if let Pooling::PerToken { dim } = cfg.pooling {
3928 let flat = if self.staged_read {
3929 self.read_staged(ctx, t * dim)?
3930 } else {
3931 let ptb = self
3932 .ptb
3933 .as_ref()
3934 .context("PerToken plan built no output buffer")?;
3935 ctx.read(ptb, t * dim)?
3936 };
3937 let mut out: Vec<Vec<Vec<f32>>> = Vec::with_capacity(b_seqs);
3938 for w in batch.seq_starts.windows(2) {
3939 out.push(
3940 (w[0] as usize..w[1] as usize)
3941 .map(|i| flat[i * dim..(i + 1) * dim].to_vec())
3942 .collect(),
3943 );
3944 }
3945 return Ok(EmbedOut::PerToken(out));
3946 }
3947 if let Some(ce) = &self.ce_head {
3948 let flat = if self.staged_read {
3950 self.read_staged(ctx, b_seqs * ce.n_labels)?
3951 } else {
3952 ctx.read(&ce.out, b_seqs * ce.n_labels)?
3953 };
3954 if _trace {
3955 eprintln!(
3956 " [enc gpu] readback (cross-encoder logits) {:?}",
3957 _t1.elapsed()
3958 );
3959 }
3960 return Ok(EmbedOut::Pooled(
3961 flat.chunks_exact(ce.n_labels)
3962 .map(<[f32]>::to_vec)
3963 .collect(),
3964 ));
3965 }
3966 let flat = if self.staged_read {
3967 self.read_staged(ctx, b_seqs * h)?
3968 } else {
3969 ctx.read(&self.pooled, b_seqs * h)?
3970 };
3971 if _trace {
3972 eprintln!(
3973 " [enc gpu] readback ({} floats) {:?}",
3974 b_seqs * h,
3975 _t1.elapsed()
3976 );
3977 }
3978 Ok(EmbedOut::Pooled(
3979 flat.chunks_exact(h).map(<[f32]>::to_vec).collect(),
3980 ))
3981 }
3982
3983 fn dispatch(&mut self, ctx: &GpuCtx, batch: &EncBatch) -> Result<usize> {
3987 let cfg = &self.cfg;
3988 let (t, b_seqs) = (batch.tokens.len(), batch.n_seqs());
3989 anyhow::ensure!(b_seqs > 0, "empty batch");
3990 anyhow::ensure!(
3991 t <= self.max_tokens,
3992 "batch of {t} tokens exceeds max_tokens {} — chunk at sequence boundaries",
3993 self.max_tokens
3994 );
3995 let offset = match cfg.pos_kind {
3996 PosKind::Learned { offset } => offset,
3997 PosKind::Rope { .. } => 0,
3998 };
3999 for w in batch.seq_starts.windows(2) {
4000 let len = (w[1] - w[0]) as usize;
4001 anyhow::ensure!(len > 0, "empty sequence in batch");
4002 anyhow::ensure!(
4003 len + offset <= cfg.max_pos,
4004 "sequence of {len} tokens exceeds max positions {} (offset {offset})",
4005 cfg.max_pos
4006 );
4007 }
4008 for &tok in &batch.tokens {
4009 anyhow::ensure!((tok as usize) < cfg.vocab, "token id {tok} out of vocab");
4010 }
4011
4012 let h = cfg.hidden;
4014 let mut emb = vec![0f32; t * h];
4015 let seq_of = row_to_seq(&batch.seq_starts, t);
4016 for (i, row) in emb.chunks_exact_mut(h).enumerate() {
4017 let tok = batch.tokens[i] as usize;
4018 row.copy_from_slice(&self.word[tok * h..(tok + 1) * h]);
4019 if let Some(pos_table) = &self.pos_table {
4020 let local = i - batch.seq_starts[seq_of[i] as usize] as usize;
4021 for (r, p) in row
4022 .iter_mut()
4023 .zip(&pos_table[(offset + local) * h..(offset + local + 1) * h])
4024 {
4025 *r += p;
4026 }
4027 }
4028 if let Some(tt) = &self.ttype {
4029 let ty = batch.type_ids.as_ref().map_or(0, |t| t[i] as usize);
4035 anyhow::ensure!(
4036 (ty + 1) * h <= tt.len(),
4037 "token_type id {ty} exceeds the checkpoint's segment table"
4038 );
4039 for (r, v) in row.iter_mut().zip(&tt[ty * h..(ty + 1) * h]) {
4040 *r += v;
4041 }
4042 }
4043 }
4044 ctx.queue
4045 .write_buffer(&self.emb_in, 0, bytemuck::cast_slice(&emb));
4046 ctx.queue.write_buffer(
4047 &self.seq_starts_buf,
4048 0,
4049 bytemuck::cast_slice(&batch.seq_starts),
4050 );
4051 ctx.queue
4052 .write_buffer(&self.seq_of_buf, 0, bytemuck::cast_slice(&seq_of));
4053 let ones;
4056 let valid_slice: &[u32] = match &batch.valid {
4057 Some(v) => {
4058 anyhow::ensure!(
4059 v.len() == t,
4060 "validity mask length {} != tokens {t}",
4061 v.len()
4062 );
4063 v
4064 }
4065 None => {
4066 ones = vec![1u32; t];
4067 &ones
4068 }
4069 };
4070 ctx.queue
4071 .write_buffer(&self.valid_buf, 0, bytemuck::cast_slice(valid_slice));
4072 for m in &self.metas_tokens {
4073 ctx.queue
4074 .write_buffer(m, 0, bytemuck::cast_slice(&[t as u32]));
4075 }
4076 for (m, n, k) in &self.metas_sk {
4077 let z = gemm3_sk_plan(t, *n as usize, *k as usize).1;
4078 ctx.queue.write_buffer(m, 12, bytemuck::cast_slice(&[z]));
4079 }
4080 for m in &self.metas_elems {
4081 ctx.queue
4082 .write_buffer(m, 0, bytemuck::cast_slice(&[(t * h) as u32]));
4083 }
4084 if let Some(ce) = &self.ce_head {
4086 for m in &ce.metas {
4087 ctx.queue
4088 .write_buffer(m, 0, bytemuck::cast_slice(&[b_seqs as u32]));
4089 }
4090 }
4091
4092 let mut enc = ctx
4094 .device
4095 .create_command_encoder(&wgpu::CommandEncoderDescriptor { label: None });
4096 {
4097 let mut pass = enc.begin_compute_pass(&wgpu::ComputePassDescriptor::default());
4098 for step in &self.steps {
4099 let (d, d2) = self.resolve_step(step, t)?;
4100 for (pl, bg, gx, gy, gz, _) in std::iter::once(d).chain(d2) {
4101 pass.set_pipeline(pl);
4102 pass.set_bind_group(0, bg, &[]);
4103 pass.dispatch_workgroups(gx, gy, gz);
4104 }
4105 }
4106 let gemm_pl = self.kernels.gemm_pipeline(self.use_f16)?;
4107 match cfg.pooling {
4108 Pooling::Mean | Pooling::LastToken | Pooling::Cls => {
4109 let pool_pl = match cfg.pooling {
4110 Pooling::LastToken => &self.kernels.last_pool,
4111 Pooling::Cls => &self.kernels.cls_pool,
4112 _ => &self.kernels.mean_pool,
4113 };
4114 let (bg_pool, bg_l2) = (
4115 self.bg_pool.as_ref().expect("validated at load"),
4116 self.bg_l2.as_ref().expect("validated at load"),
4117 );
4118 pass.set_pipeline(pool_pl);
4119 pass.set_bind_group(0, bg_pool, &[]);
4120 pass.dispatch_workgroups(b_seqs as u32, 1, 1);
4121 pass.set_pipeline(&self.kernels.l2norm);
4122 pass.set_bind_group(0, bg_l2, &[]);
4123 pass.dispatch_workgroups(b_seqs as u32, 1, 1);
4124 }
4125 Pooling::CrossEncoder { .. } => {
4126 let ce = self
4130 .ce_head
4131 .as_ref()
4132 .expect("CrossEncoder head built at load");
4133 if let Some((pl, bg)) = &ce.fused {
4134 pass.set_pipeline(pl);
4136 pass.set_bind_group(0, bg, &[]);
4137 pass.dispatch_workgroups(b_seqs as u32, 1, 1);
4138 } else {
4139 let bg_pool = self.bg_pool.as_ref().expect("validated at load");
4140 pass.set_pipeline(&self.kernels.cls_pool);
4141 pass.set_bind_group(0, bg_pool, &[]);
4142 pass.dispatch_workgroups(b_seqs as u32, 1, 1);
4143
4144 let m = b_seqs;
4145 for (bg, n) in [
4146 (&ce.bg_pooler, cfg.hidden),
4147 (&ce.bg_classifier, ce.n_labels),
4148 ] {
4149 let (pl, bg, gx, gy) = match bg {
4150 GemmBg::V1(bg) => (
4151 gemm_pl,
4152 bg,
4153 (n as u32).div_ceil(16),
4154 (m as u32).div_ceil(16),
4155 ),
4156 GemmBg::V2(bgs) => {
4157 let tile = gemm2_tile(m, n);
4161 (
4162 self.kernels.gemm2_pipeline_tile(self.use_f16, tile)?,
4163 &bgs[gemm2_tier(tile)],
4164 (n as u32).div_ceil(tile.1 as u32),
4165 (m as u32).div_ceil(tile.0 as u32),
4166 )
4167 }
4168 GemmBg::V3(_) | GemmBg::V3Sk { .. } | GemmBg::F16A { .. } => {
4169 unreachable!(
4170 "head_gemm never uploads v3/sk/f16a (see its `keep`)"
4171 )
4172 }
4173 };
4174 pass.set_pipeline(pl);
4175 pass.set_bind_group(0, bg, &[]);
4176 pass.dispatch_workgroups(gx, gy, 1);
4177 }
4178 }
4179 }
4180 Pooling::PerToken { .. } => {}
4182 other => anyhow::bail!("pooling {other:?} lands with its phase"),
4183 }
4184 }
4185 if self.staged_read {
4189 let (src, floats) = if let Pooling::PerToken { dim } = cfg.pooling {
4190 (
4191 self.ptb
4192 .as_ref()
4193 .context("PerToken plan built no output buffer")?,
4194 t * dim,
4195 )
4196 } else if let Some(ce) = &self.ce_head {
4197 (&ce.out, b_seqs * ce.n_labels)
4198 } else {
4199 (&self.pooled, b_seqs * cfg.hidden)
4200 };
4201 enc.copy_buffer_to_buffer(src, 0, &self.staging, 0, (floats * 4) as u64);
4202 }
4203 ctx.queue.submit([enc.finish()]);
4204 Ok(t)
4205 }
4206
4207 #[allow(clippy::type_complexity)]
4213 fn resolve_step<'s>(
4214 &'s self,
4215 step: &'s EStep,
4216 t: usize,
4217 ) -> Result<(StepDisp<'s>, Option<StepDisp<'s>>)> {
4218 let cfg = &self.cfg;
4219 let t32 = t as u32;
4220 let h = cfg.hidden;
4221 Ok(match step {
4222 EStep::Gemm {
4223 bg:
4224 GemmBg::V3Sk {
4225 plain,
4226 part,
4227 part32,
4228 reduce,
4229 k,
4230 f16a,
4231 },
4232 n,
4233 } => {
4234 let tile = gemm3_tile(t, *n as usize);
4235 let wgs = (*n as usize).div_ceil(tile.1) * t.div_ceil(tile.0);
4236 let plain_ok = if *n <= 512 {
4240 wgs >= 80
4241 } else {
4242 gemm3_grid_ok(t, tile, wgs)
4243 };
4244 if plain_ok {
4245 let pl = if *f16a {
4246 self.kernels
4247 .gemm3_f16a_pipeline_tile(tile)
4248 .ok_or_else(|| anyhow::anyhow!("f16a plan on a non-f16 adapter"))?
4249 } else {
4250 self.kernels.gemm3_pipeline_tile(self.use_f16, tile)?
4251 };
4252 (
4253 (
4254 pl,
4255 &plain[gemm3_tier(tile)],
4256 n.div_ceil(tile.1 as u32),
4257 t32.div_ceil(tile.0 as u32),
4258 1,
4259 if *f16a { "gemm3f16a" } else { "gemm3" },
4260 ),
4261 None,
4262 )
4263 } else {
4264 let (use32, z) = gemm3_sk_plan(t, *n as usize, *k as usize);
4267 let pl = match (*f16a, use32) {
4268 (true, false) => self
4269 .kernels
4270 .gemm3_sk_f16a
4271 .as_ref()
4272 .ok_or_else(|| anyhow::anyhow!("f16a plan on a non-f16 adapter"))?,
4273 (true, true) => self
4274 .kernels
4275 .gemm3_sk32_f16a
4276 .as_ref()
4277 .ok_or_else(|| anyhow::anyhow!("f16a plan on a non-f16 adapter"))?,
4278 (false, false) => self.kernels.gemm3_sk_pipeline(self.use_f16)?,
4279 (false, true) => self.kernels.gemm3_sk32_pipeline(self.use_f16)?,
4280 };
4281 let (pbg, rows) = if use32 {
4282 (part32, GEMM3_SK_TILE32.0 as u32)
4283 } else {
4284 (part, GEMM3_SK_TILE.0 as u32)
4285 };
4286 anyhow::ensure!(
4287 z as usize * t * (*n as usize) <= SK_PART_F32,
4288 "split-K partials overflow: z={z} t={t} n={n} — the occupancy knee \
4289 should make this impossible"
4290 );
4291 (
4292 (
4293 pl,
4294 pbg,
4295 n.div_ceil(GEMM3_SK_TILE.1 as u32),
4296 t32.div_ceil(rows),
4297 z,
4298 "gemm3sk",
4299 ),
4300 Some((
4301 &self.kernels.gemm3_sk_reduce,
4302 reduce,
4303 (t32 * n).div_ceil(256),
4304 1,
4305 1,
4306 "skreduce",
4307 )),
4308 )
4309 }
4310 }
4311 EStep::Gemm { bg, n } => (
4312 match bg {
4313 GemmBg::F16A { tiles, coop } => {
4314 let wgs128 = (*n as usize).div_ceil(128) * t.div_ceil(128);
4318 let m_pad = t.div_ceil(128) * 128;
4319 if let (Some(bg), true, true) =
4320 (coop.as_ref(), wgs128 >= 48, m_pad <= self.max_tokens)
4321 {
4322 (
4323 self.kernels
4324 .gemm4_f16w
4325 .as_ref()
4326 .ok_or_else(|| anyhow::anyhow!("coop arm without kernel"))?,
4327 bg,
4328 n.div_ceil(128),
4329 (m_pad as u32) / 128,
4330 1,
4331 "gemm4f16",
4332 )
4333 } else {
4334 let tile = gemm3_tile(t, *n as usize);
4335 (
4336 self.kernels.gemm3_f16a_pipeline_tile(tile).ok_or_else(|| {
4337 anyhow::anyhow!("f16a plan on a non-f16 adapter")
4338 })?,
4339 &tiles[gemm3_tier(tile)],
4340 n.div_ceil(tile.1 as u32),
4341 t32.div_ceil(tile.0 as u32),
4342 1,
4343 "gemm3f16a",
4344 )
4345 }
4346 }
4347 GemmBg::V1(bg) => (
4348 self.kernels.gemm_pipeline(self.use_f16)?,
4349 bg,
4350 n.div_ceil(16),
4351 t32.div_ceil(16),
4352 1,
4353 "gemm1",
4354 ),
4355 GemmBg::V2(bgs) => {
4356 let tile = gemm2_tile(t, *n as usize);
4360 (
4361 self.kernels.gemm2_pipeline_tile(self.use_f16, tile)?,
4362 &bgs[gemm2_tier(tile)],
4363 n.div_ceil(tile.1 as u32),
4364 t32.div_ceil(tile.0 as u32),
4365 1,
4366 "gemm2",
4367 )
4368 }
4369 GemmBg::V3(bgs) => {
4370 let tile = gemm3_tile(t, *n as usize);
4371 (
4372 self.kernels.gemm3_pipeline_tile(self.use_f16, tile)?,
4373 &bgs[gemm3_tier(tile)],
4374 n.div_ceil(tile.1 as u32),
4375 t32.div_ceil(tile.0 as u32),
4376 1,
4377 "gemm3",
4378 )
4379 }
4380 GemmBg::V3Sk { .. } => unreachable!("handled by the arm above"),
4381 },
4382 None,
4383 ),
4384 EStep::Ln { bg } => ((&self.kernels.layernorm, bg, t32, 1, 1, "ln"), None),
4385 EStep::Rope { bg } => ((&self.kernels.rope, bg, t32, 1, 1, "rope"), None),
4386 EStep::QkNorm { bg } => ((&self.kernels.qk_norm, bg, t32, 1, 1, "qknorm"), None),
4387 EStep::Glu { bg } => ((&self.kernels.glu, bg, t32, 1, 1, "glu"), None),
4388 EStep::Attn { bg } => (
4389 (&self.kernels.attn, bg, t32, cfg.n_heads as u32, 1, "attn"),
4390 None,
4391 ),
4392 EStep::Attn4 { bg } => (
4393 (&self.kernels.attn4, bg, t32, cfg.n_heads as u32, 1, "attn4"),
4394 None,
4395 ),
4396 EStep::AttnRb { bg } => (
4397 (
4398 &self.kernels.attn_rb4,
4399 bg,
4400 t32.div_ceil(4),
4401 cfg.n_heads as u32,
4402 1,
4403 "attnRB",
4404 ),
4405 None,
4406 ),
4407 EStep::AttnDisent { bg } => (
4408 (
4409 &self.kernels.attn_disent,
4410 bg,
4411 t32,
4412 cfg.n_heads as u32,
4413 1,
4414 "disent",
4415 ),
4416 None,
4417 ),
4418 EStep::AttnDisentRb { bg } => (
4419 (
4420 &self.kernels.attn_disent_rb4,
4421 bg,
4422 t32.div_ceil(4),
4423 cfg.n_heads as u32,
4424 1,
4425 "disentRB",
4426 ),
4427 None,
4428 ),
4429 EStep::Add { bg } => (
4430 (
4431 &self.kernels.add,
4432 bg,
4433 ((t * h) as u32).div_ceil(256),
4434 1,
4435 1,
4436 "add",
4437 ),
4438 None,
4439 ),
4440 EStep::Copy { bg } => (
4441 (
4442 &self.kernels.copy,
4443 bg,
4444 ((t * h) as u32).div_ceil(256),
4445 1,
4446 1,
4447 "copy",
4448 ),
4449 None,
4450 ),
4451 EStep::Conv { bg } => ((&self.kernels.conv, bg, t32, 1, 1, "conv"), None),
4452 EStep::L2Rows { bg } => ((&self.kernels.l2norm, bg, t32, 1, 1, "l2"), None),
4453 })
4454 }
4455
4456 pub fn profile_encode(
4463 &mut self,
4464 ctx: &GpuCtx,
4465 batch: &EncBatch,
4466 ) -> Result<Vec<(String, f64, u32)>> {
4467 anyhow::ensure!(ctx.timestamps, "adapter lacks TIMESTAMP_QUERY");
4468 let t = self.dispatch(ctx, batch)?; let ndisp: usize = self
4470 .steps
4471 .iter()
4472 .map(|s| match self.resolve_step(s, t) {
4473 Ok((_, Some(_))) => 2,
4474 _ => 1,
4475 })
4476 .sum();
4477 let nq = (ndisp * 2) as u32;
4478 let qs = ctx.device.create_query_set(&wgpu::QuerySetDescriptor {
4479 label: Some("enc_prof"),
4480 ty: wgpu::QueryType::Timestamp,
4481 count: nq,
4482 });
4483 let qbuf = ctx.device.create_buffer(&wgpu::BufferDescriptor {
4484 label: Some("enc_prof_resolve"),
4485 size: u64::from(nq) * 8,
4486 usage: wgpu::BufferUsages::QUERY_RESOLVE | wgpu::BufferUsages::COPY_SRC,
4487 mapped_at_creation: false,
4488 });
4489 let qstage = ctx.device.create_buffer(&wgpu::BufferDescriptor {
4490 label: Some("enc_prof_stage"),
4491 size: u64::from(nq) * 8,
4492 usage: wgpu::BufferUsages::MAP_READ | wgpu::BufferUsages::COPY_DST,
4493 mapped_at_creation: false,
4494 });
4495 let mut enc = ctx
4496 .device
4497 .create_command_encoder(&wgpu::CommandEncoderDescriptor::default());
4498 let mut labels: Vec<(&'static str, u32)> = Vec::with_capacity(ndisp);
4499 {
4500 let mut qi = 0u32;
4501 for step in &self.steps {
4502 let n_of = |s: &EStep| match s {
4503 EStep::Gemm { n, .. } => *n,
4504 _ => 0,
4505 };
4506 let (d, d2) = self.resolve_step(step, t)?;
4507 for (pl, bg, gx, gy, gz, label) in std::iter::once(d).chain(d2) {
4508 let mut p = enc.begin_compute_pass(&wgpu::ComputePassDescriptor {
4509 label: None,
4510 timestamp_writes: Some(wgpu::ComputePassTimestampWrites {
4511 query_set: &qs,
4512 beginning_of_pass_write_index: Some(qi * 2),
4513 end_of_pass_write_index: Some(qi * 2 + 1),
4514 }),
4515 });
4516 p.set_pipeline(pl);
4517 p.set_bind_group(0, bg, &[]);
4518 p.dispatch_workgroups(gx, gy, gz);
4519 labels.push((label, n_of(step)));
4520 qi += 1;
4521 }
4522 }
4523 }
4524 enc.resolve_query_set(&qs, 0..nq, &qbuf, 0);
4525 enc.copy_buffer_to_buffer(&qbuf, 0, &qstage, 0, u64::from(nq) * 8);
4526 ctx.queue.submit([enc.finish()]);
4527 let slice = qstage.slice(..);
4528 let (tx, rx) = std::sync::mpsc::channel();
4529 slice.map_async(wgpu::MapMode::Read, move |r| {
4530 let _ = tx.send(r);
4531 });
4532 let _ = ctx.device.poll(wgpu::PollType::wait_indefinitely());
4533 rx.recv()
4534 .map_err(|_| anyhow::anyhow!("prof staging dropped"))?
4535 .map_err(|e| anyhow::anyhow!("map_async: {e:?}"))?;
4536 let raw: Vec<u64> =
4537 bytemuck::cast_slice(&slice.get_mapped_range().expect("mapped range")).to_vec();
4538 qstage.unmap();
4539 let mut agg: std::collections::BTreeMap<(&'static str, u32), (f64, u32)> =
4540 std::collections::BTreeMap::new();
4541 for (qi, (label, n)) in labels.iter().enumerate() {
4542 let dt =
4543 raw[qi * 2 + 1].saturating_sub(raw[qi * 2]) as f64 * f64::from(ctx.ts_period) / 1e3;
4544 let e = agg.entry((label, *n)).or_insert((0.0, 0));
4545 e.0 += dt;
4546 e.1 += 1;
4547 }
4548 let mut rows: Vec<(String, f64, u32)> = agg
4549 .into_iter()
4550 .map(|((label, n), (us, cnt))| {
4551 let name = if n > 0 {
4552 format!("{label} n={n}")
4553 } else {
4554 label.to_string()
4555 };
4556 (name, us, cnt)
4557 })
4558 .collect();
4559 rows.sort_by(|a, b| b.1.partial_cmp(&a.1).unwrap_or(std::cmp::Ordering::Equal));
4560 Ok(rows)
4561 }
4562
4563 fn read_staged(&self, ctx: &GpuCtx, len: usize) -> Result<Vec<f32>> {
4567 let slice = self.staging.slice(..(len * 4) as u64);
4568 let (tx, rx) = std::sync::mpsc::channel();
4569 slice.map_async(wgpu::MapMode::Read, move |r| {
4570 let _ = tx.send(r);
4571 });
4572 if ctx.spin_poll {
4573 loop {
4574 let r = ctx.device.poll(wgpu::PollType::Poll);
4575 if rx.try_recv().is_ok() {
4576 break;
4577 }
4578 if let Ok(status) = &r
4579 && status.is_queue_empty()
4580 {
4581 rx.recv()
4582 .map_err(|_| anyhow::anyhow!("map_async callback dropped"))?
4583 .map_err(|e| anyhow::anyhow!("map_async: {e:?}"))?;
4584 break;
4585 }
4586 std::hint::spin_loop();
4587 }
4588 } else {
4589 let _ = ctx.device.poll(wgpu::PollType::wait_indefinitely());
4590 rx.recv()
4591 .map_err(|_| anyhow::anyhow!("map_async callback dropped"))?
4592 .map_err(|e| anyhow::anyhow!("map_async: {e:?}"))?;
4593 }
4594 let data = slice.get_mapped_range().expect("mapped range");
4595 let out: Vec<f32> = bytemuck::cast_slice(&data).to_vec();
4596 drop(data);
4597 self.staging.unmap();
4598 Ok(out)
4599 }
4600}
4601
4602struct PlanBuilder<'a> {
4607 ctx: &'a GpuCtx,
4608 kernels: &'a EncKernels,
4609 cfg: &'a EncoderConfig,
4610 use_f16: bool,
4611 eps: f32,
4612 emb_in: wgpu::Buffer,
4613 cur: wgpu::Buffer,
4614 alt: wgpu::Buffer,
4615 qb: wgpu::Buffer,
4616 kb: wgpu::Buffer,
4617 vb: wgpu::Buffer,
4618 cb: wgpu::Buffer,
4619 mid: wgpu::Buffer,
4620 glu: wgpu::Buffer,
4621 pooled: wgpu::Buffer,
4622 ptb: Option<wgpu::Buffer>,
4624 pool_src_seq_buf: wgpu::Buffer,
4625 seq_starts_buf: wgpu::Buffer,
4626 seq_of_buf: wgpu::Buffer,
4627 valid_buf: wgpu::Buffer,
4628 steps: Vec<EStep>,
4629 bg_pool: Option<wgpu::BindGroup>,
4630 bg_l2: Option<wgpu::BindGroup>,
4631 metas_tokens: Vec<wgpu::Buffer>,
4632 metas_sk: Vec<(wgpu::Buffer, u32, u32)>,
4635 metas_elems: Vec<wgpu::Buffer>,
4636 weights: Vec<EncGpuLinear>,
4637 norms: Vec<EncGpuNorm>,
4638 scratch: Vec<wgpu::Buffer>,
4639 sk_part: Option<wgpu::Buffer>,
4641 max_tokens: usize,
4642}
4643
4644impl<'a> PlanBuilder<'a> {
4645 fn new(
4646 ctx: &'a GpuCtx,
4647 kernels: &'a EncKernels,
4648 cfg: &'a EncoderConfig,
4649 max_tokens: usize,
4650 use_f16: bool,
4651 ) -> Self {
4652 let h = cfg.hidden;
4653 let mut up_width = match cfg.mlp_kind {
4654 MlpKind::Dense { .. } => cfg.intermediate,
4655 MlpKind::Glu { .. } => 2 * cfg.intermediate,
4656 };
4657 if cfg.layer_is_attn.iter().any(|a| !a) {
4658 up_width = up_width.max(3 * h);
4660 }
4661 let qw = (cfg.n_heads * cfg.head_dim).max(h);
4664 let zeros = |n: usize| vec![0f32; n];
4665 Self {
4666 eps: cfg.eps,
4667 emb_in: ctx.storage(&zeros(max_tokens * h)),
4668 cur: ctx.storage(&zeros(max_tokens * h)),
4669 alt: ctx.storage(&zeros(max_tokens * h)),
4670 qb: ctx.storage(&zeros(max_tokens * qw)),
4671 kb: ctx.storage(&zeros(max_tokens * qw)),
4672 vb: ctx.storage(&zeros(max_tokens * qw)),
4673 cb: ctx.storage(&zeros(max_tokens * qw)),
4674 mid: ctx.storage(&zeros(max_tokens * up_width)),
4675 glu: ctx.storage(&zeros(max_tokens * cfg.intermediate)),
4676 pooled: ctx.storage(&zeros(max_tokens * h)),
4677 ptb: None,
4678 pool_src_seq_buf: ctx.storage_bytes(&[0u8; 8]),
4679 seq_starts_buf: ctx.storage_bytes(&vec![0u8; (max_tokens + 1) * 4]),
4680 seq_of_buf: ctx.storage_bytes(&vec![0u8; max_tokens * 4]),
4681 valid_buf: ctx.storage_bytes(bytemuck::cast_slice::<u32, u8>(&vec![1u32; max_tokens])),
4682 steps: Vec::new(),
4683 bg_pool: None,
4684 bg_l2: None,
4685 metas_tokens: Vec::new(),
4686 metas_sk: Vec::new(),
4687 metas_elems: Vec::new(),
4688 weights: Vec::new(),
4689 norms: Vec::new(),
4690 scratch: Vec::new(),
4691 sk_part: None,
4692 max_tokens,
4693 ctx,
4694 kernels,
4695 cfg,
4696 use_f16,
4697 }
4698 }
4699
4700 fn upload_linear(&mut self, w: &[f32], b: Option<Vec<f32>>, n: u32, k: u32) -> Result<usize> {
4701 anyhow::ensure!(
4702 w.len() == (n as usize) * (k as usize),
4703 "weight shape {} != {n}×{k}",
4704 w.len()
4705 );
4706 let v3 = gemm3_eligible(n as usize);
4711 let sk = (!v3 && gemm3_sk_band(n as usize)) || (v3 && gemm3_smalln_sk(n as usize));
4714 let f16a = (v3 || sk) && f16a_enabled() && self.ctx.f16;
4717 let wt_store;
4718 let w = if v3 || sk {
4719 let (n, k) = (n as usize, k as usize);
4720 let mut wt = vec![0f32; n * k];
4721 for nn in 0..n {
4722 for kk in 0..k {
4723 wt[kk * n + nn] = w[nn * k + kk];
4724 }
4725 }
4726 wt_store = wt;
4727 &wt_store[..]
4728 } else {
4729 w
4730 };
4731 let mat = if f16a || self.use_f16 {
4732 EncMatBuf::F16(self.ctx.storage_bytes(&f32_to_f16_bytes(w)))
4733 } else {
4734 EncMatBuf::F32(self.ctx.storage(w))
4735 };
4736 let b_host = if f16a { b.clone() } else { None };
4737 self.weights.push(EncGpuLinear {
4738 w: mat,
4739 b: b.map(|bv| self.ctx.storage(&bv)),
4740 n,
4741 k,
4742 v3,
4743 sk,
4744 f16a,
4745 b_host,
4746 });
4747 Ok(self.weights.len() - 1)
4748 }
4749
4750 fn sk_part_buf(&mut self) -> wgpu::Buffer {
4755 if self.sk_part.is_none() {
4756 self.sk_part = Some(self.ctx.storage(&vec![0f32; SK_PART_F32]));
4757 }
4758 self.sk_part.clone().expect("just filled")
4759 }
4760
4761 fn push_gemm(&mut self, wi: usize, x: &wgpu::Buffer, y: &wgpu::Buffer, act: Option<Act>) {
4763 let lw = &self.weights[wi];
4764 let (n, k, v3, sk, has_b) = (lw.n, lw.k, lw.v3, lw.sk, lw.b.is_some());
4765 let lw_f16a = lw.f16a;
4766 let wbuf = match &lw.w {
4767 EncMatBuf::F16(b) | EncMatBuf::F32(b) => b.clone(),
4768 };
4769 let bbuf = lw.b.clone().unwrap_or_else(|| self.ctx.storage(&[0.0]));
4771 let flags = u32::from(has_b) | (act_code(act) << 8);
4772 let meta = uni(self.ctx, bytemuck::cast_slice(&[0u32, n, k, flags]));
4773 let bufs = [x, &wbuf, &bbuf, y];
4774 let bg = if lw_f16a && sk {
4775 let part = self.sk_part_buf();
4779 let plain = GEMM3_TILES
4780 .iter()
4781 .map(|&tl| {
4782 make_bg(
4783 self.ctx,
4784 self.kernels
4785 .gemm3_f16a_pipeline_tile(tl)
4786 .expect("f16a implies SHADER_F16 (checked at upload)"),
4787 &bufs,
4788 &meta,
4789 )
4790 })
4791 .collect();
4792 let part_bg = make_bg(
4793 self.ctx,
4794 self.kernels
4795 .gemm3_sk_f16a
4796 .as_ref()
4797 .expect("f16a implies SHADER_F16 (checked at upload)"),
4798 &[x, &wbuf, &bbuf, &part],
4799 &meta,
4800 );
4801 let part32_bg = make_bg(
4802 self.ctx,
4803 self.kernels
4804 .gemm3_sk32_f16a
4805 .as_ref()
4806 .expect("f16a implies SHADER_F16 (checked at upload)"),
4807 &[x, &wbuf, &bbuf, &part],
4808 &meta,
4809 );
4810 let skn = gemm3_sk_chunks(k as usize); let rmeta = uni(self.ctx, bytemuck::cast_slice(&[0u32, n, flags, skn]));
4812 let reduce = make_bg(
4813 self.ctx,
4814 &self.kernels.gemm3_sk_reduce,
4815 &[&part, &bbuf, y],
4816 &rmeta,
4817 );
4818 self.metas_sk.push((rmeta.clone(), n, k));
4819 self.metas_tokens.push(rmeta);
4820 GemmBg::V3Sk {
4821 plain,
4822 part: part_bg,
4823 part32: part32_bg,
4824 reduce,
4825 k,
4826 f16a: true,
4827 }
4828 } else if lw_f16a {
4829 let set = GEMM3_TILES
4830 .iter()
4831 .map(|&tl| {
4832 make_bg(
4833 self.ctx,
4834 self.kernels
4835 .gemm3_f16a_pipeline_tile(tl)
4836 .expect("f16a implies SHADER_F16 (checked at upload)"),
4837 &bufs,
4838 &meta,
4839 )
4840 })
4841 .collect();
4842 let coop = match (&self.kernels.gemm4_f16w, act, &self.weights[wi].b_host) {
4845 (Some(pl), None, bh) => {
4846 let bias_vec = bh.clone().unwrap_or_else(|| vec![0f32; n as usize]);
4847 let mut b8 = Vec::with_capacity(8 * n as usize);
4848 for _ in 0..8 {
4849 b8.extend_from_slice(&bias_vec);
4850 }
4851 let b8buf = self.ctx.storage(&b8);
4852 let bg = make_bg(self.ctx, pl, &[x, &wbuf, &b8buf, y], &meta);
4853 self.scratch.push(b8buf);
4854 Some(bg)
4855 }
4856 _ => None,
4857 };
4858 GemmBg::F16A { tiles: set, coop }
4859 } else if sk {
4860 let part = self.sk_part_buf();
4863 let plain = v3_bgs(self.ctx, self.kernels, self.use_f16, &bufs, &meta)
4864 .expect("checked at load");
4865 let part_bg = make_bg(
4866 self.ctx,
4867 self.kernels
4868 .gemm3_sk_pipeline(self.use_f16)
4869 .expect("checked at load"),
4870 &[x, &wbuf, &bbuf, &part],
4871 &meta,
4872 );
4873 let part32_bg = make_bg(
4874 self.ctx,
4875 self.kernels
4876 .gemm3_sk32_pipeline(self.use_f16)
4877 .expect("checked at load"),
4878 &[x, &wbuf, &bbuf, &part],
4879 &meta,
4880 );
4881 let skn = gemm3_sk_chunks(k as usize); let rmeta = uni(self.ctx, bytemuck::cast_slice(&[0u32, n, flags, skn]));
4883 let reduce = make_bg(
4884 self.ctx,
4885 &self.kernels.gemm3_sk_reduce,
4886 &[&part, &bbuf, y],
4887 &rmeta,
4888 );
4889 self.metas_sk.push((rmeta.clone(), n, k));
4890 self.metas_tokens.push(rmeta); GemmBg::V3Sk {
4892 plain,
4893 part: part_bg,
4894 part32: part32_bg,
4895 reduce,
4896 k,
4897 f16a: false,
4898 }
4899 } else if v3 {
4900 GemmBg::V3(
4903 v3_bgs(self.ctx, self.kernels, self.use_f16, &bufs, &meta)
4904 .expect("checked at load"),
4905 )
4906 } else if gemm2_enabled() {
4907 GemmBg::V2(
4911 v2_bgs(self.ctx, self.kernels, self.use_f16, &bufs, &meta)
4912 .expect("checked at load"),
4913 )
4914 } else {
4915 GemmBg::V1(make_bg(
4916 self.ctx,
4917 self.kernels
4918 .gemm_pipeline(self.use_f16)
4919 .expect("checked at load"),
4920 &bufs,
4921 &meta,
4922 ))
4923 };
4924 if !has_b {
4925 self.scratch.push(bbuf);
4926 }
4927 self.steps.push(EStep::Gemm { bg, n });
4928 self.metas_tokens.push(meta);
4929 }
4930
4931 fn push_ln(
4933 &mut self,
4934 ni: usize,
4935 x: &wgpu::Buffer,
4936 res: Option<&wgpu::Buffer>,
4937 out: &wgpu::Buffer,
4938 ) {
4939 let n = &self.norms[ni];
4940 let rms = matches!(self.cfg.norm_kind, NormKind::RmsNorm);
4941 let flags =
4942 u32::from(res.is_some()) | (u32::from(n.b.is_some()) << 1) | (u32::from(rms) << 2);
4943 let meta = uni(
4944 self.ctx,
4945 bytemuck::cast_slice(&[self.cfg.hidden as u32, flags, self.eps.to_bits(), 0]),
4946 );
4947 let res_buf = res.unwrap_or(x); let b_placeholder;
4949 let bbuf = match &n.b {
4950 Some(b) => b,
4951 None => {
4952 b_placeholder = self.ctx.storage(&[0.0]);
4953 self.scratch.push(b_placeholder.clone());
4954 self.scratch.last().expect("just pushed")
4955 }
4956 };
4957 let bg = make_bg(
4958 self.ctx,
4959 &self.kernels.layernorm,
4960 &[x, res_buf, &n.w, bbuf, out],
4961 &meta,
4962 );
4963 self.steps.push(EStep::Ln { bg });
4964 self.scratch.push(meta); }
4966
4967 fn push_rope(&mut self, x: &wgpu::Buffer, theta: f32) {
4968 self.push_rope_heads(x, theta, self.cfg.n_heads);
4969 }
4970
4971 fn push_rope_heads(&mut self, x: &wgpu::Buffer, theta: f32, heads: usize) {
4973 let meta = uni(
4974 self.ctx,
4975 bytemuck::cast_slice(&[
4976 0u32,
4977 heads as u32,
4978 self.cfg.head_dim as u32,
4979 theta.to_bits(),
4980 ]),
4981 );
4982 let bg = make_bg(
4983 self.ctx,
4984 &self.kernels.rope,
4985 &[x, &self.seq_starts_buf, &self.seq_of_buf],
4986 &meta,
4987 );
4988 self.steps.push(EStep::Rope { bg });
4989 self.scratch.push(meta);
4990 }
4991
4992 fn push_qk_norm(&mut self, x: &wgpu::Buffer, heads: usize, w: Vec<f32>) {
4994 let wb = self.ctx.storage(&w);
4995 let meta = uni(
4996 self.ctx,
4997 bytemuck::cast_slice(&[
4998 0u32,
4999 heads as u32,
5000 self.cfg.head_dim as u32,
5001 self.eps.to_bits(),
5002 ]),
5003 );
5004 let bg = make_bg(self.ctx, &self.kernels.qk_norm, &[x, &wb], &meta);
5005 self.steps.push(EStep::QkNorm { bg });
5006 self.scratch.push(wb);
5007 self.scratch.push(meta);
5008 }
5009
5010 fn push_attn(&mut self, window: u32) {
5011 self.push_attn_src(window, None);
5012 }
5013
5014 fn push_attn_src(&mut self, window: u32, packed: Option<&wgpu::Buffer>) {
5018 let mode = match self.cfg.attn_mask {
5019 MaskKind::Bidirectional => 0u32,
5020 MaskKind::Causal => 1u32,
5021 };
5022 let meta = uni(
5023 self.ctx,
5024 bytemuck::cast_slice(&[
5025 0u32,
5026 self.cfg.n_heads as u32,
5027 self.cfg.head_dim as u32,
5028 mode,
5029 window,
5030 self.cfg.n_kv_heads as u32,
5031 u32::from(packed.is_some()),
5032 0,
5033 ]),
5034 );
5035 let (q, k, v) = match packed {
5036 Some(b) => (b, b, b),
5037 None => (&self.qb, &self.kb, &self.vb),
5038 };
5039 let v4 = self.cfg.head_dim.is_multiple_of(4)
5044 && std::env::var("OSFKB_ENC_ATTN_V4").ok().as_deref() != Some("0");
5045 let rb = v4
5046 && self.cfg.head_dim <= 64
5047 && std::env::var("OSFKB_ENC_ATTN_RB").ok().as_deref() != Some("0");
5048 let pl = if rb {
5049 &self.kernels.attn_rb4
5050 } else if v4 {
5051 &self.kernels.attn4
5052 } else {
5053 &self.kernels.attn
5054 };
5055 let bg = make_bg(
5056 self.ctx,
5057 pl,
5058 &[
5059 q,
5060 k,
5061 v,
5062 &self.cb,
5063 &self.seq_starts_buf,
5064 &self.seq_of_buf,
5065 &self.valid_buf,
5066 ],
5067 &meta,
5068 );
5069 self.steps.push(if rb {
5070 EStep::AttnRb { bg }
5071 } else if v4 {
5072 EStep::Attn4 { bg }
5073 } else {
5074 EStep::Attn { bg }
5075 });
5076 self.metas_tokens.push(meta);
5077 }
5078
5079 fn push_attn_disent(&mut self, pos_k: &wgpu::Buffer, pos_q: &wgpu::Buffer) {
5083 self.push_attn_disent_src(pos_k, pos_q, None);
5084 }
5085
5086 fn push_attn_disent_src(
5090 &mut self,
5091 pos_k: &wgpu::Buffer,
5092 pos_q: &wgpu::Buffer,
5093 packed: Option<&wgpu::Buffer>,
5094 ) {
5095 let rel = self
5096 .cfg
5097 .rel_attn
5098 .expect("push_attn_disent on a non-DeBERTa config");
5099 let meta = uni(
5100 self.ctx,
5101 bytemuck::cast_slice(&[
5102 0u32,
5103 self.cfg.n_heads as u32,
5104 self.cfg.head_dim as u32,
5105 rel.span as u32,
5106 rel.max_rel as u32,
5107 rel.scale_factor() as u32,
5108 u32::from(rel.c2p),
5109 u32::from(rel.p2c) | (u32::from(packed.is_some()) << 1),
5110 ]),
5111 );
5112 let rb = self.cfg.head_dim.is_multiple_of(4)
5115 && self.cfg.head_dim <= 64
5116 && std::env::var("OSFKB_ENC_ATTN_RB").ok().as_deref() != Some("0");
5117 let pl = if rb {
5118 &self.kernels.attn_disent_rb4
5119 } else {
5120 &self.kernels.attn_disent
5121 };
5122 let (q, k, v) = match packed {
5123 Some(b) => (b, b, b),
5124 None => (&self.qb, &self.kb, &self.vb),
5125 };
5126 let bg = make_bg(
5127 self.ctx,
5128 pl,
5129 &[
5130 q,
5131 k,
5132 v,
5133 &self.cb,
5134 &self.seq_starts_buf,
5135 &self.seq_of_buf,
5136 &self.valid_buf,
5137 pos_k,
5138 pos_q,
5139 ],
5140 &meta,
5141 );
5142 self.steps.push(if rb {
5143 EStep::AttnDisentRb { bg }
5144 } else {
5145 EStep::AttnDisent { bg }
5146 });
5147 self.metas_tokens.push(meta);
5148 }
5149
5150 fn push_glu(&mut self, act: Act) {
5151 let meta = uni(
5152 self.ctx,
5153 bytemuck::cast_slice(&[self.cfg.intermediate as u32, act_code(Some(act)), 0, 0]),
5154 );
5155 let bg = make_bg(self.ctx, &self.kernels.glu, &[&self.mid, &self.glu], &meta);
5156 self.steps.push(EStep::Glu { bg });
5157 self.scratch.push(meta);
5158 }
5159
5160 fn push_copy(&mut self, dst: &wgpu::Buffer, src: &wgpu::Buffer) {
5162 let meta = uni(self.ctx, bytemuck::cast_slice(&[0u32, 0, 0, 0]));
5163 let bg = make_bg(self.ctx, &self.kernels.copy, &[dst, src], &meta);
5164 self.steps.push(EStep::Copy { bg });
5165 self.metas_elems.push(meta);
5166 }
5167
5168 fn push_conv(&mut self, bcx: &wgpu::Buffer, conv_w: Vec<f32>, y: &wgpu::Buffer) {
5170 let wb = self.ctx.storage(&conv_w);
5171 let meta = uni(
5172 self.ctx,
5173 bytemuck::cast_slice(&[self.cfg.hidden as u32, self.cfg.conv_l as u32, 0, 0]),
5174 );
5175 let bg = make_bg(
5176 self.ctx,
5177 &self.kernels.conv,
5178 &[
5179 bcx,
5180 &wb,
5181 y,
5182 &self.seq_starts_buf,
5183 &self.seq_of_buf,
5184 &self.valid_buf,
5185 ],
5186 &meta,
5187 );
5188 self.steps.push(EStep::Conv { bg });
5189 self.scratch.push(wb);
5190 self.scratch.push(meta);
5191 }
5192
5193 fn push_l2_rows(&mut self, buf: &wgpu::Buffer, dim: usize) {
5195 let meta = uni(self.ctx, bytemuck::cast_slice(&[dim as u32, 0, 0, 0]));
5196 let bg = make_bg(self.ctx, &self.kernels.l2norm, &[buf], &meta);
5197 self.steps.push(EStep::L2Rows { bg });
5198 self.scratch.push(meta);
5199 }
5200
5201 fn push_add(&mut self, dst: &wgpu::Buffer, src: &wgpu::Buffer) {
5202 let meta = uni(self.ctx, bytemuck::cast_slice(&[0u32, 0, 0, 0]));
5203 let bg = make_bg(self.ctx, &self.kernels.add, &[dst, src], &meta);
5204 self.steps.push(EStep::Add { bg });
5205 self.metas_elems.push(meta);
5206 }
5207
5208 fn push_pool(&mut self, src: &wgpu::Buffer) {
5209 let h = self.cfg.hidden as u32;
5210 let meta = uni(self.ctx, bytemuck::cast_slice(&[h, 0, 0, 0]));
5211 let pool_pl = match self.cfg.pooling {
5215 crate::pooling::Pooling::LastToken => &self.kernels.last_pool,
5216 crate::pooling::Pooling::Cls | crate::pooling::Pooling::CrossEncoder { .. } => {
5223 &self.kernels.cls_pool
5224 }
5225 _ => &self.kernels.mean_pool,
5226 };
5227 self.bg_pool = Some(make_bg(
5228 self.ctx,
5229 pool_pl,
5230 &[src, &self.pooled, &self.seq_starts_buf],
5231 &meta,
5232 ));
5233 self.scratch.push(meta);
5234 let meta2 = uni(self.ctx, bytemuck::cast_slice(&[h, 0, 0, 0]));
5235 self.bg_l2 = Some(make_bg(
5236 self.ctx,
5237 &self.kernels.l2norm,
5238 &[&self.pooled],
5239 &meta2,
5240 ));
5241 self.scratch.push(meta2);
5242 }
5243
5244 fn norm_from(&mut self, w: Vec<f32>, b: Option<Vec<f32>>) -> usize {
5245 self.norms.push(EncGpuNorm {
5246 w: self.ctx.storage(&w),
5247 b: b.map(|bv| self.ctx.storage(&bv)),
5248 });
5249 self.norms.len() - 1
5250 }
5251
5252 #[allow(clippy::type_complexity)]
5254 fn build_bert_family(
5255 &mut self,
5256 st: &LazySt,
5257 ) -> Result<(Vec<f32>, Option<Vec<f32>>, Option<Vec<f32>>)> {
5258 let cfg = self.cfg;
5259 let prefix = ["", "bert.", "roberta."]
5260 .into_iter()
5261 .find(|p| st.has(&format!("{p}embeddings.word_embeddings.weight")))
5262 .context("word embeddings not found under known prefixes ('', 'bert.', 'roberta.')")?;
5263 let get = |name: &str| -> Result<Vec<f32>> { st.tensor_f32(&format!("{prefix}{name}")) };
5264 let word = get("embeddings.word_embeddings.weight")?;
5265 anyhow::ensure!(
5266 word.len() == cfg.vocab * cfg.hidden,
5267 "word embedding shape {} != vocab {} × hidden {}",
5268 word.len(),
5269 cfg.vocab,
5270 cfg.hidden
5271 );
5272 anyhow::ensure!(
5273 matches!(cfg.norm_kind, NormKind::LayerNorm { bias: true }),
5274 "BERT-family norms are biased LayerNorm"
5275 );
5276 let (h, im) = (cfg.hidden as u32, cfg.intermediate as u32);
5277 let mlp_act = match cfg.mlp_kind {
5278 MlpKind::Dense { act, .. } => act,
5279 MlpKind::Glu { .. } => anyhow::bail!("BERT family MLPs are dense"),
5280 };
5281 let jina = st.has(&format!("{prefix}encoder.layers.0.mixer.Wqkv.weight"));
5285 let emb_ln = if jina {
5287 self.norm_from(get("emb_ln.weight")?, Some(get("emb_ln.bias")?))
5288 } else {
5289 self.norm_from(
5290 get("embeddings.LayerNorm.weight")?,
5291 Some(get("embeddings.LayerNorm.bias")?),
5292 )
5293 };
5294 let emb_in = self.emb_in.clone();
5295 let cur = self.cur.clone();
5296 self.push_ln(emb_ln, &emb_in, None, &cur);
5297 let (alt, qb, kb, vb, cb, mid) = (
5298 self.alt.clone(),
5299 self.qb.clone(),
5300 self.kb.clone(),
5301 self.vb.clone(),
5302 self.cb.clone(),
5303 self.mid.clone(),
5304 );
5305 let qkv_fused = cfg.intermediate >= 3 * cfg.hidden
5312 && std::env::var("OSFKB_ENC_QKV_FUSED").ok().as_deref() != Some("0");
5313 let mid_qkv = self.mid.clone();
5314 for i in 0..cfg.n_layers {
5315 let p = if jina {
5316 format!("encoder.layers.{i}")
5317 } else {
5318 format!("encoder.layer.{i}")
5319 };
5320 let lin = |b: &mut Self, wn: String, bn: String, n: u32, k: u32| -> Result<usize> {
5321 b.upload_linear(
5322 &st.tensor_f32(&format!("{prefix}{wn}"))?,
5323 Some(st.tensor_f32(&format!("{prefix}{bn}"))?),
5324 n,
5325 k,
5326 )
5327 };
5328 let qkv = if jina {
5331 let w = st.tensor_f32(&format!("{prefix}{p}.mixer.Wqkv.weight"))?;
5332 let bias = st.tensor_f32(&format!("{prefix}{p}.mixer.Wqkv.bias"))?;
5333 anyhow::ensure!(
5334 w.len() == 3 * cfg.hidden * cfg.hidden,
5335 "{p}.mixer.Wqkv shape {} != 3·{}·{}",
5336 w.len(),
5337 cfg.hidden,
5338 cfg.hidden
5339 );
5340 Some(self.upload_linear(&w, Some(bias), 3 * h, h)?)
5341 } else if qkv_fused {
5342 let mut w = Vec::with_capacity(3 * cfg.hidden * cfg.hidden);
5343 let mut bias = Vec::with_capacity(3 * cfg.hidden);
5344 for part in ["query", "key", "value"] {
5345 w.extend(st.tensor_f32(&format!("{prefix}{p}.attention.self.{part}.weight"))?);
5346 bias.extend(st.tensor_f32(&format!("{prefix}{p}.attention.self.{part}.bias"))?);
5347 }
5348 Some(self.upload_linear(&w, Some(bias), 3 * h, h)?)
5349 } else {
5350 None
5351 };
5352 let (q, k, v) = if qkv.is_some() {
5353 (0, 0, 0) } else {
5355 (
5356 lin(
5357 self,
5358 format!("{p}.attention.self.query.weight"),
5359 format!("{p}.attention.self.query.bias"),
5360 h,
5361 h,
5362 )?,
5363 lin(
5364 self,
5365 format!("{p}.attention.self.key.weight"),
5366 format!("{p}.attention.self.key.bias"),
5367 h,
5368 h,
5369 )?,
5370 lin(
5371 self,
5372 format!("{p}.attention.self.value.weight"),
5373 format!("{p}.attention.self.value.bias"),
5374 h,
5375 h,
5376 )?,
5377 )
5378 };
5379 let (o_n, up_n, down_n, ln1_n, ln2_n) = if jina {
5381 ("mixer.out_proj", "mlp.fc1", "mlp.fc2", "norm1", "norm2")
5382 } else {
5383 (
5384 "attention.output.dense",
5385 "intermediate.dense",
5386 "output.dense",
5387 "attention.output.LayerNorm",
5388 "output.LayerNorm",
5389 )
5390 };
5391 let o = lin(
5392 self,
5393 format!("{p}.{o_n}.weight"),
5394 format!("{p}.{o_n}.bias"),
5395 h,
5396 h,
5397 )?;
5398 let up = lin(
5399 self,
5400 format!("{p}.{up_n}.weight"),
5401 format!("{p}.{up_n}.bias"),
5402 im,
5403 h,
5404 )?;
5405 let down = lin(
5406 self,
5407 format!("{p}.{down_n}.weight"),
5408 format!("{p}.{down_n}.bias"),
5409 h,
5410 im,
5411 )?;
5412 let ln1 = self.norm_from(
5413 get(&format!("{p}.{ln1_n}.weight"))?,
5414 Some(get(&format!("{p}.{ln1_n}.bias"))?),
5415 );
5416 let ln2 = self.norm_from(
5417 get(&format!("{p}.{ln2_n}.weight"))?,
5418 Some(get(&format!("{p}.{ln2_n}.bias"))?),
5419 );
5420 if let Some(qkv) = qkv {
5422 self.push_gemm(qkv, &cur, &mid_qkv, None);
5423 self.push_attn_src(cfg.layer_window[i], Some(&mid_qkv));
5424 } else {
5425 self.push_gemm(q, &cur, &qb, None);
5426 self.push_gemm(k, &cur, &kb, None);
5427 self.push_gemm(v, &cur, &vb, None);
5428 self.push_attn(cfg.layer_window[i]);
5429 }
5430 self.push_gemm(o, &cb, &qb, None);
5431 self.push_ln(ln1, &qb, Some(&cur), &alt);
5432 self.push_gemm(up, &alt, &mid, Some(mlp_act));
5433 self.push_gemm(down, &mid, &kb, None);
5434 self.push_ln(ln2, &kb, Some(&alt), &cur);
5435 }
5436 self.push_pool(&cur);
5437 Ok((
5438 word,
5439 Some(get("embeddings.position_embeddings.weight")?),
5440 if cfg.type_vocab > 0 {
5441 Some(get("embeddings.token_type_embeddings.weight")?)
5442 } else {
5443 None
5444 },
5445 ))
5446 }
5447
5448 #[allow(clippy::type_complexity)]
5455 fn build_deberta(
5456 &mut self,
5457 st: &LazySt,
5458 ) -> Result<(Vec<f32>, Option<Vec<f32>>, Option<Vec<f32>>)> {
5459 let cfg = self.cfg;
5460 let rel = cfg.rel_attn.context("DeBERTa config without rel_attn")?;
5461 let prefix = ["", "deberta.", "token_rep_layer.bert_layer.model."]
5466 .into_iter()
5467 .find(|p| st.has(&format!("{p}embeddings.word_embeddings.weight")))
5468 .context(
5469 "word embeddings not found under known prefixes ('', 'deberta.', \
5470 'token_rep_layer.bert_layer.model.')",
5471 )?;
5472 let get = |name: &str| -> Result<Vec<f32>> { st.tensor_f32(&format!("{prefix}{name}")) };
5473 let word = get("embeddings.word_embeddings.weight")?;
5474 anyhow::ensure!(
5475 word.len() == cfg.vocab * cfg.hidden,
5476 "word embedding shape {} != vocab {} × hidden {}",
5477 word.len(),
5478 cfg.vocab,
5479 cfg.hidden
5480 );
5481 let (h, im) = (cfg.hidden as u32, cfg.intermediate as u32);
5482 let mlp_act = match cfg.mlp_kind {
5483 MlpKind::Dense { act, .. } => act,
5484 MlpKind::Glu { .. } => anyhow::bail!("DeBERTa MLPs are dense"),
5485 };
5486
5487 let rows = 2 * rel.span;
5489 let hs = cfg.hidden;
5490 let mut rel_emb = get("encoder.rel_embeddings.weight")?;
5491 anyhow::ensure!(
5492 rel_emb.len() >= rows * hs,
5493 "rel_embeddings has {} rows, need 2·span = {rows}",
5494 rel_emb.len() / hs.max(1)
5495 );
5496 rel_emb.truncate(rows * hs);
5497 cpu_layer_norm_rows(
5498 &mut rel_emb,
5499 hs,
5500 &get("encoder.LayerNorm.weight")?,
5501 &get("encoder.LayerNorm.bias")?,
5502 cfg.eps,
5503 );
5504
5505 let emb_ln = self.norm_from(
5507 get("embeddings.LayerNorm.weight")?,
5508 Some(get("embeddings.LayerNorm.bias")?),
5509 );
5510 let emb_in = self.emb_in.clone();
5511 let cur = self.cur.clone();
5512 self.push_ln(emb_ln, &emb_in, None, &cur);
5513 let (alt, qb, kb, vb, cb, mid) = (
5514 self.alt.clone(),
5515 self.qb.clone(),
5516 self.kb.clone(),
5517 self.vb.clone(),
5518 self.cb.clone(),
5519 self.mid.clone(),
5520 );
5521 for i in 0..cfg.n_layers {
5522 let p = format!("encoder.layer.{i}");
5523 let getw = |n: String| -> Result<Vec<f32>> { st.tensor_f32(&format!("{prefix}{n}")) };
5524 let qw = getw(format!("{p}.attention.self.query_proj.weight"))?;
5525 let qb_w = getw(format!("{p}.attention.self.query_proj.bias"))?;
5526 let kw = getw(format!("{p}.attention.self.key_proj.weight"))?;
5527 let kb_w = getw(format!("{p}.attention.self.key_proj.bias"))?;
5528 let pos_k = self
5531 .ctx
5532 .storage(&cpu_linear(&rel_emb, &kw, Some(&kb_w), hs, hs));
5533 let pos_q = self
5534 .ctx
5535 .storage(&cpu_linear(&rel_emb, &qw, Some(&qb_w), hs, hs));
5536
5537 let vw = getw(format!("{p}.attention.self.value_proj.weight"))?;
5542 let vb_w = getw(format!("{p}.attention.self.value_proj.bias"))?;
5543 let fused = cfg.intermediate >= 3 * cfg.hidden
5544 && std::env::var("OSFKB_ENC_QKV_FUSED").ok().as_deref() != Some("0");
5545 let qkv = if fused {
5546 let mut w = Vec::with_capacity(3 * cfg.hidden * cfg.hidden);
5547 let mut bias = Vec::with_capacity(3 * cfg.hidden);
5548 for (ww, bb) in [(&qw, &qb_w), (&kw, &kb_w), (&vw, &vb_w)] {
5549 w.extend_from_slice(ww);
5550 bias.extend_from_slice(bb);
5551 }
5552 Some(self.upload_linear(&w, Some(bias), 3 * h, h)?)
5553 } else {
5554 None
5555 };
5556 let (q, k, v) = if qkv.is_some() {
5557 (0, 0, 0) } else {
5559 (
5560 self.upload_linear(&qw, Some(qb_w.clone()), h, h)?,
5561 self.upload_linear(&kw, Some(kb_w.clone()), h, h)?,
5562 self.upload_linear(&vw, Some(vb_w.clone()), h, h)?,
5563 )
5564 };
5565 let o = self.upload_linear(
5566 &getw(format!("{p}.attention.output.dense.weight"))?,
5567 Some(getw(format!("{p}.attention.output.dense.bias"))?),
5568 h,
5569 h,
5570 )?;
5571 let up = self.upload_linear(
5572 &getw(format!("{p}.intermediate.dense.weight"))?,
5573 Some(getw(format!("{p}.intermediate.dense.bias"))?),
5574 im,
5575 h,
5576 )?;
5577 let down = self.upload_linear(
5578 &getw(format!("{p}.output.dense.weight"))?,
5579 Some(getw(format!("{p}.output.dense.bias"))?),
5580 h,
5581 im,
5582 )?;
5583 let ln1 = self.norm_from(
5584 get(&format!("{p}.attention.output.LayerNorm.weight"))?,
5585 Some(get(&format!("{p}.attention.output.LayerNorm.bias"))?),
5586 );
5587 let ln2 = self.norm_from(
5588 get(&format!("{p}.output.LayerNorm.weight"))?,
5589 Some(get(&format!("{p}.output.LayerNorm.bias"))?),
5590 );
5591 if let Some(qkv) = qkv {
5592 self.push_gemm(qkv, &cur, &mid, None);
5593 let mid_qkv = self.mid.clone();
5594 self.push_attn_disent_src(&pos_k, &pos_q, Some(&mid_qkv));
5595 } else {
5596 self.push_gemm(q, &cur, &qb, None);
5597 self.push_gemm(k, &cur, &kb, None);
5598 self.push_gemm(v, &cur, &vb, None);
5599 self.push_attn_disent(&pos_k, &pos_q);
5600 }
5601 self.push_gemm(o, &cb, &qb, None);
5602 self.push_ln(ln1, &qb, Some(&cur), &alt);
5603 self.push_gemm(up, &alt, &mid, Some(mlp_act));
5604 self.push_gemm(down, &mid, &kb, None);
5605 self.push_ln(ln2, &kb, Some(&alt), &cur);
5606 }
5607 self.push_pool(&cur);
5608 Ok((
5609 word,
5610 None, if cfg.type_vocab > 0 {
5612 Some(get("embeddings.token_type_embeddings.weight")?)
5613 } else {
5614 None
5615 },
5616 ))
5617 }
5618
5619 #[allow(clippy::type_complexity)]
5622 fn build_lfm2_colbert(
5623 &mut self,
5624 st: &LazySt,
5625 dir: &Path,
5626 ) -> Result<(Vec<f32>, Option<Vec<f32>>, Option<Vec<f32>>)> {
5627 let cfg = self.cfg;
5628 let get = |name: &str| -> Result<Vec<f32>> { st.tensor_f32(name) };
5629 let word = get("embed_tokens.weight")?;
5630 anyhow::ensure!(
5631 word.len() == cfg.vocab * cfg.hidden,
5632 "embed_tokens shape {} != vocab {} × hidden {}",
5633 word.len(),
5634 cfg.vocab,
5635 cfg.hidden
5636 );
5637 let (hd, im) = (cfg.head_dim as u32, cfg.intermediate as u32);
5638 let hu = cfg.hidden as u32;
5639 let (qh, kvh) = (cfg.n_heads as u32 * hd, cfg.n_kv_heads as u32 * hd);
5640 let theta = match cfg.pos_kind {
5641 PosKind::Rope { theta, .. } => theta,
5642 PosKind::Learned { .. } => anyhow::bail!("LFM2 uses rotary positions"),
5643 };
5644 let glu_act = match cfg.mlp_kind {
5645 MlpKind::Glu { act } => act,
5646 MlpKind::Dense { .. } => anyhow::bail!("LFM2 MLPs are GLU"),
5647 };
5648 let dim = match cfg.pooling {
5649 Pooling::PerToken { dim } => dim,
5650 other => anyhow::bail!("ColBERT plan requires PerToken pooling, got {other:?}"),
5651 };
5652 let emb_in = self.emb_in.clone();
5653 let cur = self.cur.clone();
5654 self.push_copy(&cur, &emb_in);
5655 let (alt, qb, kb, vb, cb, mid, glu) = (
5656 self.alt.clone(),
5657 self.qb.clone(),
5658 self.kb.clone(),
5659 self.vb.clone(),
5660 self.cb.clone(),
5661 self.mid.clone(),
5662 self.glu.clone(),
5663 );
5664 for i in 0..cfg.n_layers {
5665 let p = format!("layers.{i}");
5666 let an = self.norm_from(get(&format!("{p}.operator_norm.weight"))?, None);
5667 self.push_ln(an, &cur, None, &alt);
5668 if cfg.layer_is_attn[i] {
5669 let q = self.upload_linear(
5670 &get(&format!("{p}.self_attn.q_proj.weight"))?,
5671 None,
5672 qh,
5673 hu,
5674 )?;
5675 let k = self.upload_linear(
5676 &get(&format!("{p}.self_attn.k_proj.weight"))?,
5677 None,
5678 kvh,
5679 hu,
5680 )?;
5681 let v = self.upload_linear(
5682 &get(&format!("{p}.self_attn.v_proj.weight"))?,
5683 None,
5684 kvh,
5685 hu,
5686 )?;
5687 let o = self.upload_linear(
5688 &get(&format!("{p}.self_attn.out_proj.weight"))?,
5689 None,
5690 hu,
5691 qh,
5692 )?;
5693 let q_norm_w = get(&format!("{p}.self_attn.q_layernorm.weight"))?;
5694 let k_norm_w = get(&format!("{p}.self_attn.k_layernorm.weight"))?;
5695 self.push_gemm(q, &alt, &qb, None);
5696 self.push_gemm(k, &alt, &kb, None);
5697 self.push_gemm(v, &alt, &vb, None);
5698 self.push_qk_norm(&qb, cfg.n_heads, q_norm_w);
5699 self.push_qk_norm(&kb, cfg.n_kv_heads, k_norm_w);
5700 self.push_rope_heads(&qb, theta, cfg.n_heads);
5701 self.push_rope_heads(&kb, theta, cfg.n_kv_heads);
5702 self.push_attn(0);
5703 self.push_gemm(o, &cb, &alt, None);
5704 } else {
5705 let in_proj = self.upload_linear(
5706 &get(&format!("{p}.conv.in_proj.weight"))?,
5707 None,
5708 3 * hu,
5709 hu,
5710 )?;
5711 let conv_w = get(&format!("{p}.conv.conv.weight"))?;
5712 anyhow::ensure!(
5713 conv_w.len() == cfg.hidden * cfg.conv_l,
5714 "{p} conv taps {} != hidden × conv_l",
5715 conv_w.len()
5716 );
5717 let out_proj =
5718 self.upload_linear(&get(&format!("{p}.conv.out_proj.weight"))?, None, hu, hu)?;
5719 self.push_gemm(in_proj, &alt, &mid, None);
5720 self.push_conv(&mid, conv_w, &cb);
5721 self.push_gemm(out_proj, &cb, &alt, None);
5722 }
5723 self.push_add(&cur, &alt);
5724 let mn = self.norm_from(get(&format!("{p}.ffn_norm.weight"))?, None);
5725 self.push_ln(mn, &cur, None, &alt);
5726 let gate = get(&format!("{p}.feed_forward.w1.weight"))?;
5727 let up = get(&format!("{p}.feed_forward.w3.weight"))?;
5728 let mut wi = gate;
5729 wi.extend_from_slice(&up);
5730 let wi = self.upload_linear(&wi, None, 2 * im, hu)?;
5731 let wo =
5732 self.upload_linear(&get(&format!("{p}.feed_forward.w2.weight"))?, None, hu, im)?;
5733 self.push_gemm(wi, &alt, &mid, None);
5734 self.push_glu(glu_act);
5735 self.push_gemm(wo, &glu, &alt, None);
5736 self.push_add(&cur, &alt);
5737 }
5738 let fin = self.norm_from(get("embedding_norm.weight")?, None);
5741 self.push_ln(fin, &cur, None, &alt);
5742 let dense_st = LazySt::open(&dir.join("1_Dense"))?;
5743 let proj =
5744 self.upload_linear(&dense_st.tensor_f32("linear.weight")?, None, dim as u32, hu)?;
5745 let ptb = self.ctx.storage(&vec![
5746 0f32;
5747 self.pooled.size() as usize / (4 * self.cfg.hidden)
5748 * dim
5749 ]);
5750 self.push_gemm(proj, &alt, &ptb, None);
5751 self.push_l2_rows(&ptb, dim);
5752 self.ptb = Some(ptb);
5753 Ok((word, None, None))
5754 }
5755
5756 #[allow(clippy::type_complexity)]
5760 fn build_nomic(
5761 &mut self,
5762 st: &LazySt,
5763 ) -> Result<(Vec<f32>, Option<Vec<f32>>, Option<Vec<f32>>)> {
5764 let cfg = self.cfg;
5765 let get = |name: &str| -> Result<Vec<f32>> { st.tensor_f32(name) };
5766 let word = get("embeddings.word_embeddings.weight")?;
5767 anyhow::ensure!(
5768 word.len() == cfg.vocab * cfg.hidden,
5769 "word embedding shape {} != vocab {} × hidden {}",
5770 word.len(),
5771 cfg.vocab,
5772 cfg.hidden
5773 );
5774 let hu = cfg.hidden as u32;
5775 let (im, h) = (cfg.intermediate as u32, cfg.hidden);
5776 let theta = match cfg.pos_kind {
5777 PosKind::Rope { theta, .. } => theta,
5778 PosKind::Learned { .. } => anyhow::bail!("nomic uses rotary positions"),
5779 };
5780 let glu_act = match cfg.mlp_kind {
5781 MlpKind::Glu { act } => act,
5782 MlpKind::Dense { .. } => anyhow::bail!("nomic MLPs are GLU"),
5783 };
5784 let emb_ln = self.norm_from(get("emb_ln.weight")?, Some(get("emb_ln.bias")?));
5786 let emb_in = self.emb_in.clone();
5787 let cur = self.cur.clone();
5788 self.push_ln(emb_ln, &emb_in, None, &cur);
5789 let (alt, qb, kb, vb, cb, mid, glu) = (
5790 self.alt.clone(),
5791 self.qb.clone(),
5792 self.kb.clone(),
5793 self.vb.clone(),
5794 self.cb.clone(),
5795 self.mid.clone(),
5796 self.glu.clone(),
5797 );
5798 for i in 0..cfg.n_layers {
5799 let p = format!("encoder.layers.{i}");
5800 let wqkv = get(&format!("{p}.attn.Wqkv.weight"))?;
5801 anyhow::ensure!(wqkv.len() == 3 * h * h, "{p}.attn.Wqkv shape");
5802 let q = self.upload_linear(&wqkv[..h * h], None, hu, hu)?;
5803 let k = self.upload_linear(&wqkv[h * h..2 * h * h], None, hu, hu)?;
5804 let v = self.upload_linear(&wqkv[2 * h * h..], None, hu, hu)?;
5805 let o =
5806 self.upload_linear(&get(&format!("{p}.attn.out_proj.weight"))?, None, hu, hu)?;
5807 let gate = get(&format!("{p}.mlp.fc12.weight"))?;
5810 let lin_half = get(&format!("{p}.mlp.fc11.weight"))?;
5811 let mut wi = gate;
5812 wi.extend_from_slice(&lin_half);
5813 let wi = self.upload_linear(&wi, None, 2 * im, hu)?;
5814 let wo = self.upload_linear(&get(&format!("{p}.mlp.fc2.weight"))?, None, hu, im)?;
5815 let ln1 = self.norm_from(
5816 get(&format!("{p}.norm1.weight"))?,
5817 Some(get(&format!("{p}.norm1.bias"))?),
5818 );
5819 let ln2 = self.norm_from(
5820 get(&format!("{p}.norm2.weight"))?,
5821 Some(get(&format!("{p}.norm2.bias"))?),
5822 );
5823 self.push_gemm(q, &cur, &qb, None);
5826 self.push_gemm(k, &cur, &kb, None);
5827 self.push_gemm(v, &cur, &vb, None);
5828 self.push_rope(&qb, theta);
5829 self.push_rope(&kb, theta);
5830 self.push_attn(cfg.layer_window[i]);
5831 self.push_gemm(o, &cb, &qb, None);
5832 self.push_ln(ln1, &qb, Some(&cur), &alt);
5833 self.push_gemm(wi, &alt, &mid, None);
5834 self.push_glu(glu_act);
5835 self.push_gemm(wo, &glu, &kb, None);
5836 self.push_ln(ln2, &kb, Some(&alt), &cur);
5837 }
5838 self.push_pool(&cur);
5839 Ok((
5840 word,
5841 None,
5842 Some(get("embeddings.token_type_embeddings.weight")?),
5843 ))
5844 }
5845
5846 #[allow(clippy::type_complexity)]
5851 fn build_qwen3_embed(
5852 &mut self,
5853 st: &LazySt,
5854 ) -> Result<(Vec<f32>, Option<Vec<f32>>, Option<Vec<f32>>)> {
5855 let cfg = self.cfg;
5856 let get = |name: &str| -> Result<Vec<f32>> { st.tensor_f32(&format!("model.{name}")) };
5857 let word = get("embed_tokens.weight")?;
5858 anyhow::ensure!(
5859 word.len() == cfg.vocab * cfg.hidden,
5860 "embed_tokens shape {} != vocab {} × hidden {}",
5861 word.len(),
5862 cfg.vocab,
5863 cfg.hidden
5864 );
5865 let (hd, im) = (cfg.head_dim as u32, cfg.intermediate as u32);
5866 let hu = cfg.hidden as u32;
5867 let (qh, kvh) = (cfg.n_heads as u32 * hd, cfg.n_kv_heads as u32 * hd);
5868 let glu_act = match cfg.mlp_kind {
5869 MlpKind::Glu { act } => act,
5870 MlpKind::Dense { .. } => anyhow::bail!("Qwen3 embedder MLPs are GLU"),
5871 };
5872 let emb_in = self.emb_in.clone();
5876 let cur = self.cur.clone();
5877 self.push_copy(&cur, &emb_in);
5878 let (alt, qb, kb, vb, cb, mid, glu) = (
5879 self.alt.clone(),
5880 self.qb.clone(),
5881 self.kb.clone(),
5882 self.vb.clone(),
5883 self.cb.clone(),
5884 self.mid.clone(),
5885 self.glu.clone(),
5886 );
5887 for i in 0..cfg.n_layers {
5888 let p = format!("layers.{i}");
5889 let q =
5890 self.upload_linear(&get(&format!("{p}.self_attn.q_proj.weight"))?, None, qh, hu)?;
5891 let k = self.upload_linear(
5892 &get(&format!("{p}.self_attn.k_proj.weight"))?,
5893 None,
5894 kvh,
5895 hu,
5896 )?;
5897 let v = self.upload_linear(
5898 &get(&format!("{p}.self_attn.v_proj.weight"))?,
5899 None,
5900 kvh,
5901 hu,
5902 )?;
5903 let o =
5904 self.upload_linear(&get(&format!("{p}.self_attn.o_proj.weight"))?, None, hu, qh)?;
5905 let gate = get(&format!("{p}.mlp.gate_proj.weight"))?;
5906 let up = get(&format!("{p}.mlp.up_proj.weight"))?;
5907 let mut wi = gate;
5908 wi.extend_from_slice(&up);
5909 let wi = self.upload_linear(&wi, None, 2 * im, hu)?;
5910 let wo =
5911 self.upload_linear(&get(&format!("{p}.mlp.down_proj.weight"))?, None, hu, im)?;
5912 let an = self.norm_from(get(&format!("{p}.input_layernorm.weight"))?, None);
5913 let mn = self.norm_from(get(&format!("{p}.post_attention_layernorm.weight"))?, None);
5914 let q_norm_w = get(&format!("{p}.self_attn.q_norm.weight"))?;
5915 let k_norm_w = get(&format!("{p}.self_attn.k_norm.weight"))?;
5916
5917 self.push_ln(an, &cur, None, &alt);
5918 self.push_gemm(q, &alt, &qb, None);
5919 self.push_gemm(k, &alt, &kb, None);
5920 self.push_gemm(v, &alt, &vb, None);
5921 self.push_qk_norm(&qb, cfg.n_heads, q_norm_w);
5922 self.push_qk_norm(&kb, cfg.n_kv_heads, k_norm_w);
5923 let theta = match cfg.pos_kind {
5924 PosKind::Rope { theta, .. } => theta,
5925 PosKind::Learned { .. } => anyhow::bail!("Qwen3 embedder uses rotary positions"),
5926 };
5927 self.push_rope_heads(&qb, theta, cfg.n_heads);
5928 self.push_rope_heads(&kb, theta, cfg.n_kv_heads);
5929 self.push_attn(0);
5930 self.push_gemm(o, &cb, &alt, None);
5931 self.push_add(&cur, &alt);
5932 self.push_ln(mn, &cur, None, &alt);
5933 self.push_gemm(wi, &alt, &mid, None);
5934 self.push_glu(glu_act);
5935 self.push_gemm(wo, &glu, &alt, None);
5936 self.push_add(&cur, &alt);
5937 }
5938 let fin = self.norm_from(get("norm.weight")?, None);
5939 self.push_ln(fin, &cur, None, &alt);
5940 self.push_pool(&alt);
5941 Ok((word, None, None))
5942 }
5943
5944 #[allow(clippy::type_complexity)]
5953 fn build_siglip_vision(
5954 &mut self,
5955 st: &LazySt,
5956 ) -> Result<(Vec<f32>, Option<Vec<f32>>, Option<Vec<f32>>)> {
5957 let cfg = self.cfg;
5958 let get =
5959 |name: &str| -> Result<Vec<f32>> { st.tensor_f32(&format!("vision_model.{name}")) };
5960 let (h, im) = (cfg.hidden as u32, cfg.intermediate as u32);
5961 let mlp_act = match cfg.mlp_kind {
5962 MlpKind::Dense { act, .. } => act,
5963 MlpKind::Glu { .. } => anyhow::bail!("SigLIP vision MLPs are dense"),
5964 };
5965 let emb_in = self.emb_in.clone();
5969 let cur = self.cur.clone();
5970 self.push_copy(&cur, &emb_in);
5971 let (alt, qb, kb, vb, cb, mid) = (
5972 self.alt.clone(),
5973 self.qb.clone(),
5974 self.kb.clone(),
5975 self.vb.clone(),
5976 self.cb.clone(),
5977 self.mid.clone(),
5978 );
5979 for i in 0..cfg.n_layers {
5980 let p = format!("encoder.layers.{i}");
5981 let biased = |b: &mut Self, name: &str, n: u32, k: u32| -> Result<usize> {
5983 b.upload_linear(
5984 &st.tensor_f32(&format!("vision_model.{p}.{name}.weight"))?,
5985 Some(st.tensor_f32(&format!("vision_model.{p}.{name}.bias"))?),
5986 n,
5987 k,
5988 )
5989 };
5990 let q = biased(self, "self_attn.q_proj", h, h)?;
5991 let k = biased(self, "self_attn.k_proj", h, h)?;
5992 let v = biased(self, "self_attn.v_proj", h, h)?;
5993 let o = biased(self, "self_attn.out_proj", h, h)?;
5994 let up = biased(self, "mlp.fc1", im, h)?;
5995 let down = biased(self, "mlp.fc2", h, im)?;
5996 let ln1 = self.norm_from(
5997 get(&format!("{p}.layer_norm1.weight"))?,
5998 Some(get(&format!("{p}.layer_norm1.bias"))?),
5999 );
6000 let ln2 = self.norm_from(
6001 get(&format!("{p}.layer_norm2.weight"))?,
6002 Some(get(&format!("{p}.layer_norm2.bias"))?),
6003 );
6004 self.push_ln(ln1, &cur, None, &alt);
6007 self.push_gemm(q, &alt, &qb, None);
6008 self.push_gemm(k, &alt, &kb, None);
6009 self.push_gemm(v, &alt, &vb, None);
6010 self.push_attn(0); self.push_gemm(o, &cb, &alt, None);
6012 self.push_add(&cur, &alt);
6013 self.push_ln(ln2, &cur, None, &alt);
6014 self.push_gemm(up, &alt, &mid, Some(mlp_act));
6015 self.push_gemm(down, &mid, &alt, None);
6016 self.push_add(&cur, &alt);
6017 }
6018 let fin = self.norm_from(
6019 get("post_layernorm.weight")?,
6020 Some(get("post_layernorm.bias")?),
6021 );
6022 self.push_ln(fin, &cur, None, &alt);
6023 self.push_copy(&cur, &alt);
6026 Ok((Vec::new(), None, None))
6027 }
6028
6029 #[allow(clippy::type_complexity)]
6032 fn build_modernbert(
6033 &mut self,
6034 st: &LazySt,
6035 ) -> Result<(Vec<f32>, Option<Vec<f32>>, Option<Vec<f32>>)> {
6036 let cfg = self.cfg;
6037 let prefix = ["", "model."]
6038 .into_iter()
6039 .find(|p| st.has(&format!("{p}embeddings.tok_embeddings.weight")))
6040 .context("ModernBERT tok_embeddings not found under known prefixes ('', 'model.')")?;
6041 let get = |name: &str| -> Result<Vec<f32>> { st.tensor_f32(&format!("{prefix}{name}")) };
6042 let has = |name: &str| st.has(&format!("{prefix}{name}"));
6043 let word = get("embeddings.tok_embeddings.weight")?;
6044 anyhow::ensure!(
6045 word.len() == cfg.vocab * cfg.hidden,
6046 "tok embedding shape {} != vocab {} × hidden {}",
6047 word.len(),
6048 cfg.vocab,
6049 cfg.hidden
6050 );
6051 anyhow::ensure!(!cfg.qkv_bias, "ModernBERT is bias-free");
6052 let (theta, local_theta) = match cfg.pos_kind {
6053 PosKind::Rope { theta, local_theta } => (theta, local_theta),
6054 PosKind::Learned { .. } => anyhow::bail!("ModernBERT uses rotary positions"),
6055 };
6056 let glu_act = match cfg.mlp_kind {
6057 MlpKind::Glu { act } => act,
6058 MlpKind::Dense { .. } => anyhow::bail!("ModernBERT MLPs are GLU"),
6059 };
6060 let norm_of = |b: &mut Self, name: &str| -> Result<usize> {
6061 let bias_name = format!("{}.bias", name.trim_end_matches(".weight"));
6062 let bias = if has(&bias_name) {
6063 Some(get(&bias_name)?)
6064 } else {
6065 None
6066 };
6067 Ok(b.norm_from(get(name)?, bias))
6068 };
6069 let hu = cfg.hidden as u32;
6070 let (im, h) = (cfg.intermediate as u32, cfg.hidden);
6071 let emb_ln = norm_of(self, "embeddings.norm.weight")?;
6073 let emb_in = self.emb_in.clone();
6074 let cur = self.cur.clone();
6075 self.push_ln(emb_ln, &emb_in, None, &cur);
6076 let (alt, qb, kb, vb, cb, mid, glu) = (
6077 self.alt.clone(),
6078 self.qb.clone(),
6079 self.kb.clone(),
6080 self.vb.clone(),
6081 self.cb.clone(),
6082 self.mid.clone(),
6083 self.glu.clone(),
6084 );
6085 for i in 0..cfg.n_layers {
6086 let p = format!("layers.{i}");
6087 let wqkv = get(&format!("{p}.attn.Wqkv.weight"))?;
6088 anyhow::ensure!(
6089 wqkv.len() == 3 * h * h,
6090 "{p}.attn.Wqkv shape {} != 3·{h}·{h}",
6091 wqkv.len()
6092 );
6093 let q = self.upload_linear(&wqkv[..h * h], None, hu, hu)?;
6094 let k = self.upload_linear(&wqkv[h * h..2 * h * h], None, hu, hu)?;
6095 let v = self.upload_linear(&wqkv[2 * h * h..], None, hu, hu)?;
6096 let o = self.upload_linear(&get(&format!("{p}.attn.Wo.weight"))?, None, hu, hu)?;
6097 let wi = self.upload_linear(&get(&format!("{p}.mlp.Wi.weight"))?, None, 2 * im, hu)?;
6098 let wo = self.upload_linear(&get(&format!("{p}.mlp.Wo.weight"))?, None, hu, im)?;
6099 let window = cfg.layer_window[i];
6100 let layer_theta = if window > 0 { local_theta } else { theta };
6101 let attn_src = if i == 0 && cfg.skip_first_attn_norm {
6103 anyhow::ensure!(
6104 !has(&format!("{p}.attn_norm.weight")),
6105 "layer 0 attn_norm present but config says skip — checkpoint mismatch"
6106 );
6107 &cur
6108 } else {
6109 let an = norm_of(self, &format!("{p}.attn_norm.weight"))?;
6110 self.push_ln(an, &cur, None, &alt);
6111 &alt
6112 };
6113 self.push_gemm(q, attn_src, &qb, None);
6114 self.push_gemm(k, attn_src, &kb, None);
6115 self.push_gemm(v, attn_src, &vb, None);
6116 self.push_rope(&qb, layer_theta);
6117 self.push_rope(&kb, layer_theta);
6118 self.push_attn(window);
6119 self.push_gemm(o, &cb, &alt, None);
6120 self.push_add(&cur, &alt);
6121 let mn = norm_of(self, &format!("{p}.mlp_norm.weight"))?;
6122 self.push_ln(mn, &cur, None, &alt);
6123 self.push_gemm(wi, &alt, &mid, None);
6124 self.push_glu(glu_act);
6125 self.push_gemm(wo, &glu, &alt, None);
6126 self.push_add(&cur, &alt);
6127 }
6128 let fin = norm_of(self, "final_norm.weight")?;
6129 self.push_ln(fin, &cur, None, &alt);
6130 self.push_pool(&alt);
6131 Ok((word, None, None))
6132 }
6133}