use std::collections::HashMap;
use anyhow::Result;
use crate::GpuCtx;
use crate::forward::pipeline;
use crate::mimi::{Conv1d, ConvTr1d, Mimi, Transformer};
const DIM: usize = 512;
const VQ_DIM: usize = 256;
const BINS: usize = 2048;
const CONTEXT: usize = 250;
const FRAME: usize = 1920;
const NQ: usize = 8;
pub(crate) const INVALID_POS: u32 = 0xFFFF_FFFF;
pub(crate) const ELU_IP: &str = r#"
@group(0) @binding(0) var<storage, read_write> x: array<f32>;
@group(0) @binding(1) var<uniform> p: vec4<u32>;
@compute @workgroup_size(256)
fn main(@builtin(global_invocation_id) g: vec3<u32>) {
if (g.x < p.x) { let v = x[g.x]; x[g.x] = select(exp(v) - 1.0, v, v > 0.0); }
}
"#;
pub(crate) const ELU_TO: &str = r#"
@group(0) @binding(0) var<storage, read> src: array<f32>;
@group(0) @binding(1) var<storage, read_write> dst: array<f32>;
@group(0) @binding(2) var<uniform> p: vec4<u32>;
@compute @workgroup_size(256)
fn main(@builtin(global_invocation_id) g: vec3<u32>) {
if (g.x < p.x) { let v = src[g.x]; dst[g.x] = select(exp(v) - 1.0, v, v > 0.0); }
}
"#;
pub(crate) const REPLICATE: &str = r#"
@group(0) @binding(0) var<storage, read_write> xa: array<f32>;
@group(0) @binding(1) var<uniform> p: vec4<u32>;
@group(0) @binding(2) var<uniform> u: vec4<u32>;
@compute @workgroup_size(256)
fn main(@builtin(global_invocation_id) g: vec3<u32>) {
if (u.y != 1u || g.x >= p.x * p.w) { return; }
let c = g.x / p.w;
let i = g.x % p.w;
xa[c * p.z + i] = xa[c * p.z + p.w];
}
"#;
pub(crate) const WRITE_XA: &str = r#"
@group(0) @binding(0) var<storage, read> src: array<f32>;
@group(0) @binding(1) var<storage, read_write> xa: array<f32>;
@group(0) @binding(2) var<uniform> p: vec4<u32>;
@group(0) @binding(3) var<uniform> q: vec4<u32>;
@compute @workgroup_size(256)
fn main(@builtin(global_invocation_id) g: vec3<u32>) {
if (g.x >= p.x * p.y) { return; }
let c = g.x / p.y;
let t = g.x % p.y;
let v = select(src[c * p.y + t], src[t * p.x + c], q.x == 1u);
xa[c * p.z + p.w + t] = v;
}
"#;
pub(crate) const CARRY: &str = r#"
@group(0) @binding(0) var<storage, read_write> xa: array<f32>;
@group(0) @binding(1) var<uniform> p: vec4<u32>;
@compute @workgroup_size(64)
fn main(@builtin(global_invocation_id) g: vec3<u32>) {
if (g.x >= p.x) { return; }
let base = g.x * p.z;
for (var i = 0u; i < p.w; i++) { xa[base + i] = xa[base + p.y + i]; }
}
"#;
pub(crate) const CONV1D: &str = r#"
@group(0) @binding(0) var<storage, read> w: array<f32>;
@group(0) @binding(1) var<storage, read> b: array<f32>;
@group(0) @binding(2) var<storage, read> xa: array<f32>;
@group(0) @binding(3) var<storage, read_write> y: array<f32>;
@group(0) @binding(4) var<uniform> p: vec4<u32>;
@group(0) @binding(5) var<uniform> q: vec4<u32>;
@compute @workgroup_size(64)
fn main(@builtin(global_invocation_id) g: vec3<u32>, @builtin(workgroup_id) wid: vec3<u32>) {
let ot = g.x;
let oc = wid.y;
if (ot >= p.w) { return; }
var acc = b[oc];
let base = ot * p.z;
for (var ic = 0u; ic < p.x; ic++) {
let xo = ic * q.x + base;
let wo = (oc * p.x + ic) * p.y;
for (var k = 0u; k < p.y; k++) { acc += w[wo + k] * xa[xo + k]; }
}
let o = oc * p.w + ot;
if (q.y == 1u) { y[o] += acc; } else { y[o] = acc; }
}
"#;
pub(crate) fn convtr_src(channel_wise: bool, pre_elu: bool) -> String {
let act = if pre_elu {
"if (xv < 0.0) { xv = exp(xv) - 1.0; }"
} else {
""
};
let body = if channel_wise {
format!(
"var xv = x[oc * p.w + ti];\n {act}\n acc += w[oc * p.y + k] * xv;"
)
} else {
format!(
"for (var ic = 0u; ic < p.x; ic++) {{\n var xv = x[ic * p.w + ti];\n {act}\n acc += w[(ic * q.z + oc) * p.y + k] * xv;\n }}"
)
};
format!(
r#"
@group(0) @binding(0) var<storage, read> w: array<f32>;
@group(0) @binding(1) var<storage, read> b: array<f32>;
@group(0) @binding(2) var<storage, read> x: array<f32>;
@group(0) @binding(3) var<storage, read_write> full: array<f32>;
@group(0) @binding(4) var<uniform> p: vec4<u32>;
@group(0) @binding(5) var<uniform> q: vec4<u32>;
@compute @workgroup_size(64)
fn main(@builtin(global_invocation_id) g: vec3<u32>, @builtin(workgroup_id) wid: vec3<u32>) {{
let i = g.x;
let oc = wid.y;
if (i >= q.x) {{ return; }}
var acc = b[oc];
for (var k = 0u; k < p.y; k++) {{
if (i < k) {{ continue; }}
let d = i - k;
if (d % p.z != 0u) {{ continue; }}
let ti = d / p.z;
if (ti >= p.w) {{ continue; }}
{body}
}}
full[oc * q.x + i] = acc;
}}
"#
)
}
pub(crate) const TRP_ADD: &str = r#"
@group(0) @binding(0) var<storage, read_write> full: array<f32>;
@group(0) @binding(1) var<storage, read> part: array<f32>;
@group(0) @binding(2) var<uniform> p: vec4<u32>;
@compute @workgroup_size(256)
fn main(@builtin(global_invocation_id) g: vec3<u32>) {
if (g.x >= p.x * p.z) { return; }
let oc = g.x / p.z;
let i = g.x % p.z;
full[oc * p.y + i] += part[g.x];
}
"#;
pub(crate) const TR_OUT: &str = r#"
@group(0) @binding(0) var<storage, read> full: array<f32>;
@group(0) @binding(1) var<storage, read_write> y: array<f32>;
@group(0) @binding(2) var<uniform> p: vec4<u32>;
@compute @workgroup_size(256)
fn main(@builtin(global_invocation_id) g: vec3<u32>) {
if (g.x >= p.x * p.z) { return; }
let oc = g.x / p.z;
let i = g.x % p.z;
y[g.x] = full[oc * p.y + i];
}
"#;
pub(crate) const TR_CARRY: &str = r#"
@group(0) @binding(0) var<storage, read> full: array<f32>;
@group(0) @binding(1) var<storage, read> b: array<f32>;
@group(0) @binding(2) var<storage, read_write> part: array<f32>;
@group(0) @binding(3) var<uniform> p: vec4<u32>;
@compute @workgroup_size(256)
fn main(@builtin(global_invocation_id) g: vec3<u32>) {
if (g.x >= p.x * p.z) { return; }
let oc = g.x / p.z;
let i = g.x % p.z;
part[g.x] = full[oc * p.y + p.w + i] - b[oc];
}
"#;
pub(crate) const CT2TM: &str = r#"
@group(0) @binding(0) var<storage, read> src: array<f32>;
@group(0) @binding(1) var<storage, read_write> dst: array<f32>;
@group(0) @binding(2) var<uniform> p: vec4<u32>;
@compute @workgroup_size(256)
fn main(@builtin(global_invocation_id) g: vec3<u32>) {
if (g.x >= p.x * p.y) { return; }
let c = g.x / p.y;
let t = g.x % p.y;
dst[t * p.x + c] = src[g.x];
}
"#;
pub(crate) const TM2CT: &str = r#"
@group(0) @binding(0) var<storage, read> src: array<f32>;
@group(0) @binding(1) var<storage, read_write> dst: array<f32>;
@group(0) @binding(2) var<uniform> p: vec4<u32>;
@compute @workgroup_size(256)
fn main(@builtin(global_invocation_id) g: vec3<u32>) {
if (g.x >= p.x * p.y) { return; }
let c = g.x / p.y;
let t = g.x % p.y;
dst[g.x] = src[t * p.x + c];
}
"#;
pub(crate) const LN512: &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> b: array<f32>;
@group(0) @binding(3) var<storage, read_write> y: array<f32>;
var<workgroup> red: array<f32, 256>;
@compute @workgroup_size(256)
fn main(@builtin(workgroup_id) wid: vec3<u32>, @builtin(local_invocation_id) lid: vec3<u32>) {
let base = wid.x * 512u;
let t = lid.x;
red[t] = x[base + t] + x[base + t + 256u];
workgroupBarrier();
var s = 128u;
while (s > 0u) { if (t < s) { red[t] += red[t + s]; } workgroupBarrier(); s /= 2u; }
let mean = red[0] / 512.0;
workgroupBarrier();
let d0 = x[base + t] - mean;
let d1 = x[base + t + 256u] - mean;
red[t] = d0 * d0 + d1 * d1;
workgroupBarrier();
s = 128u;
while (s > 0u) { if (t < s) { red[t] += red[t + s]; } workgroupBarrier(); s /= 2u; }
let inv = 1.0 / sqrt(red[0] / 512.0 + 1e-5);
y[base + t] = d0 * inv * w[t] + b[t];
y[base + t + 256u] = d1 * inv * w[t + 256u] + b[t + 256u];
}
"#;
pub(crate) fn ln_src(dim: usize, eps: f32) -> String {
format!(
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> b: array<f32>;
@group(0) @binding(3) var<storage, read_write> y: array<f32>;
var<workgroup> red: array<f32, 256>;
@compute @workgroup_size(256)
fn main(@builtin(workgroup_id) wid: vec3<u32>, @builtin(local_invocation_id) lid: vec3<u32>) {{
let base = wid.x * {dim}u;
let t = lid.x;
var acc = 0.0;
for (var i = t; i < {dim}u; i += 256u) {{ acc += x[base + i]; }}
red[t] = acc;
workgroupBarrier();
var s = 128u;
while (s > 0u) {{ if (t < s) {{ red[t] += red[t + s]; }} workgroupBarrier(); s /= 2u; }}
let mean = red[0] / {dim}.0;
workgroupBarrier();
var vacc = 0.0;
for (var i = t; i < {dim}u; i += 256u) {{ let d = x[base + i] - mean; vacc += d * d; }}
red[t] = vacc;
workgroupBarrier();
s = 128u;
while (s > 0u) {{ if (t < s) {{ red[t] += red[t + s]; }} workgroupBarrier(); s /= 2u; }}
let inv = 1.0 / sqrt(red[0] / {dim}.0 + {eps:e});
workgroupBarrier();
for (var i = t; i < {dim}u; i += 256u) {{
y[base + i] = (x[base + i] - mean) * inv * w[i] + b[i];
}}
}}
"#
)
}
pub(crate) const MATVEC: &str = r#"
@group(0) @binding(0) var<storage, read> w: array<f32>;
@group(0) @binding(1) var<storage, read> x: array<f32>;
@group(0) @binding(2) var<storage, read_write> y: array<f32>;
@group(0) @binding(3) var<uniform> p: vec4<u32>;
@compute @workgroup_size(64)
fn main(@builtin(global_invocation_id) g: vec3<u32>, @builtin(workgroup_id) wid: vec3<u32>) {
let n = g.x;
let t = wid.y;
if (n >= p.x) { return; }
var acc = 0.0;
let wo = n * p.y;
let xo = t * p.y;
for (var k = 0u; k < p.y; k++) { acc += w[wo + k] * x[xo + k]; }
let o = t * p.x + n;
if (p.z == 1u) { y[o] += acc; } else { y[o] = acc; }
}
"#;
pub(crate) fn rope_push_src(dim: usize, heads: usize, ring: usize) -> String {
let (hd, lanes, half) = (dim / heads, dim / 2, dim / heads / 2);
let (qkv_row, koff, voff) = (3 * dim, dim, 2 * dim);
format!(
r#"
@group(0) @binding(0) var<storage, read> qkv: array<f32>;
@group(0) @binding(1) var<storage, read_write> q: array<f32>;
@group(0) @binding(2) var<storage, read_write> kr: array<f32>;
@group(0) @binding(3) var<storage, read_write> vr: array<f32>;
@group(0) @binding(4) var<storage, read_write> pr: array<u32>;
@group(0) @binding(5) var<uniform> p: vec4<u32>;
@group(0) @binding(6) var<uniform> u: vec4<u32>;
@compute @workgroup_size(64)
fn main(@builtin(global_invocation_id) g: vec3<u32>) {{
if (g.x >= p.x * {lanes}u) {{ return; }}
let t = g.x / {lanes}u;
let pr_i = g.x % {lanes}u;
let h = pr_i / {half}u;
let j = pr_i % {half}u;
let pos = u.x + t;
let slot = pos % {ring}u;
let ts = f32(pos);
let freq = exp(-log(10000.0) * 2.0 * f32(j) / {hd}.0);
let sn = sin(freq * ts);
let cs = cos(freq * ts);
let base = t * {qkv_row}u + h * {hd}u + 2u * j;
let o = t * {dim}u + h * {hd}u + 2u * j;
let ro = slot * {dim}u + h * {hd}u + 2u * j;
let qr = qkv[base]; let qi = qkv[base + 1u];
q[o] = qr * cs - qi * sn; q[o + 1u] = qr * sn + qi * cs;
let kr_ = qkv[base + {koff}u]; let ki = qkv[base + {koff}u + 1u];
kr[ro] = kr_ * cs - ki * sn; kr[ro + 1u] = kr_ * sn + ki * cs;
vr[ro] = qkv[base + {voff}u]; vr[ro + 1u] = qkv[base + {voff}u + 1u];
if (pr_i == 0u) {{ pr[slot] = pos; }}
}}
"#
)
}
pub(crate) const MV_LS: &str = r#"
@group(0) @binding(0) var<storage, read> w: array<f32>;
@group(0) @binding(1) var<storage, read> x: array<f32>;
@group(0) @binding(2) var<storage, read> ls: array<f32>;
@group(0) @binding(3) var<storage, read_write> y: array<f32>;
@group(0) @binding(4) var<uniform> p: vec4<u32>;
@compute @workgroup_size(64)
fn main(@builtin(global_invocation_id) g: vec3<u32>, @builtin(workgroup_id) wid: vec3<u32>) {
let n = g.x;
let t = wid.y;
if (n >= p.x) { return; }
var acc = 0.0;
let wo = n * p.y;
let xo = t * p.y;
for (var k = 0u; k < p.y; k++) { acc += w[wo + k] * x[xo + k]; }
y[t * p.x + n] += ls[n] * acc;
}
"#;
pub(crate) fn attn_src(dim: usize, heads: usize, ring: usize, window: usize) -> String {
let hd = dim / heads;
let scale = 1.0f32 / (hd as f32).sqrt();
format!(
r#"
@group(0) @binding(0) var<storage, read> q: array<f32>;
@group(0) @binding(1) var<storage, read> kr: array<f32>;
@group(0) @binding(2) var<storage, read> vr: array<f32>;
@group(0) @binding(3) var<storage, read> pr: array<u32>;
@group(0) @binding(4) var<storage, read_write> attn: array<f32>;
@group(0) @binding(5) var<uniform> u: vec4<u32>;
var<workgroup> sc: array<f32, {ring}>;
var<workgroup> red: array<f32, 64>;
@compute @workgroup_size(64)
fn main(@builtin(workgroup_id) wid: vec3<u32>, @builtin(local_invocation_id) lid: vec3<u32>) {{
let t = wid.x;
let h = wid.y;
let d = lid.x;
let pos_q = u.x + t;
let qo = t * {dim}u + h * {hd}u;
var lmax = -3.0e38;
for (var j = d; j < {ring}u; j += 64u) {{
let pj = pr[j];
var s = -3.0e38;
if (pj != 0xffffffffu && pj <= pos_q && pos_q - pj < {window}u) {{
var acc = 0.0;
let ko = j * {dim}u + h * {hd}u;
for (var i = 0u; i < {hd}u; i++) {{ acc += q[qo + i] * kr[ko + i]; }}
s = acc * {scale:e};
}}
sc[j] = s;
lmax = max(lmax, s);
}}
red[d] = lmax;
workgroupBarrier();
var st = 32u;
while (st > 0u) {{ if (d < st) {{ red[d] = max(red[d], red[d + st]); }} workgroupBarrier(); st /= 2u; }}
let gmax = red[0];
workgroupBarrier();
var lsum = 0.0;
for (var j = d; j < {ring}u; j += 64u) {{
var e = 0.0;
if (sc[j] > -1.0e38) {{ e = exp(sc[j] - gmax); }}
sc[j] = e;
lsum += e;
}}
red[d] = lsum;
workgroupBarrier();
st = 32u;
while (st > 0u) {{ if (d < st) {{ red[d] += red[d + st]; }} workgroupBarrier(); st /= 2u; }}
let den = red[0];
workgroupBarrier();
var acc = 0.0;
for (var j = 0u; j < {ring}u; j++) {{
let wgt = sc[j];
if (wgt > 0.0) {{ acc += wgt * vr[j * {dim}u + h * {hd}u + d]; }}
}}
attn[t * {dim}u + h * {hd}u + d] = acc / den;
}}
"#
)
}
pub(crate) const GELU_ERF: &str = r#"
@group(0) @binding(0) var<storage, read_write> x: array<f32>;
@group(0) @binding(1) var<uniform> p: vec4<u32>;
fn erff(xin: f32) -> f32 {
let ax = abs(xin);
let sgn = xin < 0.0;
if (ax >= 6.0) { return select(1.0, -1.0, sgn); }
if (ax < 0.84375) {
let z = xin * xin;
let r = 1.2837916613e-01 + z * (-3.2504209876e-01 + z * (-2.8481749818e-02
+ z * (-5.7702702470e-03 + z * -2.3763017452e-05)));
let s = 1.0 + z * (3.9791721106e-01 + z * (6.5022252500e-02 + z * (5.0813062117e-03
+ z * (1.3249473704e-04 + z * -3.9602282413e-06))));
return xin + xin * (r / s);
}
if (ax < 1.25) {
let s = ax - 1.0;
let pn = -2.3621185683e-03 + s * (4.1485610604e-01 + s * (-3.7220788002e-01
+ s * (3.1834661961e-01 + s * (-1.1089469492e-01 + s * (3.5478305072e-02
+ s * -2.1663755178e-03)))));
let qn = 1.0 + s * (1.0642088205e-01 + s * (5.4039794207e-01 + s * (7.1828655899e-02
+ s * (1.2617121637e-01 + s * (1.3637083583e-02 + s * 1.1984500103e-02)))));
let e = 8.4506291151e-01 + pn / qn;
return select(e, -e, sgn);
}
let s = 1.0 / (ax * ax);
var rr: f32;
var ss: f32;
if (ax < 2.857142857142857) {
rr = -9.8649440333e-03 + s * (-6.9385856390e-01 + s * (-1.0558626175e+01
+ s * (-6.2375331879e+01 + s * (-1.6239666748e+02 + s * (-1.8460508728e+02
+ s * (-8.1287437439e+01 + s * -9.8143291473e+00))))));
ss = 1.0 + s * (1.9651271820e+01 + s * (1.3765776062e+02 + s * (4.3456588745e+02
+ s * (6.4538726807e+02 + s * (4.2900814819e+02 + s * (1.0863500214e+02
+ s * (6.5702495575e+00 + s * -6.0424413532e-02)))))));
} else {
rr = -9.8649431020e-03 + s * (-7.9928326607e-01 + s * (-1.7757955551e+01
+ s * (-1.6063638306e+02 + s * (-6.3756646729e+02 + s * (-1.0250950928e+03
+ s * -4.8351919556e+02)))));
ss = 1.0 + s * (3.0338060379e+01 + s * (3.2579251099e+02 + s * (1.5367296143e+03
+ s * (3.1998581543e+03 + s * (2.5530502930e+03 + s * (4.7452853394e+02
+ s * -2.2440952301e+01))))));
}
let z = bitcast<f32>(bitcast<u32>(ax) & 0xffffe000u);
let r = exp(-z * z - 0.5625) * exp((z - ax) * (z + ax) + rr / ss);
let e = 1.0 - r / ax;
return select(e, -e, sgn);
}
@compute @workgroup_size(256)
fn main(@builtin(global_invocation_id) g: vec3<u32>) {
if (g.x < p.x) {
let v = x[g.x];
x[g.x] = 0.5 * v * (1.0 + erff(v * 0.70710678118654752440));
}
}
"#;
const EMB_SUM: &str = r#"
@group(0) @binding(0) var<storage, read> cb: array<f32>;
@group(0) @binding(1) var<storage, read> codes: array<u32>;
@group(0) @binding(2) var<storage, read_write> acc: array<f32>;
@group(0) @binding(3) var<uniform> p: vec4<u32>;
@compute @workgroup_size(64)
fn main(@builtin(global_invocation_id) g: vec3<u32>) {
if (g.x >= 256u) { return; }
var a = 0.0;
for (var l = 0u; l < p.x; l++) {
let idx = codes[p.y + l];
a += cb[(l * 2048u + idx) * 256u + g.x];
}
acc[g.x] = a;
}
"#;
const NEAREST: &str = r#"
@group(0) @binding(0) var<storage, read> cb: array<f32>;
@group(0) @binding(1) var<storage, read> r: array<f32>;
@group(0) @binding(2) var<storage, read_write> codes: array<u32>;
@group(0) @binding(3) var<uniform> p: vec4<u32>;
var<workgroup> bd: array<f32, 256>;
var<workgroup> bi: array<u32, 256>;
@compute @workgroup_size(256)
fn main(@builtin(local_invocation_id) lid: vec3<u32>) {
let t = lid.x;
var best = 3.0e38;
var idx = 0u;
for (var bin = t; bin < 2048u; bin += 256u) {
let o = (p.x + bin) * 256u;
var d = 0.0;
for (var c = 0u; c < 256u; c++) { let df = cb[o + c] - r[c]; d += df * df; }
if (d < best || (d == best && bin < idx)) { best = d; idx = bin; }
}
bd[t] = best;
bi[t] = idx;
workgroupBarrier();
var s = 128u;
while (s > 0u) {
if (t < s && (bd[t + s] < bd[t] || (bd[t + s] == bd[t] && bi[t + s] < bi[t]))) {
bd[t] = bd[t + s];
bi[t] = bi[t + s];
}
workgroupBarrier();
s /= 2u;
}
if (t == 0u) { codes[p.y] = bi[0]; }
}
"#;
const SUB_EMB: &str = r#"
@group(0) @binding(0) var<storage, read> cb: array<f32>;
@group(0) @binding(1) var<storage, read> codes: array<u32>;
@group(0) @binding(2) var<storage, read_write> r: array<f32>;
@group(0) @binding(3) var<uniform> p: vec4<u32>;
@compute @workgroup_size(64)
fn main(@builtin(global_invocation_id) g: vec3<u32>) {
if (g.x >= 256u) { return; }
r[g.x] -= cb[(p.x + codes[p.y]) * 256u + g.x];
}
"#;
pub(crate) struct Step {
pl: wgpu::ComputePipeline,
bg: wgpu::BindGroup,
gx: u32,
gy: u32,
label: String,
}
#[derive(Clone, Copy)]
pub(crate) struct TrShape {
pub(crate) dim: usize,
pub(crate) heads: usize,
pub(crate) ffn: usize,
pub(crate) ring: usize,
pub(crate) window: usize,
pub(crate) ln_eps: f32,
}
pub(crate) struct ConvView<'a> {
pub(crate) w: &'a [f32],
pub(crate) b: Option<&'a [f32]>,
pub(crate) in_c: usize,
pub(crate) out_c: usize,
pub(crate) k: usize,
pub(crate) k_eff: usize,
pub(crate) stride: usize,
pub(crate) replicate: bool,
}
pub(crate) struct ConvTrView<'a> {
pub(crate) w: &'a [f32],
pub(crate) b: Option<&'a [f32]>,
pub(crate) in_c: usize,
pub(crate) out_c: usize,
pub(crate) k: usize,
pub(crate) stride: usize,
pub(crate) groups: usize,
}
pub(crate) struct LayerView<'a> {
pub(crate) norm1: (&'a [f32], &'a [f32]),
pub(crate) norm2: (&'a [f32], &'a [f32]),
pub(crate) in_proj: &'a [f32],
pub(crate) out_proj: &'a [f32],
pub(crate) lin1: &'a [f32],
pub(crate) lin2: &'a [f32],
pub(crate) ls1: &'a [f32],
pub(crate) ls2: &'a [f32],
}
impl<'a> From<&'a Conv1d> for ConvView<'a> {
fn from(c: &'a Conv1d) -> Self {
Self {
w: &c.w,
b: c.b.as_deref(),
in_c: c.in_c,
out_c: c.out_c,
k: c.k,
k_eff: c.k_eff(),
stride: c.stride,
replicate: c.replicate,
}
}
}
impl<'a> From<&'a ConvTr1d> for ConvTrView<'a> {
fn from(c: &'a ConvTr1d) -> Self {
Self {
w: &c.w,
b: c.b.as_deref(),
in_c: c.in_c,
out_c: c.out_c,
k: c.k,
stride: c.stride,
groups: c.groups,
}
}
}
impl Transformer {
fn views(&self) -> Vec<LayerView<'_>> {
self.layers
.iter()
.map(|l| LayerView {
norm1: (&l.norm1.0, &l.norm1.1),
norm2: (&l.norm2.0, &l.norm2.1),
in_proj: &l.in_proj.w,
out_proj: &l.out_proj.w,
lin1: &l.lin1.w,
lin2: &l.lin2.w,
ls1: &l.ls1,
ls2: &l.ls2,
})
.collect()
}
}
#[derive(Clone)]
pub(crate) struct NamedPl {
name: String,
pl: wgpu::ComputePipeline,
}
pub(crate) struct Builder<'a> {
ctx: &'a GpuCtx,
pls: HashMap<String, NamedPl>,
steps: Vec<Step>,
state: Vec<wgpu::Buffer>,
pos_rings: Vec<wgpu::Buffer>,
taps: HashMap<String, (wgpu::Buffer, usize)>,
dummy: wgpu::Buffer,
}
impl<'a> Builder<'a> {
pub(crate) fn new(ctx: &'a GpuCtx) -> Self {
Self {
ctx,
pls: HashMap::new(),
steps: Vec::new(),
state: Vec::new(),
pos_rings: Vec::new(),
taps: HashMap::new(),
dummy: ctx.storage(&[0f32]),
}
}
pub(crate) fn pl(&mut self, name: &str, src: &str) -> NamedPl {
if let Some(p) = self.pls.get(name) {
return p.clone();
}
let p = NamedPl {
name: name.to_string(),
pl: pipeline(self.ctx, name, src),
};
self.pls.insert(name.to_string(), p.clone());
p
}
pub(crate) fn u4(&self, a: u32, b: u32, c: u32, d: u32) -> wgpu::Buffer {
crate::forward::uni(self.ctx, bytemuck::cast_slice(&[a, b, c, d]))
}
pub(crate) fn step(&mut self, pl: &NamedPl, bufs: &[&wgpu::Buffer], gx: u32, gy: u32) {
let entries: Vec<wgpu::BindGroupEntry> = bufs
.iter()
.enumerate()
.map(|(i, b)| wgpu::BindGroupEntry {
binding: i as u32,
resource: b.as_entire_binding(),
})
.collect();
let bg = self
.ctx
.device
.create_bind_group(&wgpu::BindGroupDescriptor {
label: None,
layout: &pl.pl.get_bind_group_layout(0),
entries: &entries,
});
self.steps.push(Step {
pl: pl.pl.clone(),
bg,
gx,
gy,
label: pl.name.clone(),
});
}
#[allow(clippy::type_complexity)]
pub(crate) fn finish(
&mut self,
) -> (
Vec<Step>,
Vec<wgpu::Buffer>,
Vec<wgpu::Buffer>,
HashMap<String, (wgpu::Buffer, usize)>,
) {
(
std::mem::take(&mut self.steps),
std::mem::take(&mut self.state),
std::mem::take(&mut self.pos_rings),
std::mem::take(&mut self.taps),
)
}
pub(crate) fn tap(&mut self, name: &str, buf: &wgpu::Buffer, len: usize) {
self.taps.insert(name.to_string(), (buf.clone(), len));
}
fn bias_buf(&self, b: Option<&[f32]>, n: usize) -> wgpu::Buffer {
match b {
Some(v) => self.ctx.storage(v),
None => self.ctx.storage(&vec![0f32; n]),
}
}
pub(crate) fn conv(
&mut self,
cv: ConvView<'_>,
src: &wgpu::Buffer,
src_tm: bool,
t_in: usize,
frame_uni: Option<&wgpu::Buffer>,
acc_into: Option<&wgpu::Buffer>,
) -> (wgpu::Buffer, usize) {
let tp = cv.k_eff - cv.stride;
let t_out = t_in / cv.stride;
let ta = t_in + tp;
let w = self.ctx.storage(cv.w);
let b = self.bias_buf(cv.b, cv.out_c);
let acc = acc_into.is_some();
let y = match acc_into {
Some(buf) => buf.clone(),
None => self.ctx.empty(cv.out_c * t_out),
};
let (xa, eff_ta) = if tp == 0 && !src_tm {
(src.clone(), t_in) } else {
let xa = self.ctx.empty(cv.in_c * ta);
self.state.push(xa.clone());
let p = self.u4(cv.in_c as u32, t_in as u32, ta as u32, tp as u32);
let q = self.u4(src_tm as u32, 0, 0, 0);
let wpl = self.pl("mimi_write_xa", WRITE_XA);
self.step(
&wpl,
&[src, &xa, &p, &q],
((cv.in_c * t_in) as u32).div_ceil(256),
1,
);
if cv.replicate && tp > 0 {
let u = frame_uni.expect("replicate conv needs the frame uniform");
let rp = self.u4(cv.in_c as u32, 0, ta as u32, tp as u32);
let rpl = self.pl("mimi_replicate", REPLICATE);
self.step(
&rpl,
&[&xa, &rp, u],
((cv.in_c * tp) as u32).div_ceil(256),
1,
);
}
(xa, ta)
};
let p = self.u4(cv.in_c as u32, cv.k as u32, cv.stride as u32, t_out as u32);
let q = self.u4(eff_ta as u32, acc as u32, 0, 0);
let cpl = self.pl("mimi_conv1d", CONV1D);
self.step(
&cpl,
&[&w, &b, &xa, &y, &p, &q],
(t_out as u32).div_ceil(64),
cv.out_c as u32,
);
if tp > 0 {
let p = self.u4(cv.in_c as u32, t_in as u32, ta as u32, tp as u32);
let kpl = self.pl("mimi_carry", CARRY);
self.step(&kpl, &[&xa, &p], (cv.in_c as u32).div_ceil(64), 1);
}
(y, t_out)
}
pub(crate) fn convtr(
&mut self,
tr: ConvTrView<'_>,
src: &wgpu::Buffer,
t_in: usize,
) -> (wgpu::Buffer, usize) {
let (k, s) = (tr.k, tr.stride);
let tp = k - s;
let ts = t_in * s;
let full_t = (t_in - 1) * s + k;
let cw = tr.groups == tr.in_c && tr.groups == tr.out_c && tr.groups > 1;
assert!(cw || tr.groups == 1, "unsupported convtr grouping");
let w = self.ctx.storage(tr.w);
let b = self.bias_buf(tr.b, tr.out_c);
let full = self.ctx.empty(tr.out_c * full_t);
let part = self.ctx.empty(tr.out_c * tp);
self.state.push(part.clone());
let y = self.ctx.empty(tr.out_c * ts);
let p = self.u4(tr.in_c as u32, k as u32, s as u32, t_in as u32);
let q = self.u4(full_t as u32, 0, tr.out_c as u32, 0);
let name = format!("mimi_convtr_{}", if cw { "cw" } else { "full" });
let pl = {
let src_wgsl = convtr_src(cw, false);
self.pl(&name, &src_wgsl)
};
self.step(
&pl,
&[&w, &b, src, &full, &p, &q],
(full_t as u32).div_ceil(64),
tr.out_c as u32,
);
let p = self.u4(tr.out_c as u32, full_t as u32, tp as u32, 0);
let apl = self.pl("mimi_trp_add", TRP_ADD);
self.step(
&apl,
&[&full, &part, &p],
((tr.out_c * tp) as u32).div_ceil(256),
1,
);
let p = self.u4(tr.out_c as u32, full_t as u32, ts as u32, 0);
let opl = self.pl("mimi_tr_out", TR_OUT);
self.step(
&opl,
&[&full, &y, &p],
((tr.out_c * ts) as u32).div_ceil(256),
1,
);
let p = self.u4(tr.out_c as u32, full_t as u32, tp as u32, ts as u32);
let cpl = self.pl("mimi_tr_carry", TR_CARRY);
self.step(
&cpl,
&[&full, &b, &part, &p],
((tr.out_c * tp) as u32).div_ceil(256),
1,
);
(y, ts)
}
pub(crate) fn elu_to(&mut self, src: &wgpu::Buffer, n: usize) -> wgpu::Buffer {
let dst = self.ctx.empty(n);
let p = self.u4(n as u32, 0, 0, 0);
let pl = self.pl("mimi_elu_to", ELU_TO);
self.step(&pl, &[src, &dst, &p], (n as u32).div_ceil(256), 1);
dst
}
pub(crate) fn elu_ip(&mut self, buf: &wgpu::Buffer, n: usize) {
let p = self.u4(n as u32, 0, 0, 0);
let pl = self.pl("mimi_elu_ip", ELU_IP);
self.step(&pl, &[buf, &p], (n as u32).div_ceil(256), 1);
}
pub(crate) fn resblock(
&mut self,
c1: ConvView<'_>,
c2: ConvView<'_>,
src: &wgpu::Buffer,
c: usize,
t: usize,
) -> wgpu::Buffer {
let e1 = self.elu_to(src, c * t);
let (y1, _) = self.conv(c1, &e1, false, t, None, None);
self.elu_ip(&y1, (c / 2) * t);
let y2 = self.ctx.empty(c * t);
let p = self.u4(c as u32, t as u32, t as u32, 0);
let cpl = self.pl("mimi_tr_out", TR_OUT); self.step(&cpl, &[src, &y2, &p], ((c * t) as u32).div_ceil(256), 1);
let _ = self.conv(c2, &y1, false, t, None, Some(&y2));
y2
}
pub(crate) fn transformer(
&mut self,
layers: &[LayerView<'_>],
xt: &wgpu::Buffer,
t: usize,
uni: &wgpu::Buffer,
) {
self.transformer_shaped(
layers,
xt,
t,
uni,
TrShape {
dim: DIM,
heads: 8,
ffn: 2048,
ring: CONTEXT + t,
window: CONTEXT,
ln_eps: 1e-5,
},
);
}
pub(crate) fn transformer_shaped(
&mut self,
layers: &[LayerView<'_>],
xt: &wgpu::Buffer,
t: usize,
uni: &wgpu::Buffer,
shape: TrShape,
) -> wgpu::Buffer {
let TrShape {
dim,
heads,
ffn,
ring,
window,
ln_eps,
} = shape;
let h = self.ctx.empty(t * dim);
let qkv = self.ctx.empty(t * 3 * dim);
let q = self.ctx.empty(t * dim);
let attn = self.ctx.empty(t * dim);
let ffh = self.ctx.empty(t * ffn);
let pos_ring = {
let b = self.ctx.empty(ring);
self.pos_rings.push(b.clone());
b
};
let ln = if dim == DIM && ln_eps == 1e-5 {
self.pl("mimi_ln512", LN512)
} else {
self.pl(&format!("ln_{dim}"), &ln_src(dim, ln_eps))
};
let mv = self.pl("mimi_matvec", MATVEC);
let rp = self.pl(
&format!("rope_push_{dim}_{heads}_{ring}"),
&rope_push_src(dim, heads, ring),
);
let at = self.pl(
&format!("attn_{dim}_{heads}_{ring}_{window}"),
&attn_src(dim, heads, ring, window),
);
let mls = self.pl("mimi_mv_ls", MV_LS);
let ge = self.pl("mimi_gelu", GELU_ERF);
for layer in layers {
let kr = self.ctx.empty(ring * dim);
let vr = self.ctx.empty(ring * dim);
self.state.push(kr.clone());
self.state.push(vr.clone());
let n1w = self.ctx.storage(layer.norm1.0);
let n1b = self.ctx.storage(layer.norm1.1);
let n2w = self.ctx.storage(layer.norm2.0);
let n2b = self.ctx.storage(layer.norm2.1);
let win = self.ctx.storage(layer.in_proj);
let wout = self.ctx.storage(layer.out_proj);
let w1 = self.ctx.storage(layer.lin1);
let w2 = self.ctx.storage(layer.lin2);
let ls1 = self.ctx.storage(layer.ls1);
let ls2 = self.ctx.storage(layer.ls2);
self.step(&ln, &[xt, &n1w, &n1b, &h], t as u32, 1);
let p = self.u4(3 * dim as u32, dim as u32, 0, 0);
self.step(
&mv,
&[&win, &h, &qkv, &p],
(3 * dim as u32).div_ceil(64),
t as u32,
);
let p = self.u4(t as u32, 0, 0, 0);
self.step(
&rp,
&[&qkv, &q, &kr, &vr, &pos_ring, &p, uni],
((t * dim / 2) as u32).div_ceil(64),
1,
);
self.step(
&at,
&[&q, &kr, &vr, &pos_ring, &attn, uni],
t as u32,
heads as u32,
);
let p = self.u4(dim as u32, dim as u32, 0, 0);
self.step(
&mls,
&[&wout, &attn, &ls1, xt, &p],
(dim as u32).div_ceil(64),
t as u32,
);
self.step(&ln, &[xt, &n2w, &n2b, &h], t as u32, 1);
let p = self.u4(ffn as u32, dim as u32, 0, 0);
self.step(
&mv,
&[&w1, &h, &ffh, &p],
(ffn as u32).div_ceil(64),
t as u32,
);
let p = self.u4((t * ffn) as u32, 0, 0, 0);
self.step(&ge, &[&ffh, &p], ((t * ffn) as u32).div_ceil(256), 1);
let p = self.u4(dim as u32, ffn as u32, 0, 0);
self.step(
&mls,
&[&w2, &ffh, &ls2, xt, &p],
(dim as u32).div_ceil(64),
t as u32,
);
}
h
}
}
pub struct MimiGpu {
enc_steps: Vec<Step>,
dec_steps: Vec<Step>,
pcm_in: wgpu::Buffer,
codes_out: wgpu::Buffer,
codes_in: wgpu::Buffer,
pcm_out: wgpu::Buffer,
enc_uni: wgpu::Buffer,
dec_uni: wgpu::Buffer,
state: Vec<wgpu::Buffer>,
pos_rings: Vec<wgpu::Buffer>,
taps: HashMap<String, (wgpu::Buffer, usize)>,
emb_pl: NamedPl,
cb_first: wgpu::Buffer,
cb_rest: wgpu::Buffer,
acc_f: wgpu::Buffer,
acc_r: wgpu::Buffer,
enc_pos: u32,
dec_pos: u32,
first: bool,
}
impl MimiGpu {
pub fn new(ctx: &GpuCtx, m: &Mimi) -> Result<Self> {
let pcm_in = ctx.empty(FRAME);
let codes_out = ctx.empty(NQ);
let codes_in = ctx.empty(NQ);
let enc_uni = crate::forward::uni(ctx, bytemuck::cast_slice(&[0u32, 1, 0, 0]));
let dec_uni = crate::forward::uni(ctx, bytemuck::cast_slice(&[0u32, 0, 0, 0]));
let mut b = Builder::new(ctx);
let (mut y, mut t) = b.conv((&m.enc.init).into(), &pcm_in, false, FRAME, None, None);
let mut c = m.enc.init.out_c;
for (rb, down) in &m.enc.blocks {
y = b.resblock((&rb.c1).into(), (&rb.c2).into(), &y, c, t);
let e = b.elu_to(&y, c * t);
let (dy, dt) = b.conv(down.into(), &e, false, t, None, None);
y = dy;
t = dt;
c = down.out_c;
}
let e = b.elu_to(&y, c * t);
let (y, t) = b.conv((&m.enc.last).into(), &e, false, t, None, None);
b.tap("enc_seanet", &y, DIM * t);
let xt = ctx.empty(t * DIM);
{
let p = b.u4(DIM as u32, t as u32, 0, 0);
let pl = b.pl("mimi_ct2tm", CT2TM);
b.step(&pl, &[&y, &xt, &p], ((DIM * t) as u32).div_ceil(256), 1);
}
b.transformer(&m.enc_tr.views(), &xt, t, &enc_uni);
b.tap("enc_tr", &xt, t * DIM);
let (latent, lt) = b.conv((&m.down).into(), &xt, true, t, Some(&enc_uni), None);
assert_eq!(lt, 1, "one latent frame per 80 ms");
b.tap("enc_latent", &latent, DIM);
let cb_first = ctx.storage(&m.rvq_first.codebooks.concat());
let cb_rest = ctx.storage(&m.rvq_rest.codebooks.concat());
let w_in_f = ctx.storage(&m.rvq_first.in_proj);
let w_in_r = ctx.storage(&m.rvq_rest.in_proj);
let r_first = ctx.empty(VQ_DIM);
let r_rest = ctx.empty(VQ_DIM);
{
let mv = b.pl("mimi_matvec", MATVEC);
let p = b.u4(VQ_DIM as u32, DIM as u32, 0, 0);
b.step(
&mv,
&[&w_in_f, &latent, &r_first, &p],
(VQ_DIM as u32).div_ceil(64),
1,
);
let p = b.u4(VQ_DIM as u32, DIM as u32, 0, 0);
b.step(
&mv,
&[&w_in_r, &latent, &r_rest, &p],
(VQ_DIM as u32).div_ceil(64),
1,
);
let ne = b.pl("mimi_nearest", NEAREST);
let se = b.pl("mimi_sub_emb", SUB_EMB);
let p = b.u4(0, 0, 0, 0);
b.step(&ne, &[&cb_first, &r_first, &codes_out, &p], 1, 1);
for l in 0..NQ - 1 {
let p = b.u4((l * BINS) as u32, (1 + l) as u32, 0, 0);
b.step(&ne, &[&cb_rest, &r_rest, &codes_out, &p], 1, 1);
if l < NQ - 2 {
let p = b.u4((l * BINS) as u32, (1 + l) as u32, 0, 0);
b.step(
&se,
&[&cb_rest, &codes_out, &r_rest, &p],
(VQ_DIM as u32).div_ceil(64),
1,
);
}
}
}
let enc_steps = std::mem::take(&mut b.steps);
let mut state = std::mem::take(&mut b.state);
let mut pos_rings = std::mem::take(&mut b.pos_rings);
let mut taps = std::mem::take(&mut b.taps);
let mut b = Builder::new(ctx);
let acc_f = ctx.empty(VQ_DIM);
let acc_r = ctx.empty(VQ_DIM);
let latent_d = ctx.empty(DIM);
let emb_pl = b.pl("mimi_emb_sum", EMB_SUM);
{
let es = emb_pl.clone();
let p = b.u4(1, 0, 0, 0);
b.step(
&es,
&[&cb_first, &codes_in, &acc_f, &p],
(VQ_DIM as u32).div_ceil(64),
1,
);
let p = b.u4((NQ - 1) as u32, 1, 0, 0);
b.step(
&es,
&[&cb_rest, &codes_in, &acc_r, &p],
(VQ_DIM as u32).div_ceil(64),
1,
);
let w_out_f = ctx.storage(&m.rvq_first.out_proj);
let w_out_r = ctx.storage(&m.rvq_rest.out_proj);
let mv = b.pl("mimi_matvec", MATVEC);
let p = b.u4(DIM as u32, VQ_DIM as u32, 0, 0);
b.step(
&mv,
&[&w_out_f, &acc_f, &latent_d, &p],
(DIM as u32).div_ceil(64),
1,
);
let p = b.u4(DIM as u32, VQ_DIM as u32, 1, 0);
b.step(
&mv,
&[&w_out_r, &acc_r, &latent_d, &p],
(DIM as u32).div_ceil(64),
1,
);
}
b.tap("dec_latent", &latent_d, DIM);
let (y_up, t_up) = b.convtr((&m.up).into(), &latent_d, 1);
b.tap("dec_up", &y_up, DIM * t_up);
let xt_d = ctx.empty(t_up * DIM);
{
let p = b.u4(DIM as u32, t_up as u32, 0, 0);
let pl = b.pl("mimi_ct2tm", CT2TM);
b.step(
&pl,
&[&y_up, &xt_d, &p],
((DIM * t_up) as u32).div_ceil(256),
1,
);
}
b.transformer(&m.dec_tr.views(), &xt_d, t_up, &dec_uni);
b.tap("dec_tr", &xt_d, t_up * DIM);
let y_ct = ctx.empty(DIM * t_up);
{
let p = b.u4(DIM as u32, t_up as u32, 0, 0);
let pl = b.pl("mimi_tm2ct", TM2CT);
b.step(
&pl,
&[&xt_d, &y_ct, &p],
((DIM * t_up) as u32).div_ceil(256),
1,
);
}
let (mut y, mut t) = b.conv((&m.dec.init).into(), &y_ct, false, t_up, None, None);
let mut c = m.dec.init.out_c;
for (up, rb) in &m.dec.blocks {
b.elu_ip(&y, c * t);
let (uy, ut) = b.convtr(up.into(), &y, t);
c = up.out_c;
t = ut;
y = b.resblock((&rb.c1).into(), (&rb.c2).into(), &uy, c, t);
}
b.elu_ip(&y, c * t);
let (pcm_out, t_pcm) = b.conv((&m.dec.last).into(), &y, false, t, None, None);
assert_eq!(t_pcm, FRAME, "decode must emit exactly one frame");
let dec_steps = std::mem::take(&mut b.steps);
state.extend(std::mem::take(&mut b.state));
pos_rings.extend(std::mem::take(&mut b.pos_rings));
taps.extend(std::mem::take(&mut b.taps));
for b in &pos_rings {
let inval = vec![INVALID_POS; (b.size() / 4) as usize];
ctx.queue.write_buffer(b, 0, bytemuck::cast_slice(&inval));
}
Ok(Self {
enc_steps,
dec_steps,
pcm_in,
codes_out,
codes_in,
pcm_out,
enc_uni,
dec_uni,
state,
pos_rings,
taps,
emb_pl,
cb_first,
cb_rest,
acc_f,
acc_r,
enc_pos: 0,
dec_pos: 0,
first: true,
})
}
pub fn bind_lm_tokens(&mut self, ctx: &GpuCtx, tok: &wgpu::Buffer) {
let mk = |cb: &wgpu::Buffer, acc: &wgpu::Buffer, n_cb: u32, off: u32| {
let p = crate::forward::uni(ctx, bytemuck::cast_slice(&[n_cb, off, 0u32, 0]));
let bufs = [cb, tok, acc, &p];
let entries: Vec<wgpu::BindGroupEntry> = bufs
.iter()
.enumerate()
.map(|(i, b)| wgpu::BindGroupEntry {
binding: i as u32,
resource: b.as_entire_binding(),
})
.collect();
ctx.device.create_bind_group(&wgpu::BindGroupDescriptor {
label: None,
layout: &self.emb_pl.pl.get_bind_group_layout(0),
entries: &entries,
})
};
self.dec_steps[0].bg = mk(&self.cb_first, &self.acc_f, 1, 1);
self.dec_steps[1].bg = mk(&self.cb_rest, &self.acc_r, (NQ - 1) as u32, 2);
}
pub fn submit_decode(&mut self, ctx: &GpuCtx) {
ctx.queue.write_buffer(
&self.dec_uni,
0,
bytemuck::cast_slice(&[self.dec_pos, 0u32, 0, 0]),
);
Self::run(ctx, &self.dec_steps);
self.dec_pos += 2;
}
pub fn read_decoded(&self, ctx: &GpuCtx) -> Result<Vec<f32>> {
ctx.read(&self.pcm_out, FRAME)
}
pub(crate) fn run(ctx: &GpuCtx, steps: &[Step]) {
let mut enc = ctx
.device
.create_command_encoder(&wgpu::CommandEncoderDescriptor::default());
{
let mut p = enc.begin_compute_pass(&wgpu::ComputePassDescriptor::default());
for s in steps {
p.set_pipeline(&s.pl);
p.set_bind_group(0, &s.bg, &[]);
p.dispatch_workgroups(s.gx, s.gy, 1);
}
}
ctx.queue.submit([enc.finish()]);
}
pub fn encode_frame(&mut self, ctx: &GpuCtx, frame: &[f32]) -> Result<[u32; NQ]> {
self.submit_encode(ctx, frame);
self.read_codes(ctx)
}
pub fn submit_encode(&mut self, ctx: &GpuCtx, frame: &[f32]) {
assert_eq!(frame.len(), FRAME);
ctx.queue
.write_buffer(&self.pcm_in, 0, bytemuck::cast_slice(frame));
ctx.queue.write_buffer(
&self.enc_uni,
0,
bytemuck::cast_slice(&[self.enc_pos, self.first as u32, 0, 0]),
);
Self::run(ctx, &self.enc_steps);
self.enc_pos += 2;
self.first = false;
}
pub fn read_codes(&self, ctx: &GpuCtx) -> Result<[u32; NQ]> {
let codes = ctx.read_u32(&self.codes_out, NQ)?;
let mut out = [0u32; NQ];
out.copy_from_slice(&codes);
Ok(out)
}
pub fn decode_frame(&mut self, ctx: &GpuCtx, codes: &[u32; NQ]) -> Result<Vec<f32>> {
ctx.queue
.write_buffer(&self.codes_in, 0, bytemuck::cast_slice(codes));
self.submit_decode(ctx);
self.read_decoded(ctx)
}
pub fn reset(&mut self, ctx: &GpuCtx) {
let mut enc = ctx
.device
.create_command_encoder(&wgpu::CommandEncoderDescriptor::default());
for b in &self.state {
enc.clear_buffer(b, 0, None);
}
ctx.queue.submit([enc.finish()]);
for b in &self.pos_rings {
let inval = vec![INVALID_POS; (b.size() / 4) as usize];
ctx.queue.write_buffer(b, 0, bytemuck::cast_slice(&inval));
}
self.enc_pos = 0;
self.dec_pos = 0;
self.first = true;
}
pub fn bench_steps(
&mut self,
ctx: &GpuCtx,
encode: bool,
reps: usize,
) -> Vec<(String, f64, usize)> {
let steps = if encode {
&self.enc_steps
} else {
&self.dec_steps
};
let mut agg: HashMap<String, (f64, usize)> = HashMap::new();
for _ in 0..reps {
for s in steps {
let t0 = std::time::Instant::now();
let mut enc = ctx
.device
.create_command_encoder(&wgpu::CommandEncoderDescriptor::default());
{
let mut p = enc.begin_compute_pass(&wgpu::ComputePassDescriptor::default());
p.set_pipeline(&s.pl);
p.set_bind_group(0, &s.bg, &[]);
p.dispatch_workgroups(s.gx, s.gy, 1);
}
ctx.queue.submit([enc.finish()]);
let _ = ctx.device.poll(wgpu::PollType::wait_indefinitely());
let e = agg.entry(s.label.clone()).or_insert((0.0, 0));
e.0 += t0.elapsed().as_secs_f64() * 1e3;
e.1 += 1;
}
}
let mut v: Vec<(String, f64, usize)> = agg
.into_iter()
.map(|(k, (ms, n))| (k, ms / reps as f64, n / reps))
.collect();
v.sort_by(|a, b| b.1.total_cmp(&a.1));
v
}
pub fn debug_read(&self, ctx: &GpuCtx, name: &str) -> Result<Vec<f32>> {
let (buf, len) = self
.taps
.get(name)
.ok_or_else(|| anyhow::anyhow!("no tap {name}"))?;
ctx.read(buf, *len)
}
}