use crate::GpuCtx;
use crate::encoder_weights::{Act, MaskKind};
use crate::forward::ShaderModuleTuned as _;
use crate::forward::{make_bg, pipeline, uni};
use anyhow::Result;
pub(crate) fn act_code(act: Option<Act>) -> u32 {
match act {
None => 0,
Some(Act::GeluErf) => 1,
Some(Act::GeluTanh) => 2,
Some(Act::Silu) => 3,
Some(Act::Tanh) => 4,
Some(Act::Relu) => 5,
}
}
const ACT_FNS: &str = r#"
fn erf_as(x: f32) -> f32 {
let s = sign(x);
let a = abs(x);
let t = 1.0 / (1.0 + 0.3275911 * a);
let y = 1.0 - (((((1.061405429 * t - 1.453152027) * t) + 1.421413741) * t - 0.284496736) * t
+ 0.254829592) * t * exp(-a * a);
return s * y;
}
fn apply_act(v: f32, code: u32) -> f32 {
switch code {
case 1u: { return 0.5 * v * (1.0 + erf_as(v * 0.70710678)); }
// tanh-GELU. The argument is CLAMPED, exactly as the decoder's kernels already do
// (forward.rs / lib.rs) — WGSL's `tanh` overflows to NaN on some backends for |arg| >~ 88,
// because a naive (e^2x - 1)/(e^2x + 1) expansion divides inf by inf. `arg` reaches 105 at a
// pre-activation of only 13.8, which the text encoders never hit and a ViT's outlier tokens
// hit on the first block. tanh is ±1 to well inside f32 epsilon by |arg| = 20, so this is
// EXACT, not an approximation.
case 2u: {
let a = clamp(0.7978845608 * (v + 0.044715 * v * v * v), -20.0, 20.0);
return 0.5 * v * (1.0 + tanh(a));
}
case 3u: { return v / (1.0 + exp(-v)); }
case 4u: { return tanh(v); }
case 5u: { return max(v, 0.0); }
default: { return v; }
}
}
"#;
pub(crate) fn enc_gemm_src(f16: bool) -> String {
let (enable, wty) = if f16 {
("enable f16;", "f16")
} else {
("", "f32")
};
format!(
r#"{enable}
struct Meta {{ m: u32, n: u32, k: u32, flags: u32 }}
@group(0) @binding(0) var<storage, read> x: array<f32>;
@group(0) @binding(1) var<storage, read> w: array<{wty}>;
@group(0) @binding(2) var<storage, read> bias: array<f32>;
@group(0) @binding(3) var<storage, read_write> y: array<f32>;
@group(0) @binding(4) var<uniform> mt: Meta;
var<workgroup> xt: array<f32, 256>;
var<workgroup> wt: array<f32, 256>;
{ACT_FNS}
@compute @workgroup_size(16, 16)
fn main(@builtin(workgroup_id) wg: vec3<u32>, @builtin(local_invocation_id) li: vec3<u32>) {{
let row = wg.y * 16u + li.y;
let col = wg.x * 16u + li.x;
var acc = 0.0;
let ktiles = (mt.k + 15u) / 16u;
for (var kb = 0u; kb < ktiles; kb++) {{
let kx = kb * 16u + li.x;
var xv = 0.0;
if (row < mt.m && kx < mt.k) {{ xv = x[row * mt.k + kx]; }}
xt[li.y * 16u + li.x] = xv;
let kw = kb * 16u + li.y;
var wv = 0.0;
if (col < mt.n && kw < mt.k) {{ wv = f32(w[col * mt.k + kw]); }}
wt[li.y * 16u + li.x] = wv;
workgroupBarrier();
for (var kk = 0u; kk < 16u; kk++) {{
acc += xt[li.y * 16u + kk] * wt[kk * 16u + li.x];
}}
workgroupBarrier();
}}
if (row < mt.m && col < mt.n) {{
if ((mt.flags & 1u) != 0u) {{ acc += bias[col]; }}
y[row * mt.n + col] = apply_act(acc, mt.flags >> 8u);
}}
}}
"#
)
}
pub fn enc_gemm2_src(f16: bool, bm: usize, bn: usize) -> String {
let (enable, wty) = if f16 {
("enable f16;", "f16")
} else {
("", "f32")
};
const BK: usize = 16;
let tm = bm / 16; let cpt = bn / 16; let x_iters = bm * BK / 256; let w_iters = bn / 16; let bn_ = bn;
debug_assert!(tm >= 1 && bm.is_multiple_of(16) && x_iters * 256 == bm * BK);
debug_assert!(
matches!(cpt, 2 | 4 | 8),
"cols/thread must be 2, 4 or 8, got {cpt}"
);
let (inner, acc_decl, epi) = if cpt <= 4 {
let mut inner = String::new();
for kk in 0..BK {
let comps = (0..cpt)
.map(|j| format!("ws[{kk}u * {bn_}u + col0 + {j}u]"))
.collect::<Vec<_>>()
.join(", ");
inner.push_str(&format!(" {{ let b = vec{cpt}<f32>({comps});\n"));
for i in 0..tm {
inner.push_str(&format!(
" acc{i} += xs[(row0 + {i}u) * {BK}u + {kk}u] * b;\n"
));
}
inner.push_str(" }\n");
}
let acc_decl = (0..tm)
.map(|i| format!(" var acc{i} = vec{cpt}<f32>(0.0);"))
.collect::<Vec<_>>()
.join("\n");
let mut epi = String::new();
for i in 0..tm {
epi.push_str(&format!(
" {{ let r = mrow + {i}u; if (r < mt.m) {{\n let acc = acc{i};\n"
));
for j in 0..cpt {
epi.push_str(&format!(
" {{ let c = ncol + {j}u; if (c < mt.n) {{ var v = acc[{j}u]; \
if ((mt.flags & 1u) != 0u) {{ v += bias[c]; }} \
y[r * mt.n + c] = apply_act(v, mt.flags >> 8u); }} }}\n"
));
}
epi.push_str(" } }\n");
}
(inner, acc_decl, epi)
} else {
const NV: usize = 2; let mut inner = String::new();
for kk in 0..BK {
for b in 0..NV {
let comps = (0..4)
.map(|j| format!("ws[{kk}u * {bn_}u + col0 + {}u]", b * 4 + j))
.collect::<Vec<_>>()
.join(", ");
inner.push_str(&format!(" let b{b}_{kk} = vec4<f32>({comps});\n"));
}
for i in 0..tm {
inner.push_str(&format!(
" let a{i}_{kk} = xs[(row0 + {i}u) * {BK}u + {kk}u];\n"
));
for b in 0..NV {
inner.push_str(&format!(" acc{i}_{b} += a{i}_{kk} * b{b}_{kk};\n"));
}
}
}
let acc_decl = (0..tm)
.flat_map(|i| (0..NV).map(move |b| format!(" var acc{i}_{b} = vec4<f32>(0.0);")))
.collect::<Vec<_>>()
.join("\n");
let mut epi = String::new();
for i in 0..tm {
epi.push_str(&format!(" {{ let r = mrow + {i}u; if (r < mt.m) {{\n"));
for b in 0..NV {
for j in 0..4 {
let c = b * 4 + j;
epi.push_str(&format!(
" {{ let c = ncol + {c}u; if (c < mt.n) {{ var v = acc{i}_{b}[{j}u]; \
if ((mt.flags & 1u) != 0u) {{ v += bias[c]; }} \
y[r * mt.n + c] = apply_act(v, mt.flags >> 8u); }} }}\n"
));
}
}
epi.push_str(" } }\n");
}
(inner, acc_decl, epi)
};
format!(
r#"{enable}
struct Meta {{ m: u32, n: u32, k: u32, flags: u32 }}
@group(0) @binding(0) var<storage, read> x: array<f32>;
@group(0) @binding(1) var<storage, read> w: array<{wty}>;
@group(0) @binding(2) var<storage, read> bias: array<f32>;
@group(0) @binding(3) var<storage, read_write> y: array<f32>;
@group(0) @binding(4) var<uniform> mt: Meta;
var<workgroup> xs: array<f32, {xs_len}>; // [BM][BK]
var<workgroup> ws: array<f32, {ws_len}>; // [BK][BN], transposed at stage
{ACT_FNS}
@compute @workgroup_size(16, 16)
fn main(@builtin(workgroup_id) wg: vec3<u32>, @builtin(local_invocation_id) li: vec3<u32>) {{
let t = li.y * 16u + li.x; // 0..255
let mbase = wg.y * {bm}u; // first row of this tile
let nbase = wg.x * {bn_}u; // first col of this tile
let row0 = li.y * {tm}u; // thread's rows, tile-local
let col0 = li.x * {cpt}u; // thread's cols, tile-local
let mrow = mbase + row0;
let ncol = nbase + col0;
{acc_decl}
let ktiles = (mt.k + {BK}u - 1u) / {BK}u;
for (var kb = 0u; kb < ktiles; kb++) {{
let koff = kb * {BK}u;
// Stage x[BM][BK]: consecutive threads read consecutive k (coalesced; x is [m,k]).
for (var i = 0u; i < {x_iters}u; i++) {{
let r = i * 16u + t / {BK}u;
let c = t % {BK}u;
let gr = mbase + r;
let gc = koff + c;
var v = 0.0;
if (gr < mt.m && gc < mt.k) {{ v = x[gr * mt.k + gc]; }}
xs[r * {BK}u + c] = v;
}}
// Stage w[BN][BK] TRANSPOSED into ws[BK][BN]: consecutive threads read consecutive k of
// one output column (w is [n,k]), so the global reads coalesce and the transpose happens
// in shared, where it is free.
for (var i = 0u; i < {w_iters}u; i++) {{
let c = i * 16u + t / {BK}u;
let kk = t % {BK}u;
let gc = nbase + c;
let gk = koff + kk;
var v = 0.0;
if (gc < mt.n && gk < mt.k) {{ v = f32(w[gc * mt.k + gk]); }}
ws[kk * {bn_}u + c] = v;
}}
workgroupBarrier();
{inner}
workgroupBarrier();
}}
{epi}}}
"#,
xs_len = bm * BK,
ws_len = BK * bn,
)
}
pub fn enc_gemm3_src(f16: bool, bm: usize, bn: usize, bk: usize) -> String {
let (enable, wty) = if f16 {
("enable f16;", "f16")
} else {
("", "f32")
};
let tm = bm / 16;
let cpt = bn / 16;
let nv = cpt / 4; let kq = bk / 4; assert!(cpt.is_multiple_of(4) && bk.is_multiple_of(4) && bm.is_multiple_of(16));
let bnq = bn / 4; let xs_len = bm * kq; let ws_len = bk * bnq; let x_rounds = xs_len.div_ceil(256);
let w_rounds = ws_len.div_ceil(256);
let acc_decl = (0..tm)
.flat_map(|i| (0..nv).map(move |b| format!(" var acc{i}_{b} = vec4<f32>(0.0);")))
.collect::<Vec<_>>()
.join("\n");
let mut inner = String::new();
for q in 0..kq {
for i in 0..tm {
inner.push_str(&format!(
" let a{q}_{i} = xs[(row0 + {i}u) * {kq}u + {q}u];\n"
));
}
for kk in 0..4 {
for b in 0..nv {
inner.push_str(&format!(
" let w{q}_{kk}_{b} = ws[({}u) * {bnq}u + colq + {b}u];\n",
q * 4 + kk
));
}
for i in 0..tm {
for b in 0..nv {
inner.push_str(&format!(
" acc{i}_{b} += a{q}_{i}[{kk}u] * w{q}_{kk}_{b};\n"
));
}
}
}
}
let mut epi = String::new();
for i in 0..tm {
epi.push_str(&format!(" {{ let r = mrow + {i}u; if (r < mt.m) {{\n"));
for b in 0..nv {
for j in 0..4 {
let c = b * 4 + j;
epi.push_str(&format!(
" {{ let c = ncol + {c}u; if (c < mt.n) {{ var v = acc{i}_{b}[{j}u]; \
if ((mt.flags & 1u) != 0u) {{ v += bias[c]; }} \
y[r * mt.n + c] = apply_act(v, mt.flags >> 8u); }} }}\n"
));
}
}
epi.push_str(" } }\n");
}
format!(
r#"{enable}
struct Meta {{ m: u32, n: u32, k: u32, flags: u32 }}
@group(0) @binding(0) var<storage, read> x: array<f32>; // [m, k]
@group(0) @binding(1) var<storage, read> w: array<{wty}>; // [k, n] -- TRANSPOSED at upload
@group(0) @binding(2) var<storage, read> bias: array<f32>;
@group(0) @binding(3) var<storage, read_write> y: array<f32>;
@group(0) @binding(4) var<uniform> mt: Meta;
var<workgroup> xs: array<vec4<f32>, {xs_len}>; // [bm][BK/4]
var<workgroup> ws: array<vec4<f32>, {ws_len}>; // [BK][bn/4]
{ACT_FNS}
@compute @workgroup_size(16, 16)
fn main(@builtin(workgroup_id) wg: vec3<u32>, @builtin(local_invocation_id) li: vec3<u32>) {{
let t = li.y * 16u + li.x;
let mbase = wg.y * {bm}u;
let nbase = wg.x * {bn}u;
let row0 = li.y * {tm}u;
let colq = li.x * {nv}u; // this thread's first COLUMN-QUAD in the tile
let mrow = mbase + row0;
let ncol = nbase + colq * 4u;
{acc_decl}
let ktiles = (mt.k + {bk}u - 1u) / {bk}u;
for (var kb = 0u; kb < ktiles; kb++) {{
let koff = kb * {bk}u;
// Stage x as vec4 along k (x is [m,k]: contiguous, and consecutive threads take consecutive
// k-quads of a row, so the warp's reads are contiguous too).
for (var i = 0u; i < {x_rounds}u; i++) {{
let idx = i * 256u + t;
if (idx < {xs_len}u) {{
let r = idx / {kq}u;
let q = idx % {kq}u;
let gr = mbase + r;
let gk = koff + q * 4u;
var v = vec4<f32>(0.0);
if (gr < mt.m) {{
if (gk + 0u < mt.k) {{ v.x = x[gr * mt.k + gk + 0u]; }}
if (gk + 1u < mt.k) {{ v.y = x[gr * mt.k + gk + 1u]; }}
if (gk + 2u < mt.k) {{ v.z = x[gr * mt.k + gk + 2u]; }}
if (gk + 3u < mt.k) {{ v.w = x[gr * mt.k + gk + 3u]; }}
}}
xs[r * {kq}u + q] = v;
}}
}}
// Stage wT as vec4 along n. THIS is what the transpose bought: at a fixed k the columns are
// contiguous, so one thread loads four of them in one go and consecutive threads load the
// next four.
for (var i = 0u; i < {w_rounds}u; i++) {{
let idx = i * 256u + t;
if (idx < {ws_len}u) {{
let kk = idx / {bnq}u;
let cq = idx % {bnq}u;
let gk = koff + kk;
let gc = nbase + cq * 4u;
var v = vec4<f32>(0.0);
if (gk < mt.k) {{
if (gc + 0u < mt.n) {{ v.x = f32(w[gk * mt.n + gc + 0u]); }}
if (gc + 1u < mt.n) {{ v.y = f32(w[gk * mt.n + gc + 1u]); }}
if (gc + 2u < mt.n) {{ v.z = f32(w[gk * mt.n + gc + 2u]); }}
if (gc + 3u < mt.n) {{ v.w = f32(w[gk * mt.n + gc + 3u]); }}
}}
ws[kk * {bnq}u + cq] = v;
}}
}}
workgroupBarrier();
{inner}
workgroupBarrier();
}}
{epi}}}
"#
)
}
pub fn enc_gemm3_f16a_src(bm: usize, bn: usize, bk: usize) -> String {
let tm = bm / 16;
let cpt = bn / 16;
let nv = cpt / 4;
let kq = bk / 4;
assert!(cpt.is_multiple_of(4) && bk.is_multiple_of(4) && bm.is_multiple_of(16));
let bnq = bn / 4;
let xs_len = bm * kq;
let ws_len = bk * bnq;
let x_rounds = xs_len.div_ceil(256);
let w_rounds = ws_len.div_ceil(256);
let acc_decl = (0..tm)
.flat_map(|i| (0..nv).map(move |b| format!(" var acc{i}_{b} = vec4<f32>(0.0);")))
.collect::<Vec<_>>()
.join("\n");
let hacc_decl = (0..tm)
.flat_map(|i| (0..nv).map(move |b| format!(" var h{i}_{b} = vec4<f16>(0.0);")))
.collect::<Vec<_>>()
.join("\n");
let fold = (0..tm)
.flat_map(|i| (0..nv).map(move |b| format!(" acc{i}_{b} += vec4<f32>(h{i}_{b});")))
.collect::<Vec<_>>()
.join("\n");
let mut inner = String::new();
for q in 0..kq {
for i in 0..tm {
inner.push_str(&format!(
" let a{q}_{i} = xs[(row0 + {i}u) * {kq}u + {q}u];\n"
));
}
for kk in 0..4 {
for b in 0..nv {
inner.push_str(&format!(
" let w{q}_{kk}_{b} = ws[({}u) * {bnq}u + colq + {b}u];\n",
q * 4 + kk
));
}
for i in 0..tm {
for b in 0..nv {
inner.push_str(&format!(
" h{i}_{b} = fma(vec4<f16>(a{q}_{i}[{kk}u]), w{q}_{kk}_{b}, h{i}_{b});\n"
));
}
}
}
}
let mut epi = String::new();
for i in 0..tm {
epi.push_str(&format!(" {{ let r = mrow + {i}u; if (r < mt.m) {{\n"));
for b in 0..nv {
for j in 0..4 {
let c = b * 4 + j;
epi.push_str(&format!(
" {{ let c = ncol + {c}u; if (c < mt.n) {{ var v = acc{i}_{b}[{j}u]; \
if ((mt.flags & 1u) != 0u) {{ v += bias[c]; }} \
y[r * mt.n + c] = apply_act(v, mt.flags >> 8u); }} }}\n"
));
}
}
epi.push_str(" } }\n");
}
format!(
r#"enable f16;
struct Meta {{ m: u32, n: u32, k: u32, flags: u32 }}
@group(0) @binding(0) var<storage, read> x: array<f32>; // [m, k]
@group(0) @binding(1) var<storage, read> w: array<vec4<f16>>; // [k, n/4] -- TRANSPOSED f16
@group(0) @binding(2) var<storage, read> bias: array<f32>;
@group(0) @binding(3) var<storage, read_write> y: array<f32>;
@group(0) @binding(4) var<uniform> mt: Meta;
var<workgroup> xs: array<vec4<f16>, {xs_len}>; // [bm][BK/4], converted once at stage
var<workgroup> ws: array<vec4<f16>, {ws_len}>; // [BK][bn/4]
{ACT_FNS}
@compute @workgroup_size(16, 16)
fn main(@builtin(workgroup_id) wg: vec3<u32>, @builtin(local_invocation_id) li: vec3<u32>) {{
let t = li.y * 16u + li.x;
let mbase = wg.y * {bm}u;
let nbase = wg.x * {bn}u;
let row0 = li.y * {tm}u;
let colq = li.x * {nv}u;
let mrow = mbase + row0;
let ncol = nbase + colq * 4u;
{acc_decl}
let ktiles = (mt.k + {bk}u - 1u) / {bk}u;
for (var kb = 0u; kb < ktiles; kb++) {{
let koff = kb * {bk}u;
for (var i = 0u; i < {x_rounds}u; i++) {{
let idx = i * 256u + t;
if (idx < {xs_len}u) {{
let r = idx / {kq}u;
let q = idx % {kq}u;
let gr = mbase + r;
let gk = koff + q * 4u;
var v = vec4<f32>(0.0);
if (gr < mt.m) {{
if (gk + 0u < mt.k) {{ v.x = x[gr * mt.k + gk + 0u]; }}
if (gk + 1u < mt.k) {{ v.y = x[gr * mt.k + gk + 1u]; }}
if (gk + 2u < mt.k) {{ v.z = x[gr * mt.k + gk + 2u]; }}
if (gk + 3u < mt.k) {{ v.w = x[gr * mt.k + gk + 3u]; }}
}}
xs[idx] = vec4<f16>(v);
}}
}}
for (var i = 0u; i < {w_rounds}u; i++) {{
let idx = i * 256u + t;
if (idx < {ws_len}u) {{
let kk = idx / {bnq}u;
let cq = idx % {bnq}u;
let gk = koff + kk;
var v = vec4<f16>(0.0);
if (gk < mt.k && (nbase + cq * 4u) < mt.n) {{ v = w[(gk * mt.n) / 4u + nbase / 4u + cq]; }}
ws[kk * {bnq}u + cq] = v;
}}
}}
workgroupBarrier();
{hacc_decl}
{inner}
{fold}
workgroupBarrier();
}}
{epi}}}
"#
)
}
pub fn enc_gemm4_f16_src(bm: usize, bn: usize, bk: usize) -> String {
enc_gemm4_typed_src(bm, bn, bk, true)
}
pub fn enc_gemm4_src(bm: usize, bn: usize, bk: usize) -> String {
enc_gemm4_typed_src(bm, bn, bk, false)
}
pub fn enc_gemm4_f16w_src(bm: usize, bn: usize, bk: usize) -> String {
let src = enc_gemm4_typed_src(bm, bn, bk, true)
.replace(
"@group(0) @binding(1) var<storage, read> w: array<f32>;",
"@group(0) @binding(1) var<storage, read> w: array<f16>;",
)
.replace(
"{ v = w[gk * mt.n + gc]; }",
"{ v = f32(w[gk * mt.n + gc]); }",
);
assert!(
src.contains("array<f16>;") && src.contains("f32(w[gk"),
"v4 w-binding/stage drifted"
);
src
}
fn enc_gemm4_typed_src(bm: usize, bn: usize, bk: usize, ab_f16: bool) -> String {
assert!(bm.is_multiple_of(8) && bn.is_multiple_of(8) && bk.is_multiple_of(8));
const SG_ROWS: usize = 2; const SG_COLS: usize = 4;
let nsg = SG_ROWS * SG_COLS;
let threads = nsg * 32;
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;
let abty = if ab_f16 { "f16" } else { "f32" };
let f16_enable = if ab_f16 { "enable f16;\n" } else { "" };
let acc_decl = (0..ai)
.flat_map(|i| {
(0..aj).map(move |j| {
format!(
" var acc{i}_{j} = coopLoadT<coop_mat8x8<f32, C>>(&bias8[nbase + qc * {cols_per_sg}u + {}u], mt.n);",
j * 8
)
})
})
.collect::<Vec<_>>()
.join("\n");
let mut inner = String::new();
for kk in 0..ksteps {
for i in 0..ai {
inner.push_str(&format!(
" let a{kk}_{i} = coopLoadT<coop_mat8x8<{abty}, A>>(&xs[(qr * {rows_per_sg}u + {}u) * {bk}u + {}u], {bk}u);\n",
i * 8, kk * 8
));
}
for j in 0..aj {
inner.push_str(&format!(
" let b{kk}_{j} = coopLoadT<coop_mat8x8<{abty}, B>>(&ws[{}u * {bn}u + qc * {cols_per_sg}u + {}u], {bn}u);\n",
kk * 8, j * 8
));
}
for i in 0..ai {
for j in 0..aj {
inner.push_str(&format!(
" acc{i}_{j} = coopMultiplyAdd(a{kk}_{i}, b{kk}_{j}, acc{i}_{j});\n"
));
}
}
}
let mut store = String::new();
for i in 0..ai {
for j in 0..aj {
store.push_str(&format!(
" {{ let r = mbase + qr * {rows_per_sg}u + {}u; let c = nbase + qc * {cols_per_sg}u + {}u;\n \
if (r < mt.m && c < mt.n) {{ coopStoreT(acc{i}_{j}, &y[r * mt.n + c], mt.n); }} }}\n",
i * 8, j * 8
));
}
}
let xs_len = bm * bk;
let ws_len = bk * bn;
let x_rounds = xs_len.div_ceil(threads);
let w_rounds = ws_len.div_ceil(threads);
format!(
r#"{f16_enable}enable wgpu_cooperative_matrix;
struct Meta {{ m: u32, n: u32, k: u32, flags: u32 }}
@group(0) @binding(0) var<storage, read> x: array<f32>; // [m, k]
@group(0) @binding(1) var<storage, read> w: array<f32>; // [k, n] -- TRANSPOSED at upload
@group(0) @binding(2) var<storage, read> bias8: array<f32>; // [8, n] -- bias broadcast to 8 rows
@group(0) @binding(3) var<storage, read_write> y: array<f32>; // [m, n]
@group(0) @binding(4) var<uniform> mt: Meta;
var<workgroup> xs: array<{abty}, {xs_len}>; // [bm][bk]
var<workgroup> ws: array<{abty}, {ws_len}>; // [bk][bn]
@compute @workgroup_size({threads})
fn main(
@builtin(workgroup_id) wg: vec3<u32>,
@builtin(local_invocation_index) t: u32,
@builtin(subgroup_id) sg: u32,
) {{
let mbase = wg.y * {bm}u;
let nbase = wg.x * {bn}u;
let qr = sg / {SG_COLS}u;
let qc = sg % {SG_COLS}u;
{acc_decl}
let ktiles = (mt.k + {bk}u - 1u) / {bk}u;
for (var kb = 0u; kb < ktiles; kb++) {{
let koff = kb * {bk}u;
for (var i = 0u; i < {x_rounds}u; i++) {{
let idx = i * {threads}u + t;
if (idx < {xs_len}u) {{
let r = idx / {bk}u;
let c = idx % {bk}u;
let gr = mbase + r;
let gc = koff + c;
var v = 0.0;
if (gr < mt.m && gc < mt.k) {{ v = x[gr * mt.k + gc]; }}
xs[idx] = {abty}(v);
}}
}}
for (var i = 0u; i < {w_rounds}u; i++) {{
let idx = i * {threads}u + t;
if (idx < {ws_len}u) {{
let kk = idx / {bn}u;
let c = idx % {bn}u;
let gk = koff + kk;
let gc = nbase + c;
var v = 0.0;
if (gk < mt.k && gc < mt.n) {{ v = w[gk * mt.n + gc]; }}
ws[idx] = {abty}(v);
}}
}}
workgroupBarrier();
{inner}
workgroupBarrier();
}}
{store}}}
"#
)
}
pub fn enc_gemm3_sk_src(f16: bool, bm: usize, bn: usize, bk: usize) -> String {
gemm3_sk_patch(enc_gemm3_src(f16, bm, bn, bk), bk)
}
pub fn enc_gemm3_sk_f16a_src(bm: usize, bn: usize, bk: usize) -> String {
gemm3_sk_patch(enc_gemm3_f16a_src(bm, bn, bk), bk)
}
fn gemm3_sk_patch(base: String, bk: usize) -> String {
let frags = [
"fn main(@builtin(workgroup_id) wg: vec3<u32>, @builtin(local_invocation_id) li: vec3<u32>) {",
" let ktiles = (mt.k + ",
"@group(0) @binding(3) var<storage, read_write> y: array<f32>;",
];
for f in frags {
assert!(base.contains(f), "v3 source drifted: {f}");
}
let src = base
.replace(
"fn main(@builtin(workgroup_id) wg: vec3<u32>, @builtin(local_invocation_id) li: vec3<u32>) {",
"fn main(@builtin(workgroup_id) wg: vec3<u32>, @builtin(local_invocation_id) li: vec3<u32>,\n @builtin(num_workgroups) nwg: vec3<u32>) {",
)
.replace(
"@group(0) @binding(3) var<storage, read_write> y: array<f32>;",
"@group(0) @binding(3) var<storage, read_write> y: array<f32>; // [nz, m, n] partials",
)
.replace(
" let t = li.y * 16u + li.x;",
" let t = li.y * 16u + li.x;\n if (mt.flags == 0xFFFFFFFFu) { y[0] = bias[0]; } // keep bias bound (never true)",
);
let (patched, klooped) = {
let needle = format!(
" let ktiles = (mt.k + {bk}u - 1u) / {bk}u;\n for (var kb = 0u; kb < ktiles; kb++) {{\n let koff = kb * {bk}u;"
);
let replacement = format!(
" 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;"
);
let ok = src.contains(&needle);
(src.replace(&needle, &replacement), ok)
};
assert!(klooped, "v3 k-loop shape drifted");
let mut out = patched;
assert!(
out.contains("var v = acc") || out.contains("var v = f32(acc"),
"v3/f16a epilogue drifted"
);
out = rewrite_v3_epilogue_for_sk(&out);
out
}
fn rewrite_v3_epilogue_for_sk(src: &str) -> String {
let mut out = String::with_capacity(src.len());
for line in src.lines() {
let t = line.trim_start();
if let Some(rest) = t.strip_prefix("{ let c = ncol + ") {
let j = rest.split("u;").next().expect("column offset");
let acc = rest
.split("var v = ")
.nth(1)
.and_then(|s| s.split(';').next())
.expect("accumulator expr");
let indent = &line[..line.len() - t.len()];
out.push_str(&format!(
"{indent}{{ let c = ncol + {j}u; if (c < mt.n) {{ y[wg.z * mt.m * mt.n + r * mt.n + c] = {acc}; }} }}\n"
));
} else {
out.push_str(line);
out.push('\n');
}
}
out
}
fn enc_gemm3_sk_reduce_src() -> String {
format!(
r#"
struct Meta {{ m: u32, n: u32, flags: u32, sk: u32 }}
@group(0) @binding(0) var<storage, read> part: array<f32>;
@group(0) @binding(1) var<storage, read> bias: array<f32>;
@group(0) @binding(2) var<storage, read_write> y: array<f32>;
@group(0) @binding(3) var<uniform> mt: Meta;
{ACT_FNS}
@compute @workgroup_size(256)
fn main(@builtin(global_invocation_id) gid: vec3<u32>) {{
let i = gid.x;
let total = mt.m * mt.n;
if (i >= total) {{ return; }}
var acc = 0.0;
for (var c = 0u; c < mt.sk; c++) {{ acc += part[c * total + i]; }}
if ((mt.flags & 1u) != 0u) {{ acc += bias[i % mt.n]; }}
y[i] = apply_act(acc, mt.flags >> 8u);
}}
"#
)
}
pub(crate) const GEMM2_TILES: [(usize, usize); 9] = [
(128, 128),
(128, 64),
(64, 128),
(64, 64),
(32, 64),
(64, 32),
(32, 32),
(16, 64),
(16, 32),
];
const GEMM2_MIN_WGS: usize = 96;
pub fn gemm2_tile(m: usize, n: usize) -> (usize, usize) {
if let Ok(v) = std::env::var("OSFKB_ENC_TILE")
&& let Some((a, b)) = v.split_once('x')
&& let (Ok(a), Ok(b)) = (a.parse(), b.parse())
&& GEMM2_TILES.contains(&(a, b))
{
return (a, b);
}
let wgs = |(bm, bn): (usize, usize)| n.div_ceil(bn) * m.div_ceil(bm);
GEMM2_TILES
.into_iter()
.find(|&t| wgs(t) >= GEMM2_MIN_WGS)
.unwrap_or_else(|| {
GEMM2_TILES
.into_iter()
.rev()
.max_by_key(|&t| wgs(t))
.expect("non-empty tile table")
})
}
pub(crate) fn gemm2_tier(tile: (usize, usize)) -> usize {
GEMM2_TILES
.iter()
.position(|&t| t == tile)
.expect("tile came from GEMM2_TILES")
}
pub(crate) fn gemm2_enabled() -> bool {
!matches!(
std::env::var("OSFKB_ENC_GEMM2").ok().as_deref(),
Some("0") | Some("off")
)
}
pub(crate) const GEMM3_TILES: [(usize, usize, usize); 5] = [
(128, 128, 16),
(64, 128, 16),
(64, 64, 16),
(32, 64, 8),
(16, 64, 8),
];
fn gemm3_min_wgs() -> usize {
static V: std::sync::OnceLock<usize> = std::sync::OnceLock::new();
*V.get_or_init(|| {
std::env::var("OSFKB_ENC_G3_MINWGS")
.ok()
.and_then(|v| v.parse().ok())
.unwrap_or(64)
})
}
pub(crate) fn gemm3_grid_ok(m: usize, tile: (usize, usize, usize), wgs: usize) -> bool {
wgs >= gemm3_min_wgs() || (wgs >= 48 && tile.0 >= 64 && m / tile.0 >= 3)
}
pub(crate) fn gemm3_tile(m: usize, n: usize) -> (usize, usize, usize) {
let wgs = |(bm, bn, _): (usize, usize, usize)| n.div_ceil(bn) * m.div_ceil(bm);
let mut pick = GEMM3_TILES
.into_iter()
.find(|&t| gemm3_grid_ok(m, t, wgs(t)))
.unwrap_or_else(|| {
GEMM3_TILES
.into_iter()
.rev()
.max_by_key(|&t| wgs(t))
.expect("non-empty tile table")
});
if std::env::var("OSFKB_ENC_G3_DEMOTE").ok().as_deref() == Some("0") {
return pick;
}
let mut idx = gemm3_tier(pick);
while pick.0 >= 64 && wgs(pick) < 80 && idx + 1 < GEMM3_TILES.len() {
let next = GEMM3_TILES[idx + 1];
if wgs(next) >= 128 {
pick = next;
idx += 1;
} else {
break;
}
}
pick
}
pub(crate) fn gemm3_tier(tile: (usize, usize, usize)) -> usize {
GEMM3_TILES
.iter()
.position(|&t| t == tile)
.expect("tile came from GEMM3_TILES")
}
pub(crate) fn gemm3_eligible(n: usize) -> bool {
(n >= 1152 || n <= 512)
&& gemm2_enabled() && !matches!(
std::env::var("OSFKB_ENC_GEMM3").ok().as_deref(),
Some("0") | Some("off")
)
}
pub(crate) fn gemm3_sk_band(n: usize) -> bool {
n > 512
&& n < 1152
&& n.is_multiple_of(64)
&& gemm2_enabled()
&& !matches!(
std::env::var("OSFKB_ENC_GEMM3").ok().as_deref(),
Some("0") | Some("off")
)
&& std::env::var("OSFKB_ENC_GEMM3_SK").ok().as_deref() != Some("0")
}
pub(crate) fn gemm3_smalln_sk(n: usize) -> bool {
n <= 512
&& n.is_multiple_of(64)
&& std::env::var("OSFKB_ENC_SK_SMALLN").ok().as_deref() != Some("0")
&& std::env::var("OSFKB_ENC_GEMM3_SK").ok().as_deref() != Some("0")
}
pub(crate) fn f16a_enabled() -> bool {
std::env::var("OSFKB_ENC_F16A").ok().as_deref() == Some("1")
}
pub(crate) fn gemm3_sk_chunks(k: usize) -> u32 {
if k >= 2048 { 4 } else { 2 }
}
pub(crate) fn gemm3_sk_plan(m: usize, n: usize, k: usize) -> (bool, u32) {
if std::env::var("OSFKB_ENC_SK_DYNZ").ok().as_deref() == Some("0") {
return (false, gemm3_sk_chunks(k));
}
let cols = n.div_ceil(GEMM3_SK_TILE.1);
let zpick = |base: usize| {
let mut z = 8u32;
while z > 2 && base * z as usize > 160 {
z >>= 1;
}
z.max(gemm3_sk_chunks(k))
};
let base32 = cols * m.div_ceil(GEMM3_SK_TILE32.0);
let z32 = zpick(base32);
if base32 * z32 as usize >= 96 && std::env::var("OSFKB_ENC_SK32").ok().as_deref() != Some("0") {
return (true, z32);
}
(false, zpick(cols * m.div_ceil(GEMM3_SK_TILE.0)))
}
pub(crate) const SK_PART_F32: usize = 8 * 81 * 1024;
pub(crate) const GEMM3_SK_TILE: (usize, usize, usize) = (16, 64, 8);
pub(crate) const GEMM3_SK_TILE32: (usize, usize, usize) = (32, 64, 8);
const ENC_LAYERNORM: &str = r#"
struct Meta { h: u32, flags: u32, eps: f32, pad: u32 }
@group(0) @binding(0) var<storage, read> x: array<f32>;
@group(0) @binding(1) var<storage, read> res: array<f32>;
@group(0) @binding(2) var<storage, read> w: array<f32>;
@group(0) @binding(3) var<storage, read> b: array<f32>;
@group(0) @binding(4) var<storage, read_write> out: array<f32>;
@group(0) @binding(5) var<uniform> mt: Meta;
var<workgroup> sh: array<f32, 256>;
fn val(base: u32, i: u32) -> f32 {
var v = x[base + i];
if ((mt.flags & 1u) != 0u) { v += res[base + i]; }
return v;
}
@compute @workgroup_size(256)
fn main(@builtin(workgroup_id) wg: vec3<u32>, @builtin(local_invocation_index) t: u32) {
let base = wg.x * mt.h;
var sum = 0.0;
for (var i = t; i < mt.h; i += 256u) { sum += val(base, i); }
sh[t] = sum;
workgroupBarrier();
for (var s = 128u; s > 0u; s >>= 1u) {
if (t < s) { sh[t] += sh[t + s]; }
workgroupBarrier();
}
// flags bit 2: RMS mode — no mean-centering (variance below becomes mean-of-squares).
var mean = sh[0] / f32(mt.h);
if ((mt.flags & 4u) != 0u) { mean = 0.0; }
workgroupBarrier();
var sq = 0.0;
for (var i = t; i < mt.h; i += 256u) {
let d = val(base, i) - mean;
sq += d * d;
}
sh[t] = sq;
workgroupBarrier();
for (var s = 128u; s > 0u; s >>= 1u) {
if (t < s) { sh[t] += sh[t + s]; }
workgroupBarrier();
}
let inv = 1.0 / sqrt(sh[0] / f32(mt.h) + mt.eps);
for (var i = t; i < mt.h; i += 256u) {
var o = (val(base, i) - mean) * inv * w[i];
if ((mt.flags & 2u) != 0u) { o += b[i]; }
out[base + i] = o;
}
}
"#;
const ENC_ATTN: &str = r#"
struct Meta { nrows: u32, n_heads: u32, hd: u32, mode: u32, window: u32, n_kv_heads: u32, packed: u32, p2: u32 }
@group(0) @binding(0) var<storage, read> q: array<f32>;
@group(0) @binding(1) var<storage, read> k: array<f32>;
@group(0) @binding(2) var<storage, read> v: array<f32>;
@group(0) @binding(3) var<storage, read_write> out: array<f32>;
@group(0) @binding(4) var<storage, read> seq_starts: array<u32>;
@group(0) @binding(5) var<storage, read> seq_of: array<u32>;
@group(0) @binding(6) var<storage, read> valid: array<u32>;
@group(0) @binding(7) var<uniform> mt: Meta;
var<workgroup> qsh: array<f32, 128>;
var<workgroup> psh: array<f32, 64>;
var<workgroup> red: array<f32, 64>;
@compute @workgroup_size(64)
fn main(@builtin(workgroup_id) wg: vec3<u32>, @builtin(local_invocation_index) t: u32) {
let i = wg.x;
let head = wg.y;
let h = mt.n_heads * mt.hd;
// GQA: q spans n_heads, k/v span n_kv_heads; each kv head serves n_heads/n_kv_heads q heads.
let kvh = mt.n_kv_heads * mt.hd;
let kv_head = head / (mt.n_heads / mt.n_kv_heads);
// mt.packed == 1: q/k/v are three VIEWS of one [T, q|k|v] buffer (the fused-QKV GEMM's
// output, bound to all three slots) — row stride h+2·kvh, k at offset h, v at h+kvh.
// mt.packed == 0: the historical three separate [T, ·] buffers.
var qstride = h;
var kstride = kvh;
var koff = 0u;
var voff = 0u;
if (mt.packed == 1u) {
qstride = h + 2u * kvh;
kstride = h + 2u * kvh;
koff = h;
voff = h + kvh;
}
let qbase = i * qstride + head * mt.hd;
let obase = i * h + head * mt.hd;
for (var d = t; d < mt.hd; d += 64u) { qsh[d] = q[qbase + d]; }
workgroupBarrier();
let s = seq_of[i];
let a = seq_starts[s];
let b = seq_starts[s + 1u];
var jstart = a;
var jend = b;
if (mt.mode == 1u) { jend = min(jend, i + 1u); }
if (mt.window > 0u) {
let half_w = mt.window / 2u;
if (i - a > half_w) { jstart = max(jstart, i - half_w); }
jend = min(jend, i + half_w + 1u);
}
let scale = 1.0 / sqrt(f32(mt.hd));
var m_run = -3.0e38;
var l_run = 0.0;
// Per-thread output dims: d = t, t + 64 (hd ≤ 128).
var acc0 = 0.0;
var acc1 = 0.0;
let ntiles = (jend - jstart + 63u) / 64u;
for (var tile = 0u; tile < ntiles; tile++) {
let j = jstart + tile * 64u + t;
var score = -3.0e38;
// valid[j] == 0 (ColBERT expansion pads): never an attention key.
if (j < jend && valid[j] != 0u) {
var dot = 0.0;
let kbase = j * kstride + koff + kv_head * mt.hd;
for (var d = 0u; d < mt.hd; d++) { dot += qsh[d] * k[kbase + d]; }
score = dot * scale;
}
psh[t] = score;
red[t] = score;
workgroupBarrier();
for (var r = 32u; r > 0u; r >>= 1u) {
if (t < r) { red[t] = max(red[t], red[t + r]); }
workgroupBarrier();
}
let tile_max = red[0];
workgroupBarrier();
let new_m = max(m_run, tile_max);
let rescale = exp(m_run - new_m);
// exp of masked lanes is exp(-inf) = 0 — they contribute nothing.
var p = 0.0;
if (psh[t] > -3.0e37) { p = exp(psh[t] - new_m); }
psh[t] = p;
red[t] = p;
workgroupBarrier();
for (var r = 32u; r > 0u; r >>= 1u) {
if (t < r) { red[t] += red[t + r]; }
workgroupBarrier();
}
l_run = l_run * rescale + red[0];
workgroupBarrier();
acc0 *= rescale;
acc1 *= rescale;
let tile_len = min(64u, jend - (jstart + tile * 64u));
let d0 = t;
let d1 = t + 64u;
for (var jj = 0u; jj < tile_len; jj++) {
let vbase = (jstart + tile * 64u + jj) * kstride + voff + kv_head * mt.hd;
let p_j = psh[jj];
if (d0 < mt.hd) { acc0 += p_j * v[vbase + d0]; }
if (d1 < mt.hd) { acc1 += p_j * v[vbase + d1]; }
}
m_run = new_m;
workgroupBarrier();
}
var inv = 0.0;
if (l_run > 0.0) { inv = 1.0 / l_run; }
if (t < mt.hd) { out[obase + t] = acc0 * inv; }
if (t + 64u < mt.hd) { out[obase + t + 64u] = acc1 * inv; }
}
"#;
const ENC_ATTN_V4: &str = r#"
struct Meta { nrows: u32, n_heads: u32, hd: u32, mode: u32, window: u32, n_kv_heads: u32, packed: u32, p2: u32 }
@group(0) @binding(0) var<storage, read> q: array<vec4<f32>>;
@group(0) @binding(1) var<storage, read> k: array<vec4<f32>>;
@group(0) @binding(2) var<storage, read> v: array<vec4<f32>>;
@group(0) @binding(3) var<storage, read_write> out: array<vec4<f32>>;
@group(0) @binding(4) var<storage, read> seq_starts: array<u32>;
@group(0) @binding(5) var<storage, read> seq_of: array<u32>;
@group(0) @binding(6) var<storage, read> valid: array<u32>;
@group(0) @binding(7) var<uniform> mt: Meta;
var<workgroup> qsh: array<vec4<f32>, 32>;
var<workgroup> psh: array<f32, 64>;
var<workgroup> red: array<f32, 64>;
@compute @workgroup_size(64)
fn main(@builtin(workgroup_id) wg: vec3<u32>, @builtin(local_invocation_index) t: u32) {
let i = wg.x;
let head = wg.y;
let h = mt.n_heads * mt.hd;
let kvh = mt.n_kv_heads * mt.hd;
let kv_head = head / (mt.n_heads / mt.n_kv_heads);
let hd4 = mt.hd / 4u;
var qstride = h;
var kstride = kvh;
var koff = 0u;
var voff = 0u;
if (mt.packed == 1u) {
qstride = h + 2u * kvh;
kstride = h + 2u * kvh;
koff = h;
voff = h + kvh;
}
let qbase4 = (i * qstride + head * mt.hd) / 4u;
let obase4 = (i * h + head * mt.hd) / 4u;
for (var d = t; d < hd4; d += 64u) { qsh[d] = q[qbase4 + d]; }
workgroupBarrier();
let s = seq_of[i];
let a = seq_starts[s];
let b = seq_starts[s + 1u];
var jstart = a;
var jend = b;
if (mt.mode == 1u) { jend = min(jend, i + 1u); }
if (mt.window > 0u) {
let half_w = mt.window / 2u;
if (i - a > half_w) { jstart = max(jstart, i - half_w); }
jend = min(jend, i + half_w + 1u);
}
let scale = 1.0 / sqrt(f32(mt.hd));
var m_run = -3.0e38;
var l_run = 0.0;
// Per-thread output: dims [4t, 4t+4) — one vec4 (hd ≤ 128 ⇒ hd4 ≤ 32 active threads).
var acc = vec4<f32>(0.0);
let ntiles = (jend - jstart + 63u) / 64u;
for (var tile = 0u; tile < ntiles; tile++) {
let j = jstart + tile * 64u + t;
var score = -3.0e38;
if (j < jend && valid[j] != 0u) {
var dot4 = 0.0;
let kbase4 = (j * kstride + koff + kv_head * mt.hd) / 4u;
for (var d = 0u; d < hd4; d++) { dot4 += dot(qsh[d], k[kbase4 + d]); }
score = dot4 * scale;
}
psh[t] = score;
red[t] = score;
workgroupBarrier();
for (var r = 32u; r > 0u; r >>= 1u) {
if (t < r) { red[t] = max(red[t], red[t + r]); }
workgroupBarrier();
}
let tile_max = red[0];
workgroupBarrier();
let new_m = max(m_run, tile_max);
let rescale = exp(m_run - new_m);
var p = 0.0;
if (psh[t] > -3.0e37) { p = exp(psh[t] - new_m); }
psh[t] = p;
red[t] = p;
workgroupBarrier();
for (var r = 32u; r > 0u; r >>= 1u) {
if (t < r) { red[t] += red[t + r]; }
workgroupBarrier();
}
l_run = l_run * rescale + red[0];
workgroupBarrier();
acc *= rescale;
let tile_len = min(64u, jend - (jstart + tile * 64u));
for (var jj = 0u; jj < tile_len; jj++) {
let vbase4 = ((jstart + tile * 64u + jj) * kstride + voff + kv_head * mt.hd) / 4u;
let p_j = psh[jj];
if (t < hd4) { acc += p_j * v[vbase4 + t]; }
}
m_run = new_m;
workgroupBarrier();
}
var inv = 0.0;
if (l_run > 0.0) { inv = 1.0 / l_run; }
if (t < hd4) { out[obase4 + t] = acc * inv; }
}
"#;
const ENC_ATTN_RB4: &str = r#"
struct Meta { nrows: u32, n_heads: u32, hd: u32, mode: u32, window: u32, n_kv_heads: u32, packed: u32, p2: u32 }
@group(0) @binding(0) var<storage, read> q: array<vec4<f32>>;
@group(0) @binding(1) var<storage, read> k: array<vec4<f32>>;
@group(0) @binding(2) var<storage, read> v: array<vec4<f32>>;
@group(0) @binding(3) var<storage, read_write> out: array<vec4<f32>>;
@group(0) @binding(4) var<storage, read> seq_starts: array<u32>;
@group(0) @binding(5) var<storage, read> seq_of: array<u32>;
@group(0) @binding(6) var<storage, read> valid: array<u32>;
@group(0) @binding(7) var<uniform> mt: Meta;
const RB: u32 = 4u;
var<workgroup> qsh: array<vec4<f32>, 64>; // [RB][hd/4 <= 16]
var<workgroup> ssh: array<f32, 256>; // [RB][64] scores -> probabilities
var<workgroup> red: array<f32, 64>; // [RB][16] reduction ladder
var<workgroup> ja_r: array<u32, 4>;
var<workgroup> jb_r: array<u32, 4>;
var<workgroup> m_run: array<f32, 4>;
var<workgroup> l_run: array<f32, 4>;
var<workgroup> m_new: array<f32, 4>;
var<workgroup> resc: array<f32, 4>;
@compute @workgroup_size(64)
fn main(@builtin(workgroup_id) wg: vec3<u32>, @builtin(local_invocation_index) t: u32) {
let head = wg.y;
let r0 = wg.x * RB;
let h = mt.n_heads * mt.hd;
let kvh = mt.n_kv_heads * mt.hd;
let kv_head = head / (mt.n_heads / mt.n_kv_heads);
let hd4 = mt.hd / 4u;
var qstride = h;
var kstride = kvh;
var koff = 0u;
var voff = 0u;
if (mt.packed == 1u) {
qstride = h + 2u * kvh;
kstride = h + 2u * kvh;
koff = h;
voff = h + kvh;
}
// Per-row bounds (ragged: rows may sit in different sequences; window/causal are per row).
if (t < RB) {
let i = r0 + t;
var ja = 0u;
var jb = 0u;
if (i < mt.nrows) {
let s = seq_of[i];
ja = seq_starts[s];
jb = seq_starts[s + 1u];
if (mt.mode == 1u) { jb = min(jb, i + 1u); }
if (mt.window > 0u) {
let half_w = mt.window / 2u;
if (i - ja > half_w) { ja = max(ja, i - half_w); }
jb = min(jb, i + half_w + 1u);
}
}
ja_r[t] = ja;
jb_r[t] = jb;
m_run[t] = -3.0e38;
l_run[t] = 0.0;
}
// Stage the RB query rows (cooperative: 64 threads over RB·hd4 <= 64 vec4s).
if (t < RB * hd4) {
let r = t / hd4;
let d = t % hd4;
let i = r0 + r;
var qv = vec4<f32>(0.0);
if (i < mt.nrows) { qv = q[(i * qstride + head * mt.hd) / 4u + d]; }
qsh[r * hd4 + d] = qv;
}
workgroupBarrier();
let jlo = min(min(ja_r[0], ja_r[1]), min(ja_r[2], ja_r[3]));
let jhi = max(max(jb_r[0], jb_r[1]), max(jb_r[2], jb_r[3]));
let scale = 1.0 / sqrt(f32(mt.hd));
// Thread (rr, dd) owns output dims [4·dd, 4·dd+4) of row r0+rr.
let rr = t / 16u;
let dd = t % 16u;
var acc = vec4<f32>(0.0);
let span = jhi - jlo;
let ntiles = (span + 63u) / 64u;
for (var tile = 0u; tile < ntiles; tile++) {
let j = jlo + tile * 64u + t;
// Score key j against ALL RB rows: k[j] is read ONCE and dotted four times.
var d0 = -3.0e38;
var d1 = -3.0e38;
var d2 = -3.0e38;
var d3 = -3.0e38;
if (j < jhi && valid[j] != 0u) {
var s0 = 0.0;
var s1 = 0.0;
var s2 = 0.0;
var s3 = 0.0;
let kbase4 = (j * kstride + koff + kv_head * mt.hd) / 4u;
for (var d = 0u; d < hd4; d++) {
let kv = k[kbase4 + d];
s0 += dot(kv, qsh[d]);
s1 += dot(kv, qsh[hd4 + d]);
s2 += dot(kv, qsh[2u * hd4 + d]);
s3 += dot(kv, qsh[3u * hd4 + d]);
}
if (j >= ja_r[0] && j < jb_r[0]) { d0 = s0 * scale; }
if (j >= ja_r[1] && j < jb_r[1]) { d1 = s1 * scale; }
if (j >= ja_r[2] && j < jb_r[2]) { d2 = s2 * scale; }
if (j >= ja_r[3] && j < jb_r[3]) { d3 = s3 * scale; }
}
ssh[t] = d0;
ssh[64u + t] = d1;
ssh[128u + t] = d2;
ssh[192u + t] = d3;
workgroupBarrier();
// Row max, all four rows at once: lane L of row R reduces keys L, L+16, L+32, L+48.
let lane = dd;
var mx = ssh[rr * 64u + lane];
mx = max(mx, ssh[rr * 64u + lane + 16u]);
mx = max(mx, ssh[rr * 64u + lane + 32u]);
mx = max(mx, ssh[rr * 64u + lane + 48u]);
red[t] = mx;
workgroupBarrier();
for (var r = 8u; r > 0u; r >>= 1u) {
if (lane < r) { red[t] = max(red[t], red[t + r]); }
workgroupBarrier();
}
if (t < RB) {
let nm = max(m_run[t], red[t * 16u]);
m_new[t] = nm;
resc[t] = exp(m_run[t] - nm);
}
workgroupBarrier();
// exp + row sums, same ladder.
var ps = 0.0;
for (var rr2 = 0u; rr2 < RB; rr2++) {
let sc = ssh[rr2 * 64u + t];
var p = 0.0;
if (sc > -3.0e37) { p = exp(sc - m_new[rr2]); }
ssh[rr2 * 64u + t] = p;
}
workgroupBarrier();
ps = ssh[rr * 64u + lane] + ssh[rr * 64u + lane + 16u]
+ ssh[rr * 64u + lane + 32u] + ssh[rr * 64u + lane + 48u];
red[t] = ps;
workgroupBarrier();
for (var r = 8u; r > 0u; r >>= 1u) {
if (lane < r) { red[t] += red[t + r]; }
workgroupBarrier();
}
if (t < RB) { l_run[t] = l_run[t] * resc[t] + red[t * 16u]; }
// P·V: thread (rr, dd) accumulates its vec4 of row rr.
acc *= resc[rr];
let tile_len = min(64u, jhi - (jlo + tile * 64u));
if (dd < hd4) {
for (var jj = 0u; jj < tile_len; jj++) {
let p_j = ssh[rr * 64u + jj];
let vbase4 = ((jlo + tile * 64u + jj) * kstride + voff + kv_head * mt.hd) / 4u;
acc += p_j * v[vbase4 + dd];
}
}
if (t < RB) { m_run[t] = m_new[t]; }
workgroupBarrier();
}
let i = r0 + rr;
if (i < mt.nrows && dd < hd4) {
var inv = 0.0;
if (l_run[rr] > 0.0) { inv = 1.0 / l_run[rr]; }
out[(i * h + head * mt.hd) / 4u + dd] = acc * inv;
}
}
"#;
const ENC_ATTN_DISENT: &str = r#"
struct Meta { nrows: u32, n_heads: u32, hd: u32, span: u32, max_rel: u32, sf: u32, c2p: u32, p2c: u32 }
@group(0) @binding(0) var<storage, read> q: array<f32>;
@group(0) @binding(1) var<storage, read> k: array<f32>;
@group(0) @binding(2) var<storage, read> v: array<f32>;
@group(0) @binding(3) var<storage, read_write> out: array<f32>;
@group(0) @binding(4) var<storage, read> seq_starts: array<u32>;
@group(0) @binding(5) var<storage, read> seq_of: array<u32>;
@group(0) @binding(6) var<storage, read> valid: array<u32>;
@group(0) @binding(7) var<storage, read> pos_k: array<f32>;
@group(0) @binding(8) var<storage, read> pos_q: array<f32>;
@group(0) @binding(9) var<uniform> mt: Meta;
var<workgroup> qsh: array<f32, 128>;
var<workgroup> psh: array<f32, 64>;
var<workgroup> red: array<f32, 64>;
// HF make_log_bucket_position: identity within ±span/2, logarithmic (ceil) beyond. Odd in rel,
// which is why one index m serves both the c2p and p2c gathers.
fn bucket(rel: i32, span: u32, max_rel: u32) -> i32 {
let mid = i32(span / 2u);
let a = abs(rel);
if (a <= mid) { return rel; }
let mid_f = f32(mid);
let lp = ceil(log(f32(a) / mid_f) / log((f32(max_rel) - 1.0) / mid_f) * (mid_f - 1.0)) + mid_f;
let sgn = select(1.0, -1.0, rel < 0);
return i32(sgn * lp);
}
@compute @workgroup_size(64)
fn main(@builtin(workgroup_id) wg: vec3<u32>, @builtin(local_invocation_index) t: u32) {
let i = wg.x;
let head = wg.y;
let h = mt.n_heads * mt.hd;
let p2c_on = (mt.p2c & 1u) != 0u;
var qkstride = h;
var koff = 0u;
var voff = 0u;
if ((mt.p2c & 2u) != 0u) {
qkstride = 3u * h;
koff = h;
voff = 2u * h;
}
let qbase = i * qkstride + head * mt.hd;
let obase = i * h + head * mt.hd;
for (var d = t; d < mt.hd; d += 64u) { qsh[d] = q[qbase + d]; }
workgroupBarrier();
let s = seq_of[i];
let a = seq_starts[s];
let b = seq_starts[s + 1u];
let qi_local = i32(i - a);
let two_span = i32(2u * mt.span);
let scale = 1.0 / sqrt(f32(mt.hd) * f32(mt.sf));
var m_run = -3.0e38;
var l_run = 0.0;
var acc0 = 0.0;
var acc1 = 0.0;
let ntiles = (b - a + 63u) / 64u;
for (var tile = 0u; tile < ntiles; tile++) {
let j = a + tile * 64u + t;
var score = -3.0e38;
if (j < b && valid[j] != 0u) {
let kbase = j * qkstride + koff + head * mt.hd;
// Relative-table row for (i, j): shared by the c2p and p2c terms.
let m = u32(clamp(bucket(qi_local - i32(j - a), mt.span, mt.max_rel) + i32(mt.span),
0, two_span - 1));
let rbase = m * h + head * mt.hd;
var dot = 0.0;
for (var d = 0u; d < mt.hd; d++) {
let kd = k[kbase + d];
dot += qsh[d] * kd; // content → content
if (mt.c2p != 0u) { dot += qsh[d] * pos_k[rbase + d]; } // content q → position k
if (p2c_on) { dot += kd * pos_q[rbase + d]; } // position q -> content k
}
score = dot * scale;
}
psh[t] = score;
red[t] = score;
workgroupBarrier();
for (var r = 32u; r > 0u; r >>= 1u) {
if (t < r) { red[t] = max(red[t], red[t + r]); }
workgroupBarrier();
}
let tile_max = red[0];
workgroupBarrier();
let new_m = max(m_run, tile_max);
let rescale = exp(m_run - new_m);
var p = 0.0;
if (psh[t] > -3.0e37) { p = exp(psh[t] - new_m); }
psh[t] = p;
red[t] = p;
workgroupBarrier();
for (var r = 32u; r > 0u; r >>= 1u) {
if (t < r) { red[t] += red[t + r]; }
workgroupBarrier();
}
l_run = l_run * rescale + red[0];
workgroupBarrier();
acc0 *= rescale;
acc1 *= rescale;
let tile_len = min(64u, b - (a + tile * 64u));
let d0 = t;
let d1 = t + 64u;
for (var jj = 0u; jj < tile_len; jj++) {
let vbase = (a + tile * 64u + jj) * qkstride + voff + head * mt.hd;
let p_j = psh[jj];
if (d0 < mt.hd) { acc0 += p_j * v[vbase + d0]; }
if (d1 < mt.hd) { acc1 += p_j * v[vbase + d1]; }
}
m_run = new_m;
workgroupBarrier();
}
var inv = 0.0;
if (l_run > 0.0) { inv = 1.0 / l_run; }
if (t < mt.hd) { out[obase + t] = acc0 * inv; }
if (t + 64u < mt.hd) { out[obase + t + 64u] = acc1 * inv; }
}
"#;
const ENC_ATTN_DISENT_RB4: &str = r#"
struct Meta { nrows: u32, n_heads: u32, hd: u32, span: u32, max_rel: u32, sf: u32, c2p: u32, p2c: u32 }
@group(0) @binding(0) var<storage, read> q: array<vec4<f32>>;
@group(0) @binding(1) var<storage, read> k: array<vec4<f32>>;
@group(0) @binding(2) var<storage, read> v: array<vec4<f32>>;
@group(0) @binding(3) var<storage, read_write> out: array<vec4<f32>>;
@group(0) @binding(4) var<storage, read> seq_starts: array<u32>;
@group(0) @binding(5) var<storage, read> seq_of: array<u32>;
@group(0) @binding(6) var<storage, read> valid: array<u32>;
@group(0) @binding(7) var<storage, read> pos_k: array<vec4<f32>>;
@group(0) @binding(8) var<storage, read> pos_q: array<vec4<f32>>;
@group(0) @binding(9) var<uniform> mt: Meta;
const RB: u32 = 4u;
var<workgroup> qsh: array<vec4<f32>, 64>; // [RB][hd/4 <= 16]
var<workgroup> ssh: array<f32, 256>; // [RB][64]
var<workgroup> red: array<f32, 64>;
var<workgroup> ja_r: array<u32, 4>;
var<workgroup> jb_r: array<u32, 4>;
var<workgroup> qi_r: array<i32, 4>;
var<workgroup> m_run: array<f32, 4>;
var<workgroup> l_run: array<f32, 4>;
var<workgroup> m_new: array<f32, 4>;
var<workgroup> resc: array<f32, 4>;
fn bucket(rel: i32, span: u32, max_rel: u32) -> i32 {
let mid = i32(span / 2u);
let a = abs(rel);
if (a <= mid) { return rel; }
let mid_f = f32(mid);
let lp = ceil(log(f32(a) / mid_f) / log((f32(max_rel) - 1.0) / mid_f) * (mid_f - 1.0)) + mid_f;
let sgn = select(1.0, -1.0, rel < 0);
return i32(sgn * lp);
}
@compute @workgroup_size(64)
fn main(@builtin(workgroup_id) wg: vec3<u32>, @builtin(local_invocation_index) t: u32) {
let head = wg.y;
let r0 = wg.x * RB;
let h = mt.n_heads * mt.hd;
let hd4 = mt.hd / 4u;
// mt.p2c: bit 0 = the p2c score term; bit 1 = PACKED q|k|v (the fused-QKV GEMM's output
// bound to all three slots — same trick as ENC_ATTN's packed mode; no GQA here).
let p2c_on = (mt.p2c & 1u) != 0u;
var qkstride = h;
var koff = 0u;
var voff = 0u;
if ((mt.p2c & 2u) != 0u) {
qkstride = 3u * h;
koff = h;
voff = 2u * h;
}
if (t < RB) {
let i = r0 + t;
var ja = 0u;
var jb = 0u;
var qi = 0;
if (i < mt.nrows) {
let s = seq_of[i];
ja = seq_starts[s];
jb = seq_starts[s + 1u];
qi = i32(i - ja);
}
ja_r[t] = ja;
jb_r[t] = jb;
qi_r[t] = qi;
m_run[t] = -3.0e38;
l_run[t] = 0.0;
}
if (t < RB * hd4) {
let r = t / hd4;
let d = t % hd4;
let i = r0 + r;
var qv = vec4<f32>(0.0);
if (i < mt.nrows) { qv = q[(i * qkstride + head * mt.hd) / 4u + d]; }
qsh[r * hd4 + d] = qv;
}
workgroupBarrier();
let jlo = min(min(ja_r[0], ja_r[1]), min(ja_r[2], ja_r[3]));
let jhi = max(max(jb_r[0], jb_r[1]), max(jb_r[2], jb_r[3]));
let scale = 1.0 / sqrt(f32(mt.hd) * f32(mt.sf));
let two_span = i32(2u * mt.span);
let rr = t / 16u;
let dd = t % 16u;
var acc = vec4<f32>(0.0);
let span_keys = jhi - jlo;
let ntiles = (span_keys + 63u) / 64u;
for (var tile = 0u; tile < ntiles; tile++) {
let j = jlo + tile * 64u + t;
var sc0 = -3.0e38;
var sc1 = -3.0e38;
var sc2 = -3.0e38;
var sc3 = -3.0e38;
if (j < jhi && valid[j] != 0u) {
let kbase4 = (j * qkstride + koff + head * mt.hd) / 4u;
// Content dots: k[j] read ONCE, dotted with all four staged q rows.
var c0 = 0.0;
var c1 = 0.0;
var c2 = 0.0;
var c3 = 0.0;
for (var d = 0u; d < hd4; d++) {
let kv = k[kbase4 + d];
c0 += dot(kv, qsh[d]);
c1 += dot(kv, qsh[hd4 + d]);
c2 += dot(kv, qsh[2u * hd4 + d]);
c3 += dot(kv, qsh[3u * hd4 + d]);
}
// Positional terms per row (the LUT row m depends on i−j — not shareable).
for (var r = 0u; r < RB; r++) {
if (j < ja_r[r] || j >= jb_r[r]) { continue; }
let m = u32(clamp(
bucket(qi_r[r] - i32(j - ja_r[r]), mt.span, mt.max_rel) + i32(mt.span),
0, two_span - 1));
let rbase4 = (m * h + head * mt.hd) / 4u;
var pterm = 0.0;
for (var d = 0u; d < hd4; d++) {
if (mt.c2p != 0u) { pterm += dot(qsh[r * hd4 + d], pos_k[rbase4 + d]); }
if (p2c_on) { pterm += dot(k[kbase4 + d], pos_q[rbase4 + d]); }
}
var content = c0;
if (r == 1u) { content = c1; }
if (r == 2u) { content = c2; }
if (r == 3u) { content = c3; }
let sc = (content + pterm) * scale;
if (r == 0u) { sc0 = sc; }
if (r == 1u) { sc1 = sc; }
if (r == 2u) { sc2 = sc; }
if (r == 3u) { sc3 = sc; }
}
}
ssh[t] = sc0;
ssh[64u + t] = sc1;
ssh[128u + t] = sc2;
ssh[192u + t] = sc3;
workgroupBarrier();
let lane = dd;
var mx = ssh[rr * 64u + lane];
mx = max(mx, ssh[rr * 64u + lane + 16u]);
mx = max(mx, ssh[rr * 64u + lane + 32u]);
mx = max(mx, ssh[rr * 64u + lane + 48u]);
red[t] = mx;
workgroupBarrier();
for (var r = 8u; r > 0u; r >>= 1u) {
if (lane < r) { red[t] = max(red[t], red[t + r]); }
workgroupBarrier();
}
if (t < RB) {
let nm = max(m_run[t], red[t * 16u]);
m_new[t] = nm;
resc[t] = exp(m_run[t] - nm);
}
workgroupBarrier();
for (var rr2 = 0u; rr2 < RB; rr2++) {
let sc = ssh[rr2 * 64u + t];
var p = 0.0;
if (sc > -3.0e37) { p = exp(sc - m_new[rr2]); }
ssh[rr2 * 64u + t] = p;
}
workgroupBarrier();
let ps = ssh[rr * 64u + lane] + ssh[rr * 64u + lane + 16u]
+ ssh[rr * 64u + lane + 32u] + ssh[rr * 64u + lane + 48u];
red[t] = ps;
workgroupBarrier();
for (var r = 8u; r > 0u; r >>= 1u) {
if (lane < r) { red[t] += red[t + r]; }
workgroupBarrier();
}
if (t < RB) { l_run[t] = l_run[t] * resc[t] + red[t * 16u]; }
acc *= resc[rr];
let tile_len = min(64u, jhi - (jlo + tile * 64u));
if (dd < hd4) {
for (var jj = 0u; jj < tile_len; jj++) {
let p_j = ssh[rr * 64u + jj];
let vbase4 = ((jlo + tile * 64u + jj) * qkstride + voff + head * mt.hd) / 4u;
acc += p_j * v[vbase4 + dd];
}
}
if (t < RB) { m_run[t] = m_new[t]; }
workgroupBarrier();
}
let i = r0 + rr;
if (i < mt.nrows && dd < hd4) {
var inv = 0.0;
if (l_run[rr] > 0.0) { inv = 1.0 / l_run[rr]; }
out[(i * h + head * mt.hd) / 4u + dd] = acc * inv;
}
}
"#;
const ENC_MEAN_POOL: &str = r#"
struct Meta { h: u32, pad0: u32, pad1: u32, pad2: u32 }
@group(0) @binding(0) var<storage, read> hidden: array<f32>;
@group(0) @binding(1) var<storage, read_write> out: array<f32>;
@group(0) @binding(2) var<storage, read> seq_starts: array<u32>;
@group(0) @binding(3) var<uniform> mt: Meta;
@compute @workgroup_size(256)
fn main(@builtin(workgroup_id) wg: vec3<u32>, @builtin(local_invocation_index) t: u32) {
let s = wg.x;
let a = seq_starts[s];
let b = seq_starts[s + 1u];
let inv = 1.0 / f32(max(b - a, 1u));
for (var i = t; i < mt.h; i += 256u) {
var sum = 0.0;
for (var r = a; r < b; r++) { sum += hidden[r * mt.h + i]; }
out[s * mt.h + i] = sum * inv;
}
}
"#;
const ENC_L2NORM: &str = r#"
struct Meta { dim: u32, pad0: u32, pad1: u32, pad2: u32 }
@group(0) @binding(0) var<storage, read_write> buf: array<f32>;
@group(0) @binding(1) var<uniform> mt: Meta;
var<workgroup> sh: array<f32, 256>;
@compute @workgroup_size(256)
fn main(@builtin(workgroup_id) wg: vec3<u32>, @builtin(local_invocation_index) t: u32) {
let base = wg.x * mt.dim;
var sq = 0.0;
for (var i = t; i < mt.dim; i += 256u) {
let v = buf[base + i];
sq += v * v;
}
sh[t] = sq;
workgroupBarrier();
for (var s = 128u; s > 0u; s >>= 1u) {
if (t < s) { sh[t] += sh[t + s]; }
workgroupBarrier();
}
let n = sqrt(sh[0]);
if (n > 0.0) {
let inv = 1.0 / n;
for (var i = t; i < mt.dim; i += 256u) { buf[base + i] *= inv; }
}
}
"#;
const ROPE_ENC: &str = r#"
struct Meta { nrows: u32, nh: u32, hd: u32, theta: f32 }
@group(0) @binding(0) var<storage, read_write> x: array<f32>;
@group(0) @binding(1) var<storage, read> seq_starts: array<u32>;
@group(0) @binding(2) var<storage, read> seq_of: array<u32>;
@group(0) @binding(3) var<uniform> mt: Meta;
@compute @workgroup_size(64)
fn main(@builtin(workgroup_id) wg: vec3<u32>, @builtin(local_invocation_index) t: u32) {
let i = wg.x;
let p = f32(i - seq_starts[seq_of[i]]);
let half = mt.hd / 2u;
let pairs = mt.nh * half;
let h = mt.nh * mt.hd;
for (var pj = t; pj < pairs; pj += 64u) {
let head = pj / half;
let j = pj % half;
let freq = pow(mt.theta, -2.0 * f32(j) / f32(mt.hd));
let ang = p * freq;
let c = cos(ang);
let s = sin(ang);
let base = i * h + head * mt.hd;
let a = x[base + j];
let b = x[base + j + half];
x[base + j] = a * c - b * s;
x[base + j + half] = a * s + b * c;
}
}
"#;
const GLU_SPLIT: &str = const_format_glu();
const fn const_format_glu() -> &'static str {
r#"
struct Meta { i_width: u32, act: u32, p0: u32, p1: u32 }
@group(0) @binding(0) var<storage, read> mid: array<f32>;
@group(0) @binding(1) var<storage, read_write> out: array<f32>;
@group(0) @binding(2) var<uniform> mt: Meta;
fn erf_as(x: f32) -> f32 {
let s = sign(x);
let a = abs(x);
let t = 1.0 / (1.0 + 0.3275911 * a);
let y = 1.0 - (((((1.061405429 * t - 1.453152027) * t) + 1.421413741) * t - 0.284496736) * t
+ 0.254829592) * t * exp(-a * a);
return s * y;
}
fn apply_act(v: f32, code: u32) -> f32 {
switch code {
case 1u: { return 0.5 * v * (1.0 + erf_as(v * 0.70710678)); }
// tanh-GELU. The argument is CLAMPED, exactly as the decoder's kernels already do
// (forward.rs / lib.rs) — WGSL's `tanh` overflows to NaN on some backends for |arg| >~ 88,
// because a naive (e^2x - 1)/(e^2x + 1) expansion divides inf by inf. `arg` reaches 105 at a
// pre-activation of only 13.8, which the text encoders never hit and a ViT's outlier tokens
// hit on the first block. tanh is ±1 to well inside f32 epsilon by |arg| = 20, so this is
// EXACT, not an approximation.
case 2u: {
let a = clamp(0.7978845608 * (v + 0.044715 * v * v * v), -20.0, 20.0);
return 0.5 * v * (1.0 + tanh(a));
}
case 3u: { return v / (1.0 + exp(-v)); }
case 4u: { return tanh(v); }
case 5u: { return max(v, 0.0); }
default: { return v; }
}
}
@compute @workgroup_size(256)
fn main(@builtin(workgroup_id) wg: vec3<u32>, @builtin(local_invocation_index) t: u32) {
let row = wg.x;
for (var j = t; j < mt.i_width; j += 256u) {
let a = mid[row * 2u * mt.i_width + j];
let g = mid[row * 2u * mt.i_width + mt.i_width + j];
out[row * mt.i_width + j] = apply_act(a, mt.act) * g;
}
}
"#
}
const ENC_ADD: &str = r#"
struct Meta { n: u32, p0: u32, p1: u32, p2: u32 }
@group(0) @binding(0) var<storage, read_write> dst: array<f32>;
@group(0) @binding(1) var<storage, read> src: array<f32>;
@group(0) @binding(2) var<uniform> mt: Meta;
@compute @workgroup_size(256)
fn main(@builtin(global_invocation_id) gid: vec3<u32>) {
if (gid.x < mt.n) { dst[gid.x] += src[gid.x]; }
}
"#;
const ENC_QK_NORM: &str = r#"
struct Meta { nrows: u32, n_heads: u32, hd: u32, eps: f32 }
@group(0) @binding(0) var<storage, read_write> x: array<f32>;
@group(0) @binding(1) var<storage, read> w: array<f32>;
@group(0) @binding(2) var<uniform> mt: Meta;
@compute @workgroup_size(64)
fn main(@builtin(workgroup_id) wg: vec3<u32>, @builtin(local_invocation_index) t: u32) {
let i = wg.x;
let h = mt.n_heads * mt.hd;
for (var head = t; head < mt.n_heads; head += 64u) {
let base = i * h + head * mt.hd;
var ms = 0.0;
for (var d = 0u; d < mt.hd; d++) { let v = x[base + d]; ms += v * v; }
let inv = 1.0 / sqrt(ms / f32(mt.hd) + mt.eps);
for (var d = 0u; d < mt.hd; d++) { x[base + d] = x[base + d] * inv * w[d]; }
}
}
"#;
const ENC_LAST_POOL: &str = r#"
struct Meta { h: u32, p0: u32, p1: u32, p2: u32 }
@group(0) @binding(0) var<storage, read> hidden: array<f32>;
@group(0) @binding(1) var<storage, read_write> out: array<f32>;
@group(0) @binding(2) var<storage, read> seq_starts: array<u32>;
@group(0) @binding(3) var<uniform> mt: Meta;
@compute @workgroup_size(256)
fn main(@builtin(workgroup_id) wg: vec3<u32>, @builtin(local_invocation_index) t: u32) {
let s = wg.x;
let last = seq_starts[s + 1u] - 1u;
for (var i = t; i < mt.h; i += 256u) {
out[s * mt.h + i] = hidden[last * mt.h + i];
}
}
"#;
const ENC_CLS_POOL: &str = r#"
struct Meta { h: u32, p0: u32, p1: u32, p2: u32 }
@group(0) @binding(0) var<storage, read> hidden: array<f32>;
@group(0) @binding(1) var<storage, read_write> out: array<f32>;
@group(0) @binding(2) var<storage, read> seq_starts: array<u32>;
@group(0) @binding(3) var<uniform> mt: Meta;
@compute @workgroup_size(256)
fn main(@builtin(workgroup_id) wg: vec3<u32>, @builtin(local_invocation_index) t: u32) {
let s = wg.x;
let first = seq_starts[s];
for (var i = t; i < mt.h; i += 256u) {
out[s * mt.h + i] = hidden[first * mt.h + i];
}
}
"#;
const ENC_COPY: &str = r#"
struct Meta { n: u32, p0: u32, p1: u32, p2: u32 }
@group(0) @binding(0) var<storage, read_write> dst: array<f32>;
@group(0) @binding(1) var<storage, read> src: array<f32>;
@group(0) @binding(2) var<uniform> mt: Meta;
@compute @workgroup_size(256)
fn main(@builtin(global_invocation_id) gid: vec3<u32>) {
if (gid.x < mt.n) { dst[gid.x] = src[gid.x]; }
}
"#;
const ENC_CONV: &str = r#"
struct Meta { h: u32, l: u32, p0: u32, p1: u32 }
@group(0) @binding(0) var<storage, read> bcx: array<f32>;
@group(0) @binding(1) var<storage, read> conv_w: array<f32>;
@group(0) @binding(2) var<storage, read_write> y: array<f32>;
@group(0) @binding(3) var<storage, read> seq_starts: array<u32>;
@group(0) @binding(4) var<storage, read> seq_of: array<u32>;
@group(0) @binding(5) var<storage, read> valid: array<u32>;
@group(0) @binding(6) var<uniform> mt: Meta;
@compute @workgroup_size(256)
fn main(@builtin(workgroup_id) wg: vec3<u32>, @builtin(local_invocation_index) t: u32) {
let i = wg.x;
let a = seq_starts[seq_of[i]];
let w3 = 3u * mt.h;
// HF zeroes the conv-block input at padded positions: invalid row ⇒ zero output.
let own_valid = valid[i] != 0u;
for (var c = t; c < mt.h; c += 256u) {
if (!own_valid) {
y[i * mt.h + c] = 0.0;
continue;
}
let cc = bcx[i * w3 + mt.h + c];
var acc = 0.0;
for (var k = 0u; k < mt.l; k++) {
let back = mt.l - 1u - k;
if (i >= a + back) {
let j = i - back;
if (valid[j] != 0u) {
let bx = bcx[j * w3 + c] * bcx[j * w3 + 2u * mt.h + c];
acc += conv_w[c * mt.l + k] * bx;
}
}
}
y[i * mt.h + c] = cc * acc;
}
}
"#;
pub struct EncKernels {
gemm_f32: wgpu::ComputePipeline,
gemm_f16: Option<wgpu::ComputePipeline>,
gemm2_f32: Vec<wgpu::ComputePipeline>,
gemm2_f16: Option<Vec<wgpu::ComputePipeline>>,
gemm3_f32: Vec<wgpu::ComputePipeline>,
gemm3_f16: Option<Vec<wgpu::ComputePipeline>>,
gemm3_sk_f32: wgpu::ComputePipeline,
gemm3_sk_f16: Option<wgpu::ComputePipeline>,
gemm3_sk32_f32: wgpu::ComputePipeline,
gemm3_sk32_f16: Option<wgpu::ComputePipeline>,
gemm3_sk_reduce: wgpu::ComputePipeline,
gemm3_f16a: Option<Vec<wgpu::ComputePipeline>>,
gemm3_sk_f16a: Option<wgpu::ComputePipeline>,
gemm3_sk32_f16a: Option<wgpu::ComputePipeline>,
gemm4_f16w: Option<wgpu::ComputePipeline>,
layernorm: wgpu::ComputePipeline,
attn: wgpu::ComputePipeline,
attn4: wgpu::ComputePipeline,
attn_rb4: wgpu::ComputePipeline,
attn_disent: wgpu::ComputePipeline,
attn_disent_rb4: wgpu::ComputePipeline,
mean_pool: wgpu::ComputePipeline,
l2norm: wgpu::ComputePipeline,
rope: wgpu::ComputePipeline,
glu: wgpu::ComputePipeline,
add: wgpu::ComputePipeline,
qk_norm: wgpu::ComputePipeline,
last_pool: wgpu::ComputePipeline,
cls_pool: wgpu::ComputePipeline,
copy: wgpu::ComputePipeline,
conv: wgpu::ComputePipeline,
}
pub fn gelu_tanh_on_gpu(ctx: &GpuCtx, xs: &[f32]) -> anyhow::Result<Vec<f32>> {
let src = format!(
"{ACT_FNS}\n\
@group(0) @binding(0) var<storage, read_write> v: array<f32>;\n\
@compute @workgroup_size(64)\n\
fn main(@builtin(global_invocation_id) g: vec3<u32>) {{\n\
\x20 if (g.x < {}u) {{ v[g.x] = apply_act(v[g.x], 2u); }}\n\
}}\n",
xs.len()
);
let module = ctx
.device
.shader_module_tuned(wgpu::ShaderModuleDescriptor {
label: Some("gelu_probe"),
source: wgpu::ShaderSource::Wgsl(src.into()),
});
let pl = ctx
.device
.create_compute_pipeline(&wgpu::ComputePipelineDescriptor {
label: Some("gelu_probe"),
layout: None,
module: &module,
entry_point: Some("main"),
compilation_options: Default::default(),
cache: None,
});
let buf = ctx.storage(xs);
let bg = make_bg1_pub(ctx, &pl, &buf);
dispatch(ctx, &pl, &bg, (xs.len() as u32).div_ceil(64), 1);
ctx.read(&buf, xs.len())
}
fn make_bg1_pub(ctx: &GpuCtx, pl: &wgpu::ComputePipeline, buf: &wgpu::Buffer) -> wgpu::BindGroup {
ctx.device.create_bind_group(&wgpu::BindGroupDescriptor {
label: None,
layout: &pl.get_bind_group_layout(0),
entries: &[wgpu::BindGroupEntry {
binding: 0,
resource: buf.as_entire_binding(),
}],
})
}
impl EncKernels {
pub fn new(ctx: &GpuCtx) -> Self {
Self {
gemm_f32: pipeline(ctx, "enc_gemm_f32", &enc_gemm_src(false)),
gemm_f16: ctx
.f16
.then(|| pipeline(ctx, "enc_gemm_f16", &enc_gemm_src(true))),
gemm2_f32: GEMM2_TILES
.iter()
.map(|&(bm, bn)| {
pipeline(
ctx,
&format!("enc_gemm2_{bm}x{bn}_f32"),
&enc_gemm2_src(false, bm, bn),
)
})
.collect(),
gemm2_f16: ctx.f16.then(|| {
GEMM2_TILES
.iter()
.map(|&(bm, bn)| {
pipeline(
ctx,
&format!("enc_gemm2_{bm}x{bn}_f16"),
&enc_gemm2_src(true, bm, bn),
)
})
.collect()
}),
gemm3_f32: GEMM3_TILES
.iter()
.map(|&(bm, bn, bk)| {
pipeline(
ctx,
&format!("enc_gemm3_{bm}x{bn}x{bk}_f32"),
&enc_gemm3_src(false, bm, bn, bk),
)
})
.collect(),
gemm3_f16: ctx.f16.then(|| {
GEMM3_TILES
.iter()
.map(|&(bm, bn, bk)| {
pipeline(
ctx,
&format!("enc_gemm3_{bm}x{bn}x{bk}_f16"),
&enc_gemm3_src(true, bm, bn, bk),
)
})
.collect()
}),
gemm3_sk_f32: pipeline(
ctx,
"enc_gemm3_sk_f32",
&enc_gemm3_sk_src(false, GEMM3_SK_TILE.0, GEMM3_SK_TILE.1, GEMM3_SK_TILE.2),
),
gemm3_sk_f16: ctx.f16.then(|| {
pipeline(
ctx,
"enc_gemm3_sk_f16",
&enc_gemm3_sk_src(true, GEMM3_SK_TILE.0, GEMM3_SK_TILE.1, GEMM3_SK_TILE.2),
)
}),
gemm3_sk32_f32: pipeline(
ctx,
"enc_gemm3_sk32_f32",
&enc_gemm3_sk_src(
false,
GEMM3_SK_TILE32.0,
GEMM3_SK_TILE32.1,
GEMM3_SK_TILE32.2,
),
),
gemm3_sk32_f16: ctx.f16.then(|| {
pipeline(
ctx,
"enc_gemm3_sk32_f16",
&enc_gemm3_sk_src(
true,
GEMM3_SK_TILE32.0,
GEMM3_SK_TILE32.1,
GEMM3_SK_TILE32.2,
),
)
}),
gemm3_sk_reduce: pipeline(ctx, "enc_gemm3_sk_reduce", &enc_gemm3_sk_reduce_src()),
gemm3_f16a: ctx.f16.then(|| {
GEMM3_TILES
.iter()
.map(|&(bm, bn, bk)| {
pipeline(
ctx,
&format!("enc_gemm3_f16a_{bm}x{bn}x{bk}"),
&enc_gemm3_f16a_src(bm, bn, bk),
)
})
.collect()
}),
gemm3_sk_f16a: ctx.f16.then(|| {
pipeline(
ctx,
"enc_gemm3_sk_f16a",
&enc_gemm3_sk_f16a_src(GEMM3_SK_TILE.0, GEMM3_SK_TILE.1, GEMM3_SK_TILE.2),
)
}),
gemm3_sk32_f16a: ctx.f16.then(|| {
pipeline(
ctx,
"enc_gemm3_sk32_f16a",
&enc_gemm3_sk_f16a_src(GEMM3_SK_TILE32.0, GEMM3_SK_TILE32.1, GEMM3_SK_TILE32.2),
)
}),
gemm4_f16w: (ctx.f16 && ctx.coop_matrix)
.then(|| pipeline(ctx, "enc_gemm4_f16w", &enc_gemm4_f16w_src(128, 128, 8))),
layernorm: pipeline(ctx, "enc_layernorm", ENC_LAYERNORM),
attn: pipeline(ctx, "enc_attn", ENC_ATTN),
attn4: pipeline(ctx, "enc_attn_v4", ENC_ATTN_V4),
attn_rb4: pipeline(ctx, "enc_attn_rb4", ENC_ATTN_RB4),
attn_disent: pipeline(ctx, "enc_attn_disent", ENC_ATTN_DISENT),
attn_disent_rb4: pipeline(ctx, "enc_attn_disent_rb4", ENC_ATTN_DISENT_RB4),
mean_pool: pipeline(ctx, "enc_mean_pool", ENC_MEAN_POOL),
l2norm: pipeline(ctx, "enc_l2norm", ENC_L2NORM),
rope: pipeline(ctx, "enc_rope", ROPE_ENC),
glu: pipeline(ctx, "enc_glu", GLU_SPLIT),
add: pipeline(ctx, "enc_add", ENC_ADD),
qk_norm: pipeline(ctx, "enc_qk_norm", ENC_QK_NORM),
last_pool: pipeline(ctx, "enc_last_pool", ENC_LAST_POOL),
cls_pool: pipeline(ctx, "enc_cls_pool", ENC_CLS_POOL),
copy: pipeline(ctx, "enc_copy", ENC_COPY),
conv: pipeline(ctx, "enc_conv", ENC_CONV),
}
}
pub fn rope_for_tests(
&self,
ctx: &GpuCtx,
x: &[f32],
seq_starts: &[u32],
n_heads: usize,
hd: usize,
theta: f32,
) -> Result<Vec<f32>> {
let h = n_heads * hd;
let nrows = x.len() / h;
let xb = ctx.storage(x);
let sb = ctx.storage_bytes(bytemuck::cast_slice(seq_starts));
let mb = ctx.storage_bytes(bytemuck::cast_slice(&row_to_seq(seq_starts, nrows)));
let meta = uni(
ctx,
bytemuck::cast_slice(&[nrows as u32, n_heads as u32, hd as u32, theta.to_bits()]),
);
let bg = make_bg(ctx, &self.rope, &[&xb, &sb, &mb], &meta);
dispatch(ctx, &self.rope, &bg, nrows as u32, 1);
ctx.read(&xb, x.len())
}
pub fn glu_for_tests(
&self,
ctx: &GpuCtx,
mid: &[f32],
i_width: usize,
act: Act,
) -> Result<Vec<f32>> {
let rows = mid.len() / (2 * i_width);
let mb = ctx.storage(mid);
let ob = ctx.storage(&vec![0f32; rows * i_width]);
let meta = uni(
ctx,
bytemuck::cast_slice(&[i_width as u32, act_code(Some(act)), 0, 0]),
);
let bg = make_bg(ctx, &self.glu, &[&mb, &ob], &meta);
dispatch(ctx, &self.glu, &bg, rows as u32, 1);
ctx.read(&ob, rows * i_width)
}
pub(crate) fn ln_pl(&self) -> &wgpu::ComputePipeline {
&self.layernorm
}
pub(crate) fn attn_pl(&self) -> &wgpu::ComputePipeline {
&self.attn
}
pub(crate) fn add_pl(&self) -> &wgpu::ComputePipeline {
&self.add
}
pub(crate) fn gemm2_pipeline_tile(
&self,
f16: bool,
tile: (usize, usize),
) -> Result<&wgpu::ComputePipeline> {
let set = if f16 {
self.gemm2_f16
.as_ref()
.ok_or_else(|| anyhow::anyhow!("adapter has no SHADER_F16"))?
} else {
&self.gemm2_f32
};
Ok(&set[gemm2_tier(tile)])
}
pub(crate) fn gemm3_f16a_pipeline_tile(
&self,
tile: (usize, usize, usize),
) -> Option<&wgpu::ComputePipeline> {
self.gemm3_f16a.as_ref().map(|set| &set[gemm3_tier(tile)])
}
pub(crate) fn gemm3_sk32_pipeline(&self, f16: bool) -> Result<&wgpu::ComputePipeline> {
if f16 {
self.gemm3_sk32_f16
.as_ref()
.ok_or_else(|| anyhow::anyhow!("adapter has no SHADER_F16"))
} else {
Ok(&self.gemm3_sk32_f32)
}
}
pub(crate) fn gemm3_sk_pipeline(&self, f16: bool) -> Result<&wgpu::ComputePipeline> {
if f16 {
self.gemm3_sk_f16
.as_ref()
.ok_or_else(|| anyhow::anyhow!("adapter has no SHADER_F16"))
} else {
Ok(&self.gemm3_sk_f32)
}
}
pub(crate) fn gemm3_pipeline_tile(
&self,
f16: bool,
tile: (usize, usize, usize),
) -> Result<&wgpu::ComputePipeline> {
let set = if f16 {
self.gemm3_f16
.as_ref()
.ok_or_else(|| anyhow::anyhow!("adapter has no SHADER_F16"))?
} else {
&self.gemm3_f32
};
Ok(&set[gemm3_tier(tile)])
}
pub(crate) fn gemm_pipeline(&self, f16: bool) -> Result<&wgpu::ComputePipeline> {
if f16 {
self.gemm_f16
.as_ref()
.ok_or_else(|| anyhow::anyhow!("adapter has no SHADER_F16"))
} else {
Ok(&self.gemm_f32)
}
}
#[allow(clippy::too_many_arguments)]
pub fn gemm_for_tests(
&self,
ctx: &GpuCtx,
x: &[f32],
w: &[f32],
bias: Option<&[f32]>,
m: usize,
n: usize,
k: usize,
act: Option<Act>,
w_f16: bool,
) -> Result<Vec<f32>> {
let (pl, bm, bn) = if gemm2_enabled() {
let t = gemm2_tile(m, n);
(self.gemm2_pipeline_tile(w_f16, t)?, t.0, t.1)
} else {
(self.gemm_pipeline(w_f16)?, 16, 16)
};
let xb = ctx.storage(x);
let wb = if w_f16 {
let h: Vec<u8> = w
.iter()
.flat_map(|&v| half::f16::from_f32(v).to_le_bytes())
.collect();
ctx.storage_bytes(&h)
} else {
ctx.storage(w)
};
let bb = ctx.storage(bias.unwrap_or(&[0.0]));
let yb = ctx.storage(&vec![0f32; m * n]);
let flags = u32::from(bias.is_some()) | (act_code(act) << 8);
let meta = uni(
ctx,
bytemuck::cast_slice(&[m as u32, n as u32, k as u32, flags]),
);
let bg = make_bg(ctx, pl, &[&xb, &wb, &bb, &yb], &meta);
dispatch(
ctx,
pl,
&bg,
(n as u32).div_ceil(bn as u32),
(m as u32).div_ceil(bm as u32),
);
ctx.read(&yb, m * n)
}
#[allow(clippy::too_many_arguments)]
pub fn gemm_resident(
&self,
ctx: &GpuCtx,
x: &[f32],
wb: &wgpu::Buffer,
bb: &wgpu::Buffer,
m: usize,
n: usize,
k: usize,
act: Option<Act>,
) -> Result<Vec<f32>> {
let (pl, bm, bn) = if gemm2_enabled() {
let t = gemm2_tile(m, n);
(self.gemm2_pipeline_tile(false, t)?, t.0, t.1)
} else {
(self.gemm_pipeline(false)?, 16, 16)
};
let xb = ctx.storage(x);
let yb = ctx.storage(&vec![0f32; m * n]);
let flags = 1u32 | (act_code(act) << 8); let meta = uni(
ctx,
bytemuck::cast_slice(&[m as u32, n as u32, k as u32, flags]),
);
let bg = make_bg(ctx, pl, &[&xb, wb, bb, &yb], &meta);
dispatch(
ctx,
pl,
&bg,
(n as u32).div_ceil(bn as u32),
(m as u32).div_ceil(bm as u32),
);
ctx.read(&yb, m * n)
}
#[allow(clippy::too_many_arguments)]
pub fn gemm_resident_into(
&self,
ctx: &GpuCtx,
xb: &wgpu::Buffer,
wb: &wgpu::Buffer,
bb: &wgpu::Buffer,
yb: &wgpu::Buffer,
m: usize,
n: usize,
k: usize,
act: Option<Act>,
) -> Result<()> {
let (pl, bm, bn) = if gemm2_enabled() {
let t = gemm2_tile(m, n);
(self.gemm2_pipeline_tile(false, t)?, t.0, t.1)
} else {
(self.gemm_pipeline(false)?, 16, 16)
};
let flags = 1u32 | (act_code(act) << 8); let meta = uni(
ctx,
bytemuck::cast_slice(&[m as u32, n as u32, k as u32, flags]),
);
let bg = make_bg(ctx, pl, &[xb, wb, bb, yb], &meta);
dispatch(
ctx,
pl,
&bg,
(n as u32).div_ceil(bn as u32),
(m as u32).div_ceil(bm as u32),
);
Ok(())
}
#[allow(clippy::too_many_arguments)]
pub fn gemm_bench(
&self,
ctx: &GpuCtx,
m: usize,
n: usize,
k: usize,
w_f16: bool,
v2: bool,
reps: usize,
) -> Result<f64> {
let (pl, bm, bn) = if v2 {
let t = gemm2_tile(m, n);
(self.gemm2_pipeline_tile(w_f16, t)?, t.0, t.1)
} else {
(self.gemm_pipeline(w_f16)?, 16, 16)
};
let xb = ctx.storage(&vec![0.5f32; m * k]);
let wb = if w_f16 {
ctx.storage_bytes(&vec![0u8; n * k * 2])
} else {
ctx.storage(&vec![0.25f32; n * k])
};
let bb = ctx.storage(&[0.0f32]);
let yb = ctx.storage(&vec![0f32; m * n]);
let meta = uni(
ctx,
bytemuck::cast_slice(&[m as u32, n as u32, k as u32, 0u32]),
);
let bg = make_bg(ctx, pl, &[&xb, &wb, &bb, &yb], &meta);
let (gx, gy) = (
(n as u32).div_ceil(bn as u32),
(m as u32).div_ceil(bm as u32),
);
let run = |reps: usize| -> f64 {
let mut enc = ctx
.device
.create_command_encoder(&wgpu::CommandEncoderDescriptor { label: None });
{
let mut pass = enc.begin_compute_pass(&wgpu::ComputePassDescriptor::default());
pass.set_pipeline(pl);
pass.set_bind_group(0, &bg, &[]);
for _ in 0..reps {
pass.dispatch_workgroups(gx, gy, 1);
}
}
let t0 = std::time::Instant::now();
ctx.queue.submit([enc.finish()]);
let _ = ctx.device.poll(wgpu::PollType::wait_indefinitely());
t0.elapsed().as_secs_f64()
};
run(2); let mut best = f64::MAX;
for _ in 0..5 {
best = best.min(run(reps));
}
let flop = 2.0 * m as f64 * n as f64 * k as f64 * reps as f64;
Ok(flop / best / 1e9)
}
#[allow(clippy::too_many_arguments)]
pub fn gemm3_bench(
&self,
ctx: &GpuCtx,
m: usize,
n: usize,
k: usize,
bm: usize,
bn: usize,
bk: usize,
reps: usize,
) -> Result<f64> {
let pl = pipeline(
ctx,
&format!("enc_gemm3_{bm}x{bn}x{bk}_lab"),
&enc_gemm3_src(false, bm, bn, bk),
);
let xb = ctx.storage(&vec![0.5f32; m * k]);
let wb = ctx.storage(&vec![0.25f32; k * n]); let bb = ctx.storage(&[0.0f32]);
let yb = ctx.storage(&vec![0f32; m * n]);
let meta = uni(
ctx,
bytemuck::cast_slice(&[m as u32, n as u32, k as u32, 0u32]),
);
let bg = make_bg(ctx, &pl, &[&xb, &wb, &bb, &yb], &meta);
let (gx, gy) = (
(n as u32).div_ceil(bn as u32),
(m as u32).div_ceil(bm as u32),
);
let run = |reps: usize| -> f64 {
let mut enc = ctx
.device
.create_command_encoder(&wgpu::CommandEncoderDescriptor { label: None });
{
let mut pass = enc.begin_compute_pass(&wgpu::ComputePassDescriptor::default());
pass.set_pipeline(&pl);
pass.set_bind_group(0, &bg, &[]);
for _ in 0..reps {
pass.dispatch_workgroups(gx, gy, 1);
}
}
let t0 = std::time::Instant::now();
ctx.queue.submit([enc.finish()]);
let _ = ctx.device.poll(wgpu::PollType::wait_indefinitely());
t0.elapsed().as_secs_f64()
};
run(2);
let mut best = f64::MAX;
for _ in 0..5 {
best = best.min(run(reps));
}
let flop = 2.0 * m as f64 * n as f64 * k as f64 * reps as f64;
Ok(flop / best / 1e9)
}
#[allow(clippy::too_many_arguments)]
pub fn gemm3_sk_bench(
&self,
ctx: &GpuCtx,
m: usize,
n: usize,
k: usize,
bm: usize,
bn: usize,
bk: usize,
chunks: u32,
reps: usize,
) -> Result<f64> {
let pl = pipeline(
ctx,
&format!("enc_gemm3_sk_{bm}x{bn}x{bk}_lab"),
&enc_gemm3_sk_src(false, bm, bn, bk),
);
let rpl = pipeline(ctx, "enc_gemm3_sk_reduce_lab", &enc_gemm3_sk_reduce_src());
let xb = ctx.storage(&vec![0.5f32; m * k]);
let wb = ctx.storage(&vec![0.25f32; k * n]); let bb = ctx.storage(&[0.0f32]);
let part = ctx.storage(&vec![0f32; chunks as usize * m * n]);
let yb = ctx.storage(&vec![0f32; m * n]);
let meta = uni(
ctx,
bytemuck::cast_slice(&[m as u32, n as u32, k as u32, 0u32]),
);
let meta_r = uni(
ctx,
bytemuck::cast_slice(&[m as u32, n as u32, 0u32, chunks]),
);
let bg = make_bg(ctx, &pl, &[&xb, &wb, &bb, &part], &meta);
let rbg = make_bg(ctx, &rpl, &[&part, &bb, &yb], &meta_r);
let (gx, gy) = (
(n as u32).div_ceil(bn as u32),
(m as u32).div_ceil(bm as u32),
);
let rgx = ((m * n) as u32).div_ceil(256);
let run = |reps: usize| -> f64 {
let mut enc = ctx
.device
.create_command_encoder(&wgpu::CommandEncoderDescriptor { label: None });
{
let mut pass = enc.begin_compute_pass(&wgpu::ComputePassDescriptor::default());
for _ in 0..reps {
pass.set_pipeline(&pl);
pass.set_bind_group(0, &bg, &[]);
pass.dispatch_workgroups(gx, gy, chunks);
pass.set_pipeline(&rpl);
pass.set_bind_group(0, &rbg, &[]);
pass.dispatch_workgroups(rgx, 1, 1);
}
}
let t0 = std::time::Instant::now();
ctx.queue.submit([enc.finish()]);
let _ = ctx.device.poll(wgpu::PollType::wait_indefinitely());
t0.elapsed().as_secs_f64()
};
run(2);
let mut best = f64::MAX;
for _ in 0..5 {
best = best.min(run(reps));
}
let flop = 2.0 * m as f64 * n as f64 * k as f64 * reps as f64;
Ok(flop / best / 1e9)
}
#[allow(clippy::too_many_arguments)]
pub fn layernorm_for_tests(
&self,
ctx: &GpuCtx,
x: &[f32],
res: Option<&[f32]>,
w: &[f32],
b: Option<&[f32]>,
t: usize,
h: usize,
eps: f32,
) -> Result<Vec<f32>> {
let xb = ctx.storage(x);
let rb = ctx.storage(res.unwrap_or(&[0.0]));
let wb = ctx.storage(w);
let bb = ctx.storage(b.unwrap_or(&[0.0]));
let ob = ctx.storage(&vec![0f32; t * h]);
let flags = u32::from(res.is_some()) | (u32::from(b.is_some()) << 1);
let meta = uni(
ctx,
bytemuck::cast_slice(&[h as u32, flags, eps.to_bits(), 0]),
);
let bg = make_bg(ctx, &self.layernorm, &[&xb, &rb, &wb, &bb, &ob], &meta);
dispatch(ctx, &self.layernorm, &bg, t as u32, 1);
ctx.read(&ob, t * h)
}
#[allow(clippy::too_many_arguments)]
#[allow(clippy::too_many_arguments)]
pub fn attention_for_tests(
&self,
ctx: &GpuCtx,
q: &[f32],
k: &[f32],
v: &[f32],
seq_starts: &[u32],
n_heads: usize,
n_kv_heads: usize,
hd: usize,
mask: MaskKind,
window: u32,
) -> Result<Vec<f32>> {
anyhow::ensure!(hd <= 128, "encoder attention supports head_dim ≤ 128");
let h = n_heads * hd;
let nrows = q.len() / h;
let seq_of = row_to_seq(seq_starts, nrows);
let qb = ctx.storage(q);
let kb = ctx.storage(k);
let vb = ctx.storage(v);
let ob = ctx.storage(&vec![0f32; nrows * h]);
let sb = ctx.storage_bytes(bytemuck::cast_slice(seq_starts));
let mb = ctx.storage_bytes(bytemuck::cast_slice(&seq_of));
let mode = match mask {
MaskKind::Bidirectional => 0u32,
MaskKind::Causal => 1u32,
};
let meta = uni(
ctx,
bytemuck::cast_slice(&[
nrows as u32,
n_heads as u32,
hd as u32,
mode,
window,
n_kv_heads as u32,
0,
0,
]),
);
let valid = ctx.storage_bytes(bytemuck::cast_slice(&vec![1u32; nrows]));
let bg = make_bg(
ctx,
&self.attn,
&[&qb, &kb, &vb, &ob, &sb, &mb, &valid],
&meta,
);
dispatch(ctx, &self.attn, &bg, nrows as u32, n_heads as u32);
ctx.read(&ob, nrows * h)
}
pub fn mean_pool_l2_for_tests(
&self,
ctx: &GpuCtx,
hidden: &[f32],
h: usize,
seq_starts: &[u32],
) -> Result<Vec<Vec<f32>>> {
let n_seqs = seq_starts.len() - 1;
let hb = ctx.storage(hidden);
let ob = ctx.storage(&vec![0f32; n_seqs * h]);
let sb = ctx.storage_bytes(bytemuck::cast_slice(seq_starts));
let meta = uni(ctx, bytemuck::cast_slice(&[h as u32, 0, 0, 0]));
let bg = make_bg(ctx, &self.mean_pool, &[&hb, &ob, &sb], &meta);
dispatch(ctx, &self.mean_pool, &bg, n_seqs as u32, 1);
let meta2 = uni(ctx, bytemuck::cast_slice(&[h as u32, 0, 0, 0]));
let bg2 = make_bg(ctx, &self.l2norm, &[&ob], &meta2);
dispatch(ctx, &self.l2norm, &bg2, n_seqs as u32, 1);
let flat = ctx.read(&ob, n_seqs * h)?;
Ok(flat.chunks_exact(h).map(<[f32]>::to_vec).collect())
}
}
fn cpu_layer_norm_rows(x: &mut [f32], h: usize, w: &[f32], b: &[f32], eps: f32) {
for row in x.chunks_exact_mut(h) {
let mean = row.iter().sum::<f32>() / h as f32;
let var = row.iter().map(|v| (v - mean) * (v - mean)).sum::<f32>() / h as f32;
let inv = 1.0 / (var + eps).sqrt();
for (j, v) in row.iter_mut().enumerate() {
*v = (*v - mean) * inv * w[j] + b[j];
}
}
}
fn cpu_linear(x: &[f32], w: &[f32], bias: Option<&[f32]>, n: usize, k: usize) -> Vec<f32> {
let rows = x.len() / k;
let mut out = vec![0f32; rows * n];
for (o, xr) in out.chunks_exact_mut(n).zip(x.chunks_exact(k)) {
for (nn, o_n) in o.iter_mut().enumerate() {
let wr = &w[nn * k..(nn + 1) * k];
let mut acc = bias.map_or(0.0, |b| b[nn]);
for (xk, wk) in xr.iter().zip(wr) {
acc += xk * wk;
}
*o_n = acc;
}
}
out
}
pub(crate) fn dispatch(
ctx: &GpuCtx,
pl: &wgpu::ComputePipeline,
bg: &wgpu::BindGroup,
gx: u32,
gy: u32,
) {
let mut enc = ctx
.device
.create_command_encoder(&wgpu::CommandEncoderDescriptor { label: None });
{
let mut pass = enc.begin_compute_pass(&wgpu::ComputePassDescriptor::default());
pass.set_pipeline(pl);
pass.set_bind_group(0, bg, &[]);
pass.dispatch_workgroups(gx, gy, 1);
}
ctx.queue.submit([enc.finish()]);
}
pub(crate) fn row_to_seq(seq_starts: &[u32], t: usize) -> Vec<u32> {
let mut map = vec![0u32; t];
for (s, w) in seq_starts.windows(2).enumerate() {
for r in w[0]..w[1] {
map[r as usize] = s as u32;
}
}
map
}
use crate::encoder_weights::{EncArch, EncBatch, EncoderConfig, MlpKind, NormKind, PosKind};
use crate::pooling::{EmbedOut, Pooling};
use crate::weights::{LazySt, f32_to_f16_bytes};
use anyhow::Context;
use std::path::Path;
enum EncMatBuf {
F16(wgpu::Buffer),
F32(wgpu::Buffer),
}
struct EncGpuLinear {
w: EncMatBuf,
b: Option<wgpu::Buffer>,
n: u32,
k: u32,
v3: bool,
sk: bool,
f16a: bool,
b_host: Option<Vec<f32>>,
}
struct CeHeadGpu {
bg_pooler: GemmBg,
bg_classifier: GemmBg,
metas: Vec<wgpu::Buffer>,
out: wgpu::Buffer,
n_labels: usize,
fused: Option<(wgpu::ComputePipeline, wgpu::BindGroup)>,
_mid: wgpu::Buffer,
_weights: Vec<EncGpuLinear>,
}
const ENC_CE_FUSED: &str = r#"
struct Meta { h: u32, n_labels: u32, p0: u32, p1: u32 }
@group(0) @binding(0) var<storage, read> hidden: array<f32>; // [T, h]
@group(0) @binding(1) var<storage, read> seq_starts: array<u32>;
@group(0) @binding(2) var<storage, read> pw: array<f32>; // pooler [h, h] (HF [n,k])
@group(0) @binding(3) var<storage, read> pb: array<f32>; // pooler bias [h]
@group(0) @binding(4) var<storage, read> cw: array<f32>; // classifier [n_labels, h]
@group(0) @binding(5) var<storage, read> cb: array<f32>; // classifier bias [n_labels]
@group(0) @binding(6) var<storage, read_write> out: array<f32>; // [b, n_labels]
@group(0) @binding(7) var<uniform> mt: Meta;
var<workgroup> cls: array<f32, 2048>;
var<workgroup> pooled: array<f32, 2048>;
var<workgroup> red: array<f32, 256>;
@compute @workgroup_size(256)
fn main(@builtin(workgroup_id) wg: vec3<u32>, @builtin(local_invocation_index) t: u32) {
let s = wg.x;
let h = mt.h;
let base = seq_starts[s] * h; // the sequence's FIRST token = CLS
for (var j = t; j < h; j += 256u) { cls[j] = hidden[base + j]; }
workgroupBarrier();
// pooled[j] = tanh(pw[j]·cls + pb[j]) — each thread owns rows j, j+256, …
for (var j = t; j < h; j += 256u) {
var acc = 0.0;
let rb = j * h;
for (var d = 0u; d < h; d++) { acc += pw[rb + d] * cls[d]; }
pooled[j] = tanh(clamp(acc + pb[j], -20.0, 20.0));
}
workgroupBarrier();
// classifier: one ladder-reduced dot per label (n_labels is 1-2 in practice).
for (var l = 0u; l < mt.n_labels; l++) {
var part = 0.0;
let rb = l * h;
for (var d = t; d < h; d += 256u) { part += cw[rb + d] * pooled[d]; }
red[t] = part;
workgroupBarrier();
for (var r = 128u; r > 0u; r >>= 1u) {
if (t < r) { red[t] += red[t + r]; }
workgroupBarrier();
}
if (t == 0u) { out[s * mt.n_labels + l] = red[0] + cb[l]; }
workgroupBarrier();
}
}
"#;
pub(crate) fn v2_bgs(
ctx: &GpuCtx,
kernels: &EncKernels,
f16: bool,
bufs: &[&wgpu::Buffer],
meta: &wgpu::Buffer,
) -> Result<Vec<wgpu::BindGroup>> {
GEMM2_TILES
.into_iter()
.map(|t| {
Ok(make_bg(
ctx,
kernels.gemm2_pipeline_tile(f16, t)?,
bufs,
meta,
))
})
.collect()
}
pub(crate) fn v3_bgs(
ctx: &GpuCtx,
kernels: &EncKernels,
f16: bool,
bufs: &[&wgpu::Buffer],
meta: &wgpu::Buffer,
) -> Result<Vec<wgpu::BindGroup>> {
GEMM3_TILES
.into_iter()
.map(|t| {
Ok(make_bg(
ctx,
kernels.gemm3_pipeline_tile(f16, t)?,
bufs,
meta,
))
})
.collect()
}
#[allow(clippy::too_many_arguments)]
fn head_gemm(
ctx: &GpuCtx,
kernels: &EncKernels,
use_f16: bool,
w: &[f32],
bias: &[f32],
n: u32,
k: u32,
x: &wgpu::Buffer,
y: &wgpu::Buffer,
act: Option<Act>,
) -> Result<(GemmBg, wgpu::Buffer, EncGpuLinear)> {
anyhow::ensure!(
w.len() == (n as usize) * (k as usize),
"head weight shape {} != {n}×{k}",
w.len()
);
let mat = if use_f16 {
EncMatBuf::F16(ctx.storage_bytes(&f32_to_f16_bytes(w)))
} else {
EncMatBuf::F32(ctx.storage(w))
};
let bbuf = ctx.storage(bias);
let flags = 1u32 | (act_code(act) << 8);
let meta = uni(ctx, bytemuck::cast_slice(&[0u32, n, k, flags]));
let wbuf = match &mat {
EncMatBuf::F16(b) | EncMatBuf::F32(b) => b,
};
let bufs = [x, wbuf, &bbuf, y];
let bg = if gemm2_enabled() {
GemmBg::V2(v2_bgs(ctx, kernels, use_f16, &bufs, &meta)?)
} else {
GemmBg::V1(make_bg(ctx, kernels.gemm_pipeline(use_f16)?, &bufs, &meta))
};
let keep = EncGpuLinear {
w: mat,
b: Some(bbuf),
n,
k,
v3: false, sk: false,
f16a: false,
b_host: None,
};
Ok((bg, meta, keep))
}
struct EncGpuNorm {
w: wgpu::Buffer,
b: Option<wgpu::Buffer>,
}
enum GemmBg {
V1(wgpu::BindGroup),
V2(Vec<wgpu::BindGroup>),
V3(Vec<wgpu::BindGroup>),
V3Sk {
plain: Vec<wgpu::BindGroup>,
part: wgpu::BindGroup,
part32: wgpu::BindGroup,
reduce: wgpu::BindGroup,
k: u32,
f16a: bool,
},
F16A {
tiles: Vec<wgpu::BindGroup>,
coop: Option<wgpu::BindGroup>,
},
}
type StepDisp<'s> = (
&'s wgpu::ComputePipeline,
&'s wgpu::BindGroup,
u32,
u32,
u32,
&'static str,
);
enum EStep {
Gemm { bg: GemmBg, n: u32 },
Ln { bg: wgpu::BindGroup },
Rope { bg: wgpu::BindGroup },
QkNorm { bg: wgpu::BindGroup },
Glu { bg: wgpu::BindGroup },
Attn { bg: wgpu::BindGroup },
Attn4 { bg: wgpu::BindGroup },
AttnRb { bg: wgpu::BindGroup },
AttnDisent { bg: wgpu::BindGroup },
AttnDisentRb { bg: wgpu::BindGroup },
Add { bg: wgpu::BindGroup },
Copy { bg: wgpu::BindGroup },
Conv { bg: wgpu::BindGroup },
L2Rows { bg: wgpu::BindGroup },
}
pub struct EncoderGpu {
cfg: EncoderConfig,
kernels: EncKernels,
use_f16: bool,
max_tokens: usize,
word: Vec<f32>,
pos_table: Option<Vec<f32>>,
ttype: Option<Vec<f32>>,
emb_in: wgpu::Buffer,
hidden_buf: wgpu::Buffer,
_pool_src_seq_buf: wgpu::Buffer,
pooled: wgpu::Buffer,
seq_starts_buf: wgpu::Buffer,
seq_of_buf: wgpu::Buffer,
valid_buf: wgpu::Buffer,
steps: Vec<EStep>,
bg_pool: Option<wgpu::BindGroup>,
bg_l2: Option<wgpu::BindGroup>,
ptb: Option<wgpu::Buffer>,
metas_tokens: Vec<wgpu::Buffer>,
metas_sk: Vec<(wgpu::Buffer, u32, u32)>,
metas_elems: Vec<wgpu::Buffer>,
ce_head: Option<CeHeadGpu>,
staging: wgpu::Buffer,
staged_read: bool,
_weights: Vec<EncGpuLinear>,
_norms: Vec<EncGpuNorm>,
_scratch: Vec<wgpu::Buffer>,
}
pub const MAX_VISION_BATCH: usize = 32;
impl EncoderGpu {
pub fn kernels(&self) -> &EncKernels {
&self.kernels
}
pub fn load(ctx: &GpuCtx, dir: &Path, max_tokens: usize) -> Result<Self> {
let want_f16 = matches!(std::env::var("OSFKB_ENC_F16").ok().as_deref(), Some("1"));
Self::load_with(ctx, dir, max_tokens, want_f16 && ctx.f16)
}
pub fn load_f32(ctx: &GpuCtx, dir: &Path, max_tokens: usize) -> Result<Self> {
Self::load_with(ctx, dir, max_tokens, false)
}
pub fn load_siglip_vision(ctx: &GpuCtx, dir: &Path) -> Result<Self> {
let want_f16 = matches!(std::env::var("OSFKB_ENC_F16").ok().as_deref(), Some("1"));
let (_, spec) = crate::encoder_weights::siglip_configs_from_dir(dir)?;
let n = (spec.image_size / spec.patch_size).pow(2);
Self::load_with_cfg(ctx, dir, spec.config, MAX_VISION_BATCH * n, want_f16 && ctx.f16)
}
pub fn load_siglip_vision_f32(ctx: &GpuCtx, dir: &Path) -> Result<Self> {
let (_, spec) = crate::encoder_weights::siglip_configs_from_dir(dir)?;
let n = (spec.image_size / spec.patch_size).pow(2);
Self::load_with_cfg(ctx, dir, spec.config, MAX_VISION_BATCH * n, false)
}
fn load_with(ctx: &GpuCtx, dir: &Path, max_tokens: usize, use_f16: bool) -> Result<Self> {
let cfg = crate::encoder_weights::encoder_config_from_dir(dir)?;
Self::load_with_cfg(ctx, dir, cfg, max_tokens, use_f16)
}
fn load_with_cfg(
ctx: &GpuCtx,
dir: &Path,
cfg: EncoderConfig,
max_tokens: usize,
use_f16: bool,
) -> Result<Self> {
anyhow::ensure!(
cfg.head_dim <= 128,
"encoder attention supports head_dim ≤ 128"
);
let st = LazySt::open(dir)?;
let kernels = EncKernels::new(ctx);
if use_f16 {
kernels.gemm_pipeline(true)?; }
let mut b = PlanBuilder::new(ctx, &kernels, &cfg, max_tokens, use_f16);
let (word, pos_table, ttype) = match cfg.arch {
EncArch::Bert | EncArch::XlmRoberta => b.build_bert_family(&st)?,
EncArch::DebertaV2 => b.build_deberta(&st)?,
EncArch::ModernBert => b.build_modernbert(&st)?,
EncArch::Qwen3Embed => b.build_qwen3_embed(&st)?,
EncArch::Lfm2Colbert => b.build_lfm2_colbert(&st, dir)?,
EncArch::NomicBert => b.build_nomic(&st)?,
EncArch::SiglipVision => b.build_siglip_vision(&st)?,
other => anyhow::bail!(
"GPU encoder tensor table for {other:?} pending verification against a real \
checkpoint"
),
};
let PlanBuilder {
emb_in,
cur,
pool_src_seq_buf,
pooled,
ptb,
valid_buf,
seq_starts_buf,
seq_of_buf,
steps,
bg_pool,
bg_l2,
metas_tokens,
metas_sk,
metas_elems,
weights,
norms,
scratch,
..
} = b;
if !matches!(cfg.pooling, Pooling::PerToken { .. } | Pooling::MapHead) {
anyhow::ensure!(
bg_pool.is_some() && bg_l2.is_some(),
"plan built no pooling bind groups"
);
}
let ce_head = if let Pooling::CrossEncoder { n_labels } = cfg.pooling {
let h = cfg.hidden;
let prefix = ["", "bert.", "roberta."]
.into_iter()
.find(|p| st.has(&format!("{p}embeddings.word_embeddings.weight")))
.context("cross-encoder: word embeddings not under '', 'bert.' or 'roberta.'")?;
let cap = max_tokens;
let mid = ctx.storage(&vec![0f32; cap * h]);
let out = ctx.storage(&vec![0f32; cap * n_labels]);
let (pool_w, pool_b, cls_w, cls_b) = if st.has(&format!("{prefix}pooler.dense.weight"))
{
(
format!("{prefix}pooler.dense.weight"),
format!("{prefix}pooler.dense.bias"),
"classifier.weight".to_string(),
"classifier.bias".to_string(),
)
} else {
(
"classifier.dense.weight".to_string(),
"classifier.dense.bias".to_string(),
"classifier.out_proj.weight".to_string(),
"classifier.out_proj.bias".to_string(),
)
};
let (bg_pooler, m1, w1) = head_gemm(
ctx,
&kernels,
use_f16,
&st.tensor_f32(&pool_w)?,
&st.tensor_f32(&pool_b)?,
h as u32,
h as u32,
&pooled,
&mid,
Some(Act::Tanh),
)?;
let (bg_classifier, m2, w2) = head_gemm(
ctx,
&kernels,
use_f16,
&st.tensor_f32(&cls_w)?,
&st.tensor_f32(&cls_b)?,
n_labels as u32,
h as u32,
&mid,
&out,
None,
)?;
let fused = if !use_f16
&& h <= 2048
&& std::env::var("OSFKB_ENC_CE_FUSED").ok().as_deref() != Some("0")
{
let (EncMatBuf::F32(pwb) | EncMatBuf::F16(pwb)) = &w1.w;
let (EncMatBuf::F32(cwb) | EncMatBuf::F16(cwb)) = &w2.w;
let pl = pipeline(ctx, "enc_ce_fused", ENC_CE_FUSED);
let meta = uni(
ctx,
bytemuck::cast_slice(&[h as u32, n_labels as u32, 0, 0]),
);
let bg = make_bg(
ctx,
&pl,
&[
&cur,
&seq_starts_buf,
pwb,
w1.b.as_ref().expect("pooler bias"),
cwb,
w2.b.as_ref().expect("classifier bias"),
&out,
],
&meta,
);
Some((pl, bg))
} else {
None
};
Some(CeHeadGpu {
bg_pooler,
bg_classifier,
metas: vec![m1, m2],
out,
n_labels,
fused,
_mid: mid,
_weights: vec![w1, w2],
})
} else {
None
};
let stage_floats = max_tokens
* cfg.hidden.max(match cfg.pooling {
Pooling::PerToken { dim } => dim,
_ => 0,
});
let staging = ctx.device.create_buffer(&wgpu::BufferDescriptor {
label: Some("enc_readback_staging"),
size: (stage_floats * 4) as u64,
usage: wgpu::BufferUsages::MAP_READ | wgpu::BufferUsages::COPY_DST,
mapped_at_creation: false,
});
let staged_read = std::env::var("OSFKB_ENC_STAGED_READ").ok().as_deref() == Some("1");
Ok(Self {
cfg,
kernels,
use_f16,
max_tokens,
word,
pos_table,
ttype,
emb_in,
hidden_buf: cur,
_pool_src_seq_buf: pool_src_seq_buf,
pooled,
seq_starts_buf,
seq_of_buf,
valid_buf,
steps,
bg_pool,
bg_l2,
ptb,
metas_tokens,
metas_sk,
metas_elems,
ce_head,
staging,
staged_read,
_weights: weights,
_norms: norms,
_scratch: scratch,
})
}
pub fn config(&self) -> &EncoderConfig {
&self.cfg
}
pub fn precision(&self) -> &'static str {
if self.use_f16 { "f16" } else { "f32" }
}
pub fn forward_hidden(&mut self, ctx: &GpuCtx, batch: &EncBatch) -> Result<Vec<f32>> {
if std::env::var("OSFKB_ENC_TS").ok().as_deref() == Some("1") {
static ONCE: std::sync::Once = std::sync::Once::new();
let mut rows = None;
ONCE.call_once(|| rows = Some(self.profile_encode(ctx, batch)));
if let Some(Ok(rows)) = rows {
let total: f64 = rows.iter().map(|r| r.1).sum();
eprintln!(" [enc ts] plan kernels {total:8.1} µs (head/readback excluded)");
for (name, us, cnt) in rows {
eprintln!(
" [enc ts] {name:16} {us:8.1} µs ({:4.1}%) ×{cnt}",
us / total * 100.0
);
}
}
}
let t = self.dispatch(ctx, batch)?;
ctx.read(&self.hidden_buf, t * self.cfg.hidden)
}
pub fn run_layers_gpu(&mut self, ctx: &GpuCtx, rows: &[f32], n: usize) -> Result<Vec<f32>> {
self.run_layers_gpu_batched(ctx, rows, n, 1)
}
pub fn run_layers_gpu_batched(
&mut self,
ctx: &GpuCtx,
rows: &[f32],
seq_len: usize,
batch: usize,
) -> Result<Vec<f32>> {
let h = self.cfg.hidden;
let t = batch * seq_len;
anyhow::ensure!(
rows.len() == t * h,
"rows {} != batch {batch} × seq_len {seq_len} × hidden {h}",
rows.len()
);
anyhow::ensure!(
t <= self.max_tokens,
"batch·seq_len {t} exceeds max_tokens {}",
self.max_tokens
);
ctx.queue
.write_buffer(&self.emb_in, 0, bytemuck::cast_slice(rows));
let seq_starts: Vec<u32> = (0..=batch).map(|i| (i * seq_len) as u32).collect();
let seq_of = row_to_seq(&seq_starts, t);
let valid = vec![1u32; t];
ctx.queue
.write_buffer(&self.seq_starts_buf, 0, bytemuck::cast_slice(&seq_starts));
ctx.queue
.write_buffer(&self.seq_of_buf, 0, bytemuck::cast_slice(&seq_of));
ctx.queue
.write_buffer(&self.valid_buf, 0, bytemuck::cast_slice(&valid));
for m in &self.metas_tokens {
ctx.queue
.write_buffer(m, 0, bytemuck::cast_slice(&[t as u32]));
}
for (m, nn, k) in &self.metas_sk {
let z = gemm3_sk_plan(t, *nn as usize, *k as usize).1;
ctx.queue.write_buffer(m, 12, bytemuck::cast_slice(&[z]));
}
for m in &self.metas_elems {
ctx.queue
.write_buffer(m, 0, bytemuck::cast_slice(&[(t * h) as u32]));
}
let mut cmd = ctx
.device
.create_command_encoder(&wgpu::CommandEncoderDescriptor { label: None });
{
let mut pass = cmd.begin_compute_pass(&wgpu::ComputePassDescriptor::default());
for step in &self.steps {
let (d, d2) = self.resolve_step(step, t)?;
for (pl, bg, gx, gy, gz, _) in std::iter::once(d).chain(d2) {
pass.set_pipeline(pl);
pass.set_bind_group(0, bg, &[]);
pass.dispatch_workgroups(gx, gy, gz);
}
}
}
ctx.queue.submit([cmd.finish()]);
ctx.read(&self.hidden_buf, t * h)
}
pub fn encode(&mut self, ctx: &GpuCtx, batch: &EncBatch) -> Result<EmbedOut> {
let _trace = std::env::var("OSFKB_ENC_GPU_TRACE").is_ok();
if std::env::var("OSFKB_ENC_TS").ok().as_deref() == Some("1") {
static ONCE: std::sync::Once = std::sync::Once::new();
let mut rows = None;
ONCE.call_once(|| rows = Some(self.profile_encode(ctx, batch)));
if let Some(rows) = rows {
match rows {
Ok(rows) => {
let total: f64 = rows.iter().map(|r| r.1).sum();
eprintln!(
" [enc ts] plan kernels {total:8.1} µs (pool/head/readback excluded)"
);
for (name, us, cnt) in rows {
eprintln!(
" [enc ts] {name:16} {us:8.1} µs ({:4.1}%) ×{cnt}",
us / total * 100.0
);
}
}
Err(e) => eprintln!(" [enc ts] profile failed: {e}"),
}
}
}
let _t0 = std::time::Instant::now();
let t = self.dispatch(ctx, batch)?;
if _trace {
eprintln!(
" [enc gpu] dispatch (gather+upload+GPU) {:?}",
_t0.elapsed()
);
}
let _t1 = std::time::Instant::now();
let (cfg, h, b_seqs) = (&self.cfg, self.cfg.hidden, batch.n_seqs());
if let Pooling::PerToken { dim } = cfg.pooling {
let flat = if self.staged_read {
self.read_staged(ctx, t * dim)?
} else {
let ptb = self
.ptb
.as_ref()
.context("PerToken plan built no output buffer")?;
ctx.read(ptb, t * dim)?
};
let mut out: Vec<Vec<Vec<f32>>> = Vec::with_capacity(b_seqs);
for w in batch.seq_starts.windows(2) {
out.push(
(w[0] as usize..w[1] as usize)
.map(|i| flat[i * dim..(i + 1) * dim].to_vec())
.collect(),
);
}
return Ok(EmbedOut::PerToken(out));
}
if let Some(ce) = &self.ce_head {
let flat = if self.staged_read {
self.read_staged(ctx, b_seqs * ce.n_labels)?
} else {
ctx.read(&ce.out, b_seqs * ce.n_labels)?
};
if _trace {
eprintln!(
" [enc gpu] readback (cross-encoder logits) {:?}",
_t1.elapsed()
);
}
return Ok(EmbedOut::Pooled(
flat.chunks_exact(ce.n_labels)
.map(<[f32]>::to_vec)
.collect(),
));
}
let flat = if self.staged_read {
self.read_staged(ctx, b_seqs * h)?
} else {
ctx.read(&self.pooled, b_seqs * h)?
};
if _trace {
eprintln!(
" [enc gpu] readback ({} floats) {:?}",
b_seqs * h,
_t1.elapsed()
);
}
Ok(EmbedOut::Pooled(
flat.chunks_exact(h).map(<[f32]>::to_vec).collect(),
))
}
fn dispatch(&mut self, ctx: &GpuCtx, batch: &EncBatch) -> Result<usize> {
let cfg = &self.cfg;
let (t, b_seqs) = (batch.tokens.len(), batch.n_seqs());
anyhow::ensure!(b_seqs > 0, "empty batch");
anyhow::ensure!(
t <= self.max_tokens,
"batch of {t} tokens exceeds max_tokens {} — chunk at sequence boundaries",
self.max_tokens
);
let offset = match cfg.pos_kind {
PosKind::Learned { offset } => offset,
PosKind::Rope { .. } => 0,
};
for w in batch.seq_starts.windows(2) {
let len = (w[1] - w[0]) as usize;
anyhow::ensure!(len > 0, "empty sequence in batch");
anyhow::ensure!(
len + offset <= cfg.max_pos,
"sequence of {len} tokens exceeds max positions {} (offset {offset})",
cfg.max_pos
);
}
for &tok in &batch.tokens {
anyhow::ensure!((tok as usize) < cfg.vocab, "token id {tok} out of vocab");
}
let h = cfg.hidden;
let mut emb = vec![0f32; t * h];
let seq_of = row_to_seq(&batch.seq_starts, t);
for (i, row) in emb.chunks_exact_mut(h).enumerate() {
let tok = batch.tokens[i] as usize;
row.copy_from_slice(&self.word[tok * h..(tok + 1) * h]);
if let Some(pos_table) = &self.pos_table {
let local = i - batch.seq_starts[seq_of[i] as usize] as usize;
for (r, p) in row
.iter_mut()
.zip(&pos_table[(offset + local) * h..(offset + local + 1) * h])
{
*r += p;
}
}
if let Some(tt) = &self.ttype {
let ty = batch.type_ids.as_ref().map_or(0, |t| t[i] as usize);
anyhow::ensure!(
(ty + 1) * h <= tt.len(),
"token_type id {ty} exceeds the checkpoint's segment table"
);
for (r, v) in row.iter_mut().zip(&tt[ty * h..(ty + 1) * h]) {
*r += v;
}
}
}
ctx.queue
.write_buffer(&self.emb_in, 0, bytemuck::cast_slice(&emb));
ctx.queue.write_buffer(
&self.seq_starts_buf,
0,
bytemuck::cast_slice(&batch.seq_starts),
);
ctx.queue
.write_buffer(&self.seq_of_buf, 0, bytemuck::cast_slice(&seq_of));
let ones;
let valid_slice: &[u32] = match &batch.valid {
Some(v) => {
anyhow::ensure!(
v.len() == t,
"validity mask length {} != tokens {t}",
v.len()
);
v
}
None => {
ones = vec![1u32; t];
&ones
}
};
ctx.queue
.write_buffer(&self.valid_buf, 0, bytemuck::cast_slice(valid_slice));
for m in &self.metas_tokens {
ctx.queue
.write_buffer(m, 0, bytemuck::cast_slice(&[t as u32]));
}
for (m, n, k) in &self.metas_sk {
let z = gemm3_sk_plan(t, *n as usize, *k as usize).1;
ctx.queue.write_buffer(m, 12, bytemuck::cast_slice(&[z]));
}
for m in &self.metas_elems {
ctx.queue
.write_buffer(m, 0, bytemuck::cast_slice(&[(t * h) as u32]));
}
if let Some(ce) = &self.ce_head {
for m in &ce.metas {
ctx.queue
.write_buffer(m, 0, bytemuck::cast_slice(&[b_seqs as u32]));
}
}
let mut enc = ctx
.device
.create_command_encoder(&wgpu::CommandEncoderDescriptor { label: None });
{
let mut pass = enc.begin_compute_pass(&wgpu::ComputePassDescriptor::default());
for step in &self.steps {
let (d, d2) = self.resolve_step(step, t)?;
for (pl, bg, gx, gy, gz, _) in std::iter::once(d).chain(d2) {
pass.set_pipeline(pl);
pass.set_bind_group(0, bg, &[]);
pass.dispatch_workgroups(gx, gy, gz);
}
}
let gemm_pl = self.kernels.gemm_pipeline(self.use_f16)?;
match cfg.pooling {
Pooling::Mean | Pooling::LastToken | Pooling::Cls => {
let pool_pl = match cfg.pooling {
Pooling::LastToken => &self.kernels.last_pool,
Pooling::Cls => &self.kernels.cls_pool,
_ => &self.kernels.mean_pool,
};
let (bg_pool, bg_l2) = (
self.bg_pool.as_ref().expect("validated at load"),
self.bg_l2.as_ref().expect("validated at load"),
);
pass.set_pipeline(pool_pl);
pass.set_bind_group(0, bg_pool, &[]);
pass.dispatch_workgroups(b_seqs as u32, 1, 1);
pass.set_pipeline(&self.kernels.l2norm);
pass.set_bind_group(0, bg_l2, &[]);
pass.dispatch_workgroups(b_seqs as u32, 1, 1);
}
Pooling::CrossEncoder { .. } => {
let ce = self
.ce_head
.as_ref()
.expect("CrossEncoder head built at load");
if let Some((pl, bg)) = &ce.fused {
pass.set_pipeline(pl);
pass.set_bind_group(0, bg, &[]);
pass.dispatch_workgroups(b_seqs as u32, 1, 1);
} else {
let bg_pool = self.bg_pool.as_ref().expect("validated at load");
pass.set_pipeline(&self.kernels.cls_pool);
pass.set_bind_group(0, bg_pool, &[]);
pass.dispatch_workgroups(b_seqs as u32, 1, 1);
let m = b_seqs;
for (bg, n) in [
(&ce.bg_pooler, cfg.hidden),
(&ce.bg_classifier, ce.n_labels),
] {
let (pl, bg, gx, gy) = match bg {
GemmBg::V1(bg) => (
gemm_pl,
bg,
(n as u32).div_ceil(16),
(m as u32).div_ceil(16),
),
GemmBg::V2(bgs) => {
let tile = gemm2_tile(m, n);
(
self.kernels.gemm2_pipeline_tile(self.use_f16, tile)?,
&bgs[gemm2_tier(tile)],
(n as u32).div_ceil(tile.1 as u32),
(m as u32).div_ceil(tile.0 as u32),
)
}
GemmBg::V3(_) | GemmBg::V3Sk { .. } | GemmBg::F16A { .. } => {
unreachable!(
"head_gemm never uploads v3/sk/f16a (see its `keep`)"
)
}
};
pass.set_pipeline(pl);
pass.set_bind_group(0, bg, &[]);
pass.dispatch_workgroups(gx, gy, 1);
}
}
}
Pooling::PerToken { .. } => {}
other => anyhow::bail!("pooling {other:?} lands with its phase"),
}
}
if self.staged_read {
let (src, floats) = if let Pooling::PerToken { dim } = cfg.pooling {
(
self.ptb
.as_ref()
.context("PerToken plan built no output buffer")?,
t * dim,
)
} else if let Some(ce) = &self.ce_head {
(&ce.out, b_seqs * ce.n_labels)
} else {
(&self.pooled, b_seqs * cfg.hidden)
};
enc.copy_buffer_to_buffer(src, 0, &self.staging, 0, (floats * 4) as u64);
}
ctx.queue.submit([enc.finish()]);
Ok(t)
}
#[allow(clippy::type_complexity)]
fn resolve_step<'s>(
&'s self,
step: &'s EStep,
t: usize,
) -> Result<(StepDisp<'s>, Option<StepDisp<'s>>)> {
let cfg = &self.cfg;
let t32 = t as u32;
let h = cfg.hidden;
Ok(match step {
EStep::Gemm {
bg:
GemmBg::V3Sk {
plain,
part,
part32,
reduce,
k,
f16a,
},
n,
} => {
let tile = gemm3_tile(t, *n as usize);
let wgs = (*n as usize).div_ceil(tile.1) * t.div_ceil(tile.0);
let plain_ok = if *n <= 512 {
wgs >= 80
} else {
gemm3_grid_ok(t, tile, wgs)
};
if plain_ok {
let pl = if *f16a {
self.kernels
.gemm3_f16a_pipeline_tile(tile)
.ok_or_else(|| anyhow::anyhow!("f16a plan on a non-f16 adapter"))?
} else {
self.kernels.gemm3_pipeline_tile(self.use_f16, tile)?
};
(
(
pl,
&plain[gemm3_tier(tile)],
n.div_ceil(tile.1 as u32),
t32.div_ceil(tile.0 as u32),
1,
if *f16a { "gemm3f16a" } else { "gemm3" },
),
None,
)
} else {
let (use32, z) = gemm3_sk_plan(t, *n as usize, *k as usize);
let pl = match (*f16a, use32) {
(true, false) => self
.kernels
.gemm3_sk_f16a
.as_ref()
.ok_or_else(|| anyhow::anyhow!("f16a plan on a non-f16 adapter"))?,
(true, true) => self
.kernels
.gemm3_sk32_f16a
.as_ref()
.ok_or_else(|| anyhow::anyhow!("f16a plan on a non-f16 adapter"))?,
(false, false) => self.kernels.gemm3_sk_pipeline(self.use_f16)?,
(false, true) => self.kernels.gemm3_sk32_pipeline(self.use_f16)?,
};
let (pbg, rows) = if use32 {
(part32, GEMM3_SK_TILE32.0 as u32)
} else {
(part, GEMM3_SK_TILE.0 as u32)
};
anyhow::ensure!(
z as usize * t * (*n as usize) <= SK_PART_F32,
"split-K partials overflow: z={z} t={t} n={n} — the occupancy knee \
should make this impossible"
);
(
(
pl,
pbg,
n.div_ceil(GEMM3_SK_TILE.1 as u32),
t32.div_ceil(rows),
z,
"gemm3sk",
),
Some((
&self.kernels.gemm3_sk_reduce,
reduce,
(t32 * n).div_ceil(256),
1,
1,
"skreduce",
)),
)
}
}
EStep::Gemm { bg, n } => (
match bg {
GemmBg::F16A { tiles, coop } => {
let wgs128 = (*n as usize).div_ceil(128) * t.div_ceil(128);
let m_pad = t.div_ceil(128) * 128;
if let (Some(bg), true, true) =
(coop.as_ref(), wgs128 >= 48, m_pad <= self.max_tokens)
{
(
self.kernels
.gemm4_f16w
.as_ref()
.ok_or_else(|| anyhow::anyhow!("coop arm without kernel"))?,
bg,
n.div_ceil(128),
(m_pad as u32) / 128,
1,
"gemm4f16",
)
} else {
let tile = gemm3_tile(t, *n as usize);
(
self.kernels.gemm3_f16a_pipeline_tile(tile).ok_or_else(|| {
anyhow::anyhow!("f16a plan on a non-f16 adapter")
})?,
&tiles[gemm3_tier(tile)],
n.div_ceil(tile.1 as u32),
t32.div_ceil(tile.0 as u32),
1,
"gemm3f16a",
)
}
}
GemmBg::V1(bg) => (
self.kernels.gemm_pipeline(self.use_f16)?,
bg,
n.div_ceil(16),
t32.div_ceil(16),
1,
"gemm1",
),
GemmBg::V2(bgs) => {
let tile = gemm2_tile(t, *n as usize);
(
self.kernels.gemm2_pipeline_tile(self.use_f16, tile)?,
&bgs[gemm2_tier(tile)],
n.div_ceil(tile.1 as u32),
t32.div_ceil(tile.0 as u32),
1,
"gemm2",
)
}
GemmBg::V3(bgs) => {
let tile = gemm3_tile(t, *n as usize);
(
self.kernels.gemm3_pipeline_tile(self.use_f16, tile)?,
&bgs[gemm3_tier(tile)],
n.div_ceil(tile.1 as u32),
t32.div_ceil(tile.0 as u32),
1,
"gemm3",
)
}
GemmBg::V3Sk { .. } => unreachable!("handled by the arm above"),
},
None,
),
EStep::Ln { bg } => ((&self.kernels.layernorm, bg, t32, 1, 1, "ln"), None),
EStep::Rope { bg } => ((&self.kernels.rope, bg, t32, 1, 1, "rope"), None),
EStep::QkNorm { bg } => ((&self.kernels.qk_norm, bg, t32, 1, 1, "qknorm"), None),
EStep::Glu { bg } => ((&self.kernels.glu, bg, t32, 1, 1, "glu"), None),
EStep::Attn { bg } => (
(&self.kernels.attn, bg, t32, cfg.n_heads as u32, 1, "attn"),
None,
),
EStep::Attn4 { bg } => (
(&self.kernels.attn4, bg, t32, cfg.n_heads as u32, 1, "attn4"),
None,
),
EStep::AttnRb { bg } => (
(
&self.kernels.attn_rb4,
bg,
t32.div_ceil(4),
cfg.n_heads as u32,
1,
"attnRB",
),
None,
),
EStep::AttnDisent { bg } => (
(
&self.kernels.attn_disent,
bg,
t32,
cfg.n_heads as u32,
1,
"disent",
),
None,
),
EStep::AttnDisentRb { bg } => (
(
&self.kernels.attn_disent_rb4,
bg,
t32.div_ceil(4),
cfg.n_heads as u32,
1,
"disentRB",
),
None,
),
EStep::Add { bg } => (
(
&self.kernels.add,
bg,
((t * h) as u32).div_ceil(256),
1,
1,
"add",
),
None,
),
EStep::Copy { bg } => (
(
&self.kernels.copy,
bg,
((t * h) as u32).div_ceil(256),
1,
1,
"copy",
),
None,
),
EStep::Conv { bg } => ((&self.kernels.conv, bg, t32, 1, 1, "conv"), None),
EStep::L2Rows { bg } => ((&self.kernels.l2norm, bg, t32, 1, 1, "l2"), None),
})
}
pub fn profile_encode(
&mut self,
ctx: &GpuCtx,
batch: &EncBatch,
) -> Result<Vec<(String, f64, u32)>> {
anyhow::ensure!(ctx.timestamps, "adapter lacks TIMESTAMP_QUERY");
let t = self.dispatch(ctx, batch)?; let ndisp: usize = self
.steps
.iter()
.map(|s| match self.resolve_step(s, t) {
Ok((_, Some(_))) => 2,
_ => 1,
})
.sum();
let nq = (ndisp * 2) as u32;
let qs = ctx.device.create_query_set(&wgpu::QuerySetDescriptor {
label: Some("enc_prof"),
ty: wgpu::QueryType::Timestamp,
count: nq,
});
let qbuf = ctx.device.create_buffer(&wgpu::BufferDescriptor {
label: Some("enc_prof_resolve"),
size: u64::from(nq) * 8,
usage: wgpu::BufferUsages::QUERY_RESOLVE | wgpu::BufferUsages::COPY_SRC,
mapped_at_creation: false,
});
let qstage = ctx.device.create_buffer(&wgpu::BufferDescriptor {
label: Some("enc_prof_stage"),
size: u64::from(nq) * 8,
usage: wgpu::BufferUsages::MAP_READ | wgpu::BufferUsages::COPY_DST,
mapped_at_creation: false,
});
let mut enc = ctx
.device
.create_command_encoder(&wgpu::CommandEncoderDescriptor::default());
let mut labels: Vec<(&'static str, u32)> = Vec::with_capacity(ndisp);
{
let mut qi = 0u32;
for step in &self.steps {
let n_of = |s: &EStep| match s {
EStep::Gemm { n, .. } => *n,
_ => 0,
};
let (d, d2) = self.resolve_step(step, t)?;
for (pl, bg, gx, gy, gz, label) in std::iter::once(d).chain(d2) {
let mut p = enc.begin_compute_pass(&wgpu::ComputePassDescriptor {
label: None,
timestamp_writes: Some(wgpu::ComputePassTimestampWrites {
query_set: &qs,
beginning_of_pass_write_index: Some(qi * 2),
end_of_pass_write_index: Some(qi * 2 + 1),
}),
});
p.set_pipeline(pl);
p.set_bind_group(0, bg, &[]);
p.dispatch_workgroups(gx, gy, gz);
labels.push((label, n_of(step)));
qi += 1;
}
}
}
enc.resolve_query_set(&qs, 0..nq, &qbuf, 0);
enc.copy_buffer_to_buffer(&qbuf, 0, &qstage, 0, u64::from(nq) * 8);
ctx.queue.submit([enc.finish()]);
let slice = qstage.slice(..);
let (tx, rx) = std::sync::mpsc::channel();
slice.map_async(wgpu::MapMode::Read, move |r| {
let _ = tx.send(r);
});
let _ = ctx.device.poll(wgpu::PollType::wait_indefinitely());
rx.recv()
.map_err(|_| anyhow::anyhow!("prof staging dropped"))?
.map_err(|e| anyhow::anyhow!("map_async: {e:?}"))?;
let raw: Vec<u64> =
bytemuck::cast_slice(&slice.get_mapped_range().expect("mapped range")).to_vec();
qstage.unmap();
let mut agg: std::collections::BTreeMap<(&'static str, u32), (f64, u32)> =
std::collections::BTreeMap::new();
for (qi, (label, n)) in labels.iter().enumerate() {
let dt =
raw[qi * 2 + 1].saturating_sub(raw[qi * 2]) as f64 * f64::from(ctx.ts_period) / 1e3;
let e = agg.entry((label, *n)).or_insert((0.0, 0));
e.0 += dt;
e.1 += 1;
}
let mut rows: Vec<(String, f64, u32)> = agg
.into_iter()
.map(|((label, n), (us, cnt))| {
let name = if n > 0 {
format!("{label} n={n}")
} else {
label.to_string()
};
(name, us, cnt)
})
.collect();
rows.sort_by(|a, b| b.1.partial_cmp(&a.1).unwrap_or(std::cmp::Ordering::Equal));
Ok(rows)
}
fn read_staged(&self, ctx: &GpuCtx, len: usize) -> Result<Vec<f32>> {
let slice = self.staging.slice(..(len * 4) as u64);
let (tx, rx) = std::sync::mpsc::channel();
slice.map_async(wgpu::MapMode::Read, move |r| {
let _ = tx.send(r);
});
if ctx.spin_poll {
loop {
let r = ctx.device.poll(wgpu::PollType::Poll);
if rx.try_recv().is_ok() {
break;
}
if let Ok(status) = &r
&& status.is_queue_empty()
{
rx.recv()
.map_err(|_| anyhow::anyhow!("map_async callback dropped"))?
.map_err(|e| anyhow::anyhow!("map_async: {e:?}"))?;
break;
}
std::hint::spin_loop();
}
} else {
let _ = ctx.device.poll(wgpu::PollType::wait_indefinitely());
rx.recv()
.map_err(|_| anyhow::anyhow!("map_async callback dropped"))?
.map_err(|e| anyhow::anyhow!("map_async: {e:?}"))?;
}
let data = slice.get_mapped_range().expect("mapped range");
let out: Vec<f32> = bytemuck::cast_slice(&data).to_vec();
drop(data);
self.staging.unmap();
Ok(out)
}
}
struct PlanBuilder<'a> {
ctx: &'a GpuCtx,
kernels: &'a EncKernels,
cfg: &'a EncoderConfig,
use_f16: bool,
eps: f32,
emb_in: wgpu::Buffer,
cur: wgpu::Buffer,
alt: wgpu::Buffer,
qb: wgpu::Buffer,
kb: wgpu::Buffer,
vb: wgpu::Buffer,
cb: wgpu::Buffer,
mid: wgpu::Buffer,
glu: wgpu::Buffer,
pooled: wgpu::Buffer,
ptb: Option<wgpu::Buffer>,
pool_src_seq_buf: wgpu::Buffer,
seq_starts_buf: wgpu::Buffer,
seq_of_buf: wgpu::Buffer,
valid_buf: wgpu::Buffer,
steps: Vec<EStep>,
bg_pool: Option<wgpu::BindGroup>,
bg_l2: Option<wgpu::BindGroup>,
metas_tokens: Vec<wgpu::Buffer>,
metas_sk: Vec<(wgpu::Buffer, u32, u32)>,
metas_elems: Vec<wgpu::Buffer>,
weights: Vec<EncGpuLinear>,
norms: Vec<EncGpuNorm>,
scratch: Vec<wgpu::Buffer>,
sk_part: Option<wgpu::Buffer>,
max_tokens: usize,
}
impl<'a> PlanBuilder<'a> {
fn new(
ctx: &'a GpuCtx,
kernels: &'a EncKernels,
cfg: &'a EncoderConfig,
max_tokens: usize,
use_f16: bool,
) -> Self {
let h = cfg.hidden;
let mut up_width = match cfg.mlp_kind {
MlpKind::Dense { .. } => cfg.intermediate,
MlpKind::Glu { .. } => 2 * cfg.intermediate,
};
if cfg.layer_is_attn.iter().any(|a| !a) {
up_width = up_width.max(3 * h);
}
let qw = (cfg.n_heads * cfg.head_dim).max(h);
let zeros = |n: usize| vec![0f32; n];
Self {
eps: cfg.eps,
emb_in: ctx.storage(&zeros(max_tokens * h)),
cur: ctx.storage(&zeros(max_tokens * h)),
alt: ctx.storage(&zeros(max_tokens * h)),
qb: ctx.storage(&zeros(max_tokens * qw)),
kb: ctx.storage(&zeros(max_tokens * qw)),
vb: ctx.storage(&zeros(max_tokens * qw)),
cb: ctx.storage(&zeros(max_tokens * qw)),
mid: ctx.storage(&zeros(max_tokens * up_width)),
glu: ctx.storage(&zeros(max_tokens * cfg.intermediate)),
pooled: ctx.storage(&zeros(max_tokens * h)),
ptb: None,
pool_src_seq_buf: ctx.storage_bytes(&[0u8; 8]),
seq_starts_buf: ctx.storage_bytes(&vec![0u8; (max_tokens + 1) * 4]),
seq_of_buf: ctx.storage_bytes(&vec![0u8; max_tokens * 4]),
valid_buf: ctx.storage_bytes(bytemuck::cast_slice::<u32, u8>(&vec![1u32; max_tokens])),
steps: Vec::new(),
bg_pool: None,
bg_l2: None,
metas_tokens: Vec::new(),
metas_sk: Vec::new(),
metas_elems: Vec::new(),
weights: Vec::new(),
norms: Vec::new(),
scratch: Vec::new(),
sk_part: None,
max_tokens,
ctx,
kernels,
cfg,
use_f16,
}
}
fn upload_linear(&mut self, w: &[f32], b: Option<Vec<f32>>, n: u32, k: u32) -> Result<usize> {
anyhow::ensure!(
w.len() == (n as usize) * (k as usize),
"weight shape {} != {n}×{k}",
w.len()
);
let v3 = gemm3_eligible(n as usize);
let sk = (!v3 && gemm3_sk_band(n as usize)) || (v3 && gemm3_smalln_sk(n as usize));
let f16a = (v3 || sk) && f16a_enabled() && self.ctx.f16;
let wt_store;
let w = if v3 || sk {
let (n, k) = (n as usize, k as usize);
let mut wt = vec![0f32; n * k];
for nn in 0..n {
for kk in 0..k {
wt[kk * n + nn] = w[nn * k + kk];
}
}
wt_store = wt;
&wt_store[..]
} else {
w
};
let mat = if f16a || self.use_f16 {
EncMatBuf::F16(self.ctx.storage_bytes(&f32_to_f16_bytes(w)))
} else {
EncMatBuf::F32(self.ctx.storage(w))
};
let b_host = if f16a { b.clone() } else { None };
self.weights.push(EncGpuLinear {
w: mat,
b: b.map(|bv| self.ctx.storage(&bv)),
n,
k,
v3,
sk,
f16a,
b_host,
});
Ok(self.weights.len() - 1)
}
fn sk_part_buf(&mut self) -> wgpu::Buffer {
if self.sk_part.is_none() {
self.sk_part = Some(self.ctx.storage(&vec![0f32; SK_PART_F32]));
}
self.sk_part.clone().expect("just filled")
}
fn push_gemm(&mut self, wi: usize, x: &wgpu::Buffer, y: &wgpu::Buffer, act: Option<Act>) {
let lw = &self.weights[wi];
let (n, k, v3, sk, has_b) = (lw.n, lw.k, lw.v3, lw.sk, lw.b.is_some());
let lw_f16a = lw.f16a;
let wbuf = match &lw.w {
EncMatBuf::F16(b) | EncMatBuf::F32(b) => b.clone(),
};
let bbuf = lw.b.clone().unwrap_or_else(|| self.ctx.storage(&[0.0]));
let flags = u32::from(has_b) | (act_code(act) << 8);
let meta = uni(self.ctx, bytemuck::cast_slice(&[0u32, n, k, flags]));
let bufs = [x, &wbuf, &bbuf, y];
let bg = if lw_f16a && sk {
let part = self.sk_part_buf();
let plain = GEMM3_TILES
.iter()
.map(|&tl| {
make_bg(
self.ctx,
self.kernels
.gemm3_f16a_pipeline_tile(tl)
.expect("f16a implies SHADER_F16 (checked at upload)"),
&bufs,
&meta,
)
})
.collect();
let part_bg = make_bg(
self.ctx,
self.kernels
.gemm3_sk_f16a
.as_ref()
.expect("f16a implies SHADER_F16 (checked at upload)"),
&[x, &wbuf, &bbuf, &part],
&meta,
);
let part32_bg = make_bg(
self.ctx,
self.kernels
.gemm3_sk32_f16a
.as_ref()
.expect("f16a implies SHADER_F16 (checked at upload)"),
&[x, &wbuf, &bbuf, &part],
&meta,
);
let skn = gemm3_sk_chunks(k as usize); let rmeta = uni(self.ctx, bytemuck::cast_slice(&[0u32, n, flags, skn]));
let reduce = make_bg(
self.ctx,
&self.kernels.gemm3_sk_reduce,
&[&part, &bbuf, y],
&rmeta,
);
self.metas_sk.push((rmeta.clone(), n, k));
self.metas_tokens.push(rmeta);
GemmBg::V3Sk {
plain,
part: part_bg,
part32: part32_bg,
reduce,
k,
f16a: true,
}
} else if lw_f16a {
let set = GEMM3_TILES
.iter()
.map(|&tl| {
make_bg(
self.ctx,
self.kernels
.gemm3_f16a_pipeline_tile(tl)
.expect("f16a implies SHADER_F16 (checked at upload)"),
&bufs,
&meta,
)
})
.collect();
let coop = match (&self.kernels.gemm4_f16w, act, &self.weights[wi].b_host) {
(Some(pl), None, bh) => {
let bias_vec = bh.clone().unwrap_or_else(|| vec![0f32; n as usize]);
let mut b8 = Vec::with_capacity(8 * n as usize);
for _ in 0..8 {
b8.extend_from_slice(&bias_vec);
}
let b8buf = self.ctx.storage(&b8);
let bg = make_bg(self.ctx, pl, &[x, &wbuf, &b8buf, y], &meta);
self.scratch.push(b8buf);
Some(bg)
}
_ => None,
};
GemmBg::F16A { tiles: set, coop }
} else if sk {
let part = self.sk_part_buf();
let plain = v3_bgs(self.ctx, self.kernels, self.use_f16, &bufs, &meta)
.expect("checked at load");
let part_bg = make_bg(
self.ctx,
self.kernels
.gemm3_sk_pipeline(self.use_f16)
.expect("checked at load"),
&[x, &wbuf, &bbuf, &part],
&meta,
);
let part32_bg = make_bg(
self.ctx,
self.kernels
.gemm3_sk32_pipeline(self.use_f16)
.expect("checked at load"),
&[x, &wbuf, &bbuf, &part],
&meta,
);
let skn = gemm3_sk_chunks(k as usize); let rmeta = uni(self.ctx, bytemuck::cast_slice(&[0u32, n, flags, skn]));
let reduce = make_bg(
self.ctx,
&self.kernels.gemm3_sk_reduce,
&[&part, &bbuf, y],
&rmeta,
);
self.metas_sk.push((rmeta.clone(), n, k));
self.metas_tokens.push(rmeta); GemmBg::V3Sk {
plain,
part: part_bg,
part32: part32_bg,
reduce,
k,
f16a: false,
}
} else if v3 {
GemmBg::V3(
v3_bgs(self.ctx, self.kernels, self.use_f16, &bufs, &meta)
.expect("checked at load"),
)
} else if gemm2_enabled() {
GemmBg::V2(
v2_bgs(self.ctx, self.kernels, self.use_f16, &bufs, &meta)
.expect("checked at load"),
)
} else {
GemmBg::V1(make_bg(
self.ctx,
self.kernels
.gemm_pipeline(self.use_f16)
.expect("checked at load"),
&bufs,
&meta,
))
};
if !has_b {
self.scratch.push(bbuf);
}
self.steps.push(EStep::Gemm { bg, n });
self.metas_tokens.push(meta);
}
fn push_ln(
&mut self,
ni: usize,
x: &wgpu::Buffer,
res: Option<&wgpu::Buffer>,
out: &wgpu::Buffer,
) {
let n = &self.norms[ni];
let rms = matches!(self.cfg.norm_kind, NormKind::RmsNorm);
let flags =
u32::from(res.is_some()) | (u32::from(n.b.is_some()) << 1) | (u32::from(rms) << 2);
let meta = uni(
self.ctx,
bytemuck::cast_slice(&[self.cfg.hidden as u32, flags, self.eps.to_bits(), 0]),
);
let res_buf = res.unwrap_or(x); let b_placeholder;
let bbuf = match &n.b {
Some(b) => b,
None => {
b_placeholder = self.ctx.storage(&[0.0]);
self.scratch.push(b_placeholder.clone());
self.scratch.last().expect("just pushed")
}
};
let bg = make_bg(
self.ctx,
&self.kernels.layernorm,
&[x, res_buf, &n.w, bbuf, out],
&meta,
);
self.steps.push(EStep::Ln { bg });
self.scratch.push(meta); }
fn push_rope(&mut self, x: &wgpu::Buffer, theta: f32) {
self.push_rope_heads(x, theta, self.cfg.n_heads);
}
fn push_rope_heads(&mut self, x: &wgpu::Buffer, theta: f32, heads: usize) {
let meta = uni(
self.ctx,
bytemuck::cast_slice(&[
0u32,
heads as u32,
self.cfg.head_dim as u32,
theta.to_bits(),
]),
);
let bg = make_bg(
self.ctx,
&self.kernels.rope,
&[x, &self.seq_starts_buf, &self.seq_of_buf],
&meta,
);
self.steps.push(EStep::Rope { bg });
self.scratch.push(meta);
}
fn push_qk_norm(&mut self, x: &wgpu::Buffer, heads: usize, w: Vec<f32>) {
let wb = self.ctx.storage(&w);
let meta = uni(
self.ctx,
bytemuck::cast_slice(&[
0u32,
heads as u32,
self.cfg.head_dim as u32,
self.eps.to_bits(),
]),
);
let bg = make_bg(self.ctx, &self.kernels.qk_norm, &[x, &wb], &meta);
self.steps.push(EStep::QkNorm { bg });
self.scratch.push(wb);
self.scratch.push(meta);
}
fn push_attn(&mut self, window: u32) {
self.push_attn_src(window, None);
}
fn push_attn_src(&mut self, window: u32, packed: Option<&wgpu::Buffer>) {
let mode = match self.cfg.attn_mask {
MaskKind::Bidirectional => 0u32,
MaskKind::Causal => 1u32,
};
let meta = uni(
self.ctx,
bytemuck::cast_slice(&[
0u32,
self.cfg.n_heads as u32,
self.cfg.head_dim as u32,
mode,
window,
self.cfg.n_kv_heads as u32,
u32::from(packed.is_some()),
0,
]),
);
let (q, k, v) = match packed {
Some(b) => (b, b, b),
None => (&self.qb, &self.kb, &self.vb),
};
let v4 = self.cfg.head_dim.is_multiple_of(4)
&& std::env::var("OSFKB_ENC_ATTN_V4").ok().as_deref() != Some("0");
let rb = v4
&& self.cfg.head_dim <= 64
&& std::env::var("OSFKB_ENC_ATTN_RB").ok().as_deref() != Some("0");
let pl = if rb {
&self.kernels.attn_rb4
} else if v4 {
&self.kernels.attn4
} else {
&self.kernels.attn
};
let bg = make_bg(
self.ctx,
pl,
&[
q,
k,
v,
&self.cb,
&self.seq_starts_buf,
&self.seq_of_buf,
&self.valid_buf,
],
&meta,
);
self.steps.push(if rb {
EStep::AttnRb { bg }
} else if v4 {
EStep::Attn4 { bg }
} else {
EStep::Attn { bg }
});
self.metas_tokens.push(meta);
}
fn push_attn_disent(&mut self, pos_k: &wgpu::Buffer, pos_q: &wgpu::Buffer) {
self.push_attn_disent_src(pos_k, pos_q, None);
}
fn push_attn_disent_src(
&mut self,
pos_k: &wgpu::Buffer,
pos_q: &wgpu::Buffer,
packed: Option<&wgpu::Buffer>,
) {
let rel = self
.cfg
.rel_attn
.expect("push_attn_disent on a non-DeBERTa config");
let meta = uni(
self.ctx,
bytemuck::cast_slice(&[
0u32,
self.cfg.n_heads as u32,
self.cfg.head_dim as u32,
rel.span as u32,
rel.max_rel as u32,
rel.scale_factor() as u32,
u32::from(rel.c2p),
u32::from(rel.p2c) | (u32::from(packed.is_some()) << 1),
]),
);
let rb = self.cfg.head_dim.is_multiple_of(4)
&& self.cfg.head_dim <= 64
&& std::env::var("OSFKB_ENC_ATTN_RB").ok().as_deref() != Some("0");
let pl = if rb {
&self.kernels.attn_disent_rb4
} else {
&self.kernels.attn_disent
};
let (q, k, v) = match packed {
Some(b) => (b, b, b),
None => (&self.qb, &self.kb, &self.vb),
};
let bg = make_bg(
self.ctx,
pl,
&[
q,
k,
v,
&self.cb,
&self.seq_starts_buf,
&self.seq_of_buf,
&self.valid_buf,
pos_k,
pos_q,
],
&meta,
);
self.steps.push(if rb {
EStep::AttnDisentRb { bg }
} else {
EStep::AttnDisent { bg }
});
self.metas_tokens.push(meta);
}
fn push_glu(&mut self, act: Act) {
let meta = uni(
self.ctx,
bytemuck::cast_slice(&[self.cfg.intermediate as u32, act_code(Some(act)), 0, 0]),
);
let bg = make_bg(self.ctx, &self.kernels.glu, &[&self.mid, &self.glu], &meta);
self.steps.push(EStep::Glu { bg });
self.scratch.push(meta);
}
fn push_copy(&mut self, dst: &wgpu::Buffer, src: &wgpu::Buffer) {
let meta = uni(self.ctx, bytemuck::cast_slice(&[0u32, 0, 0, 0]));
let bg = make_bg(self.ctx, &self.kernels.copy, &[dst, src], &meta);
self.steps.push(EStep::Copy { bg });
self.metas_elems.push(meta);
}
fn push_conv(&mut self, bcx: &wgpu::Buffer, conv_w: Vec<f32>, y: &wgpu::Buffer) {
let wb = self.ctx.storage(&conv_w);
let meta = uni(
self.ctx,
bytemuck::cast_slice(&[self.cfg.hidden as u32, self.cfg.conv_l as u32, 0, 0]),
);
let bg = make_bg(
self.ctx,
&self.kernels.conv,
&[
bcx,
&wb,
y,
&self.seq_starts_buf,
&self.seq_of_buf,
&self.valid_buf,
],
&meta,
);
self.steps.push(EStep::Conv { bg });
self.scratch.push(wb);
self.scratch.push(meta);
}
fn push_l2_rows(&mut self, buf: &wgpu::Buffer, dim: usize) {
let meta = uni(self.ctx, bytemuck::cast_slice(&[dim as u32, 0, 0, 0]));
let bg = make_bg(self.ctx, &self.kernels.l2norm, &[buf], &meta);
self.steps.push(EStep::L2Rows { bg });
self.scratch.push(meta);
}
fn push_add(&mut self, dst: &wgpu::Buffer, src: &wgpu::Buffer) {
let meta = uni(self.ctx, bytemuck::cast_slice(&[0u32, 0, 0, 0]));
let bg = make_bg(self.ctx, &self.kernels.add, &[dst, src], &meta);
self.steps.push(EStep::Add { bg });
self.metas_elems.push(meta);
}
fn push_pool(&mut self, src: &wgpu::Buffer) {
let h = self.cfg.hidden as u32;
let meta = uni(self.ctx, bytemuck::cast_slice(&[h, 0, 0, 0]));
let pool_pl = match self.cfg.pooling {
crate::pooling::Pooling::LastToken => &self.kernels.last_pool,
crate::pooling::Pooling::Cls | crate::pooling::Pooling::CrossEncoder { .. } => {
&self.kernels.cls_pool
}
_ => &self.kernels.mean_pool,
};
self.bg_pool = Some(make_bg(
self.ctx,
pool_pl,
&[src, &self.pooled, &self.seq_starts_buf],
&meta,
));
self.scratch.push(meta);
let meta2 = uni(self.ctx, bytemuck::cast_slice(&[h, 0, 0, 0]));
self.bg_l2 = Some(make_bg(
self.ctx,
&self.kernels.l2norm,
&[&self.pooled],
&meta2,
));
self.scratch.push(meta2);
}
fn norm_from(&mut self, w: Vec<f32>, b: Option<Vec<f32>>) -> usize {
self.norms.push(EncGpuNorm {
w: self.ctx.storage(&w),
b: b.map(|bv| self.ctx.storage(&bv)),
});
self.norms.len() - 1
}
#[allow(clippy::type_complexity)]
fn build_bert_family(
&mut self,
st: &LazySt,
) -> Result<(Vec<f32>, Option<Vec<f32>>, Option<Vec<f32>>)> {
let cfg = self.cfg;
let prefix = ["", "bert.", "roberta."]
.into_iter()
.find(|p| st.has(&format!("{p}embeddings.word_embeddings.weight")))
.context("word embeddings not found under known prefixes ('', 'bert.', 'roberta.')")?;
let get = |name: &str| -> Result<Vec<f32>> { st.tensor_f32(&format!("{prefix}{name}")) };
let word = get("embeddings.word_embeddings.weight")?;
anyhow::ensure!(
word.len() == cfg.vocab * cfg.hidden,
"word embedding shape {} != vocab {} × hidden {}",
word.len(),
cfg.vocab,
cfg.hidden
);
anyhow::ensure!(
matches!(cfg.norm_kind, NormKind::LayerNorm { bias: true }),
"BERT-family norms are biased LayerNorm"
);
let (h, im) = (cfg.hidden as u32, cfg.intermediate as u32);
let mlp_act = match cfg.mlp_kind {
MlpKind::Dense { act, .. } => act,
MlpKind::Glu { .. } => anyhow::bail!("BERT family MLPs are dense"),
};
let jina = st.has(&format!("{prefix}encoder.layers.0.mixer.Wqkv.weight"));
let emb_ln = if jina {
self.norm_from(get("emb_ln.weight")?, Some(get("emb_ln.bias")?))
} else {
self.norm_from(
get("embeddings.LayerNorm.weight")?,
Some(get("embeddings.LayerNorm.bias")?),
)
};
let emb_in = self.emb_in.clone();
let cur = self.cur.clone();
self.push_ln(emb_ln, &emb_in, None, &cur);
let (alt, qb, kb, vb, cb, mid) = (
self.alt.clone(),
self.qb.clone(),
self.kb.clone(),
self.vb.clone(),
self.cb.clone(),
self.mid.clone(),
);
let qkv_fused = cfg.intermediate >= 3 * cfg.hidden
&& std::env::var("OSFKB_ENC_QKV_FUSED").ok().as_deref() != Some("0");
let mid_qkv = self.mid.clone();
for i in 0..cfg.n_layers {
let p = if jina {
format!("encoder.layers.{i}")
} else {
format!("encoder.layer.{i}")
};
let lin = |b: &mut Self, wn: String, bn: String, n: u32, k: u32| -> Result<usize> {
b.upload_linear(
&st.tensor_f32(&format!("{prefix}{wn}"))?,
Some(st.tensor_f32(&format!("{prefix}{bn}"))?),
n,
k,
)
};
let qkv = if jina {
let w = st.tensor_f32(&format!("{prefix}{p}.mixer.Wqkv.weight"))?;
let bias = st.tensor_f32(&format!("{prefix}{p}.mixer.Wqkv.bias"))?;
anyhow::ensure!(
w.len() == 3 * cfg.hidden * cfg.hidden,
"{p}.mixer.Wqkv shape {} != 3·{}·{}",
w.len(),
cfg.hidden,
cfg.hidden
);
Some(self.upload_linear(&w, Some(bias), 3 * h, h)?)
} else if qkv_fused {
let mut w = Vec::with_capacity(3 * cfg.hidden * cfg.hidden);
let mut bias = Vec::with_capacity(3 * cfg.hidden);
for part in ["query", "key", "value"] {
w.extend(st.tensor_f32(&format!("{prefix}{p}.attention.self.{part}.weight"))?);
bias.extend(st.tensor_f32(&format!("{prefix}{p}.attention.self.{part}.bias"))?);
}
Some(self.upload_linear(&w, Some(bias), 3 * h, h)?)
} else {
None
};
let (q, k, v) = if qkv.is_some() {
(0, 0, 0) } else {
(
lin(
self,
format!("{p}.attention.self.query.weight"),
format!("{p}.attention.self.query.bias"),
h,
h,
)?,
lin(
self,
format!("{p}.attention.self.key.weight"),
format!("{p}.attention.self.key.bias"),
h,
h,
)?,
lin(
self,
format!("{p}.attention.self.value.weight"),
format!("{p}.attention.self.value.bias"),
h,
h,
)?,
)
};
let (o_n, up_n, down_n, ln1_n, ln2_n) = if jina {
("mixer.out_proj", "mlp.fc1", "mlp.fc2", "norm1", "norm2")
} else {
(
"attention.output.dense",
"intermediate.dense",
"output.dense",
"attention.output.LayerNorm",
"output.LayerNorm",
)
};
let o = lin(
self,
format!("{p}.{o_n}.weight"),
format!("{p}.{o_n}.bias"),
h,
h,
)?;
let up = lin(
self,
format!("{p}.{up_n}.weight"),
format!("{p}.{up_n}.bias"),
im,
h,
)?;
let down = lin(
self,
format!("{p}.{down_n}.weight"),
format!("{p}.{down_n}.bias"),
h,
im,
)?;
let ln1 = self.norm_from(
get(&format!("{p}.{ln1_n}.weight"))?,
Some(get(&format!("{p}.{ln1_n}.bias"))?),
);
let ln2 = self.norm_from(
get(&format!("{p}.{ln2_n}.weight"))?,
Some(get(&format!("{p}.{ln2_n}.bias"))?),
);
if let Some(qkv) = qkv {
self.push_gemm(qkv, &cur, &mid_qkv, None);
self.push_attn_src(cfg.layer_window[i], Some(&mid_qkv));
} else {
self.push_gemm(q, &cur, &qb, None);
self.push_gemm(k, &cur, &kb, None);
self.push_gemm(v, &cur, &vb, None);
self.push_attn(cfg.layer_window[i]);
}
self.push_gemm(o, &cb, &qb, None);
self.push_ln(ln1, &qb, Some(&cur), &alt);
self.push_gemm(up, &alt, &mid, Some(mlp_act));
self.push_gemm(down, &mid, &kb, None);
self.push_ln(ln2, &kb, Some(&alt), &cur);
}
self.push_pool(&cur);
Ok((
word,
Some(get("embeddings.position_embeddings.weight")?),
if cfg.type_vocab > 0 {
Some(get("embeddings.token_type_embeddings.weight")?)
} else {
None
},
))
}
#[allow(clippy::type_complexity)]
fn build_deberta(
&mut self,
st: &LazySt,
) -> Result<(Vec<f32>, Option<Vec<f32>>, Option<Vec<f32>>)> {
let cfg = self.cfg;
let rel = cfg.rel_attn.context("DeBERTa config without rel_attn")?;
let prefix = ["", "deberta.", "token_rep_layer.bert_layer.model."]
.into_iter()
.find(|p| st.has(&format!("{p}embeddings.word_embeddings.weight")))
.context(
"word embeddings not found under known prefixes ('', 'deberta.', \
'token_rep_layer.bert_layer.model.')",
)?;
let get = |name: &str| -> Result<Vec<f32>> { st.tensor_f32(&format!("{prefix}{name}")) };
let word = get("embeddings.word_embeddings.weight")?;
anyhow::ensure!(
word.len() == cfg.vocab * cfg.hidden,
"word embedding shape {} != vocab {} × hidden {}",
word.len(),
cfg.vocab,
cfg.hidden
);
let (h, im) = (cfg.hidden as u32, cfg.intermediate as u32);
let mlp_act = match cfg.mlp_kind {
MlpKind::Dense { act, .. } => act,
MlpKind::Glu { .. } => anyhow::bail!("DeBERTa MLPs are dense"),
};
let rows = 2 * rel.span;
let hs = cfg.hidden;
let mut rel_emb = get("encoder.rel_embeddings.weight")?;
anyhow::ensure!(
rel_emb.len() >= rows * hs,
"rel_embeddings has {} rows, need 2·span = {rows}",
rel_emb.len() / hs.max(1)
);
rel_emb.truncate(rows * hs);
cpu_layer_norm_rows(
&mut rel_emb,
hs,
&get("encoder.LayerNorm.weight")?,
&get("encoder.LayerNorm.bias")?,
cfg.eps,
);
let emb_ln = self.norm_from(
get("embeddings.LayerNorm.weight")?,
Some(get("embeddings.LayerNorm.bias")?),
);
let emb_in = self.emb_in.clone();
let cur = self.cur.clone();
self.push_ln(emb_ln, &emb_in, None, &cur);
let (alt, qb, kb, vb, cb, mid) = (
self.alt.clone(),
self.qb.clone(),
self.kb.clone(),
self.vb.clone(),
self.cb.clone(),
self.mid.clone(),
);
for i in 0..cfg.n_layers {
let p = format!("encoder.layer.{i}");
let getw = |n: String| -> Result<Vec<f32>> { st.tensor_f32(&format!("{prefix}{n}")) };
let qw = getw(format!("{p}.attention.self.query_proj.weight"))?;
let qb_w = getw(format!("{p}.attention.self.query_proj.bias"))?;
let kw = getw(format!("{p}.attention.self.key_proj.weight"))?;
let kb_w = getw(format!("{p}.attention.self.key_proj.bias"))?;
let pos_k = self
.ctx
.storage(&cpu_linear(&rel_emb, &kw, Some(&kb_w), hs, hs));
let pos_q = self
.ctx
.storage(&cpu_linear(&rel_emb, &qw, Some(&qb_w), hs, hs));
let vw = getw(format!("{p}.attention.self.value_proj.weight"))?;
let vb_w = getw(format!("{p}.attention.self.value_proj.bias"))?;
let fused = cfg.intermediate >= 3 * cfg.hidden
&& std::env::var("OSFKB_ENC_QKV_FUSED").ok().as_deref() != Some("0");
let qkv = if fused {
let mut w = Vec::with_capacity(3 * cfg.hidden * cfg.hidden);
let mut bias = Vec::with_capacity(3 * cfg.hidden);
for (ww, bb) in [(&qw, &qb_w), (&kw, &kb_w), (&vw, &vb_w)] {
w.extend_from_slice(ww);
bias.extend_from_slice(bb);
}
Some(self.upload_linear(&w, Some(bias), 3 * h, h)?)
} else {
None
};
let (q, k, v) = if qkv.is_some() {
(0, 0, 0) } else {
(
self.upload_linear(&qw, Some(qb_w.clone()), h, h)?,
self.upload_linear(&kw, Some(kb_w.clone()), h, h)?,
self.upload_linear(&vw, Some(vb_w.clone()), h, h)?,
)
};
let o = self.upload_linear(
&getw(format!("{p}.attention.output.dense.weight"))?,
Some(getw(format!("{p}.attention.output.dense.bias"))?),
h,
h,
)?;
let up = self.upload_linear(
&getw(format!("{p}.intermediate.dense.weight"))?,
Some(getw(format!("{p}.intermediate.dense.bias"))?),
im,
h,
)?;
let down = self.upload_linear(
&getw(format!("{p}.output.dense.weight"))?,
Some(getw(format!("{p}.output.dense.bias"))?),
h,
im,
)?;
let ln1 = self.norm_from(
get(&format!("{p}.attention.output.LayerNorm.weight"))?,
Some(get(&format!("{p}.attention.output.LayerNorm.bias"))?),
);
let ln2 = self.norm_from(
get(&format!("{p}.output.LayerNorm.weight"))?,
Some(get(&format!("{p}.output.LayerNorm.bias"))?),
);
if let Some(qkv) = qkv {
self.push_gemm(qkv, &cur, &mid, None);
let mid_qkv = self.mid.clone();
self.push_attn_disent_src(&pos_k, &pos_q, Some(&mid_qkv));
} else {
self.push_gemm(q, &cur, &qb, None);
self.push_gemm(k, &cur, &kb, None);
self.push_gemm(v, &cur, &vb, None);
self.push_attn_disent(&pos_k, &pos_q);
}
self.push_gemm(o, &cb, &qb, None);
self.push_ln(ln1, &qb, Some(&cur), &alt);
self.push_gemm(up, &alt, &mid, Some(mlp_act));
self.push_gemm(down, &mid, &kb, None);
self.push_ln(ln2, &kb, Some(&alt), &cur);
}
self.push_pool(&cur);
Ok((
word,
None, if cfg.type_vocab > 0 {
Some(get("embeddings.token_type_embeddings.weight")?)
} else {
None
},
))
}
#[allow(clippy::type_complexity)]
fn build_lfm2_colbert(
&mut self,
st: &LazySt,
dir: &Path,
) -> Result<(Vec<f32>, Option<Vec<f32>>, Option<Vec<f32>>)> {
let cfg = self.cfg;
let get = |name: &str| -> Result<Vec<f32>> { st.tensor_f32(name) };
let word = get("embed_tokens.weight")?;
anyhow::ensure!(
word.len() == cfg.vocab * cfg.hidden,
"embed_tokens shape {} != vocab {} × hidden {}",
word.len(),
cfg.vocab,
cfg.hidden
);
let (hd, im) = (cfg.head_dim as u32, cfg.intermediate as u32);
let hu = cfg.hidden as u32;
let (qh, kvh) = (cfg.n_heads as u32 * hd, cfg.n_kv_heads as u32 * hd);
let theta = match cfg.pos_kind {
PosKind::Rope { theta, .. } => theta,
PosKind::Learned { .. } => anyhow::bail!("LFM2 uses rotary positions"),
};
let glu_act = match cfg.mlp_kind {
MlpKind::Glu { act } => act,
MlpKind::Dense { .. } => anyhow::bail!("LFM2 MLPs are GLU"),
};
let dim = match cfg.pooling {
Pooling::PerToken { dim } => dim,
other => anyhow::bail!("ColBERT plan requires PerToken pooling, got {other:?}"),
};
let emb_in = self.emb_in.clone();
let cur = self.cur.clone();
self.push_copy(&cur, &emb_in);
let (alt, qb, kb, vb, cb, mid, glu) = (
self.alt.clone(),
self.qb.clone(),
self.kb.clone(),
self.vb.clone(),
self.cb.clone(),
self.mid.clone(),
self.glu.clone(),
);
for i in 0..cfg.n_layers {
let p = format!("layers.{i}");
let an = self.norm_from(get(&format!("{p}.operator_norm.weight"))?, None);
self.push_ln(an, &cur, None, &alt);
if cfg.layer_is_attn[i] {
let q = self.upload_linear(
&get(&format!("{p}.self_attn.q_proj.weight"))?,
None,
qh,
hu,
)?;
let k = self.upload_linear(
&get(&format!("{p}.self_attn.k_proj.weight"))?,
None,
kvh,
hu,
)?;
let v = self.upload_linear(
&get(&format!("{p}.self_attn.v_proj.weight"))?,
None,
kvh,
hu,
)?;
let o = self.upload_linear(
&get(&format!("{p}.self_attn.out_proj.weight"))?,
None,
hu,
qh,
)?;
let q_norm_w = get(&format!("{p}.self_attn.q_layernorm.weight"))?;
let k_norm_w = get(&format!("{p}.self_attn.k_layernorm.weight"))?;
self.push_gemm(q, &alt, &qb, None);
self.push_gemm(k, &alt, &kb, None);
self.push_gemm(v, &alt, &vb, None);
self.push_qk_norm(&qb, cfg.n_heads, q_norm_w);
self.push_qk_norm(&kb, cfg.n_kv_heads, k_norm_w);
self.push_rope_heads(&qb, theta, cfg.n_heads);
self.push_rope_heads(&kb, theta, cfg.n_kv_heads);
self.push_attn(0);
self.push_gemm(o, &cb, &alt, None);
} else {
let in_proj = self.upload_linear(
&get(&format!("{p}.conv.in_proj.weight"))?,
None,
3 * hu,
hu,
)?;
let conv_w = get(&format!("{p}.conv.conv.weight"))?;
anyhow::ensure!(
conv_w.len() == cfg.hidden * cfg.conv_l,
"{p} conv taps {} != hidden × conv_l",
conv_w.len()
);
let out_proj =
self.upload_linear(&get(&format!("{p}.conv.out_proj.weight"))?, None, hu, hu)?;
self.push_gemm(in_proj, &alt, &mid, None);
self.push_conv(&mid, conv_w, &cb);
self.push_gemm(out_proj, &cb, &alt, None);
}
self.push_add(&cur, &alt);
let mn = self.norm_from(get(&format!("{p}.ffn_norm.weight"))?, None);
self.push_ln(mn, &cur, None, &alt);
let gate = get(&format!("{p}.feed_forward.w1.weight"))?;
let up = get(&format!("{p}.feed_forward.w3.weight"))?;
let mut wi = gate;
wi.extend_from_slice(&up);
let wi = self.upload_linear(&wi, None, 2 * im, hu)?;
let wo =
self.upload_linear(&get(&format!("{p}.feed_forward.w2.weight"))?, None, hu, im)?;
self.push_gemm(wi, &alt, &mid, None);
self.push_glu(glu_act);
self.push_gemm(wo, &glu, &alt, None);
self.push_add(&cur, &alt);
}
let fin = self.norm_from(get("embedding_norm.weight")?, None);
self.push_ln(fin, &cur, None, &alt);
let dense_st = LazySt::open(&dir.join("1_Dense"))?;
let proj =
self.upload_linear(&dense_st.tensor_f32("linear.weight")?, None, dim as u32, hu)?;
let ptb = self.ctx.storage(&vec![
0f32;
self.pooled.size() as usize / (4 * self.cfg.hidden)
* dim
]);
self.push_gemm(proj, &alt, &ptb, None);
self.push_l2_rows(&ptb, dim);
self.ptb = Some(ptb);
Ok((word, None, None))
}
#[allow(clippy::type_complexity)]
fn build_nomic(
&mut self,
st: &LazySt,
) -> Result<(Vec<f32>, Option<Vec<f32>>, Option<Vec<f32>>)> {
let cfg = self.cfg;
let get = |name: &str| -> Result<Vec<f32>> { st.tensor_f32(name) };
let word = get("embeddings.word_embeddings.weight")?;
anyhow::ensure!(
word.len() == cfg.vocab * cfg.hidden,
"word embedding shape {} != vocab {} × hidden {}",
word.len(),
cfg.vocab,
cfg.hidden
);
let hu = cfg.hidden as u32;
let (im, h) = (cfg.intermediate as u32, cfg.hidden);
let theta = match cfg.pos_kind {
PosKind::Rope { theta, .. } => theta,
PosKind::Learned { .. } => anyhow::bail!("nomic uses rotary positions"),
};
let glu_act = match cfg.mlp_kind {
MlpKind::Glu { act } => act,
MlpKind::Dense { .. } => anyhow::bail!("nomic MLPs are GLU"),
};
let emb_ln = self.norm_from(get("emb_ln.weight")?, Some(get("emb_ln.bias")?));
let emb_in = self.emb_in.clone();
let cur = self.cur.clone();
self.push_ln(emb_ln, &emb_in, None, &cur);
let (alt, qb, kb, vb, cb, mid, glu) = (
self.alt.clone(),
self.qb.clone(),
self.kb.clone(),
self.vb.clone(),
self.cb.clone(),
self.mid.clone(),
self.glu.clone(),
);
for i in 0..cfg.n_layers {
let p = format!("encoder.layers.{i}");
let wqkv = get(&format!("{p}.attn.Wqkv.weight"))?;
anyhow::ensure!(wqkv.len() == 3 * h * h, "{p}.attn.Wqkv shape");
let q = self.upload_linear(&wqkv[..h * h], None, hu, hu)?;
let k = self.upload_linear(&wqkv[h * h..2 * h * h], None, hu, hu)?;
let v = self.upload_linear(&wqkv[2 * h * h..], None, hu, hu)?;
let o =
self.upload_linear(&get(&format!("{p}.attn.out_proj.weight"))?, None, hu, hu)?;
let gate = get(&format!("{p}.mlp.fc12.weight"))?;
let lin_half = get(&format!("{p}.mlp.fc11.weight"))?;
let mut wi = gate;
wi.extend_from_slice(&lin_half);
let wi = self.upload_linear(&wi, None, 2 * im, hu)?;
let wo = self.upload_linear(&get(&format!("{p}.mlp.fc2.weight"))?, None, hu, im)?;
let ln1 = self.norm_from(
get(&format!("{p}.norm1.weight"))?,
Some(get(&format!("{p}.norm1.bias"))?),
);
let ln2 = self.norm_from(
get(&format!("{p}.norm2.weight"))?,
Some(get(&format!("{p}.norm2.bias"))?),
);
self.push_gemm(q, &cur, &qb, None);
self.push_gemm(k, &cur, &kb, None);
self.push_gemm(v, &cur, &vb, None);
self.push_rope(&qb, theta);
self.push_rope(&kb, theta);
self.push_attn(cfg.layer_window[i]);
self.push_gemm(o, &cb, &qb, None);
self.push_ln(ln1, &qb, Some(&cur), &alt);
self.push_gemm(wi, &alt, &mid, None);
self.push_glu(glu_act);
self.push_gemm(wo, &glu, &kb, None);
self.push_ln(ln2, &kb, Some(&alt), &cur);
}
self.push_pool(&cur);
Ok((
word,
None,
Some(get("embeddings.token_type_embeddings.weight")?),
))
}
#[allow(clippy::type_complexity)]
fn build_qwen3_embed(
&mut self,
st: &LazySt,
) -> Result<(Vec<f32>, Option<Vec<f32>>, Option<Vec<f32>>)> {
let cfg = self.cfg;
let get = |name: &str| -> Result<Vec<f32>> { st.tensor_f32(&format!("model.{name}")) };
let word = get("embed_tokens.weight")?;
anyhow::ensure!(
word.len() == cfg.vocab * cfg.hidden,
"embed_tokens shape {} != vocab {} × hidden {}",
word.len(),
cfg.vocab,
cfg.hidden
);
let (hd, im) = (cfg.head_dim as u32, cfg.intermediate as u32);
let hu = cfg.hidden as u32;
let (qh, kvh) = (cfg.n_heads as u32 * hd, cfg.n_kv_heads as u32 * hd);
let glu_act = match cfg.mlp_kind {
MlpKind::Glu { act } => act,
MlpKind::Dense { .. } => anyhow::bail!("Qwen3 embedder MLPs are GLU"),
};
let emb_in = self.emb_in.clone();
let cur = self.cur.clone();
self.push_copy(&cur, &emb_in);
let (alt, qb, kb, vb, cb, mid, glu) = (
self.alt.clone(),
self.qb.clone(),
self.kb.clone(),
self.vb.clone(),
self.cb.clone(),
self.mid.clone(),
self.glu.clone(),
);
for i in 0..cfg.n_layers {
let p = format!("layers.{i}");
let q =
self.upload_linear(&get(&format!("{p}.self_attn.q_proj.weight"))?, None, qh, hu)?;
let k = self.upload_linear(
&get(&format!("{p}.self_attn.k_proj.weight"))?,
None,
kvh,
hu,
)?;
let v = self.upload_linear(
&get(&format!("{p}.self_attn.v_proj.weight"))?,
None,
kvh,
hu,
)?;
let o =
self.upload_linear(&get(&format!("{p}.self_attn.o_proj.weight"))?, None, hu, qh)?;
let gate = get(&format!("{p}.mlp.gate_proj.weight"))?;
let up = get(&format!("{p}.mlp.up_proj.weight"))?;
let mut wi = gate;
wi.extend_from_slice(&up);
let wi = self.upload_linear(&wi, None, 2 * im, hu)?;
let wo =
self.upload_linear(&get(&format!("{p}.mlp.down_proj.weight"))?, None, hu, im)?;
let an = self.norm_from(get(&format!("{p}.input_layernorm.weight"))?, None);
let mn = self.norm_from(get(&format!("{p}.post_attention_layernorm.weight"))?, None);
let q_norm_w = get(&format!("{p}.self_attn.q_norm.weight"))?;
let k_norm_w = get(&format!("{p}.self_attn.k_norm.weight"))?;
self.push_ln(an, &cur, None, &alt);
self.push_gemm(q, &alt, &qb, None);
self.push_gemm(k, &alt, &kb, None);
self.push_gemm(v, &alt, &vb, None);
self.push_qk_norm(&qb, cfg.n_heads, q_norm_w);
self.push_qk_norm(&kb, cfg.n_kv_heads, k_norm_w);
let theta = match cfg.pos_kind {
PosKind::Rope { theta, .. } => theta,
PosKind::Learned { .. } => anyhow::bail!("Qwen3 embedder uses rotary positions"),
};
self.push_rope_heads(&qb, theta, cfg.n_heads);
self.push_rope_heads(&kb, theta, cfg.n_kv_heads);
self.push_attn(0);
self.push_gemm(o, &cb, &alt, None);
self.push_add(&cur, &alt);
self.push_ln(mn, &cur, None, &alt);
self.push_gemm(wi, &alt, &mid, None);
self.push_glu(glu_act);
self.push_gemm(wo, &glu, &alt, None);
self.push_add(&cur, &alt);
}
let fin = self.norm_from(get("norm.weight")?, None);
self.push_ln(fin, &cur, None, &alt);
self.push_pool(&alt);
Ok((word, None, None))
}
#[allow(clippy::type_complexity)]
fn build_siglip_vision(
&mut self,
st: &LazySt,
) -> Result<(Vec<f32>, Option<Vec<f32>>, Option<Vec<f32>>)> {
let cfg = self.cfg;
let get =
|name: &str| -> Result<Vec<f32>> { st.tensor_f32(&format!("vision_model.{name}")) };
let (h, im) = (cfg.hidden as u32, cfg.intermediate as u32);
let mlp_act = match cfg.mlp_kind {
MlpKind::Dense { act, .. } => act,
MlpKind::Glu { .. } => anyhow::bail!("SigLIP vision MLPs are dense"),
};
let emb_in = self.emb_in.clone();
let cur = self.cur.clone();
self.push_copy(&cur, &emb_in);
let (alt, qb, kb, vb, cb, mid) = (
self.alt.clone(),
self.qb.clone(),
self.kb.clone(),
self.vb.clone(),
self.cb.clone(),
self.mid.clone(),
);
for i in 0..cfg.n_layers {
let p = format!("encoder.layers.{i}");
let biased = |b: &mut Self, name: &str, n: u32, k: u32| -> Result<usize> {
b.upload_linear(
&st.tensor_f32(&format!("vision_model.{p}.{name}.weight"))?,
Some(st.tensor_f32(&format!("vision_model.{p}.{name}.bias"))?),
n,
k,
)
};
let q = biased(self, "self_attn.q_proj", h, h)?;
let k = biased(self, "self_attn.k_proj", h, h)?;
let v = biased(self, "self_attn.v_proj", h, h)?;
let o = biased(self, "self_attn.out_proj", h, h)?;
let up = biased(self, "mlp.fc1", im, h)?;
let down = biased(self, "mlp.fc2", h, im)?;
let ln1 = self.norm_from(
get(&format!("{p}.layer_norm1.weight"))?,
Some(get(&format!("{p}.layer_norm1.bias"))?),
);
let ln2 = self.norm_from(
get(&format!("{p}.layer_norm2.weight"))?,
Some(get(&format!("{p}.layer_norm2.bias"))?),
);
self.push_ln(ln1, &cur, None, &alt);
self.push_gemm(q, &alt, &qb, None);
self.push_gemm(k, &alt, &kb, None);
self.push_gemm(v, &alt, &vb, None);
self.push_attn(0); self.push_gemm(o, &cb, &alt, None);
self.push_add(&cur, &alt);
self.push_ln(ln2, &cur, None, &alt);
self.push_gemm(up, &alt, &mid, Some(mlp_act));
self.push_gemm(down, &mid, &alt, None);
self.push_add(&cur, &alt);
}
let fin = self.norm_from(
get("post_layernorm.weight")?,
Some(get("post_layernorm.bias")?),
);
self.push_ln(fin, &cur, None, &alt);
self.push_copy(&cur, &alt);
Ok((Vec::new(), None, None))
}
#[allow(clippy::type_complexity)]
fn build_modernbert(
&mut self,
st: &LazySt,
) -> Result<(Vec<f32>, Option<Vec<f32>>, Option<Vec<f32>>)> {
let cfg = self.cfg;
let prefix = ["", "model."]
.into_iter()
.find(|p| st.has(&format!("{p}embeddings.tok_embeddings.weight")))
.context("ModernBERT tok_embeddings not found under known prefixes ('', 'model.')")?;
let get = |name: &str| -> Result<Vec<f32>> { st.tensor_f32(&format!("{prefix}{name}")) };
let has = |name: &str| st.has(&format!("{prefix}{name}"));
let word = get("embeddings.tok_embeddings.weight")?;
anyhow::ensure!(
word.len() == cfg.vocab * cfg.hidden,
"tok embedding shape {} != vocab {} × hidden {}",
word.len(),
cfg.vocab,
cfg.hidden
);
anyhow::ensure!(!cfg.qkv_bias, "ModernBERT is bias-free");
let (theta, local_theta) = match cfg.pos_kind {
PosKind::Rope { theta, local_theta } => (theta, local_theta),
PosKind::Learned { .. } => anyhow::bail!("ModernBERT uses rotary positions"),
};
let glu_act = match cfg.mlp_kind {
MlpKind::Glu { act } => act,
MlpKind::Dense { .. } => anyhow::bail!("ModernBERT MLPs are GLU"),
};
let norm_of = |b: &mut Self, name: &str| -> Result<usize> {
let bias_name = format!("{}.bias", name.trim_end_matches(".weight"));
let bias = if has(&bias_name) {
Some(get(&bias_name)?)
} else {
None
};
Ok(b.norm_from(get(name)?, bias))
};
let hu = cfg.hidden as u32;
let (im, h) = (cfg.intermediate as u32, cfg.hidden);
let emb_ln = norm_of(self, "embeddings.norm.weight")?;
let emb_in = self.emb_in.clone();
let cur = self.cur.clone();
self.push_ln(emb_ln, &emb_in, None, &cur);
let (alt, qb, kb, vb, cb, mid, glu) = (
self.alt.clone(),
self.qb.clone(),
self.kb.clone(),
self.vb.clone(),
self.cb.clone(),
self.mid.clone(),
self.glu.clone(),
);
for i in 0..cfg.n_layers {
let p = format!("layers.{i}");
let wqkv = get(&format!("{p}.attn.Wqkv.weight"))?;
anyhow::ensure!(
wqkv.len() == 3 * h * h,
"{p}.attn.Wqkv shape {} != 3·{h}·{h}",
wqkv.len()
);
let q = self.upload_linear(&wqkv[..h * h], None, hu, hu)?;
let k = self.upload_linear(&wqkv[h * h..2 * h * h], None, hu, hu)?;
let v = self.upload_linear(&wqkv[2 * h * h..], None, hu, hu)?;
let o = self.upload_linear(&get(&format!("{p}.attn.Wo.weight"))?, None, hu, hu)?;
let wi = self.upload_linear(&get(&format!("{p}.mlp.Wi.weight"))?, None, 2 * im, hu)?;
let wo = self.upload_linear(&get(&format!("{p}.mlp.Wo.weight"))?, None, hu, im)?;
let window = cfg.layer_window[i];
let layer_theta = if window > 0 { local_theta } else { theta };
let attn_src = if i == 0 && cfg.skip_first_attn_norm {
anyhow::ensure!(
!has(&format!("{p}.attn_norm.weight")),
"layer 0 attn_norm present but config says skip — checkpoint mismatch"
);
&cur
} else {
let an = norm_of(self, &format!("{p}.attn_norm.weight"))?;
self.push_ln(an, &cur, None, &alt);
&alt
};
self.push_gemm(q, attn_src, &qb, None);
self.push_gemm(k, attn_src, &kb, None);
self.push_gemm(v, attn_src, &vb, None);
self.push_rope(&qb, layer_theta);
self.push_rope(&kb, layer_theta);
self.push_attn(window);
self.push_gemm(o, &cb, &alt, None);
self.push_add(&cur, &alt);
let mn = norm_of(self, &format!("{p}.mlp_norm.weight"))?;
self.push_ln(mn, &cur, None, &alt);
self.push_gemm(wi, &alt, &mid, None);
self.push_glu(glu_act);
self.push_gemm(wo, &glu, &alt, None);
self.push_add(&cur, &alt);
}
let fin = norm_of(self, "final_norm.weight")?;
self.push_ln(fin, &cur, None, &alt);
self.push_pool(&alt);
Ok((word, None, None))
}
}