#[derive(Clone, Copy, Debug, Eq, PartialEq)]
pub enum F32LinearAccumulation {
Scalar,
Lanes4,
Lanes8,
FusedLanes4,
FusedLanes8,
Accelerate,
AccelerateRowInvariant,
AccelerateBiasSeeded,
AccelerateBiasSeededRowInvariant,
WidenedF64,
}
#[derive(Clone, Copy, Debug, Eq, PartialEq)]
pub enum F32RmsNormArithmetic {
ScalarReciprocalSqrt,
ScalarDivideSqrt,
Lanes4ReciprocalSqrt,
Lanes8ReciprocalSqrt,
Lanes16ReciprocalSqrt,
Lanes32ReciprocalSqrt,
TorchCascade4ReciprocalSqrt,
TorchCascade8ReciprocalSqrt,
F64ReciprocalSqrt,
}
impl F32RmsNormArithmetic {
pub const WIDENED_F64: Self = Self::F64ReciprocalSqrt;
}
#[derive(Clone, Copy, Debug, Eq, PartialEq)]
pub enum F32SiluArithmetic {
Divide,
MultiplyReciprocal,
WidenedF64,
}
#[derive(Clone, Copy, Debug, Eq, PartialEq)]
pub enum F32SoftmaxArithmetic {
ReciprocalMultiply,
Divide,
WidenedF64,
}
pub fn linear(
x: &[f32],
weight: &[f32],
bias: Option<&[f32]>,
m: usize,
k: usize,
n: usize,
out: &mut [f32],
) {
linear_with_accumulation(x, weight, bias, m, k, n, F32LinearAccumulation::Scalar, out);
}
#[allow(clippy::too_many_arguments)]
pub fn linear_with_accumulation(
x: &[f32],
weight: &[f32],
bias: Option<&[f32]>,
m: usize,
k: usize,
n: usize,
accumulation: F32LinearAccumulation,
out: &mut [f32],
) {
assert_eq!(x.len(), m * k, "x must be [m, k]");
assert_eq!(weight.len(), n * k, "weight must be [n, k]");
assert_eq!(out.len(), m * n, "out must be [m, n]");
if let Some(bias) = bias {
assert_eq!(bias.len(), n, "bias must be [n]");
}
if matches!(
accumulation,
F32LinearAccumulation::AccelerateBiasSeeded
| F32LinearAccumulation::AccelerateBiasSeededRowInvariant
) {
match bias {
Some(bias) => {
for row in out.chunks_exact_mut(n) {
row.copy_from_slice(bias);
}
}
None => out.fill(0.0),
}
let row_invariant = accumulation == F32LinearAccumulation::AccelerateBiasSeededRowInvariant;
if accelerate_sgemm(x, weight, m, k, n, 1.0, row_invariant, out) {
return;
}
out.fill(0.0);
}
if matches!(
accumulation,
F32LinearAccumulation::Accelerate | F32LinearAccumulation::AccelerateRowInvariant
) && accelerate_sgemm(
x,
weight,
m,
k,
n,
0.0,
accumulation == F32LinearAccumulation::AccelerateRowInvariant,
out,
) {
if let Some(bias) = bias {
for row in out.chunks_exact_mut(n) {
for (value, offset) in row.iter_mut().zip(bias) {
*value += offset;
}
}
}
return;
}
for row in 0..m {
let x_row = &x[row * k..row * k + k];
for col in 0..n {
let w_row = &weight[col * k..col * k + k];
let sum = dot_with_accumulation(x_row, w_row, accumulation);
out[row * n + col] = bias.map_or(sum, |b| sum + b[col]);
}
}
}
fn dot_with_accumulation(x: &[f32], weight: &[f32], accumulation: F32LinearAccumulation) -> f32 {
assert_eq!(x.len(), weight.len(), "dot-product inputs must match");
match accumulation {
F32LinearAccumulation::Scalar => {
let mut sum = 0.0f32;
for index in 0..x.len() {
sum += x[index] * weight[index];
}
sum
}
F32LinearAccumulation::WidenedF64 => {
let mut sum = 0.0f64;
for index in 0..x.len() {
sum += f64::from(x[index]) * f64::from(weight[index]);
}
sum as f32
}
F32LinearAccumulation::Lanes4
| F32LinearAccumulation::Lanes8
| F32LinearAccumulation::FusedLanes4
| F32LinearAccumulation::FusedLanes8
| F32LinearAccumulation::Accelerate
| F32LinearAccumulation::AccelerateRowInvariant
| F32LinearAccumulation::AccelerateBiasSeeded
| F32LinearAccumulation::AccelerateBiasSeededRowInvariant => {
let lanes = match accumulation {
F32LinearAccumulation::Lanes4 => 4,
F32LinearAccumulation::Lanes8 => 8,
F32LinearAccumulation::FusedLanes4 => 4,
F32LinearAccumulation::FusedLanes8 => 8,
F32LinearAccumulation::Accelerate
| F32LinearAccumulation::AccelerateRowInvariant
| F32LinearAccumulation::AccelerateBiasSeeded
| F32LinearAccumulation::AccelerateBiasSeededRowInvariant => 1,
F32LinearAccumulation::Scalar | F32LinearAccumulation::WidenedF64 => {
unreachable!("scalar and widened orders are handled above")
}
};
let mut partial = [0.0f32; 8];
for index in 0..x.len() {
let lane = index % lanes;
partial[lane] = match accumulation {
F32LinearAccumulation::FusedLanes4 | F32LinearAccumulation::FusedLanes8 => {
x[index].mul_add(weight[index], partial[lane])
}
F32LinearAccumulation::Scalar
| F32LinearAccumulation::Lanes4
| F32LinearAccumulation::Lanes8
| F32LinearAccumulation::Accelerate
| F32LinearAccumulation::AccelerateRowInvariant
| F32LinearAccumulation::AccelerateBiasSeeded
| F32LinearAccumulation::AccelerateBiasSeededRowInvariant
| F32LinearAccumulation::WidenedF64 => partial[lane] + x[index] * weight[index],
};
}
let mut sum = 0.0f32;
for value in &partial[..lanes] {
sum += *value;
}
sum
}
}
}
#[cfg(all(feature = "accelerate-sgemm", target_os = "macos"))]
#[allow(clippy::too_many_arguments)]
fn accelerate_sgemm(
x: &[f32],
weight: &[f32],
m: usize,
k: usize,
n: usize,
beta: f32,
row_invariant: bool,
out: &mut [f32],
) -> bool {
if row_invariant && m == 1 {
let mut doubled_x = Vec::with_capacity(2 * k);
doubled_x.extend_from_slice(x);
doubled_x.extend_from_slice(x);
let mut doubled_out = Vec::with_capacity(2 * n);
doubled_out.extend_from_slice(out);
doubled_out.extend_from_slice(out);
if !accelerate_sgemm(&doubled_x, weight, 2, k, n, beta, false, &mut doubled_out) {
return false;
}
out.copy_from_slice(&doubled_out[..n]);
return true;
}
let m = i32::try_from(m).expect("SGEMM rows fit CBLAS i32 dimensions");
let k = i32::try_from(k).expect("SGEMM reduction fits CBLAS i32 dimensions");
let n = i32::try_from(n).expect("SGEMM columns fit CBLAS i32 dimensions");
unsafe {
cblas_sgemm(
CBLAS_ROW_MAJOR,
CBLAS_NO_TRANSPOSE,
CBLAS_TRANSPOSE,
m,
n,
k,
1.0,
x.as_ptr(),
k,
weight.as_ptr(),
k,
beta,
out.as_mut_ptr(),
n,
);
}
true
}
#[cfg(not(all(feature = "accelerate-sgemm", target_os = "macos")))]
fn accelerate_sgemm(
_x: &[f32],
_weight: &[f32],
_m: usize,
_k: usize,
_n: usize,
_beta: f32,
_row_invariant: bool,
_out: &mut [f32],
) -> bool {
false
}
#[cfg(all(feature = "accelerate-sgemm", target_os = "macos"))]
const CBLAS_ROW_MAJOR: i32 = 101;
#[cfg(all(feature = "accelerate-sgemm", target_os = "macos"))]
const CBLAS_NO_TRANSPOSE: i32 = 111;
#[cfg(all(feature = "accelerate-sgemm", target_os = "macos"))]
const CBLAS_TRANSPOSE: i32 = 112;
#[cfg(all(feature = "accelerate-sgemm", target_os = "macos"))]
#[link(name = "Accelerate", kind = "framework")]
unsafe extern "C" {
fn cblas_sgemm(
order: i32,
trans_a: i32,
trans_b: i32,
m: i32,
n: i32,
k: i32,
alpha: f32,
a: *const f32,
lda: i32,
b: *const f32,
ldb: i32,
beta: f32,
c: *mut f32,
ldc: i32,
);
}
#[derive(Clone, Copy, Debug, Eq, PartialEq)]
pub enum F32Transcendental {
ScalarLibm,
AccelerateVForce,
SleefU10,
}
pub fn sin_with(x: &[f32], implementation: F32Transcendental, out: &mut [f32]) {
assert_eq!(x.len(), out.len(), "sin output must match its input");
if implementation == F32Transcendental::AccelerateVForce && vforce_sin(x, out) {
return;
}
if implementation == F32Transcendental::SleefU10 {
for (value, target) in x.iter().zip(out.iter_mut()) {
*target = crate::sleef::sinf_u10(*value);
}
return;
}
for (value, target) in x.iter().zip(out.iter_mut()) {
*target = value.sin();
}
}
pub fn exp_with(x: &[f32], implementation: F32Transcendental, out: &mut [f32]) {
assert_eq!(x.len(), out.len(), "exp output must match its input");
if implementation == F32Transcendental::AccelerateVForce && vforce_exp(x, out) {
return;
}
if implementation == F32Transcendental::SleefU10 {
for (value, target) in x.iter().zip(out.iter_mut()) {
*target = crate::sleef::expf_u10(*value);
}
return;
}
for (value, target) in x.iter().zip(out.iter_mut()) {
*target = value.exp();
}
}
#[cfg(all(feature = "accelerate-sgemm", target_os = "macos"))]
fn vforce_sin(x: &[f32], out: &mut [f32]) -> bool {
let count = i32::try_from(x.len()).expect("vForce length fits i32");
unsafe { vvsinf(out.as_mut_ptr(), x.as_ptr(), &raw const count) };
true
}
#[cfg(all(feature = "accelerate-sgemm", target_os = "macos"))]
fn vforce_exp(x: &[f32], out: &mut [f32]) -> bool {
let count = i32::try_from(x.len()).expect("vForce length fits i32");
unsafe { vvexpf(out.as_mut_ptr(), x.as_ptr(), &raw const count) };
true
}
#[cfg(not(all(feature = "accelerate-sgemm", target_os = "macos")))]
fn vforce_sin(_x: &[f32], _out: &mut [f32]) -> bool {
false
}
#[cfg(not(all(feature = "accelerate-sgemm", target_os = "macos")))]
fn vforce_exp(_x: &[f32], _out: &mut [f32]) -> bool {
false
}
#[cfg(all(feature = "accelerate-sgemm", target_os = "macos"))]
#[link(name = "Accelerate", kind = "framework")]
unsafe extern "C" {
fn vvsinf(out: *mut f32, x: *const f32, count: *const i32);
fn vvexpf(out: *mut f32, x: *const f32, count: *const i32);
}
pub fn rms_norm(x: &[f32], weight: &[f32], eps: f32, rows: usize, dim: usize, out: &mut [f32]) {
rms_norm_with_arithmetic(
x,
weight,
eps,
rows,
dim,
F32RmsNormArithmetic::ScalarReciprocalSqrt,
out,
);
}
pub fn rms_norm_with_arithmetic(
x: &[f32],
weight: &[f32],
eps: f32,
rows: usize,
dim: usize,
arithmetic: F32RmsNormArithmetic,
out: &mut [f32],
) {
assert_eq!(x.len(), rows * dim, "x must be [rows, dim]");
assert_eq!(weight.len(), dim, "weight must be [dim]");
assert_eq!(out.len(), rows * dim, "out must be [rows, dim]");
for row in 0..rows {
let src = &x[row * dim..row * dim + dim];
let scale = rms_scale(src, eps, arithmetic);
for index in 0..dim {
out[row * dim + index] = src[index] * scale * weight[index];
}
}
}
fn rms_scale(src: &[f32], eps: f32, arithmetic: F32RmsNormArithmetic) -> f32 {
match arithmetic {
F32RmsNormArithmetic::ScalarReciprocalSqrt => {
let sum = sum_squares_f32(src, 1);
(sum / src.len() as f32 + eps).sqrt().recip()
}
F32RmsNormArithmetic::ScalarDivideSqrt => {
let sum = sum_squares_f32(src, 1);
1.0f32 / (sum / src.len() as f32 + eps).sqrt()
}
F32RmsNormArithmetic::Lanes4ReciprocalSqrt => {
let sum = sum_squares_f32(src, 4);
(sum / src.len() as f32 + eps).sqrt().recip()
}
F32RmsNormArithmetic::Lanes8ReciprocalSqrt => {
let sum = sum_squares_f32(src, 8);
(sum / src.len() as f32 + eps).sqrt().recip()
}
F32RmsNormArithmetic::Lanes16ReciprocalSqrt => {
let sum = sum_squares_f32(src, 16);
(sum / src.len() as f32 + eps).sqrt().recip()
}
F32RmsNormArithmetic::Lanes32ReciprocalSqrt => {
let sum = sum_squares_f32(src, 32);
(sum / src.len() as f32 + eps).sqrt().recip()
}
F32RmsNormArithmetic::TorchCascade4ReciprocalSqrt => {
let sum = torch_cascade_sum(src, 4, |value| value * value);
(sum / src.len() as f32 + eps).sqrt().recip()
}
F32RmsNormArithmetic::TorchCascade8ReciprocalSqrt => {
let sum = torch_cascade_sum(src, 8, |value| value * value);
(sum / src.len() as f32 + eps).sqrt().recip()
}
F32RmsNormArithmetic::F64ReciprocalSqrt => {
let mut sum = 0.0f64;
for value in src {
let value = f64::from(*value);
sum += value * value;
}
(sum / src.len() as f64 + f64::from(eps)).sqrt().recip() as f32
}
}
}
fn sum_squares_f32(src: &[f32], lanes: usize) -> f32 {
let mut partial = [0.0f32; 32];
for (index, value) in src.iter().enumerate() {
partial[index % lanes] += *value * *value;
}
let mut sum = 0.0f32;
for value in &partial[..lanes] {
sum += *value;
}
sum
}
#[allow(clippy::needless_range_loop)]
pub fn torch_cascade_sum(src: &[f32], width: usize, transform: impl Fn(f32) -> f32) -> f32 {
assert!(width > 0, "vector width must be positive");
const ILP: usize = 4;
const LEVELS: usize = 4;
let vector_count = src.len() / width;
let vector = |index: usize, lane: usize| transform(src[index * width + lane]);
let size = vector_count / ILP;
let level_power = ceil_log2(size).div_euclid(LEVELS).max(4);
let level_step = 1usize << level_power;
let level_mask = level_step - 1;
let mut acc = vec![[0.0f32; ILP].map(|_| vec![0.0f32; width]); LEVELS];
let mut index = 0usize;
while index + level_step <= size {
for _ in 0..level_step {
for chain in 0..ILP {
for lane in 0..width {
acc[0][chain][lane] += vector(index * ILP + chain, lane);
}
}
index += 1;
}
for level in 1..LEVELS {
for chain in 0..ILP {
for lane in 0..width {
acc[level][chain][lane] += acc[level - 1][chain][lane];
acc[level - 1][chain][lane] = 0.0;
}
}
if index & (level_mask << (level * level_power)) != 0 {
break;
}
}
}
while index < size {
for chain in 0..ILP {
for lane in 0..width {
acc[0][chain][lane] += vector(index * ILP + chain, lane);
}
}
index += 1;
}
for level in 1..LEVELS {
for chain in 0..ILP {
for lane in 0..width {
acc[level][chain][lane] += acc[level - 1][chain][lane];
}
}
}
let mut partial = acc.swap_remove(LEVELS - 1);
for leftover in size * ILP..vector_count {
for lane in 0..width {
partial[0][lane] += vector(leftover, lane);
}
}
for chain in 1..ILP {
for lane in 0..width {
partial[0][lane] += partial[chain][lane];
}
}
let mut sum = 0.0f32;
for index in vector_count * width..src.len() {
sum += transform(src[index]);
}
for lane in 0..width {
sum += partial[0][lane];
}
sum
}
fn ceil_log2(value: usize) -> usize {
if value <= 1 {
return 0;
}
usize::BITS as usize - (value - 1).leading_zeros() as usize
}
pub fn silu_mul_in_place(gate: &mut [f32], up: &[f32]) {
silu_mul_in_place_with_arithmetic(gate, up, F32SiluArithmetic::Divide);
}
pub fn silu_mul_in_place_with_arithmetic(
gate: &mut [f32],
up: &[f32],
arithmetic: F32SiluArithmetic,
) {
assert_eq!(gate.len(), up.len(), "gate and up must match");
for (g, u) in gate.iter_mut().zip(up) {
let x = *g;
if arithmetic == F32SiluArithmetic::WidenedF64 {
let wide = f64::from(x);
*g = (wide / (1.0 + (-wide).exp()) * f64::from(*u)) as f32;
continue;
}
let denominator = 1.0 + (-x).exp();
let silu = match arithmetic {
F32SiluArithmetic::Divide => x / denominator,
F32SiluArithmetic::MultiplyReciprocal => x * denominator.recip(),
F32SiluArithmetic::WidenedF64 => unreachable!("handled above"),
};
*g = silu * u;
}
}
pub fn softmax_rows(x: &mut [f32], rows: usize, cols: usize) {
softmax_rows_with_arithmetic(x, rows, cols, F32SoftmaxArithmetic::ReciprocalMultiply);
}
pub fn softmax_rows_with_arithmetic(
x: &mut [f32],
rows: usize,
cols: usize,
arithmetic: F32SoftmaxArithmetic,
) {
assert_eq!(x.len(), rows * cols, "x must be [rows, cols]");
for row in 0..rows {
let slice = &mut x[row * cols..row * cols + cols];
let mut max = f32::NEG_INFINITY;
for value in slice.iter() {
if *value > max {
max = *value;
}
}
if arithmetic == F32SoftmaxArithmetic::WidenedF64 {
let max = f64::from(max);
let mut wide = Vec::with_capacity(slice.len());
let mut sum = 0.0f64;
for value in slice.iter() {
let exponent = (f64::from(*value) - max).exp();
sum += exponent;
wide.push(exponent);
}
for (value, exponent) in slice.iter_mut().zip(wide) {
*value = (exponent / sum) as f32;
}
continue;
}
let mut sum = 0.0f32;
for value in slice.iter_mut() {
*value = (*value - max).exp();
sum += *value;
}
for value in slice.iter_mut() {
*value = match arithmetic {
F32SoftmaxArithmetic::ReciprocalMultiply => *value * sum.recip(),
F32SoftmaxArithmetic::Divide => *value / sum,
F32SoftmaxArithmetic::WidenedF64 => unreachable!("handled above"),
};
}
}
}
#[allow(clippy::too_many_arguments)]
pub fn gqa_attention(
queries: &[f32],
keys: &[f32],
values: &[f32],
additive_mask: &[f32],
query_positions: usize,
key_positions: usize,
q_heads: usize,
kv_heads: usize,
head_dim: usize,
out: &mut [f32],
) {
gqa_attention_with_softmax(
queries,
keys,
values,
additive_mask,
query_positions,
key_positions,
q_heads,
kv_heads,
head_dim,
F32SoftmaxArithmetic::ReciprocalMultiply,
out,
);
}
#[allow(clippy::too_many_arguments)]
pub fn gqa_attention_with_softmax(
queries: &[f32],
keys: &[f32],
values: &[f32],
additive_mask: &[f32],
query_positions: usize,
key_positions: usize,
q_heads: usize,
kv_heads: usize,
head_dim: usize,
softmax_arithmetic: F32SoftmaxArithmetic,
out: &mut [f32],
) {
gqa_attention_with_arithmetic(
queries,
keys,
values,
additive_mask,
query_positions,
key_positions,
q_heads,
kv_heads,
head_dim,
softmax_arithmetic,
F32LinearAccumulation::Scalar,
out,
);
}
#[allow(clippy::too_many_arguments)]
pub fn gqa_attention_with_arithmetic(
queries: &[f32],
keys: &[f32],
values: &[f32],
additive_mask: &[f32],
query_positions: usize,
key_positions: usize,
q_heads: usize,
kv_heads: usize,
head_dim: usize,
softmax_arithmetic: F32SoftmaxArithmetic,
accumulation: F32LinearAccumulation,
out: &mut [f32],
) {
assert!(kv_heads > 0, "at least one KV head is required");
assert_eq!(
q_heads % kv_heads,
0,
"query heads must divide evenly into KV groups"
);
assert_eq!(
queries.len(),
query_positions * q_heads * head_dim,
"queries must be [query_positions, q_heads, head_dim]"
);
assert_eq!(
keys.len(),
key_positions * kv_heads * head_dim,
"keys must be [key_positions, kv_heads, head_dim]"
);
assert_eq!(
values.len(),
key_positions * kv_heads * head_dim,
"values must be [key_positions, kv_heads, head_dim]"
);
assert_eq!(
additive_mask.len(),
query_positions * key_positions,
"mask must be [query_positions, key_positions]"
);
assert_eq!(
out.len(),
query_positions * q_heads * head_dim,
"out must be [query_positions, q_heads, head_dim]"
);
if accumulation == F32LinearAccumulation::Accelerate
&& accelerate_gqa_attention(
queries,
keys,
values,
additive_mask,
query_positions,
key_positions,
q_heads,
kv_heads,
head_dim,
softmax_arithmetic,
out,
)
{
return;
}
let scale = (head_dim as f32).sqrt().recip();
let kv_group = q_heads / kv_heads;
let mut scores = vec![0.0f32; key_positions];
for query_position in 0..query_positions {
let mask =
&additive_mask[query_position * key_positions..(query_position + 1) * key_positions];
for q_head in 0..q_heads {
let kv_head = q_head / kv_group;
let query_base = (query_position * q_heads + q_head) * head_dim;
let query = &queries[query_base..query_base + head_dim];
for (key_position, score) in scores.iter_mut().enumerate() {
let key_base = (key_position * kv_heads + kv_head) * head_dim;
let key = &keys[key_base..key_base + head_dim];
let dot = dot_with_accumulation(query, key, accumulation);
*score = dot * scale + mask[key_position];
}
softmax_rows_with_arithmetic(&mut scores, 1, key_positions, softmax_arithmetic);
let out_base = query_base;
attention_weighted_sum(
&scores,
values,
kv_head,
kv_heads,
head_dim,
accumulation,
&mut out[out_base..out_base + head_dim],
);
}
}
}
#[cfg(all(feature = "accelerate-sgemm", target_os = "macos"))]
#[allow(clippy::too_many_arguments)]
fn accelerate_gqa_attention(
queries: &[f32],
keys: &[f32],
values: &[f32],
additive_mask: &[f32],
query_positions: usize,
key_positions: usize,
q_heads: usize,
kv_heads: usize,
head_dim: usize,
softmax_arithmetic: F32SoftmaxArithmetic,
out: &mut [f32],
) -> bool {
let scale = (head_dim as f32).sqrt().recip();
let kv_group = q_heads / kv_heads;
let mut query_matrix = vec![0.0f32; query_positions * head_dim];
let mut key_matrix = vec![0.0f32; key_positions * head_dim];
let mut value_transpose = vec![0.0f32; head_dim * key_positions];
let mut scores = vec![0.0f32; query_positions * key_positions];
let mut context = vec![0.0f32; query_positions * head_dim];
for q_head in 0..q_heads {
let kv_head = q_head / kv_group;
for query_position in 0..query_positions {
let query_base = (query_position * q_heads + q_head) * head_dim;
query_matrix[query_position * head_dim..(query_position + 1) * head_dim]
.copy_from_slice(&queries[query_base..query_base + head_dim]);
}
for key_position in 0..key_positions {
let key_base = (key_position * kv_heads + kv_head) * head_dim;
key_matrix[key_position * head_dim..(key_position + 1) * head_dim]
.copy_from_slice(&keys[key_base..key_base + head_dim]);
for lane in 0..head_dim {
value_transpose[lane * key_positions + key_position] = values[key_base + lane];
}
}
if !accelerate_sgemm(
&query_matrix,
&key_matrix,
query_positions,
head_dim,
key_positions,
0.0,
false,
&mut scores,
) {
return false;
}
for query_position in 0..query_positions {
let score_row =
&mut scores[query_position * key_positions..(query_position + 1) * key_positions];
let mask = &additive_mask
[query_position * key_positions..(query_position + 1) * key_positions];
for (score, mask_value) in score_row.iter_mut().zip(mask) {
*score = *score * scale + mask_value;
}
softmax_rows_with_arithmetic(score_row, 1, key_positions, softmax_arithmetic);
}
if !accelerate_sgemm(
&scores,
&value_transpose,
query_positions,
key_positions,
head_dim,
0.0,
false,
&mut context,
) {
return false;
}
for query_position in 0..query_positions {
let out_base = (query_position * q_heads + q_head) * head_dim;
out[out_base..out_base + head_dim].copy_from_slice(
&context[query_position * head_dim..(query_position + 1) * head_dim],
);
}
}
true
}
#[cfg(not(all(feature = "accelerate-sgemm", target_os = "macos")))]
#[allow(clippy::too_many_arguments)]
fn accelerate_gqa_attention(
_queries: &[f32],
_keys: &[f32],
_values: &[f32],
_additive_mask: &[f32],
_query_positions: usize,
_key_positions: usize,
_q_heads: usize,
_kv_heads: usize,
_head_dim: usize,
_softmax_arithmetic: F32SoftmaxArithmetic,
_out: &mut [f32],
) -> bool {
false
}
#[allow(clippy::too_many_arguments)]
fn attention_weighted_sum(
scores: &[f32],
values: &[f32],
kv_head: usize,
kv_heads: usize,
head_dim: usize,
accumulation: F32LinearAccumulation,
out: &mut [f32],
) {
if accumulation == F32LinearAccumulation::WidenedF64 {
for lane in 0..head_dim {
let mut sum = 0.0f64;
for (key_position, weight) in scores.iter().copied().enumerate() {
let value = values[(key_position * kv_heads + kv_head) * head_dim + lane];
sum += f64::from(weight) * f64::from(value);
}
out[lane] = sum as f32;
}
return;
}
let lanes = match accumulation {
F32LinearAccumulation::Scalar => 1,
F32LinearAccumulation::Lanes4 | F32LinearAccumulation::FusedLanes4 => 4,
F32LinearAccumulation::Lanes8 | F32LinearAccumulation::FusedLanes8 => 8,
F32LinearAccumulation::Accelerate
| F32LinearAccumulation::AccelerateRowInvariant
| F32LinearAccumulation::AccelerateBiasSeeded
| F32LinearAccumulation::AccelerateBiasSeededRowInvariant => 1,
F32LinearAccumulation::WidenedF64 => unreachable!("handled above"),
};
for lane in 0..head_dim {
let mut partial = [0.0f32; 8];
for (key_position, weight) in scores.iter().copied().enumerate() {
let value = values[(key_position * kv_heads + kv_head) * head_dim + lane];
let partial_index = key_position % lanes;
partial[partial_index] = match accumulation {
F32LinearAccumulation::FusedLanes4 | F32LinearAccumulation::FusedLanes8 => {
weight.mul_add(value, partial[partial_index])
}
F32LinearAccumulation::Scalar
| F32LinearAccumulation::Lanes4
| F32LinearAccumulation::Lanes8
| F32LinearAccumulation::Accelerate
| F32LinearAccumulation::AccelerateRowInvariant
| F32LinearAccumulation::AccelerateBiasSeeded
| F32LinearAccumulation::AccelerateBiasSeededRowInvariant
| F32LinearAccumulation::WidenedF64 => partial[partial_index] + weight * value,
};
}
let mut sum = 0.0f32;
for value in &partial[..lanes] {
sum += *value;
}
out[lane] = sum;
}
}
pub fn mrope_interleave(axes: [&[f32]; 3], sections: [usize; 3], out: &mut [f32]) {
let half = out.len();
for axis in axes {
assert!(
axis.len() >= half,
"axis row shorter than the half-dimension"
);
}
out.copy_from_slice(&axes[0][..half]);
let modality_num = 3usize;
for (axis_index, section) in sections.iter().enumerate().skip(1) {
let end = section * modality_num;
let mut lane = axis_index;
while lane < end && lane < half {
out[lane] = axes[axis_index][lane];
lane += modality_num;
}
}
}
pub fn apply_rope_in_place(row: &mut [f32], cos: &[f32], sin: &[f32]) {
let dim = row.len();
assert_eq!(cos.len(), dim, "cos must match head_dim");
assert_eq!(sin.len(), dim, "sin must match head_dim");
assert!(dim.is_multiple_of(2), "head_dim must be even");
let half = dim / 2;
let original: Vec<f32> = row.to_vec();
for index in 0..dim {
let rotated = if index < half {
-original[index + half]
} else {
original[index - half]
};
row[index] = original[index] * cos[index] + rotated * sin[index];
}
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn linear_matches_a_hand_computed_product() {
let x = [1.0, 2.0, 3.0];
let weight = [1.0, 0.0, -1.0, 2.0, 2.0, 2.0];
let mut out = [0.0; 2];
linear(&x, &weight, None, 1, 3, 2, &mut out);
assert_eq!(out, [-2.0, 12.0]);
let mut biased = [0.0; 2];
linear(&x, &weight, Some(&[10.0, -12.0]), 1, 3, 2, &mut biased);
assert_eq!(biased, [8.0, 0.0]);
}
#[test]
fn torch_cascade_sum_agrees_with_the_flat_sum_when_rounding_cannot_intervene() {
for length in [1usize, 7, 8, 15, 16, 31, 128, 1024, 3072] {
for width in [4usize, 8] {
let values: Vec<f32> = (0..length).map(|index| (index % 8) as f32).collect();
let flat: f32 = values.iter().sum();
assert_eq!(
torch_cascade_sum(&values, width, |value| value),
flat,
"length {length}, width {width}"
);
}
}
}
#[test]
fn torch_cascade_sum_applies_its_transform_before_accumulating() {
let values = [1.0f32, 2.0, 3.0, 4.0, 5.0];
assert_eq!(torch_cascade_sum(&values, 4, |value| value * value), 55.0);
}
#[test]
fn torch_cascade_sum_differs_from_a_flat_sum_once_rounding_matters() {
let mut values = vec![1.0f32; 1024];
values[0] = 1.0e8;
let flat = values.iter().fold(0.0f32, |sum, value| sum + value);
assert_ne!(torch_cascade_sum(&values, 8, |value| value), flat);
}
#[test]
fn ceil_log2_matches_its_definition() {
assert_eq!(ceil_log2(0), 0);
assert_eq!(ceil_log2(1), 0);
assert_eq!(ceil_log2(2), 1);
assert_eq!(ceil_log2(3), 2);
assert_eq!(ceil_log2(32), 5);
assert_eq!(ceil_log2(33), 6);
}
#[test]
fn rms_norm_normalizes_and_scales() {
let x = [3.0f32, 4.0];
let weight = [1.0f32, 1.0];
let mut out = [0.0; 2];
rms_norm(&x, &weight, 0.0, 1, 2, &mut out);
let expected = 12.5f32.sqrt().recip();
assert!((out[0] - 3.0 * expected).abs() < 1e-6);
assert!((out[1] - 4.0 * expected).abs() < 1e-6);
let mut weighted = [0.0; 2];
rms_norm(&x, &[2.0, 0.5], 0.0, 1, 2, &mut weighted);
assert!((weighted[0] - 3.0 * expected * 2.0).abs() < 1e-6);
assert!((weighted[1] - 4.0 * expected * 0.5).abs() < 1e-6);
}
#[test]
fn silu_mul_matches_the_definition() {
let mut gate = [0.0f32, 1.0, -1.0];
let up = [1.0f32, 2.0, 3.0];
silu_mul_in_place(&mut gate, &up);
assert_eq!(gate[0], 0.0);
let silu_one = 1.0f32 / (1.0 + (-1.0f32).exp());
assert!((gate[1] - silu_one * 2.0).abs() < 1e-6);
let silu_neg = -1.0f32 / (1.0 + 1.0f32.exp());
assert!((gate[2] - silu_neg * 3.0).abs() < 1e-6);
}
#[test]
fn softmax_rows_sums_to_one_and_is_shift_invariant() {
let mut x = [1.0f32, 2.0, 3.0, 101.0, 102.0, 103.0];
softmax_rows(&mut x, 2, 3);
let first: f32 = x[..3].iter().sum();
let second: f32 = x[3..].iter().sum();
assert!((first - 1.0).abs() < 1e-6);
assert!((second - 1.0).abs() < 1e-6);
for index in 0..3 {
assert!((x[index] - x[index + 3]).abs() < 1e-6);
}
}
#[test]
fn gqa_maps_each_query_head_to_its_kv_group() {
let (query_positions, key_positions, q_heads, kv_heads, head_dim) = (1, 1, 4, 2, 2);
let queries = vec![0.0f32; query_positions * q_heads * head_dim];
let keys = vec![0.0f32; key_positions * kv_heads * head_dim];
let values = [10.0f32, 11.0, 20.0, 21.0];
let mut out = vec![0.0f32; query_positions * q_heads * head_dim];
gqa_attention(
&queries,
&keys,
&values,
&[0.0],
query_positions,
key_positions,
q_heads,
kv_heads,
head_dim,
&mut out,
);
assert_eq!(&out[0..2], &[10.0, 11.0]);
assert_eq!(&out[2..4], &[10.0, 11.0]);
assert_eq!(&out[4..6], &[20.0, 21.0]);
assert_eq!(&out[6..8], &[20.0, 21.0]);
}
#[test]
fn gqa_honors_the_additive_causal_mask() {
let (query_positions, key_positions, q_heads, kv_heads, head_dim) = (2, 2, 1, 1, 2);
let queries = vec![0.0f32; query_positions * q_heads * head_dim];
let keys = vec![0.0f32; key_positions * kv_heads * head_dim];
let values = [2.0f32, 4.0, 10.0, 20.0];
let mask = [0.0f32, f32::NEG_INFINITY, 0.0, 0.0];
let mut out = vec![0.0f32; query_positions * q_heads * head_dim];
gqa_attention(
&queries,
&keys,
&values,
&mask,
query_positions,
key_positions,
q_heads,
kv_heads,
head_dim,
&mut out,
);
assert_eq!(&out[0..2], &[2.0, 4.0]);
assert_eq!(&out[2..4], &[6.0, 12.0]);
}
#[test]
fn rope_rotates_a_known_pair() {
let mut row = [3.0f32, 5.0];
apply_rope_in_place(&mut row, &[0.0, 0.0], &[1.0, 1.0]);
assert_eq!(row, [-5.0, 3.0]);
let mut same = [3.0f32, 5.0];
apply_rope_in_place(&mut same, &[1.0, 1.0], &[0.0, 0.0]);
assert_eq!(same, [3.0, 5.0]);
}
#[test]
fn mrope_interleave_is_identity_when_all_axes_agree() {
let axis: Vec<f32> = (0..64).map(|value| value as f32).collect();
let mut out = vec![0.0f32; 64];
mrope_interleave([&axis, &axis, &axis], [24, 20, 20], &mut out);
assert_eq!(out, axis);
}
#[test]
fn mrope_interleave_selects_the_documented_lanes() {
let zeros = vec![0.0f32; 64];
let ones = vec![1.0f32; 64];
let twos = vec![2.0f32; 64];
let mut out = vec![0.0f32; 64];
mrope_interleave([&zeros, &ones, &twos], [24, 20, 20], &mut out);
for (lane, value) in out.iter().enumerate() {
let expected = if lane < 60 && lane % 3 == 1 {
1.0
} else if lane < 60 && lane % 3 == 2 {
2.0
} else {
0.0
};
assert_eq!(*value, expected, "lane {lane}");
}
}
}