use super::*;
#[inline(always)]
pub fn softmax(x: &mut [f32]) {
if x.is_empty() {
return;
}
let max_val = crate::simd::simd_max_f32(x);
crate::simd::simd_add_scalar_inplace(x, -max_val);
let sum: f32 = crate::simd::simd_exp_sum_inplace(x);
let inv_sum = 1.0 / sum;
crate::simd::simd_scale_inplace(x, inv_sum);
}
#[inline(always)]
pub fn softmax_scaled(x: &mut [f32], inv_temp: f32) {
if x.is_empty() {
return;
}
let max_val = crate::simd::simd_max_f32(x);
crate::simd::simd_fused_sub_scale_inplace(x, max_val, inv_temp);
let sum: f32 = crate::simd::simd_exp_sum_inplace(x);
let inv_sum = 1.0 / sum;
crate::simd::simd_scale_inplace(x, inv_sum);
}
#[inline(always)]
pub fn rmsnorm(x: &mut [f32]) {
if x.is_empty() {
return;
}
let sum_sq = crate::simd::simd_sum_sq(x, x.len());
let inv_rms = 1.0 / (sum_sq / x.len() as f32 + 1e-5f32).sqrt();
crate::simd::simd_scale_inplace(x, inv_rms);
}
#[inline(always)]
pub fn gegelu(hidden: &mut [f32], gate: &[f32], up: &[f32]) {
const CHUNK: usize = 64;
let mut buf = [0.0f32; CHUNK];
let mut i = 0;
while i + CHUNK <= hidden.len() {
crate::simd::simd_fused_decay_write(&mut buf, 0.0, &gate[i..i + CHUNK], -1.702);
crate::simd::simd_exp_inplace(&mut buf);
crate::simd::simd_add_scalar_inplace(&mut buf, 1.0);
crate::simd::simd_reciprocal_inplace(&mut buf);
for j in 0..CHUNK {
hidden[i + j] = gate[i + j] * up[i + j];
}
crate::simd::simd_scale_mul_inplace(&mut hidden[i..i + CHUNK], &buf, 1.0);
i += CHUNK;
}
for i in i..hidden.len() {
let g = gate[i];
let sigmoid = 1.0 / (1.0 + (-1.702 * g).exp());
hidden[i] = g * sigmoid * up[i];
}
}
#[inline(always)]
pub fn gegelu_tanh(hidden: &mut [f32], gate: &[f32], up: &[f32]) {
const CHUNK: usize = 64;
const SQRT_2_OVER_PI: f32 = 0.797_884_6; const SCALE_2: f32 = 1.595_769_2; let mut buf = [0.0f32; CHUNK];
let mut buf2 = [0.0f32; CHUNK];
let mut i = 0;
while i + CHUNK <= hidden.len() {
crate::simd::simd_fused_decay_write(&mut buf, 0.0, &gate[i..i + CHUNK], 0.044715);
crate::simd::simd_scale_mul_inplace(&mut buf, &gate[i..i + CHUNK], 1.0); crate::simd::simd_add_scalar_inplace(&mut buf, 1.0);
crate::simd::simd_scale_mul_inplace(&mut buf, &gate[i..i + CHUNK], SCALE_2);
crate::simd::simd_exp_inplace(&mut buf);
buf2[..CHUNK].copy_from_slice(&buf);
crate::simd::simd_add_scalar_inplace(&mut buf2, 1.0); for j in 0..CHUNK {
buf[j] /= buf2[j];
hidden[i + j] = gate[i + j] * up[i + j];
}
crate::simd::simd_scale_mul_inplace(&mut hidden[i..i + CHUNK], &buf, 1.0);
i += CHUNK;
}
for i in i..hidden.len() {
let g = gate[i];
let inner = SQRT_2_OVER_PI * (g + 0.044715 * g * g * g);
let gelu_val = 0.5 * g * (1.0 + inner.tanh());
hidden[i] = gelu_val * up[i];
}
}
#[inline(always)]
pub fn silu(x: &mut [f32]) {
const CHUNK: usize = 64;
let mut buf = [0.0f32; CHUNK];
let mut i = 0;
while i + CHUNK <= x.len() {
crate::simd::simd_fused_decay_write(&mut buf, 0.0, &x[i..i + CHUNK], -1.0);
crate::simd::simd_exp_inplace(&mut buf);
crate::simd::simd_add_scalar_inplace(&mut buf, 1.0);
crate::simd::simd_reciprocal_inplace(&mut buf);
crate::simd::simd_scale_mul_inplace(&mut x[i..i + CHUNK], &buf, 1.0);
i += CHUNK;
}
for v in x[i..].iter_mut() {
*v = *v / (1.0 + (-*v).exp());
}
}
#[inline(always)]
pub fn swiglu(hidden: &mut [f32], gate: &[f32], up: &[f32]) {
const CHUNK: usize = 64;
let mut buf = [0.0f32; CHUNK];
let mut i = 0;
while i + CHUNK <= hidden.len() {
crate::simd::simd_fused_decay_write(&mut buf, 0.0, &gate[i..i + CHUNK], -1.0);
crate::simd::simd_exp_inplace(&mut buf);
crate::simd::simd_add_scalar_inplace(&mut buf, 1.0);
crate::simd::simd_reciprocal_inplace(&mut buf);
for j in 0..CHUNK {
hidden[i + j] = gate[i + j] * up[i + j];
}
crate::simd::simd_scale_mul_inplace(&mut hidden[i..i + CHUNK], &buf, 1.0);
i += CHUNK;
}
for i in i..hidden.len() {
let g = gate[i];
hidden[i] = g / (1.0 + (-g).exp()) * up[i];
}
}
#[inline(always)]
pub fn rmsnorm_with_gamma(x: &mut [f32], gamma: &[f32]) {
rmsnorm_with_gamma_eps(x, gamma, 1e-5)
}
#[inline(always)]
pub fn rmsnorm_with_gamma_eps(x: &mut [f32], gamma: &[f32], eps: f64) {
let n = x.len();
if n == 0 {
return;
}
let sum_sq = crate::simd::simd_sum_sq(x, n);
let inv_rms = 1.0 / (sum_sq / n as f32 + eps as f32).sqrt();
crate::simd::simd_scale_mul_inplace(x, gamma, inv_rms);
}
#[inline(always)]
pub fn matmul(output: &mut [f32], weight: &[f32], input: &[f32], rows: usize, cols: usize) {
crate::simd::simd_matmul_rows(output, weight, input, rows, cols);
}
#[inline(always)]
pub fn matmul_parallel(
output: &mut [f32],
weight: &[f32],
input: &[f32],
rows: usize,
cols: usize,
) {
crate::simd::simd_matmul_rows_parallel(output, weight, input, rows, cols);
}
#[inline(always)]
pub fn matmul_relu(output: &mut [f32], weight: &[f32], input: &[f32], rows: usize, cols: usize) {
crate::simd::simd_matmul_relu_rows(output, weight, input, rows, cols);
}
#[inline(always)]
pub fn matmul_f16(
output: &mut [f32],
weight: &[half::f16],
input: &[f32],
rows: usize,
cols: usize,
) {
crate::simd::simd_matmul_f16_f32_rows(output, weight, input, rows, cols);
}
#[inline(always)]
pub fn matmul_f16_parallel(
output: &mut [f32],
weight: &[half::f16],
input: &[f32],
rows: usize,
cols: usize,
) {
crate::simd::simd_matmul_f16_f32_rows_parallel(output, weight, input, rows, cols);
}
#[cfg(feature = "sparse_mlp")]
#[inline(always)]
pub fn sparse_matmul(
output: &mut [f32],
weight: &[f32],
input: &[f32],
rows: usize,
cols: usize,
active_indices: &mut [usize],
active_values: &mut [f32],
) -> usize {
let mut alive = 0;
for c in 0..cols {
let val = unsafe { *input.get_unchecked(c) };
if val > 0.0 {
unsafe {
*active_indices.get_unchecked_mut(alive) = c;
*active_values.get_unchecked_mut(alive) = val;
}
alive += 1;
}
}
crate::simd::simd_sparse_matmul_rows(
output,
weight,
active_indices,
active_values,
rows,
cols,
alive,
);
alive
}
#[deprecated(
since = "0.1.0",
note = "allocates a vocab-sized Vec per call; use `sample_token_into` on hot paths"
)]
pub fn sample_token(probs: &[f32], rng: &mut Rng) -> usize {
let mut r = rng.uniform();
while r == 0.0 {
r = rng.uniform();
}
let n = probs.len();
if n == 0 {
return 0;
}
let mut cdf = vec![0.0f32; n];
let mut sum = 0.0f32;
for (i, &p) in probs.iter().enumerate() {
sum += p;
unsafe {
*cdf.get_unchecked_mut(i) = sum;
}
}
let idx = cdf[..n].partition_point(|&c| c <= r);
idx.min(n - 1)
}
pub fn sample_token_into(probs: &[f32], rng: &mut Rng, cdf: &mut Vec<f32>) -> usize {
let mut r = rng.uniform();
while r == 0.0 {
r = rng.uniform();
}
let n = probs.len();
if n == 0 {
return 0;
}
cdf.resize(n, 0.0);
let mut sum = 0.0f32;
for (i, &p) in probs.iter().enumerate() {
sum += p;
unsafe {
*cdf.get_unchecked_mut(i) = sum;
}
}
let idx = cdf[..n].partition_point(|&c| c <= r);
idx.min(n - 1)
}