#[cfg(not(target_arch = "wasm32"))]
use rayon::prelude::*;
pub fn add(a: &[f32], b: &[f32]) -> Vec<f32> {
a.iter().zip(b).map(|(x, y)| x + y).collect()
}
pub fn gelu(x: &[f32]) -> Vec<f32> {
const C: f32 = 0.797_884_6; x.iter()
.map(|&v| 0.5 * v * (1.0 + (C * (v + 0.044715 * v * v * v)).tanh()))
.collect()
}
#[allow(clippy::too_many_arguments)]
pub fn matmul(
a: &[f32],
b: &[f32],
bias: Option<&[f32]>,
m: usize,
k: usize,
n: usize,
batch: usize,
a_stride: usize,
b_stride: usize,
trans_a: bool,
trans_b: bool,
alpha: f32,
) -> Vec<f32> {
let mut out = vec![0.0f32; batch * m * n];
#[cfg(not(target_arch = "wasm32"))]
let rows = out.par_chunks_mut(n);
#[cfg(target_arch = "wasm32")]
let rows = out.chunks_mut(n);
rows.enumerate().for_each(|(row_idx, orow)| {
let bat = row_idx / m;
let i = row_idx % m;
let a_base = bat * a_stride;
let b_base = bat * b_stride;
let a_at = |kk: usize| {
if trans_a {
a[a_base + kk * m + i]
} else {
a[a_base + i * k + kk]
}
};
if trans_b {
for (j, o) in orow.iter_mut().enumerate() {
let b_row = &b[b_base + j * k..b_base + (j + 1) * k];
let dot: f32 = b_row
.iter()
.enumerate()
.map(|(kk, &bv)| a_at(kk) * bv)
.sum();
*o = dot * alpha;
}
} else {
for kk in 0..k {
let av = a_at(kk);
let b_row = &b[b_base + kk * n..b_base + (kk + 1) * n];
for (o, &bv) in orow.iter_mut().zip(b_row) {
*o += av * bv;
}
}
if alpha != 1.0 {
for o in orow.iter_mut() {
*o *= alpha;
}
}
}
if let Some(bias) = bias {
for (o, &bb) in orow.iter_mut().zip(bias) {
*o += bb;
}
}
});
out
}
pub fn kv_append(
dst: &mut [f32],
src: &[f32],
h: usize,
t: usize,
hd: usize,
cap: usize,
len: usize,
) {
for hh in 0..h {
for tt in 0..t {
let d0 = hh * cap * hd + (len + tt) * hd;
let s0 = hh * t * hd + tt * hd;
dst[d0..d0 + hd].copy_from_slice(&src[s0..s0 + hd]);
}
}
}
pub fn softmax(
x: &[f32],
rows: usize,
cols: usize,
q_len: usize,
causal: bool,
off: usize,
) -> Vec<f32> {
let mut out = vec![0.0f32; rows * cols];
for r in 0..rows {
let base = r * cols;
let visible = |j: usize| !causal || j <= (r % q_len) + off;
let mut max = f32::NEG_INFINITY;
for j in 0..cols {
if visible(j) {
max = max.max(x[base + j]);
}
}
let mut sum = 0.0f32;
for j in 0..cols {
if visible(j) {
sum += (x[base + j] - max).exp();
}
}
for j in 0..cols {
if visible(j) {
out[base + j] = (x[base + j] - max).exp() / sum;
}
}
}
out
}
pub fn layernorm(
x: &[f32],
gamma: &[f32],
beta: &[f32],
rows: usize,
cols: usize,
eps: f32,
) -> Vec<f32> {
let mut out = vec![0.0f32; rows * cols];
let nf = cols as f32;
for r in 0..rows {
let row = &x[r * cols..(r + 1) * cols];
let mean: f32 = row.iter().sum::<f32>() / nf;
let var: f32 = row.iter().map(|&v| (v - mean) * (v - mean)).sum::<f32>() / nf;
let inv_std = 1.0 / (var + eps).sqrt();
for j in 0..cols {
out[r * cols + j] = (row[j] - mean) * inv_std * gamma[j] + beta[j];
}
}
out
}
pub fn embedding(ids: &[u32], wte: &[f32], wpe: Option<&[f32]>, c: usize, pos: usize) -> Vec<f32> {
let mut out = vec![0.0f32; ids.len() * c];
for (t, &id) in ids.iter().enumerate() {
let dst = &mut out[t * c..(t + 1) * c];
dst.copy_from_slice(&wte[id as usize * c..(id as usize + 1) * c]);
if let Some(wpe) = wpe {
for (d, &pv) in dst.iter_mut().zip(&wpe[(t + pos) * c..(t + pos + 1) * c]) {
*d += pv;
}
}
}
out
}
pub fn split_heads(qkv: &[f32], t: usize, c: usize, h: usize) -> (Vec<f32>, Vec<f32>, Vec<f32>) {
let hd = c / h;
let mut q = vec![0.0f32; t * c];
let mut k = vec![0.0f32; t * c];
let mut v = vec![0.0f32; t * c];
for hh in 0..h {
for tt in 0..t {
for d in 0..hd {
let dst = hh * t * hd + tt * hd + d;
let src_row = tt * 3 * c;
let col = hh * hd + d;
q[dst] = qkv[src_row + col];
k[dst] = qkv[src_row + c + col];
v[dst] = qkv[src_row + 2 * c + col];
}
}
}
(q, k, v)
}
pub fn merge_heads(x: &[f32], t: usize, c: usize, h: usize) -> Vec<f32> {
let hd = c / h;
let mut out = vec![0.0f32; t * c];
for tt in 0..t {
for hh in 0..h {
for d in 0..hd {
out[tt * c + hh * hd + d] = x[hh * t * hd + tt * hd + d];
}
}
}
out
}
pub fn gelu_bwd(x: &[f32], dy: &[f32]) -> Vec<f32> {
const C: f32 = 0.797_884_6; const A: f32 = 0.044715;
x.iter()
.zip(dy)
.map(|(&v, &g)| {
let u = C * (v + A * v * v * v);
let th = u.tanh();
let sech2 = 1.0 - th * th;
let d = 0.5 * (1.0 + th) + 0.5 * v * sech2 * C * (1.0 + 3.0 * A * v * v);
d * g
})
.collect()
}
pub fn softmax_bwd(y: &[f32], dy: &[f32], rows: usize, cols: usize) -> Vec<f32> {
let mut out = vec![0.0f32; rows * cols];
for r in 0..rows {
let base = r * cols;
let s: f32 = (0..cols).map(|j| y[base + j] * dy[base + j]).sum();
for j in 0..cols {
out[base + j] = y[base + j] * (dy[base + j] - s);
}
}
out
}
pub fn layernorm_bwd_dx(
x: &[f32],
gamma: &[f32],
dy: &[f32],
rows: usize,
cols: usize,
eps: f32,
) -> Vec<f32> {
let nf = cols as f32;
let mut out = vec![0.0f32; rows * cols];
for r in 0..rows {
let row = &x[r * cols..(r + 1) * cols];
let mean: f32 = row.iter().sum::<f32>() / nf;
let var: f32 = row.iter().map(|&v| (v - mean) * (v - mean)).sum::<f32>() / nf;
let inv_std = 1.0 / (var + eps).sqrt();
let mut s1 = 0.0f32; let mut s2 = 0.0f32; for j in 0..cols {
let gd = gamma[j] * dy[r * cols + j];
s1 += gd;
s2 += gd * (row[j] - mean) * inv_std;
}
for j in 0..cols {
let xhat = (row[j] - mean) * inv_std;
let gd = gamma[j] * dy[r * cols + j];
out[r * cols + j] = inv_std * (gd - s1 / nf - xhat * s2 / nf);
}
}
out
}
pub fn layernorm_bwd_dparams(
x: &[f32],
dy: &[f32],
rows: usize,
cols: usize,
eps: f32,
) -> (Vec<f32>, Vec<f32>) {
let nf = cols as f32;
let mut dgamma = vec![0.0f32; cols];
let mut dbeta = vec![0.0f32; cols];
for r in 0..rows {
let row = &x[r * cols..(r + 1) * cols];
let mean: f32 = row.iter().sum::<f32>() / nf;
let var: f32 = row.iter().map(|&v| (v - mean) * (v - mean)).sum::<f32>() / nf;
let inv_std = 1.0 / (var + eps).sqrt();
for j in 0..cols {
let d = dy[r * cols + j];
dgamma[j] += d * (row[j] - mean) * inv_std;
dbeta[j] += d;
}
}
(dgamma, dbeta)
}
pub fn sum_rows(x: &[f32], rows: usize, cols: usize) -> Vec<f32> {
let mut out = vec![0.0f32; cols];
for r in 0..rows {
for j in 0..cols {
out[j] += x[r * cols + j];
}
}
out
}
pub fn scatter_add_rows(dst: &mut [f32], ids: &[u32], src: &[f32], c: usize) {
for (r, &id) in ids.iter().enumerate() {
let d0 = id as usize * c;
for j in 0..c {
dst[d0 + j] += src[r * c + j];
}
}
}
pub fn gather_nll(probs: &[f32], ids: &[u32], cols: usize) -> Vec<f32> {
ids.iter()
.enumerate()
.map(|(r, &id)| -probs[r * cols + id as usize].max(f32::MIN_POSITIVE).ln())
.collect()
}
pub fn ce_bwd(probs: &[f32], ids: &[u32], rows: usize, cols: usize, scale: f32) -> Vec<f32> {
let mut out = vec![0.0f32; rows * cols];
for r in 0..rows {
for j in 0..cols {
let onehot = (j as u32 == ids[r]) as u32 as f32;
out[r * cols + j] = (probs[r * cols + j] - onehot) * scale;
}
}
out
}
pub fn pcg_hash(x: u32) -> u32 {
let state = x.wrapping_mul(747796405).wrapping_add(2891336453);
let word = ((state >> ((state >> 28) + 4)) ^ state).wrapping_mul(277803737);
(word >> 22) ^ word
}
pub fn dropout(x: &[f32], p: f32, scale: f32, seed: u32) -> Vec<f32> {
x.iter()
.enumerate()
.map(|(i, &v)| {
let r = pcg_hash(seed ^ (i as u32).wrapping_mul(0x9E37_79B9));
let u = (r >> 8) as f32 / 16_777_216.0; if u >= p { v * scale } else { 0.0 }
})
.collect()
}
pub fn unsplit_head(d: &[f32], t: usize, c: usize, h: usize, which: usize) -> Vec<f32> {
let hd = c / h;
let mut out = vec![0.0f32; t * 3 * c];
for hh in 0..h {
for tt in 0..t {
for dd in 0..hd {
out[tt * 3 * c + which * c + hh * hd + dd] = d[hh * t * hd + tt * hd + dd];
}
}
}
out
}
pub fn unmerge_heads(dy: &[f32], t: usize, c: usize, h: usize) -> Vec<f32> {
let hd = c / h;
let mut out = vec![0.0f32; t * c];
for hh in 0..h {
for tt in 0..t {
for dd in 0..hd {
out[hh * t * hd + tt * hd + dd] = dy[tt * c + hh * hd + dd];
}
}
}
out
}
#[allow(clippy::too_many_arguments)]
pub fn adamw(
param: &mut [f32],
grad: &[f32],
m: &mut [f32],
v: &mut [f32],
lr: f32,
beta1: f32,
beta2: f32,
eps: f32,
weight_decay: f32,
step: u32,
) {
let bc1 = 1.0 - beta1.powi(step as i32);
let bc2 = 1.0 - beta2.powi(step as i32);
for i in 0..param.len() {
m[i] = beta1 * m[i] + (1.0 - beta1) * grad[i];
v[i] = beta2 * v[i] + (1.0 - beta2) * grad[i] * grad[i];
let mhat = m[i] / bc1;
let vhat = v[i] / bc2;
param[i] -= lr * (mhat / (vhat.sqrt() + eps) + weight_decay * param[i]);
}
}