use anyhow::Result;
use crate::GpuCtx;
use crate::diffusion_gemma::{DgCache, DgDecoder, DgGenConfig, DgRng, rope_cos_sin, softmax_f32};
use crate::encoder::{GEMM3_TILES, enc_gemm3_src, gemm3_tier, gemm3_tile};
use crate::forward::{make_bg, pipeline, uni};
const DG_RMSNORM: &str = r#"
struct Meta { h: u32, weighted: u32, eps: u32, p: u32 }
@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> outp: array<f32>;
@group(0) @binding(3) 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.h;
var sq = 0.0;
for (var i = t; i < mt.h; i = i + 256u) { let v = x[base + i]; sq = sq + v * v; }
sh[t] = sq;
workgroupBarrier();
for (var s = 128u; s > 0u; s = s >> 1u) { if (t < s) { sh[t] = sh[t] + sh[t + s]; } workgroupBarrier(); }
let inv = 1.0 / sqrt(sh[0] / f32(mt.h) + bitcast<f32>(mt.eps));
for (var i = t; i < mt.h; i = i + 256u) {
var v = x[base + i] * inv;
if (mt.weighted != 0u) { v = v * w[i]; }
outp[base + i] = v;
}
}
"#;
const DG_HEADNORM_ROPE: &str = r#"
struct Meta { heads: u32, hd: u32, flags: u32, eps: u32 }
@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<storage, read> cs: array<f32>;
@group(0) @binding(3) var<uniform> mt: Meta;
var<workgroup> sh: array<f32, 128>;
var<workgroup> nv: array<f32, 128>;
@compute @workgroup_size(128)
fn main(@builtin(workgroup_id) wg: vec3<u32>, @builtin(local_invocation_index) t: u32) {
let tok = wg.x;
let head = wg.y;
let hd = mt.hd;
let base = tok * mt.heads * hd + head * hd;
var sq = 0.0;
if (t < hd) { let v = x[base + t]; sq = v * v; }
sh[t] = sq;
workgroupBarrier();
for (var s = 64u; s > 0u; s = s >> 1u) { if (t < s) { sh[t] = sh[t] + sh[t + s]; } workgroupBarrier(); }
let inv = 1.0 / sqrt(sh[0] / f32(hd) + bitcast<f32>(mt.eps));
if (t < hd) {
var v = x[base + t] * inv;
if ((mt.flags & 1u) != 0u) { v = v * w[t]; }
nv[t] = v;
}
workgroupBarrier();
if (t < hd) {
var v = nv[t];
if ((mt.flags & 2u) != 0u) {
let half = hd / 2u;
let cosv = cs[tok * 2u * hd + t];
let sinv = cs[tok * 2u * hd + hd + t];
var rot: f32;
if (t < half) { rot = -nv[t + half]; } else { rot = nv[t - half]; }
v = v * cosv + rot * sinv;
}
x[base + t] = v;
}
}
"#;
const DG_ATTN: &str = r#"
struct Meta { t: u32, nh: u32, hd: u32, ctx: u32, nkv: u32, p1: u32, p2: u32, p3: u32 }
@group(0) @binding(0) var<storage, read> q: array<f32>;
@group(0) @binding(1) var<storage, read> kc: array<f32>;
@group(0) @binding(2) var<storage, read> vc: array<f32>;
@group(0) @binding(3) var<storage, read> k: array<f32>;
@group(0) @binding(4) var<storage, read> v: array<f32>;
@group(0) @binding(5) var<storage, read_write> outp: array<f32>;
@group(0) @binding(6) var<uniform> mt: Meta;
const TILE: u32 = 128u;
var<workgroup> q_sh: array<f32, 128>;
var<workgroup> sc: array<f32, 128>;
var<workgroup> red: array<f32, 128>;
@compute @workgroup_size(128)
fn main(@builtin(workgroup_id) wg: vec3<u32>, @builtin(local_invocation_index) tid: u32) {
let t = mt.t;
let hd = mt.hd;
let ctx = mt.ctx;
let s_total = ctx + t;
let i = wg.x;
let head = wg.y;
if (i >= t || head >= mt.nh) { return; }
let kv = head / (mt.nh / mt.nkv);
let qbase = i * mt.nh * hd + head * hd;
if (tid < hd) { q_sh[tid] = q[qbase + tid]; }
workgroupBarrier();
var acc = 0.0;
var run_max = -1e30;
var run_den = 0.0;
var tile0 = 0u;
loop {
if (tile0 >= s_total) { break; }
let j = tile0 + tid;
var s = -1e30;
if (j < s_total) {
var kbase: u32;
if (j < ctx) { kbase = (kv * ctx + j) * hd; } else { kbase = ((j - ctx) * mt.nkv + kv) * hd; }
var dot = 0.0;
if (j < ctx) {
for (var c = 0u; c < hd; c = c + 1u) { dot = dot + q_sh[c] * kc[kbase + c]; }
} else {
for (var c = 0u; c < hd; c = c + 1u) { dot = dot + q_sh[c] * k[kbase + c]; }
}
s = dot; // scaling = 1.0
}
sc[tid] = s;
red[tid] = s;
workgroupBarrier();
for (var st = 64u; st > 0u; st = st >> 1u) {
if (tid < st) { red[tid] = max(red[tid], red[tid + st]); }
workgroupBarrier();
}
let nm = max(run_max, red[0]);
let corr = exp(run_max - nm);
workgroupBarrier();
var p = 0.0;
if (j < s_total) { p = exp(sc[tid] - nm); }
sc[tid] = p;
red[tid] = p;
workgroupBarrier();
for (var st = 64u; st > 0u; st = st >> 1u) {
if (tid < st) { red[tid] = red[tid] + red[tid + st]; }
workgroupBarrier();
}
run_den = run_den * corr + red[0];
if (tid < hd) {
var a = acc * corr;
let n = min(TILE, s_total - tile0);
for (var r = 0u; r < n; r = r + 1u) {
let jj = tile0 + r;
var vbase: u32;
if (jj < ctx) { vbase = (kv * ctx + jj) * hd; } else { vbase = ((jj - ctx) * mt.nkv + kv) * hd; }
if (jj < ctx) { a = a + sc[r] * vc[vbase + tid]; } else { a = a + sc[r] * v[vbase + tid]; }
}
acc = a;
}
run_max = nm;
workgroupBarrier();
tile0 = tile0 + TILE;
}
if (tid < hd) { outp[qbase + tid] = acc / run_den; }
}
"#;
const DG_GLU: &str = r#"
struct Meta { n: u32, mi: u32, fused: u32, p: u32 }
@group(0) @binding(0) var<storage, read> g: array<f32>;
@group(0) @binding(1) var<storage, read> u: array<f32>;
@group(0) @binding(2) var<storage, read_write> outp: array<f32>;
@group(0) @binding(3) var<uniform> mt: Meta;
fn gelu_tanh(x: f32) -> f32 {
return 0.5 * x * (1.0 + tanh(0.7978845608028654 * (x + 0.044715 * x * x * x)));
}
@compute @workgroup_size(256)
fn main(@builtin(global_invocation_id) gid: vec3<u32>) {
let i = gid.x;
if (i >= mt.n) { return; }
if (mt.fused != 0u) {
let r = i / mt.mi;
let c = i % mt.mi;
outp[i] = gelu_tanh(g[r * 2u * mt.mi + c]) * g[r * 2u * mt.mi + mt.mi + c];
} else {
outp[i] = gelu_tanh(g[i]) * u[i];
}
}
"#;
const DG_SOFTMAX: &str = r#"
struct Meta { n: u32, p0: u32, p1: u32, p2: u32 }
@group(0) @binding(0) var<storage, read_write> x: 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.n;
var m = -1e30;
for (var i = t; i < mt.n; i = i + 256u) { m = max(m, x[base + i]); }
sh[t] = m;
workgroupBarrier();
for (var s = 128u; s > 0u; s = s >> 1u) { if (t < s) { sh[t] = max(sh[t], sh[t + s]); } workgroupBarrier(); }
let mx = sh[0];
workgroupBarrier();
var sum = 0.0;
for (var i = t; i < mt.n; i = i + 256u) { let e = exp(x[base + i] - mx); x[base + i] = e; sum = sum + e; }
sh[t] = sum;
workgroupBarrier();
for (var s = 128u; s > 0u; s = s >> 1u) { if (t < s) { sh[t] = sh[t] + sh[t + s]; } workgroupBarrier(); }
let inv = 1.0 / sh[0];
for (var i = t; i < mt.n; i = i + 256u) { x[base + i] = x[base + i] * inv; }
}
"#;
const DG_GATHER: &str = r#"
struct Meta { n: u32, h: u32, p0: u32, p1: u32 }
@group(0) @binding(0) var<storage, read> src: array<f32>;
@group(0) @binding(1) var<storage, read> idx: array<u32>;
@group(0) @binding(2) var<storage, read_write> dst: array<f32>;
@group(0) @binding(3) var<uniform> mt: Meta;
@compute @workgroup_size(256)
fn main(@builtin(global_invocation_id) gid: vec3<u32>) {
let i = gid.x;
if (i >= mt.n * mt.h) { return; }
let r = i / mt.h;
let c = i % mt.h;
dst[i] = src[idx[r] * mt.h + c];
}
"#;
const DG_SCATTER_ADD: &str = r#"
struct Meta { n: u32, h: u32, p0: u32, p1: u32 }
@group(0) @binding(0) var<storage, read> src: array<f32>;
@group(0) @binding(1) var<storage, read> idx: array<u32>;
@group(0) @binding(2) var<storage, read> wgt: array<f32>;
@group(0) @binding(3) var<storage, read_write> dst: array<f32>;
@group(0) @binding(4) var<uniform> mt: Meta;
@compute @workgroup_size(256)
fn main(@builtin(global_invocation_id) gid: vec3<u32>) {
let i = gid.x;
if (i >= mt.n * mt.h) { return; }
let r = i / mt.h;
let c = i % mt.h;
dst[idx[r] * mt.h + c] = dst[idx[r] * mt.h + c] + wgt[r] * src[i];
}
"#;
const DG_ADD_SCALE: &str = r#"
struct Meta { n: u32, s: u32, p0: u32, p1: 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] = (dst[gid.x] + src[gid.x]) * bitcast<f32>(mt.s); }
}
"#;
struct GLin {
w: wgpu::Buffer,
b: wgpu::Buffer,
n: u32,
k: u32,
}
impl GLin {
fn new(ctx: &GpuCtx, w: &[f32], n: usize, k: usize) -> Self {
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];
}
}
Self {
w: ctx.storage(&wt),
b: ctx.storage(&vec![0f32; n]),
n: n as u32,
k: k as u32,
}
}
}
struct GLayer {
sliding: bool,
n_kv: u32,
hd: u32,
q: GLin,
k: GLin,
v: Option<GLin>,
o: GLin,
q_norm: wgpu::Buffer,
k_norm: wgpu::Buffer,
ln_in: wgpu::Buffer,
ln_post_attn: wgpu::Buffer,
ln_pre_ff: wgpu::Buffer,
ln_post_ff: wgpu::Buffer,
ln_post_ff1: wgpu::Buffer,
ln_post_ff2: wgpu::Buffer,
ln_pre_ff2: wgpu::Buffer,
router_w: wgpu::Buffer,
router_proj: GLin,
layer_scalar: f32,
gate: GLin,
up: GLin,
down: GLin,
experts_gu: Vec<GLin>,
experts_down: Vec<GLin>,
}
pub struct DgGpuCache {
k: Vec<wgpu::Buffer>,
v: Vec<wgpu::Buffer>,
ctx: Vec<usize>,
seq_len: usize,
}
pub struct DgGpu {
cpu: DgDecoder,
layers: Vec<GLayer>,
head: GLin,
sc_embed: GLin,
sc_pre_norm: wgpu::Buffer,
sc_gate: GLin,
sc_up: GLin,
sc_down: GLin,
final_norm: wgpu::Buffer,
gemm3: Vec<wgpu::ComputePipeline>,
rmsnorm: wgpu::ComputePipeline,
headnorm: wgpu::ComputePipeline,
attn: wgpu::ComputePipeline,
glu: wgpu::ComputePipeline,
softmax: wgpu::ComputePipeline,
gather: wgpu::ComputePipeline,
scatter: wgpu::ComputePipeline,
add_scale: wgpu::ComputePipeline,
ones: wgpu::Buffer,
}
impl DgGpu {
pub fn new(ctx: &GpuCtx, cpu: DgDecoder) -> Result<Self> {
let cfg = &cpu.cfg;
anyhow::ensure!(
cfg.head_dim <= 128 && cfg.head_dim_global <= 128,
"attention supports head_dim ≤ 128"
);
let (h, e, mi) = (cfg.hidden, cfg.num_experts, cfg.moe_inter);
let hscale = (h as f32).powf(-0.5);
let layers = cpu
.layers
.iter()
.map(|l| GLayer {
sliding: l.sliding,
n_kv: l.n_kv as u32,
hd: l.hd as u32,
q: GLin::new(ctx, &l.q.w, l.q.n, l.q.k),
k: GLin::new(ctx, &l.k.w, l.k.n, l.k.k),
v: l.v.as_ref().map(|v| GLin::new(ctx, &v.w, v.n, v.k)),
o: GLin::new(ctx, &l.o.w, l.o.n, l.o.k),
q_norm: ctx.storage(&l.q_norm),
k_norm: ctx.storage(&l.k_norm),
ln_in: ctx.storage(&l.ln_in),
ln_post_attn: ctx.storage(&l.ln_post_attn),
ln_pre_ff: ctx.storage(&l.ln_pre_ff),
ln_post_ff: ctx.storage(&l.ln_post_ff),
ln_post_ff1: ctx.storage(&l.ln_post_ff1),
ln_post_ff2: ctx.storage(&l.ln_post_ff2),
ln_pre_ff2: ctx.storage(&l.ln_pre_ff2),
router_w: ctx.storage(
&l.router_scale
.iter()
.map(|s| s * hscale)
.collect::<Vec<_>>(),
),
router_proj: GLin::new(ctx, &l.router_proj.w, e, h),
layer_scalar: l.layer_scalar_dec,
gate: GLin::new(ctx, &l.gate.w, l.gate.n, l.gate.k),
up: GLin::new(ctx, &l.up.w, l.up.n, l.up.k),
down: GLin::new(ctx, &l.down.w, l.down.n, l.down.k),
experts_gu: (0..e)
.map(|ei| {
GLin::new(
ctx,
&l.experts_gate_up[ei * 2 * mi * h..(ei + 1) * 2 * mi * h],
2 * mi,
h,
)
})
.collect(),
experts_down: (0..e)
.map(|ei| {
GLin::new(ctx, &l.experts_down[ei * h * mi..(ei + 1) * h * mi], h, mi)
})
.collect(),
})
.collect();
let scale = (h as f32).sqrt();
let mut sc_w = vec![0f32; h * cfg.vocab];
for v in 0..cfg.vocab {
for j in 0..h {
sc_w[j * cfg.vocab + v] = cpu.embed[v * h + j] * scale;
}
}
Ok(Self {
layers,
head: GLin::new(ctx, &cpu.embed, cfg.vocab, h),
sc_embed: GLin::new(ctx, &sc_w, h, cfg.vocab),
sc_pre_norm: ctx.storage(&cpu.sc.pre_norm),
sc_gate: GLin::new(ctx, &cpu.sc.gate.w, cpu.sc.gate.n, cpu.sc.gate.k),
sc_up: GLin::new(ctx, &cpu.sc.up.w, cpu.sc.up.n, cpu.sc.up.k),
sc_down: GLin::new(ctx, &cpu.sc.down.w, cpu.sc.down.n, cpu.sc.down.k),
final_norm: ctx.storage(&cpu.final_norm),
gemm3: GEMM3_TILES
.iter()
.map(|&(bm, bn, bk)| pipeline(ctx, "dg_gemm3", &enc_gemm3_src(false, bm, bn, bk)))
.collect(),
rmsnorm: pipeline(ctx, "dg_rmsnorm", DG_RMSNORM),
headnorm: pipeline(ctx, "dg_headnorm", DG_HEADNORM_ROPE),
attn: pipeline(ctx, "dg_attn", DG_ATTN),
glu: pipeline(ctx, "dg_glu", DG_GLU),
softmax: pipeline(ctx, "dg_softmax", DG_SOFTMAX),
gather: pipeline(ctx, "dg_gather", DG_GATHER),
scatter: pipeline(ctx, "dg_scatter", DG_SCATTER_ADD),
add_scale: pipeline(ctx, "dg_add_scale", DG_ADD_SCALE),
ones: ctx.storage(&vec![1.0f32; cpu.cfg.hidden.max(128)]),
cpu,
})
}
pub fn cpu(&self) -> &DgDecoder {
&self.cpu
}
pub fn upload_cache(&self, ctx: &GpuCtx, cache: &DgCache) -> DgGpuCache {
let one = vec![0f32; 1]; DgGpuCache {
k: cache
.k
.iter()
.map(|k| ctx.storage(if k.is_empty() { &one } else { k }))
.collect(),
v: cache
.v
.iter()
.map(|v| ctx.storage(if v.is_empty() { &one } else { v }))
.collect(),
ctx: cache.ctx.clone(),
seq_len: cache.seq_len,
}
}
pub fn canvas_forward(
&self,
ctx: &GpuCtx,
ids: &[u32],
cache: &DgGpuCache,
sc_logits: Option<&[f32]>,
) -> Result<Vec<f32>> {
let cfg = &self.cpu.cfg;
let (t_len, h, vocab) = (ids.len(), cfg.hidden, cfg.vocab);
let eps = cfg.eps.to_bits();
let scale = (h as f32).sqrt();
let mut embeds = vec![0f32; t_len * h];
for (t, &id) in ids.iter().enumerate() {
for j in 0..h {
embeds[t * h + j] = self.cpu.embed[id as usize * h + j] * scale;
}
}
let xb = ctx.storage(&embeds);
let nb = ctx.empty(t_len * h);
let sb = ctx.empty(t_len * h.max(cfg.n_heads * cfg.head_dim.max(cfg.head_dim_global)));
let qb = ctx.empty(t_len * cfg.n_heads * cfg.head_dim.max(cfg.head_dim_global));
let kb = ctx
.empty(t_len * cfg.n_kv.max(cfg.n_kv_global) * cfg.head_dim.max(cfg.head_dim_global));
let vb = ctx
.empty(t_len * cfg.n_kv.max(cfg.n_kv_global) * cfg.head_dim.max(cfg.head_dim_global));
let gb = ctx.empty(t_len * cfg.inter);
let ub = ctx.empty(t_len * cfg.inter);
let ab = ctx.empty(t_len * cfg.inter);
let h1b = ctx.empty(t_len * h);
let h1n = ctx.empty(t_len * h);
let h2n = ctx.empty(t_len * h);
let reb = ctx.empty(t_len * cfg.num_experts);
let heb = ctx.empty(t_len * h);
let mut passes: Vec<(&wgpu::ComputePipeline, wgpu::BindGroup, u32, u32)> = Vec::new();
let mut keep: Vec<wgpu::Buffer> = Vec::new();
macro_rules! flush {
() => {{
let mut enc = ctx
.device
.create_command_encoder(&wgpu::CommandEncoderDescriptor { label: Some("dg") });
{
let mut cp = enc.begin_compute_pass(&wgpu::ComputePassDescriptor {
label: Some("dg"),
timestamp_writes: None,
});
for (pl, bg, gx, gy) in passes.drain(..) {
cp.set_pipeline(pl);
cp.set_bind_group(0, &bg, &[]);
cp.dispatch_workgroups(gx, gy, 1);
}
}
ctx.queue.submit(Some(enc.finish()));
}};
}
macro_rules! gemm {
($x:expr, $lw:expr, $rows:expr, $y:expr) => {{
let lw: &GLin = $lw;
let rows: usize = $rows;
let meta = uni(ctx, bytemuck::cast_slice(&[rows as u32, lw.n, lw.k, 1u32]));
let tile = gemm3_tile(rows, lw.n as usize);
let pl = &self.gemm3[gemm3_tier(tile)];
let bg = make_bg(ctx, pl, &[$x, &lw.w, &lw.b, $y], &meta);
passes.push((
pl,
bg,
lw.n.div_ceil(tile.1 as u32),
(rows as u32).div_ceil(tile.0 as u32),
));
keep.push(meta);
}};
}
macro_rules! rms {
($x:expr, $w:expr, $weighted:expr, $y:expr) => {{
let meta = uni(
ctx,
bytemuck::cast_slice(&[h as u32, $weighted as u32, eps, 0u32]),
);
let bg = make_bg(ctx, &self.rmsnorm, &[$x, $w, $y], &meta);
passes.push((&self.rmsnorm, bg, t_len as u32, 1));
keep.push(meta);
}};
}
macro_rules! add_scale {
($dst:expr, $src:expr, $s:expr) => {{
let meta = uni(
ctx,
bytemuck::cast_slice(&[(t_len * h) as u32, ($s as f32).to_bits(), 0u32, 0u32]),
);
let bg = make_bg(ctx, &self.add_scale, &[$dst, $src], &meta);
passes.push((&self.add_scale, bg, ((t_len * h) as u32).div_ceil(256), 1));
keep.push(meta);
}};
}
match sc_logits {
Some(logits) => {
let lb = ctx.storage(logits);
let meta = uni(ctx, bytemuck::cast_slice(&[vocab as u32, 0u32, 0u32, 0u32]));
let bg = make_bg(ctx, &self.softmax, &[&lb], &meta);
passes.push((&self.softmax, bg, t_len as u32, 1));
keep.push(meta);
gemm!(&lb, &self.sc_embed, t_len, &sb); keep.push(lb);
}
None => {
let z = ctx.empty(t_len * h);
add_scale!(&sb, &z, 1.0); keep.push(z);
}
}
rms!(&sb, &self.sc_pre_norm, 1, &nb);
gemm!(&nb, &self.sc_gate, t_len, &gb);
gemm!(&nb, &self.sc_up, t_len, &ub);
{
let meta = uni(
ctx,
bytemuck::cast_slice(&[(t_len * cfg.inter) as u32, cfg.inter as u32, 0u32, 0u32]),
);
let bg = make_bg(ctx, &self.glu, &[&gb, &ub, &ab], &meta);
passes.push((&self.glu, bg, ((t_len * cfg.inter) as u32).div_ceil(256), 1));
keep.push(meta);
}
gemm!(&ab, &self.sc_down, t_len, &sb);
add_scale!(&xb, &sb, 1.0); rms!(&xb, &self.ones, 0, &nb); let cur = nb; let nb = ctx.empty(t_len * h);
let cs_sliding: Vec<f32> = (0..t_len)
.flat_map(|t| rope_cos_sin(cache.seq_len + t, cfg.head_dim, cfg.rope_sliding))
.collect();
let cs_full: Vec<f32> = (0..t_len)
.flat_map(|t| rope_cos_sin(cache.seq_len + t, cfg.head_dim_global, cfg.rope_full))
.collect();
let cs_sliding = ctx.storage(&cs_sliding);
let cs_full = ctx.storage(&cs_full);
for (li, l) in self.layers.iter().enumerate() {
let (n_kv, hd) = (l.n_kv, l.hd);
let cs = if l.sliding { &cs_sliding } else { &cs_full };
rms!(&cur, &l.ln_in, 1, &nb);
gemm!(&nb, &l.q, t_len, &qb);
gemm!(&nb, &l.k, t_len, &kb);
match &l.v {
Some(vp) => gemm!(&nb, vp, t_len, &vb),
None => gemm!(&nb, &l.k, t_len, &vb), }
let hn = |buf: &wgpu::Buffer, w: &wgpu::Buffer, heads: u32, flags: u32| {
let meta = uni(ctx, bytemuck::cast_slice(&[heads, hd, flags, eps]));
let bg = make_bg(ctx, &self.headnorm, &[buf, w, cs], &meta);
(bg, meta, heads)
};
for (bg, meta, heads) in [
hn(&qb, &l.q_norm, cfg.n_heads as u32, 3), hn(&kb, &l.k_norm, n_kv, 3),
hn(&vb, &self.ones, n_kv, 0), ] {
passes.push((&self.headnorm, bg, t_len as u32, heads));
keep.push(meta);
}
{
let meta = uni(
ctx,
bytemuck::cast_slice(&[
t_len as u32,
cfg.n_heads as u32,
hd,
cache.ctx[li] as u32,
n_kv,
0u32,
0u32,
0u32,
]),
);
let bg = make_bg(
ctx,
&self.attn,
&[&qb, &cache.k[li], &cache.v[li], &kb, &vb, &sb],
&meta,
);
passes.push((&self.attn, bg, t_len as u32, cfg.n_heads as u32));
keep.push(meta);
}
gemm!(&sb, &l.o, t_len, &nb);
rms!(&nb, &l.ln_post_attn, 1, &sb);
add_scale!(&cur, &sb, 1.0);
rms!(&cur, &l.ln_pre_ff, 1, &nb);
gemm!(&nb, &l.gate, t_len, &gb);
gemm!(&nb, &l.up, t_len, &ub);
{
let meta = uni(
ctx,
bytemuck::cast_slice(&[
(t_len * cfg.inter) as u32,
cfg.inter as u32,
0u32,
0u32,
]),
);
let bg = make_bg(ctx, &self.glu, &[&gb, &ub, &ab], &meta);
passes.push((&self.glu, bg, ((t_len * cfg.inter) as u32).div_ceil(256), 1));
keep.push(meta);
}
gemm!(&ab, &l.down, t_len, &h1b);
rms!(&h1b, &l.ln_post_ff1, 1, &h1n);
rms!(&cur, &l.router_w, 1, &nb); gemm!(&nb, &l.router_proj, t_len, &reb);
rms!(&cur, &l.ln_pre_ff2, 1, &heb);
flush!();
let re = ctx.read(&reb, t_len * cfg.num_experts)?;
let e_n = cfg.num_experts;
let mut per_expert: Vec<(Vec<u32>, Vec<f32>)> = vec![(Vec::new(), Vec::new()); e_n];
for t in 0..t_len {
let mut pr = re[t * e_n..(t + 1) * e_n].to_vec();
softmax_f32(&mut pr);
let mut idx: Vec<usize> = (0..e_n).collect();
idx.sort_by(|&a, &b| pr[b].partial_cmp(&pr[a]).unwrap());
let top = &idx[..cfg.top_k];
let wsum: f32 = top.iter().map(|&e| pr[e]).sum();
for &ei in top {
per_expert[ei].0.push(t as u32);
per_expert[ei]
.1
.push(pr[ei] / wsum * self.cpu.layers[li].per_expert_scale[ei]);
}
}
let h2b = ctx.empty(t_len * h); for (ei, (idxs, wgts)) in per_expert.iter().enumerate() {
if idxs.is_empty() {
continue;
}
let nt = idxs.len();
let ib = ctx.storage_bytes(bytemuck::cast_slice(idxs));
let wb = ctx.storage(wgts);
let xg = ctx.empty(nt * h);
let yg = ctx.empty(nt * 2 * cfg.moe_inter);
let ag = ctx.empty(nt * cfg.moe_inter);
let og = ctx.empty(nt * h);
{
let meta = uni(
ctx,
bytemuck::cast_slice(&[nt as u32, h as u32, 0u32, 0u32]),
);
let bg = make_bg(ctx, &self.gather, &[&heb, &ib, &xg], &meta);
passes.push((&self.gather, bg, ((nt * h) as u32).div_ceil(256), 1));
keep.push(meta);
}
gemm!(&xg, &l.experts_gu[ei], nt, &yg);
{
let meta = uni(
ctx,
bytemuck::cast_slice(&[
(nt * cfg.moe_inter) as u32,
cfg.moe_inter as u32,
1u32, 0u32,
]),
);
let bg = make_bg(ctx, &self.glu, &[&yg, &yg, &ag], &meta);
passes.push((
&self.glu,
bg,
((nt * cfg.moe_inter) as u32).div_ceil(256),
1,
));
keep.push(meta);
}
gemm!(&ag, &l.experts_down[ei], nt, &og);
{
let meta = uni(
ctx,
bytemuck::cast_slice(&[nt as u32, h as u32, 0u32, 0u32]),
);
let bg = make_bg(ctx, &self.scatter, &[&og, &ib, &wb, &h2b], &meta);
passes.push((&self.scatter, bg, ((nt * h) as u32).div_ceil(256), 1));
keep.push(meta);
}
keep.extend([ib, wb, xg, yg, ag, og]);
}
rms!(&h2b, &l.ln_post_ff2, 1, &h2n);
add_scale!(&h1n, &h2n, 1.0); rms!(&h1n, &l.ln_post_ff, 1, &sb);
add_scale!(&cur, &sb, l.layer_scalar); keep.push(h2b);
}
rms!(&cur, &self.final_norm, 1, &nb);
let lgb = ctx.empty(t_len * vocab);
gemm!(&nb, &self.head, t_len, &lgb);
flush!();
let mut logits = ctx.read(&lgb, t_len * vocab)?;
for v in logits.iter_mut() {
*v = cfg.softcap * (*v / cfg.softcap).tanh();
}
drop(keep);
Ok(logits)
}
pub fn generate(
&self,
ctx: &GpuCtx,
prompt: &[u32],
gc: &DgGenConfig,
rng: &mut dyn DgRng,
) -> Result<Vec<u32>> {
let canvas = self.cpu.cfg.canvas_length;
let max_new_canvases = gc.max_new_tokens.div_ceil(canvas);
let mut cache = DgCache::empty(self.cpu.cfg.n_layers);
let mut out: Vec<u32> = prompt.to_vec();
let mut to_encode: Vec<u32> = prompt.to_vec();
for _ in 0..max_new_canvases {
self.cpu.encode(&to_encode, &mut cache);
let gcache = self.upload_cache(ctx, &cache);
let fwd = |ids: &[u32], sc: Option<&[f32]>| {
self.canvas_forward(ctx, ids, &gcache, sc)
.expect("gpu forward")
};
let (mut tokens, _) = self.cpu.denoise_block_with(gc, rng, None, fwd);
let mut finished = false;
if let Some(eos) = gc.eos_token_id
&& let Some(p) = tokens.iter().position(|&t| t == eos)
{
for t in tokens[p + 1..].iter_mut() {
*t = gc.pad_token_id;
}
finished = true;
}
out.extend_from_slice(&tokens);
if finished {
break;
}
to_encode = tokens;
}
Ok(out)
}
}