use anyhow::Result;
use crate::deepencoder_gpu::{flash_reg_f16a_src, flash_reg_kv16_src, flash_reg_sg_src, flash_reg_src, u32x4, ADDACT, LAYERNORM};
use crate::encoder::{enc_gemm3_f16a_src, enc_gemm3_src};
use crate::forward::{make_bg, pipeline, uni};
use crate::vision::{rope_tables_2d, ImagePatches, VisionConfig};
use crate::vision_glm::GlmVisionTower;
use crate::GpuCtx;
const RMSNORM: &str = r#"
@group(0) @binding(0) var<storage, read> x: array<f32>;
@group(0) @binding(1) var<storage, read> w: array<f32>;
@group(0) @binding(2) var<storage, read_write> y: array<f32>;
@group(0) @binding(3) var<uniform> d: vec4<u32>; // (rows, c, eps_bits, _)
@compute @workgroup_size(64)
fn main(@builtin(global_invocation_id) gid: vec3<u32>, @builtin(num_workgroups) nwg: vec3<u32>) {
let rows = d.x; let c = d.y; let eps = bitcast<f32>(d.z);
let r = gid.x + gid.y * nwg.x * 64u; if (r >= rows) { return; }
var ms = 0.0;
for (var j = 0u; j < c; j = j + 1u) { let v = x[r * c + j]; ms = ms + v * v; }
let inv = 1.0 / sqrt(ms / f32(c) + eps);
for (var j = 0u; j < c; j = j + 1u) { y[r * c + j] = x[r * c + j] * inv * w[j]; }
}
"#;
const QKNORM_ROPE: &str = r#"
@group(0) @binding(0) var<storage, read_write> qkv: array<f32>; // [n, 3*hid]
@group(0) @binding(1) var<storage, read> qw: array<f32>; // [64]
@group(0) @binding(2) var<storage, read> kw: array<f32>; // [64]
@group(0) @binding(3) var<storage, read> cs: array<f32>; // [n, 64]
@group(0) @binding(4) var<storage, read> sn: array<f32>; // [n, 64]
@group(0) @binding(5) var<uniform> d: vec4<u32>; // (n, heads, eps_bits, hid)
@compute @workgroup_size(64)
fn main(@builtin(global_invocation_id) gid: vec3<u32>, @builtin(num_workgroups) nwg: vec3<u32>) {
let n = d.x; let heads = d.y; let eps = bitcast<f32>(d.z); let hid = d.w;
let idx = gid.x + gid.y * nwg.x * 64u; if (idx >= n * heads) { return; }
let tok = idx / heads; let h = idx % heads;
let qb = tok * 3u * hid + h * 64u;
let kb = qb + hid;
var q: array<f32, 64>;
var k: array<f32, 64>;
var qs = 0.0; var ks = 0.0;
for (var j = 0u; j < 64u; j = j + 1u) {
q[j] = qkv[qb + j]; qs = qs + q[j] * q[j];
k[j] = qkv[kb + j]; ks = ks + k[j] * k[j];
}
let qi = 1.0 / sqrt(qs / 64.0 + eps);
let ki = 1.0 / sqrt(ks / 64.0 + eps);
for (var j = 0u; j < 64u; j = j + 1u) { q[j] = q[j] * qi * qw[j]; k[j] = k[j] * ki * kw[j]; }
// NEOX rope, half = 32: out[j] = x[j]·cos[j] + rot·sin[j], rot = -x[j+32] (j<32) else x[j-32]
let cb = tok * 64u;
for (var j = 0u; j < 64u; j = j + 1u) {
var rq = 0.0; var rk = 0.0;
if (j < 32u) { rq = -q[j + 32u]; rk = -k[j + 32u]; } else { rq = q[j - 32u]; rk = k[j - 32u]; }
qkv[qb + j] = q[j] * cs[cb + j] + rq * sn[cb + j];
qkv[kb + j] = k[j] * cs[cb + j] + rk * sn[cb + j];
}
}
"#;
const PACK16: &str = r#"
enable f16;
@group(0) @binding(0) var<storage, read> x: array<f32>;
@group(0) @binding(1) var<storage, read_write> y: array<f16>;
@group(0) @binding(2) var<uniform> d: vec4<u32>; // (len, _, _, _)
@compute @workgroup_size(64)
fn main(@builtin(global_invocation_id) gid: vec3<u32>, @builtin(num_workgroups) nwg: vec3<u32>) {
let len = d.x; let i = gid.x + gid.y * nwg.x * 64u; if (i >= len) { return; }
y[i] = f16(x[i]);
}
"#;
const MUL_SILU: &str = r#"
@group(0) @binding(0) var<storage, read_write> g: array<f32>;
@group(0) @binding(1) var<storage, read> u: array<f32>;
@group(0) @binding(2) var<uniform> d: vec4<u32>; // (len, _, _, _)
@compute @workgroup_size(64)
fn main(@builtin(global_invocation_id) gid: vec3<u32>, @builtin(num_workgroups) nwg: vec3<u32>) {
let len = d.x; let i = gid.x + gid.y * nwg.x * 64u; if (i >= len) { return; }
let v = g[i];
g[i] = (v / (1.0 + exp(-v))) * u[i];
}
"#;
const ACCUM: &str = r#"
@group(0) @binding(0) var<storage, read_write> y: array<f32>;
@group(0) @binding(1) var<storage, read> b: array<f32>;
@group(0) @binding(2) var<uniform> d: vec4<u32>; // (len, _, _, _)
@compute @workgroup_size(64)
fn main(@builtin(global_invocation_id) gid: vec3<u32>, @builtin(num_workgroups) nwg: vec3<u32>) {
let len = d.x; let i = gid.x + gid.y * nwg.x * 64u; if (i >= len) { return; }
y[i] = y[i] + b[i];
}
"#;
fn tiled_linear_t_src(f16: bool) -> String {
let (enable, wty) = if f16 { ("enable f16;", "f16") } else { ("", "f32") };
format!(r#"{enable}
@group(0) @binding(0) var<storage, read> x: array<f32>;
@group(0) @binding(1) var<storage, read> w: array<{wty}>; // [k, n]
@group(0) @binding(2) var<storage, read> b: array<f32>;
@group(0) @binding(3) var<storage, read_write> y: array<f32>;
@group(0) @binding(4) var<uniform> d: vec4<u32>; // (m, k, n, has_bias)
var<workgroup> xs: array<f32, 256>;
var<workgroup> ws: array<f32, 256>;
@compute @workgroup_size(16, 16)
fn main(@builtin(workgroup_id) wid: vec3<u32>, @builtin(local_invocation_id) lid: vec3<u32>) {{
let m = d.x; let k = d.y; let n = d.z;
let row = wid.x * 16u + lid.x;
let col = wid.y * 16u + lid.y;
var acc = 0.0;
let ntile = (k + 15u) / 16u;
for (var t = 0u; t < ntile; t = t + 1u) {{
let kx = t * 16u + lid.y;
xs[lid.x * 16u + lid.y] = select(0.0, x[row * k + kx], row < m && kx < k);
let kw = t * 16u + lid.x;
ws[lid.y * 16u + lid.x] = select(0.0, f32(w[kw * n + col]), col < n && kw < k);
workgroupBarrier();
for (var p = 0u; p < 16u; p = p + 1u) {{ acc = acc + xs[lid.x * 16u + p] * ws[lid.y * 16u + p]; }}
workgroupBarrier();
}}
if (row < m && col < n) {{
if (d.w == 1u) {{ acc = acc + b[col]; }}
y[row * n + col] = acc;
}}
}}
"#)
}
fn gemm3_tile_from_env() -> (usize, usize, usize) {
if let Ok(v) = std::env::var("OSFKB_GLM_GEMM3_TILE") {
let p: Vec<usize> = v.split('x').filter_map(|t| t.parse().ok()).collect();
if p.len() == 3 { return (p[0], p[1], p[2]); }
}
(128, 128, 16)
}
fn record(enc: &mut wgpu::CommandEncoder, pl: &wgpu::ComputePipeline, bg: &wgpu::BindGroup, threads: usize) {
let mut p = enc.begin_compute_pass(&wgpu::ComputePassDescriptor::default());
p.set_pipeline(pl);
p.set_bind_group(0, bg, &[]);
let wg = threads.div_ceil(64);
let gx = wg.min(65535) as u32;
let gy = wg.div_ceil(65535) as u32;
p.dispatch_workgroups(gx, gy, 1);
}
struct GpuLin {
wt: wgpu::Buffer,
b: Option<wgpu::Buffer>,
n: usize,
k: usize,
}
fn transpose(w: &[f32], n: usize, k: usize) -> Vec<f32> {
let mut t = vec![0f32; n * k];
for r in 0..n {
for c in 0..k {
t[c * n + r] = w[r * k + c];
}
}
t
}
impl GpuLin {
fn new(ctx: &GpuCtx, l: &crate::vision::Linear, f16: bool) -> Self {
let t = transpose(&l.w, l.n, l.k);
let wt = if f16 {
ctx.storage_bytes(&crate::weights::f32_to_f16_bytes(&t))
} else {
ctx.storage(&t)
};
Self {
wt,
b: l.b.as_deref().map(|b| ctx.storage(b)),
n: l.n,
k: l.k,
}
}
}
struct GpuBlock {
norm1_w: wgpu::Buffer,
qkv: GpuLin,
q_norm_w: wgpu::Buffer,
k_norm_w: wgpu::Buffer,
proj: GpuLin,
norm2_w: wgpu::Buffer,
gate: GpuLin,
up: GpuLin,
down: GpuLin,
}
pub struct GlmVisionGpu {
pub cfg: VisionConfig,
patch: GpuLin,
blocks: Vec<GpuBlock>,
post_ln_w: wgpu::Buffer,
downsample: GpuLin,
merger_proj: GpuLin,
merger_norm_w: wgpu::Buffer,
merger_norm_b: wgpu::Buffer,
merger_gate: GpuLin,
merger_up: GpuLin,
merger_down: GpuLin,
pl_gemm: wgpu::ComputePipeline,
pl_gemm3: wgpu::ComputePipeline,
gemm3_tile: (usize, usize),
pl_gemm3_f16a: Option<wgpu::ComputePipeline>,
pl_rms: wgpu::ComputePipeline,
pl_ln: wgpu::ComputePipeline,
pl_qkrope: wgpu::ComputePipeline,
pl_pack16: Option<wgpu::ComputePipeline>,
pl_flash: wgpu::ComputePipeline,
flash_rb: usize,
pl_addact: wgpu::ComputePipeline,
pl_mulsilu: wgpu::ComputePipeline,
pl_accum: wgpu::ComputePipeline,
}
impl GlmVisionGpu {
pub fn new(ctx: &GpuCtx, cpu: &GlmVisionTower) -> Result<Self> {
let rb = std::env::var("OSFKB_GLM_FLASH_RB")
.ok()
.and_then(|v| v.parse().ok())
.filter(|&r: &usize| r >= 1 && r <= 32)
.unwrap_or(16); let f16 = ctx.f16 && std::env::var("OSFKB_GLM_F16").as_deref() != Ok("0");
Self::new_with_opts(ctx, cpu, rb, f16)
}
pub fn new_with_rb(ctx: &GpuCtx, cpu: &GlmVisionTower, rb: usize) -> Result<Self> {
let f16 = ctx.f16 && std::env::var("OSFKB_GLM_F16").as_deref() != Ok("0");
Self::new_with_opts(ctx, cpu, rb, f16)
}
pub fn new_with_opts(ctx: &GpuCtx, cpu: &GlmVisionTower, rb: usize, f16: bool) -> Result<Self> {
let (bm, bn, bk) = gemm3_tile_from_env();
Self::new_full(ctx, cpu, rb, f16, bm, bn, bk)
}
pub fn new_full(ctx: &GpuCtx, cpu: &GlmVisionTower, rb: usize, f16: bool, bm: usize, bn: usize, bk: usize) -> Result<Self> {
let cfg = cpu.cfg.clone();
anyhow::ensure!(
cfg.head_dim() == 64,
"GPU GLM tower bakes head_dim 64, checkpoint has {}",
cfg.head_dim()
);
let blocks = cpu
.blocks
.iter()
.map(|b| GpuBlock {
norm1_w: ctx.storage(&b.norm1_w),
qkv: GpuLin::new(ctx, &b.qkv, f16),
q_norm_w: ctx.storage(&b.q_norm_w),
k_norm_w: ctx.storage(&b.k_norm_w),
proj: GpuLin::new(ctx, &b.proj, f16),
norm2_w: ctx.storage(&b.norm2_w),
gate: GpuLin::new(ctx, &b.gate, f16),
up: GpuLin::new(ctx, &b.up, f16),
down: GpuLin::new(ctx, &b.down, f16),
})
.collect();
Ok(Self {
patch: GpuLin::new(ctx, &cpu.patch, f16),
blocks,
post_ln_w: ctx.storage(&cpu.post_ln_w),
downsample: GpuLin::new(ctx, &cpu.downsample, f16),
merger_proj: GpuLin::new(ctx, &cpu.merger_proj, f16),
merger_norm_w: ctx.storage(&cpu.merger_post_norm.w),
merger_norm_b: ctx.storage(&cpu.merger_post_norm.b),
merger_gate: GpuLin::new(ctx, &cpu.merger_gate, f16),
merger_up: GpuLin::new(ctx, &cpu.merger_up, f16),
merger_down: GpuLin::new(ctx, &cpu.merger_down, f16),
pl_gemm: pipeline(ctx, "glmv_gemm", &tiled_linear_t_src(f16)),
pl_gemm3: pipeline(ctx, "glmv_gemm3", &enc_gemm3_src(f16, bm, bn, bk)),
gemm3_tile: (bm, bn),
pl_gemm3_f16a: (f16
&& std::env::var("OSFKB_GLM_F16A").as_deref() == Ok("1"))
.then(|| {
pipeline(ctx, "glmv_gemm3_f16a", &enc_gemm3_f16a_src(128, 64, 32))
}),
pl_rms: pipeline(ctx, "glmv_rms", RMSNORM),
pl_ln: pipeline(ctx, "glmv_ln", LAYERNORM),
pl_qkrope: pipeline(ctx, "glmv_qkrope", QKNORM_ROPE),
pl_pack16: (ctx.f16
&& std::env::var("OSFKB_GLM_FLASH_KV16").as_deref() == Ok("1"))
.then(|| pipeline(ctx, "glmv_pack16", PACK16)),
pl_flash: if ctx.f16
&& std::env::var("OSFKB_GLM_FLASH_KV16").as_deref() == Ok("1")
{
pipeline(ctx, "glmv_flash_kv16", &flash_reg_kv16_src(64, rb, false))
} else if ctx.f16
&& std::env::var("OSFKB_GLM_FLASH_F16A").as_deref() == Ok("1")
{
pipeline(ctx, "glmv_flash_f16a", &flash_reg_f16a_src(64, rb, false))
} else if ctx.subgroups
&& std::env::var("OSFKB_GLM_FLASH_SG").as_deref() != Ok("0")
{
pipeline(ctx, "glmv_flash_sg", &flash_reg_sg_src(64, rb, false))
} else {
pipeline(ctx, "glmv_flash", &flash_reg_src(64, rb, false))
},
flash_rb: rb,
pl_addact: pipeline(ctx, "glmv_addact", ADDACT),
pl_mulsilu: pipeline(ctx, "glmv_mulsilu", MUL_SILU),
pl_accum: pipeline(ctx, "glmv_accum", ACCUM),
cfg,
})
}
fn gemm_into(&self, ctx: &GpuCtx, enc: &mut wgpu::CommandEncoder, x: &wgpu::Buffer, l: &GpuLin, m: usize, y: &wgpu::Buffer) {
let (k, n) = (l.k, l.n);
let zero = ctx.storage(&[0.0]);
let bias = l.b.as_ref().unwrap_or(&zero);
let (bm, bn) = self.gemm3_tile;
let wgs3 = n.div_ceil(bn) * m.div_ceil(bm);
if let Some(pl) = &self.pl_gemm3_f16a
&& n.div_ceil(64) * m.div_ceil(128) >= 96
&& k % 4 == 0
&& n % 4 == 0
{
let flags = u32::from(l.b.is_some());
let meta = uni(ctx, &u32x4(m as u32, n as u32, k as u32, flags));
let bg = make_bg(ctx, pl, &[x, &l.wt, bias, y], &meta);
let mut p = enc.begin_compute_pass(&wgpu::ComputePassDescriptor::default());
p.set_pipeline(pl);
p.set_bind_group(0, &bg, &[]);
p.dispatch_workgroups((n as u32).div_ceil(64), (m as u32).div_ceil(128), 1);
return;
}
if wgs3 >= 96 && k % 4 == 0 && n % 4 == 0 {
let flags = u32::from(l.b.is_some()); let meta = uni(ctx, &u32x4(m as u32, n as u32, k as u32, flags));
let bg = make_bg(ctx, &self.pl_gemm3, &[x, &l.wt, bias, y], &meta);
let mut p = enc.begin_compute_pass(&wgpu::ComputePassDescriptor::default());
p.set_pipeline(&self.pl_gemm3);
p.set_bind_group(0, &bg, &[]);
p.dispatch_workgroups((n as u32).div_ceil(bn as u32), (m as u32).div_ceil(bm as u32), 1);
return;
}
let meta = uni(ctx, &u32x4(m as u32, k as u32, n as u32, u32::from(l.b.is_some())));
let bg = make_bg(ctx, &self.pl_gemm, &[x, &l.wt, bias, y], &meta);
let (gx, gy) = ((m as u32).div_ceil(16), (n as u32).div_ceil(16));
let mut p = enc.begin_compute_pass(&wgpu::ComputePassDescriptor::default());
p.set_pipeline(&self.pl_gemm);
p.set_bind_group(0, &bg, &[]);
p.dispatch_workgroups(gx, gy, 1);
}
fn rms_into(&self, ctx: &GpuCtx, enc: &mut wgpu::CommandEncoder, x: &wgpu::Buffer, rows: usize, c: usize, w: &wgpu::Buffer, y: &wgpu::Buffer) {
let meta = uni(ctx, &u32x4(rows as u32, c as u32, self.cfg.eps.to_bits(), 0));
let bg = make_bg(ctx, &self.pl_rms, &[x, w, y], &meta);
record(enc, &self.pl_rms, &bg, rows);
}
fn accum(&self, ctx: &GpuCtx, enc: &mut wgpu::CommandEncoder, y: &wgpu::Buffer, b: &wgpu::Buffer, len: usize) {
let meta = uni(ctx, &u32x4(len as u32, 0, 0, 0));
let bg = make_bg(ctx, &self.pl_accum, &[y, b], &meta);
record(enc, &self.pl_accum, &bg, len);
}
#[allow(dead_code)]
#[allow(dead_code)]
fn gemm(&self, ctx: &GpuCtx, enc: &mut wgpu::CommandEncoder, x: &wgpu::Buffer, l: &GpuLin, m: usize) -> wgpu::Buffer {
let (k, n) = (l.k, l.n);
let y = ctx.empty(m * n);
let zero = ctx.storage(&[0.0]);
let bias = l.b.as_ref().unwrap_or(&zero);
let meta = uni(ctx, &u32x4(m as u32, k as u32, n as u32, u32::from(l.b.is_some())));
let bg = make_bg(ctx, &self.pl_gemm, &[x, &l.wt, bias, &y], &meta);
let (gx, gy) = ((m as u32).div_ceil(16), (n as u32).div_ceil(16));
{
let mut p = enc.begin_compute_pass(&wgpu::ComputePassDescriptor::default());
p.set_pipeline(&self.pl_gemm);
p.set_bind_group(0, &bg, &[]);
p.dispatch_workgroups(gx, gy, 1);
}
y
}
fn rms(&self, ctx: &GpuCtx, enc: &mut wgpu::CommandEncoder, x: &wgpu::Buffer, rows: usize, c: usize, w: &wgpu::Buffer) -> wgpu::Buffer {
let y = ctx.empty(rows * c);
let meta = uni(ctx, &u32x4(rows as u32, c as u32, self.cfg.eps.to_bits(), 0));
let bg = make_bg(ctx, &self.pl_rms, &[x, w, &y], &meta);
record(enc, &self.pl_rms, &bg, rows);
y
}
fn addact(&self, ctx: &GpuCtx, enc: &mut wgpu::CommandEncoder, a: &wgpu::Buffer, b: Option<&wgpu::Buffer>, len: usize, act: u32) -> wgpu::Buffer {
let y = ctx.empty(len);
let zero = ctx.storage(&[0.0]);
let bb = b.unwrap_or(&zero);
let meta = uni(ctx, &u32x4(len as u32, act, u32::from(b.is_some()), 0));
let bg = make_bg(ctx, &self.pl_addact, &[a, bb, &y], &meta);
record(enc, &self.pl_addact, &bg, len);
y
}
fn run_one(&self, ctx: &GpuCtx, img: &ImagePatches) -> Result<Vec<f32>> {
let cfg = &self.cfg;
let (hid, heads, inter) = (cfg.hidden, cfg.heads, cfg.intermediate);
let n = img.num_patches();
anyhow::ensure!(img.patches.len() == n * cfg.patch_dim(), "patch buffer mismatch");
let mut enc = ctx
.device
.create_command_encoder(&wgpu::CommandEncoderDescriptor::default());
let x = ctx.empty(n * hid);
let normed = ctx.empty(n * hid);
let qkv = ctx.empty(n * 3 * hid);
let merged = ctx.empty(n * hid);
let tmp = ctx.empty(n * hid);
let g = ctx.empty(n * inter);
let u = ctx.empty(n * inter);
let patches = ctx.storage(&img.patches);
self.gemm_into(ctx, &mut enc, &patches, &self.patch, n, &x);
let (cos, sin) = rope_tables_2d(cfg, img.grid);
let cos_b = ctx.storage(&cos);
let sin_b = ctx.storage(&sin);
let qkv16 = self.pl_pack16.as_ref().map(|_| {
ctx.empty((n * 3 * hid).div_ceil(2))
});
let d_flash = uni(ctx, &u32x4(n as u32, heads as u32, 64, 1));
let scale = 1.0f32 / 8.0; let e_flash = uni(ctx, &u32x4(hid as u32, 1, scale.to_bits(), 0));
let chunk = std::env::var("OSFKB_GLM_TOWER_CHUNK")
.ok()
.and_then(|v| v.parse::<usize>().ok())
.unwrap_or(1);
macro_rules! flush {
() => {
let done = std::mem::replace(
&mut enc,
ctx.device
.create_command_encoder(&wgpu::CommandEncoderDescriptor::default()),
);
ctx.queue.submit(Some(done.finish()));
};
}
for (bi, blk) in self.blocks.iter().enumerate() {
if chunk > 0 && bi > 0 && bi.is_multiple_of(chunk) {
flush!();
}
let e = &mut enc;
self.rms_into(ctx, e, &x, n, hid, &blk.norm1_w, &normed);
self.gemm_into(ctx, e, &normed, &blk.qkv, n, &qkv);
let meta = uni(ctx, &u32x4(n as u32, heads as u32, cfg.eps.to_bits(), hid as u32));
let bg = make_bg(
ctx,
&self.pl_qkrope,
&[&qkv, &blk.q_norm_w, &blk.k_norm_w, &cos_b, &sin_b],
&meta,
);
record(e, &self.pl_qkrope, &bg, n * heads);
if let (Some(pk), Some(q16)) = (&self.pl_pack16, &qkv16) {
let meta = uni(ctx, &u32x4((n * 3 * hid) as u32, 0, 0, 0));
let bg = make_bg(ctx, pk, &[&qkv, q16], &meta);
record(e, pk, &bg, n * 3 * hid);
}
let fbg = ctx.device.create_bind_group(&wgpu::BindGroupDescriptor {
label: None,
layout: &self.pl_flash.get_bind_group_layout(0),
entries: &[
wgpu::BindGroupEntry { binding: 0, resource: qkv16.as_ref().unwrap_or(&qkv).as_entire_binding() },
wgpu::BindGroupEntry { binding: 3, resource: merged.as_entire_binding() },
wgpu::BindGroupEntry { binding: 4, resource: d_flash.as_entire_binding() },
wgpu::BindGroupEntry { binding: 5, resource: e_flash.as_entire_binding() },
],
});
let wgs = heads * n.div_ceil(self.flash_rb);
{
let mut p = e.begin_compute_pass(&wgpu::ComputePassDescriptor::default());
p.set_pipeline(&self.pl_flash);
p.set_bind_group(0, &fbg, &[]);
let (gx, gy) = ((wgs.min(65535)) as u32, wgs.div_ceil(65535) as u32);
p.dispatch_workgroups(gx, gy, 1);
}
self.gemm_into(ctx, e, &merged, &blk.proj, n, &tmp);
self.accum(ctx, e, &x, &tmp, n * hid);
self.rms_into(ctx, e, &x, n, hid, &blk.norm2_w, &normed);
self.gemm_into(ctx, e, &normed, &blk.gate, n, &g);
self.gemm_into(ctx, e, &normed, &blk.up, n, &u);
let meta = uni(ctx, &u32x4((n * inter) as u32, 0, 0, 0));
let bg = make_bg(ctx, &self.pl_mulsilu, &[&g, &u], &meta);
record(e, &self.pl_mulsilu, &bg, n * inter);
self.gemm_into(ctx, e, &g, &blk.down, n, &tmp);
self.accum(ctx, e, &x, &tmp, n * hid);
}
let e = &mut enc;
self.rms_into(ctx, e, &x, n, hid, &self.post_ln_w, &normed);
let unit = cfg.merge_unit();
anyhow::ensure!(n % unit == 0, "patch count {n} not a multiple of merge²");
let tokens = n / unit;
let oh = cfg.out_hidden;
let inner = self.merger_gate.n;
let ds = ctx.empty(tokens * oh);
self.gemm_into(ctx, e, &normed, &self.downsample, tokens, &ds);
let mp = ctx.empty(tokens * oh);
self.gemm_into(ctx, e, &ds, &self.merger_proj, tokens, &mp);
let ln = ctx.empty(tokens * oh);
let meta = uni(ctx, &u32x4(tokens as u32, oh as u32, 1e-5f32.to_bits(), 0));
let bg = make_bg(ctx, &self.pl_ln, &[&mp, &self.merger_norm_w, &self.merger_norm_b, &ln], &meta);
record(e, &self.pl_ln, &bg, tokens);
let gelu = self.addact(ctx, e, &ln, None, tokens * oh, 1); let mg = ctx.empty(tokens * inner);
let mu = ctx.empty(tokens * inner);
self.gemm_into(ctx, e, &gelu, &self.merger_gate, tokens, &mg);
self.gemm_into(ctx, e, &gelu, &self.merger_up, tokens, &mu);
let meta = uni(ctx, &u32x4((tokens * inner) as u32, 0, 0, 0));
let bg = make_bg(ctx, &self.pl_mulsilu, &[&mg, &mu], &meta);
record(e, &self.pl_mulsilu, &bg, tokens * inner);
let out = ctx.empty(tokens * oh);
self.gemm_into(ctx, e, &mg, &self.merger_down, tokens, &out);
ctx.queue.submit(Some(enc.finish()));
let _ = ctx.device.poll(wgpu::PollType::wait_indefinitely());
ctx.read(&out, tokens * oh)
}
pub fn forward(&self, ctx: &GpuCtx, images: &[ImagePatches]) -> Result<Vec<f32>> {
let mut out = Vec::new();
for img in images {
out.extend(self.run_one(ctx, img)?);
}
Ok(out)
}
}
#[cfg(test)]
mod tests {
use super::*;
use crate::vision::Linear;
use crate::vision_glm::GlmBlock;
fn fill(n: usize, seed: u32) -> Vec<f32> {
let mut s = seed.wrapping_add(1);
(0..n)
.map(|_| {
s = s.wrapping_mul(1664525).wrapping_add(1013904223);
((s >> 8) as f32 / 16_777_216.0 - 0.5) * 0.2
})
.collect()
}
fn lin(nn: usize, k: usize, bias: bool, seed: u32) -> Linear {
Linear::from_parts(
fill(nn * k, seed),
bias.then(|| fill(nn, seed + 1)),
nn,
k,
)
}
#[test]
fn gemm3_route_iso() {
let Ok(ctx) = GpuCtx::new() else { return };
let (m, k, n) = (1024usize, 1024usize, 3072usize);
let x = fill(m * k, 51);
let w = fill(n * k, 52);
let b = fill(n, 53);
let want = crate::deepencoder::linear(&x, m, k, n, &w, Some(&b));
let cpu_lin = Linear::from_parts(w, Some(b), n, k);
let l = GpuLin::new(&ctx, &cpu_lin, false);
let dummy_cfg = VisionConfig {
depth: 1, hidden: 64, heads: 1, intermediate: 64, out_hidden: 64, patch: 1, merge: 2,
temporal: 1, in_channels: 3, grid_side: 0, eps: 1e-5,
act: crate::encoder_weights::Act::Silu, rope_theta: 10_000.0, min_pixels: 0, max_pixels: usize::MAX,
};
let pl3 = pipeline(&ctx, "t_gemm3", &enc_gemm3_src(false, 128, 128, 16));
let xb = ctx.storage(&x);
let y = ctx.empty(m * n);
let bias = l.b.as_ref().expect("bias");
let meta = uni(&ctx, &u32x4(m as u32, n as u32, k as u32, 1));
let bg = make_bg(&ctx, &pl3, &[&xb, &l.wt, bias, &y], &meta);
let mut enc = ctx.device.create_command_encoder(&wgpu::CommandEncoderDescriptor::default());
{
let mut p = enc.begin_compute_pass(&wgpu::ComputePassDescriptor::default());
p.set_pipeline(&pl3);
p.set_bind_group(0, &bg, &[]);
p.dispatch_workgroups((n as u32).div_ceil(128), (m as u32).div_ceil(128), 1);
}
ctx.queue.submit(Some(enc.finish()));
let _ = ctx.device.poll(wgpu::PollType::wait_indefinitely());
let got = ctx.read(&y, m * n).expect("read");
let _ = dummy_cfg;
let mx = want.iter().zip(&got).map(|(a, b)| (a - b).abs()).fold(0f32, f32::max);
println!("enc_gemm3 vs CPU linear max_abs = {mx:.3e}");
assert!(mx < 1e-3, "gemm3 not iso: {mx:.3e}");
}
#[test]
fn glm_tower_cpu_gpu_iso() {
let Ok(ctx) = GpuCtx::new() else {
eprintln!("no wgpu adapter; skipping");
return;
};
let (heads, hid, inter, oh, depth) = (2usize, 128usize, 192usize, 96usize, 2usize);
let cfg = VisionConfig {
depth,
hidden: hid,
heads,
intermediate: inter,
out_hidden: oh,
patch: 4,
merge: 2,
temporal: 1,
in_channels: 3,
grid_side: 0,
eps: 1e-5,
act: crate::encoder_weights::Act::Silu,
rope_theta: 10_000.0,
min_pixels: 0,
max_pixels: usize::MAX,
};
let pd = cfg.patch_dim();
let blocks = (0..depth)
.map(|i| {
let s = 100 + (i as u32) * 20;
GlmBlock {
norm1_w: fill(hid, s).iter().map(|v| 1.0 + v).collect(),
qkv: lin(3 * hid, hid, true, s + 1),
q_norm_w: fill(64, s + 3).iter().map(|v| 1.0 + v).collect(),
k_norm_w: fill(64, s + 4).iter().map(|v| 1.0 + v).collect(),
proj: lin(hid, hid, true, s + 5),
norm2_w: fill(hid, s + 7).iter().map(|v| 1.0 + v).collect(),
gate: lin(inter, hid, true, s + 8),
up: lin(inter, hid, true, s + 10),
down: lin(hid, inter, true, s + 12),
}
})
.collect();
let cpu = GlmVisionTower {
cfg: cfg.clone(),
patch: lin(hid, pd, true, 1),
blocks,
post_ln_w: fill(hid, 3).iter().map(|v| 1.0 + v).collect(),
downsample: lin(oh, 4 * hid, true, 4),
merger_proj: lin(oh, oh, false, 6),
merger_post_norm: crate::vision::Norm {
w: fill(oh, 8).iter().map(|v| 1.0 + v).collect(),
b: fill(oh, 9),
},
merger_gate: lin(3 * oh, oh, false, 10),
merger_up: lin(3 * oh, oh, false, 12),
merger_down: lin(oh, 3 * oh, false, 14),
};
let img = ImagePatches {
patches: fill(16 * pd, 42),
grid: [1, 4, 4],
};
let want = cpu.forward(&[img.clone()]).expect("cpu forward");
let gpu = GlmVisionGpu::new_with_opts(&ctx, &cpu, 16, false).expect("gpu build");
let got = gpu.forward(&ctx, &[img.clone()]).expect("gpu forward");
let mx = want.iter().zip(&got).map(|(a, b)| (a - b).abs()).fold(0f32, f32::max);
println!("GLM tower CPU vs GPU(f32) max_abs = {mx:.3e}");
assert!(mx < 1e-3, "GLM tower not iso: max_abs {mx:.3e}");
if ctx.f16 {
let gpu16 = GlmVisionGpu::new_with_opts(&ctx, &cpu, 16, true).expect("gpu f16 build");
let got16 = gpu16.forward(&ctx, &[img]).expect("gpu f16 forward");
let mx16 = got.iter().zip(&got16).map(|(a, b)| (a - b).abs()).fold(0f32, f32::max);
println!("GLM tower f32 vs f16-weights max_abs = {mx16:.3e}");
assert!(mx16 < 5e-2, "f16 drift too large: {mx16:.3e}");
}
}
}