use ft_core::{DType, Device, TensorMeta};
use super::tensor::{Mat, QInt8};
use crate::error::{FocrError, FocrResult};
use crate::quant::recipe::switch_on;
fn meta_2d(rows: usize, cols: usize) -> TensorMeta {
TensorMeta::from_shape(vec![rows, cols], DType::F32, Device::Cpu)
}
fn kernel_err(e: ft_kernel_cpu::KernelError) -> FocrError {
FocrError::Other(anyhow::anyhow!("ft-kernel-cpu: {e}"))
}
fn checked_mat_len(context: &str, x: &Mat) -> FocrResult<usize> {
let expected = x.rows.checked_mul(x.cols).ok_or_else(|| {
FocrError::Other(anyhow::anyhow!(
"{context}: rows*cols overflow ({} * {})",
x.rows,
x.cols
))
})?;
if x.data.len() != expected {
return Err(FocrError::Other(anyhow::anyhow!(
"{context}: data len {} != rows*cols {} for shape [{}, {}]",
x.data.len(),
expected,
x.rows,
x.cols
)));
}
Ok(expected)
}
fn checked_qint8_len(context: &str, w: &QInt8) -> FocrResult<usize> {
let expected = w.n.checked_mul(w.k).ok_or_else(|| {
FocrError::Other(anyhow::anyhow!(
"{context}: n*k overflow ({} * {})",
w.n,
w.k
))
})?;
if w.w.len() != expected {
return Err(FocrError::Other(anyhow::anyhow!(
"{context}: weight len {} != n*k {} for shape [{}, {}]",
w.w.len(),
expected,
w.n,
w.k
)));
}
if w.scales.len() != w.n {
return Err(FocrError::Other(anyhow::anyhow!(
"{context}: scales len {} != n {}",
w.scales.len(),
w.n
)));
}
Ok(expected)
}
pub fn matmul(a: &Mat, b: &Mat) -> FocrResult<Mat> {
checked_mat_len("matmul lhs", a)?;
checked_mat_len("matmul rhs", b)?;
if a.cols != b.rows {
return Err(FocrError::Other(anyhow::anyhow!(
"matmul inner dim mismatch: [{},{}] x [{},{}]",
a.rows,
a.cols,
b.rows,
b.cols
)));
}
let (m, k, n) = (a.rows, a.cols, b.cols);
let lhs_meta = meta_2d(m, k);
let rhs_meta = meta_2d(k, n);
let data = ft_kernel_cpu::matmul_tensor_contiguous_f32(&a.data, &b.data, &lhs_meta, &rhs_meta)
.map_err(kernel_err)?;
Ok(Mat::from_vec(m, n, data))
}
#[must_use]
pub fn quantize_int8(w: &[f32], out: usize, in_: usize) -> QInt8 {
let (qw, scales) = ft_kernel_cpu::quantize_per_output_channel_i8(w, out, in_);
QInt8::new(qw, scales, out, in_)
}
pub fn linear_int8_dynamic(x: &Mat, w: &QInt8, bias: Option<&[f32]>) -> FocrResult<Mat> {
checked_mat_len("linear_int8_dynamic x", x)?;
checked_qint8_len("linear_int8_dynamic weight", w)?;
if x.cols != w.k {
return Err(FocrError::Other(anyhow::anyhow!(
"linear_int8_dynamic: x.cols {} != w.k {}",
x.cols,
w.k
)));
}
if let Some(b) = bias
&& b.len() != w.n
{
return Err(FocrError::Other(anyhow::anyhow!(
"linear_int8_dynamic: bias len {} != n {}",
b.len(),
w.n
)));
}
let (m, k, n) = (x.rows, x.cols, w.n);
let data = ft_kernel_cpu::linear_int8_dynamic_f32(&x.data, m, k, &w.w, &w.scales, n, bias);
Ok(Mat::from_vec(m, n, data))
}
#[allow(clippy::too_many_arguments)]
#[must_use]
pub fn conv2d(
input: &[f32],
weight: &[f32],
bias: Option<&[f32]>,
batch: usize,
in_ch: usize,
ph: usize,
pw: usize,
kh: usize,
kw: usize,
oh: usize,
ow: usize,
sh: usize,
sw: usize,
out_ch: usize,
) -> Vec<f32> {
ft_kernel_cpu::conv2d_forward_f32(
input, weight, bias, batch, in_ch, ph, pw, kh, kw, oh, ow, sh, sw, out_ch,
)
}
#[must_use]
pub fn tf_same_pad_amounts(i: usize, k: usize, s: usize) -> (usize, usize) {
let total = (i.div_ceil(s) - 1) * s + k;
let total = total.saturating_sub(i);
(total / 2, total - total / 2)
}
#[allow(clippy::too_many_arguments)]
#[must_use]
pub fn tf_same_pad(
input: &[f32],
batch: usize,
ch: usize,
h: usize,
w: usize,
kh: usize,
kw: usize,
sh: usize,
sw: usize,
fill: f32,
) -> (Vec<f32>, usize, usize) {
let (top, bottom) = tf_same_pad_amounts(h, kh, sh);
let (left, right) = tf_same_pad_amounts(w, kw, sw);
let (ph, pw) = (h + top + bottom, w + left + right);
let mut out = vec![fill; batch * ch * ph * pw];
for bc in 0..batch * ch {
let src = &input[bc * h * w..(bc + 1) * h * w];
let dst = &mut out[bc * ph * pw..(bc + 1) * ph * pw];
for row in 0..h {
let d = (row + top) * pw + left;
dst[d..d + w].copy_from_slice(&src[row * w..(row + 1) * w]);
}
}
(out, ph, pw)
}
#[allow(clippy::too_many_arguments)]
#[must_use]
pub fn max_pool2d(
input: &[f32],
batch: usize,
ch: usize,
ph: usize,
pw: usize,
k: usize,
s: usize,
oh: usize,
ow: usize,
) -> Vec<f32> {
let mut out = vec![f32::NEG_INFINITY; batch * ch * oh * ow];
for bc in 0..batch * ch {
let src = &input[bc * ph * pw..(bc + 1) * ph * pw];
let dst = &mut out[bc * oh * ow..(bc + 1) * oh * ow];
for oy in 0..oh {
for ox in 0..ow {
let mut m = f32::NEG_INFINITY;
for ky in 0..k {
let row = oy * s + ky;
for kx in 0..k {
m = m.max(src[row * pw + ox * s + kx]);
}
}
dst[oy * ow + ox] = m;
}
}
}
out
}
#[allow(clippy::too_many_arguments)]
pub fn group_norm(
x: &mut [f32],
batch: usize,
channels: usize,
spatial: usize,
groups: usize,
eps: f32,
weight: &[f32],
bias: &[f32],
fuse_relu: bool,
) -> FocrResult<()> {
if groups == 0 || !channels.is_multiple_of(groups) {
return Err(FocrError::Other(anyhow::anyhow!(
"group_norm: channels {channels} not divisible by groups {groups}"
)));
}
if x.len() != batch * channels * spatial {
return Err(FocrError::Other(anyhow::anyhow!(
"group_norm: x len {} != batch {batch} * channels {channels} * spatial {spatial}",
x.len()
)));
}
if weight.len() != channels || bias.len() != channels {
return Err(FocrError::Other(anyhow::anyhow!(
"group_norm: weight/bias len {}/{} != channels {channels}",
weight.len(),
bias.len()
)));
}
let cpg = channels / groups;
let group_len = cpg * spatial;
for b in 0..batch {
for g in 0..groups {
let start = (b * channels + g * cpg) * spatial;
let slice = &mut x[start..start + group_len];
let mut sum = 0.0f64;
for &v in slice.iter() {
sum += f64::from(v);
}
let mean = sum / group_len as f64;
let mut var = 0.0f64;
for &v in slice.iter() {
let d = f64::from(v) - mean;
var += d * d;
}
let var = var / group_len as f64;
let inv = 1.0 / (var + f64::from(eps)).sqrt();
let (mean, inv) = (mean as f32, inv as f32);
for c in 0..cpg {
let (gamma, beta) = (weight[g * cpg + c], bias[g * cpg + c]);
for v in &mut slice[c * spatial..(c + 1) * spatial] {
let y = (*v - mean) * inv * gamma + beta;
*v = if fuse_relu { y.max(0.0) } else { y };
}
}
}
}
Ok(())
}
#[allow(clippy::too_many_arguments)]
#[must_use]
pub fn sdpa(
q: &[f32],
k: &[f32],
v: &[f32],
num_bh: usize,
seq_q: usize,
seq_k: usize,
d_k: usize,
d_v: usize,
scale: f32,
causal: bool,
) -> Vec<f32> {
ft_kernel_cpu::sdpa_forward_f32(q, k, v, num_bh, seq_q, seq_k, d_k, d_v, scale, causal)
}
pub fn rms_norm(x: &Mat, weight: Option<&[f32]>, eps: f32) -> FocrResult<Mat> {
checked_mat_len("rms_norm x", x)?;
if let Some(w) = weight
&& w.len() != x.cols
{
return Err(FocrError::Other(anyhow::anyhow!(
"rms_norm: weight len {} != cols {}",
w.len(),
x.cols
)));
}
let data = ft_kernel_cpu::rms_norm_forward_f32(&x.data, weight, x.rows, x.cols, eps);
Ok(Mat::from_vec(x.rows, x.cols, data))
}
pub fn layer_norm(
x: &Mat,
weight: Option<&[f32]>,
bias: Option<&[f32]>,
eps: f32,
) -> FocrResult<Mat> {
checked_mat_len("layer_norm x", x)?;
if let Some(w) = weight
&& w.len() != x.cols
{
return Err(FocrError::Other(anyhow::anyhow!(
"layer_norm: weight len {} != cols {}",
w.len(),
x.cols
)));
}
if let Some(b) = bias
&& b.len() != x.cols
{
return Err(FocrError::Other(anyhow::anyhow!(
"layer_norm: bias len {} != cols {}",
b.len(),
x.cols
)));
}
let data = ft_kernel_cpu::layer_norm_forward_f32(&x.data, weight, bias, x.rows, x.cols, eps);
Ok(Mat::from_vec(x.rows, x.cols, data))
}
const SOFTMAX_COPY_ENV: &str = "FOCR_SOFTMAX_COPY";
fn softmax_copy_enabled() -> bool {
use std::sync::OnceLock;
static FLAG: OnceLock<bool> = OnceLock::new();
*FLAG.get_or_init(|| switch_on(SOFTMAX_COPY_ENV))
}
fn install_softmax_output(x: &mut Mat, out: Vec<f32>, copy: bool) {
if copy {
x.data.copy_from_slice(&out);
} else {
x.data = out;
}
}
pub fn softmax_rows(x: &mut Mat) -> FocrResult<()> {
checked_mat_len("softmax_rows x", x)?;
if x.cols == 0 || x.rows == 0 {
return Ok(());
}
let meta = meta_2d(x.rows, x.cols);
let out =
ft_kernel_cpu::softmax_dim_tensor_contiguous_f32(&x.data, &meta, 1).map_err(kernel_err)?;
if out.len() != x.data.len() {
return Err(FocrError::Other(anyhow::anyhow!(
"softmax_rows: kernel output len {} != input len {}",
out.len(),
x.data.len()
)));
}
install_softmax_output(x, out, softmax_copy_enabled());
Ok(())
}
pub fn silu(x: &mut Mat) {
for v in &mut x.data {
let s = *v;
*v = s / (1.0 + (-s).exp());
}
}
pub fn gelu(x: &mut Mat) {
const INV_SQRT2: f64 = std::f64::consts::FRAC_1_SQRT_2;
for v in &mut x.data {
let xf = f64::from(*v);
*v = (0.5 * xf * (1.0 + erf_f64(xf * INV_SQRT2))) as f32;
}
}
pub fn quick_gelu(x: &mut Mat) {
for v in &mut x.data {
let s = *v;
*v = s / (1.0 + (-1.702 * s).exp());
}
}
#[inline]
#[must_use]
pub fn quick_gelu_scalar(x: f32) -> f32 {
x / (1.0 + (-1.702 * x).exp())
}
pub fn gelu_tanh(x: &mut Mat) {
for v in &mut x.data {
*v = gelu_tanh_scalar(*v);
}
}
pub fn relu(x: &mut Mat) {
for v in &mut x.data {
*v = v.max(0.0);
}
}
#[inline]
#[must_use]
pub fn gelu_tanh_scalar(x: f32) -> f32 {
const SQRT_2_OVER_PI: f32 = 0.797_884_6;
let inner = SQRT_2_OVER_PI * (x + 0.044_715 * x * x * x);
0.5 * x * (1.0 + inner.tanh())
}
#[inline]
fn erf_f64(x: f64) -> f64 {
const A1: f64 = 0.254_829_592;
const A2: f64 = -0.284_496_736;
const A3: f64 = 1.421_413_741;
const A4: f64 = -1.453_152_027;
const A5: f64 = 1.061_405_429;
const P: f64 = 0.327_591_1;
let sign = if x < 0.0 { -1.0 } else { 1.0 };
let ax = x.abs();
let t = 1.0 / (1.0 + P * ax);
let y = 1.0 - (((((A5 * t + A4) * t) + A3) * t + A2) * t + A1) * t * (-ax * ax).exp();
sign * y
}
#[cfg(test)]
mod tests {
use super::*;
fn assert_err_contains<T>(res: FocrResult<T>, needle: &str) {
let message = match res {
Ok(_) => String::from("<ok>"),
Err(err) => err.to_string(),
};
assert!(
message.contains(needle),
"error {message:?} did not contain {needle:?}"
);
}
#[test]
fn matmul_matches_hand_computed() {
let a = Mat::from_vec(2, 3, vec![1.0, 2.0, 3.0, 4.0, 5.0, 6.0]);
let b = Mat::from_vec(3, 2, vec![7.0, 8.0, 9.0, 10.0, 11.0, 12.0]);
let c = matmul(&a, &b).unwrap();
assert_eq!(c.shape(), (2, 2));
assert_eq!(c.data, vec![58.0, 64.0, 139.0, 154.0]);
}
#[test]
fn matmul_is_row_block_invariant() {
let (m, k, n) = (301usize, 96usize, 160usize);
let a = Mat::from_vec(
m,
k,
(0..m * k)
.map(|i| ((i % 37) as f32) * 0.031 - 0.5)
.collect(),
);
let b = Mat::from_vec(
k,
n,
(0..k * n)
.map(|i| ((i % 53) as f32) * 0.017 - 0.4)
.collect(),
);
let whole = matmul(&a, &b).expect("whole matmul");
for chunk in [1usize, 7, 64, 128, 300, 301] {
let mut stitched = Vec::with_capacity(m * n);
for start in (0..m).step_by(chunk) {
let rows = chunk.min(m - start);
let part = Mat::from_vec(rows, k, a.data[start * k..(start + rows) * k].to_vec());
stitched.extend_from_slice(&matmul(&part, &b).expect("chunk matmul").data);
}
assert_eq!(
stitched, whole.data,
"row-chunked matmul (chunk {chunk}) must be bit-identical to the whole GEMM"
);
}
}
#[test]
fn matmul_rejects_inner_mismatch() {
let a = Mat::zeros(2, 3);
let b = Mat::zeros(4, 2); assert!(matmul(&a, &b).is_err());
}
#[test]
fn matmul_rejects_malformed_backing_data_without_panicking() {
let a = Mat {
rows: 1,
cols: 2,
data: vec![1.0],
};
let b = Mat::zeros(2, 1);
assert_err_contains(matmul(&a, &b), "matmul lhs: data len 1 != rows*cols 2");
}
#[test]
fn rms_norm_matches_hand_computed() {
let x = Mat::from_vec(1, 2, vec![3.0, 4.0]);
let y = rms_norm(&x, None, 0.0).unwrap();
let rstd = 1.0f32 / 12.5f32.sqrt();
assert!((y.data[0] - 3.0 * rstd).abs() < 1e-6);
assert!((y.data[1] - 4.0 * rstd).abs() < 1e-6);
}
#[test]
fn rms_norm_applies_weight() {
let x = Mat::from_vec(1, 2, vec![3.0, 4.0]);
let y = rms_norm(&x, Some(&[2.0, 0.5]), 0.0).unwrap();
let rstd = 1.0f32 / 12.5f32.sqrt();
assert!((y.data[0] - 3.0 * rstd * 2.0).abs() < 1e-6);
assert!((y.data[1] - 4.0 * rstd * 0.5).abs() < 1e-6);
}
#[test]
fn rms_norm_rejects_malformed_backing_data_without_panicking() {
let x = Mat {
rows: 2,
cols: 2,
data: vec![1.0, 2.0, 3.0],
};
assert_err_contains(rms_norm(&x, None, 1e-6), "rms_norm x: data len 3");
}
#[test]
fn layer_norm_rejects_malformed_backing_data_without_panicking() {
let x = Mat {
rows: 1,
cols: 3,
data: vec![1.0, 2.0],
};
assert_err_contains(layer_norm(&x, None, None, 1e-6), "layer_norm x: data len 2");
}
#[test]
fn quick_gelu_matches_hand_computed() {
assert!((quick_gelu_scalar(0.0)).abs() < 1e-7);
let q1 = 1.0f32 / (1.0 + (-1.702f32).exp());
assert!((quick_gelu_scalar(1.0) - q1).abs() < 1e-6);
assert!((quick_gelu_scalar(1.0) - 0.845_855).abs() < 1e-4);
let qm1 = -1.0f32 / (1.0 + (1.702f32).exp());
assert!((quick_gelu_scalar(-1.0) - qm1).abs() < 1e-6);
assert!((quick_gelu_scalar(-1.0) - (-0.154_144)).abs() < 1e-4);
}
#[test]
fn quick_gelu_mat_matches_scalar() {
let mut m = Mat::from_vec(1, 3, vec![-1.0, 0.0, 1.0]);
quick_gelu(&mut m);
assert!((m.data[0] - quick_gelu_scalar(-1.0)).abs() < 1e-7);
assert!((m.data[1] - quick_gelu_scalar(0.0)).abs() < 1e-7);
assert!((m.data[2] - quick_gelu_scalar(1.0)).abs() < 1e-7);
}
#[test]
fn silu_matches_hand_computed() {
let mut m = Mat::from_vec(1, 3, vec![0.0, 1.0, -1.0]);
silu(&mut m);
assert!(m.data[0].abs() < 1e-7);
assert!((m.data[1] - 0.731_058_6).abs() < 1e-5);
assert!((m.data[2] - (-0.268_941_4)).abs() < 1e-5);
}
#[test]
fn gelu_matches_hand_computed() {
let mut m = Mat::from_vec(1, 3, vec![0.0, 1.0, -1.0]);
gelu(&mut m);
assert!(m.data[0].abs() < 1e-6);
assert!((m.data[1] - 0.841_344_7).abs() < 1e-4);
assert!((m.data[2] - (-0.158_655_3)).abs() < 1e-4);
}
#[test]
fn softmax_rows_sums_to_one() {
let mut m = Mat::from_vec(2, 3, vec![1.0, 2.0, 3.0, 1.0, 1.0, 1.0]);
softmax_rows(&mut m).unwrap();
let r0: f32 = m.row(0).iter().sum();
let r1: f32 = m.row(1).iter().sum();
assert!((r0 - 1.0).abs() < 1e-6);
assert!((r1 - 1.0).abs() < 1e-6);
for &v in m.row(1) {
assert!((v - 1.0 / 3.0).abs() < 1e-6);
}
}
#[test]
fn softmax_output_copy_fallback_preserves_values_and_old_allocation() {
let expected = vec![0.25, 0.75];
let mut copied = Mat::from_vec(1, 2, vec![1.0, 2.0]);
let copied_ptr = copied.data.as_ptr();
install_softmax_output(&mut copied, expected.clone(), true);
assert_eq!(copied.data, expected);
assert_eq!(copied.data.as_ptr(), copied_ptr);
let mut transferred = Mat::from_vec(1, 2, vec![1.0, 2.0]);
let out = expected.clone();
let out_ptr = out.as_ptr();
install_softmax_output(&mut transferred, out, false);
assert_eq!(transferred.data, expected);
assert_eq!(transferred.data.as_ptr(), out_ptr);
}
#[test]
fn softmax_rows_rejects_malformed_backing_data_without_panicking() {
let mut x = Mat {
rows: 1,
cols: 4,
data: vec![1.0, 2.0, 3.0],
};
assert_err_contains(softmax_rows(&mut x), "softmax_rows x: data len 3");
}
#[test]
fn linear_int8_dynamic_approximates_f32() {
let x = Mat::from_vec(1, 3, vec![1.0, 2.0, 3.0]);
let w = quantize_int8(&[1.0, 0.0, 1.0, 0.0, 1.0, 0.0], 2, 3);
let y = linear_int8_dynamic(&x, &w, None).unwrap();
assert_eq!(y.shape(), (1, 2));
assert!((y.data[0] - 4.0).abs() < 0.1);
assert!((y.data[1] - 2.0).abs() < 0.1);
}
#[test]
fn linear_int8_dynamic_rejects_malformed_backing_data_without_panicking() {
let x = Mat {
rows: 1,
cols: 2,
data: vec![1.0],
};
let w = QInt8 {
w: vec![1i8, 2].into(),
scales: vec![1.0],
n: 1,
k: 2,
layout: super::super::tensor::WeightLayout::RowMajor,
};
assert_err_contains(
linear_int8_dynamic(&x, &w, None),
"linear_int8_dynamic x: data len 1",
);
let x = Mat::from_vec(1, 2, vec![1.0, 2.0]);
let short_weight = QInt8 {
w: vec![1i8].into(),
scales: vec![1.0],
n: 1,
k: 2,
layout: super::super::tensor::WeightLayout::RowMajor,
};
assert_err_contains(
linear_int8_dynamic(&x, &short_weight, None),
"linear_int8_dynamic weight: weight len 1 != n*k 2",
);
let missing_scale = QInt8 {
w: vec![1i8, 2].into(),
scales: vec![],
n: 1,
k: 2,
layout: super::super::tensor::WeightLayout::RowMajor,
};
assert_err_contains(
linear_int8_dynamic(&x, &missing_scale, None),
"linear_int8_dynamic weight: scales len 0 != n 1",
);
}
#[test]
fn linear_int8_dynamic_activation_rounds_ties_to_even() {
let x = Mat::from_vec(
2,
8,
vec![
-127.0, -2.5, -1.5, -0.5, 0.5, 1.5, 2.5, 127.0, 0.0, 0.0, 0.0, 0.0, 0.0, 0.0, 0.0, 0.0,
],
);
let w = QInt8::new(
vec![
0, 0, 0, 0, 1, 0, 0, 0, 0, 0, 0, 0, 0, 0, 1, 0,
],
vec![1.0, 1.0],
2,
8,
);
let y = linear_int8_dynamic(&x, &w, None).unwrap();
assert_eq!(y.shape(), (2, 2));
assert_eq!(y.data, vec![0.0, 2.0, 0.0, 0.0]);
let again = linear_int8_dynamic(&x, &w, None).unwrap();
assert_eq!(y, again, "dynamic activation quant must be byte-identical");
}
#[test]
fn group_norm_matches_torch_golden() {
const X: [f32; 72] = [
0.0, 0.1296341, 0.2570806, 0.3801884, 0.4968801, 0.6051864, 0.7032794, 0.7895037,
0.8624042, 0.9207506, 0.9635582, 0.9901046, 0.9999417, 0.9929036, 0.9691092, 0.9289597,
0.873133, 0.8025711, 0.7184649, 0.6222337, 0.5155014, 0.4000695, 0.277886, 0.1510129,
0.0215911, -0.1081951, -0.2361552, -0.3601299, -0.4780271, -0.5878571, -0.6877661,
-0.7760682, -0.8512733, -0.9121122, -0.957558, -0.9868438, -0.9994755, -0.9952398,
-0.9742084, -0.9367359, -0.8834547, -0.8152643, -0.7333152, -0.6389909, -0.5338824,
-0.4197641, -0.2985622, -0.1723212, -0.0431721, 0.0867056, 0.21512, 0.3399035,
0.4589513, 0.5702536, 0.6719319, 0.7622709, 0.8397455, 0.9030485, 0.9511114, 0.983123,
0.9985433, 0.997112, 0.9788533, 0.9440752, 0.8933646, 0.8275774, 0.7478238, 0.6554497,
0.5520141, 0.4392635, 0.3190989, 0.1935491,
];
const Y: [f32; 72] = [
-1.1128683, -0.9128186, -0.716145, -0.5261666, -0.3460895, -0.1789526, 0.1213925,
0.3076768, 0.4651756, 0.5912306, 0.6837148, 0.7410672, 0.9592738, 0.9367534, 0.8606158,
0.7321458, 0.5535116, 0.3277276, 0.1605169, -0.2158298, -0.6332453, -1.084684,
-1.5625268, -2.0587103, 2.4798393, 1.9679233, 1.4632101, 0.9742165, 0.509194,
0.0759915, -0.305477, -0.7073503, -1.0496179, -1.3265028, -1.5333323, -1.6666154,
-0.7448562, -0.7371473, -0.6988704, -0.6306711, -0.5337003, -0.4095948, -0.2046283,
0.0357078, 0.3035217, 0.5942926, 0.9031123, 1.2247713, -1.6645347, -1.3156482,
-0.9706925, -0.6354902, -0.3156959, -0.0167077, 0.4023003, 0.6989031, 0.9532693,
1.1611068, 1.3189077, 1.424009, 1.5083864, 1.5014455, 1.4129068, 1.2442629, 0.9983606,
0.6793492, 0.3991693, -0.1176782, -0.6964164, -1.3272738, -1.9996133, -2.7020838,
];
let weight: Vec<f32> = (0..6).map(|i| 0.5 + i as f32 * 0.2).collect();
let bias: Vec<f32> = (0..6).map(|i| -0.2 + i as f32 * 0.08).collect();
let mut x = X.to_vec();
group_norm(&mut x, 2, 6, 6, 3, 1e-5, &weight, &bias, false).unwrap();
let maxabs = x
.iter()
.zip(Y.iter())
.map(|(a, b)| (a - b).abs())
.fold(0.0f32, f32::max);
assert!(maxabs <= 2e-6, "vs torch golden: maxabs {maxabs}");
let mut fused = X.to_vec();
group_norm(&mut fused, 2, 6, 6, 3, 1e-5, &weight, &bias, true).unwrap();
let clamped: Vec<f32> = x.iter().map(|v| v.max(0.0)).collect();
assert_eq!(fused, clamped, "fused ReLU must equal unfused + clamp");
}
#[test]
fn tf_same_pad_amounts_match_timm() {
assert_eq!(tf_same_pad_amounts(128, 7, 2), (2, 3));
assert_eq!(tf_same_pad_amounts(1280, 7, 2), (2, 3));
assert_eq!(tf_same_pad_amounts(37, 3, 1), (1, 1));
assert_eq!(tf_same_pad_amounts(64, 1, 1), (0, 0));
assert_eq!(tf_same_pad_amounts(64, 1, 2), (0, 0));
assert_eq!(tf_same_pad_amounts(6, 3, 2), (0, 1));
assert_eq!(tf_same_pad_amounts(9, 3, 2), (1, 1));
}
#[test]
fn max_pool2d_same_matches_timm_golden() {
let x: Vec<f32> = (0..2 * 6 * 9).map(|i| (i as f32 * 0.37).cos()).collect();
const Y: [f32; 30] = [
1.0, 0.9323273, 0.4507553, 0.93477, 0.9999768, 0.9298415, 0.7338563, 0.7475899,
0.9999071, 0.9999071, 0.7292103, 0.4628793, 0.9395249, 0.999791, 0.9247403, 0.4262586,
0.7565722, 0.9996285, 0.9996285, 0.7198165, 0.4749172, 0.9441053, 0.9994195, 0.9194672,
0.4138885, 0.7654141, 0.9991641, 0.9991641, 0.7102889, 0.1313064,
];
let (padded, ph, pw) = tf_same_pad(&x, 1, 2, 6, 9, 3, 3, 2, 2, f32::NEG_INFINITY);
assert_eq!((ph, pw), (7, 11), "padded dims match timm pad_same");
let y = max_pool2d(&padded, 1, 2, ph, pw, 3, 2, 3, 5);
assert_eq!(y.len(), Y.len());
let maxabs = y
.iter()
.zip(Y.iter())
.map(|(a, b)| (a - b).abs())
.fold(0.0f32, f32::max);
assert!(maxabs <= 1e-6, "vs timm golden: maxabs {maxabs}");
assert!(
y.iter().all(|v| v.is_finite()),
"pool output must be finite"
);
}
#[test]
fn tf_same_pad_stem_geometry_composes_with_conv2d() {
let (h, w) = (128usize, 1280usize);
let x = vec![1.0f32; h * w];
let (padded, ph, pw) = tf_same_pad(&x, 1, 1, h, w, 7, 7, 2, 2, 0.0);
assert_eq!((ph, pw), (133, 1285), "timm pad_same golden shape");
let weight = vec![1.0f32; 7 * 7];
let (oh, ow) = (h.div_ceil(2), w.div_ceil(2));
let y = conv2d(&padded, &weight, None, 1, 1, ph, pw, 7, 7, oh, ow, 2, 2, 1);
assert_eq!(y.len(), oh * ow);
assert_eq!(y[(oh / 2) * ow + ow / 2], 49.0, "interior window sums 7×7");
assert_eq!(y[0], 25.0, "corner window is 5×5 real after 2/3-2/3 pads");
}
#[test]
fn group_norm_rejects_bad_shapes() {
let w = vec![1.0f32; 6];
let b = vec![0.0f32; 6];
let mut x = vec![0.0f32; 2 * 6 * 6];
assert!(group_norm(&mut x, 2, 6, 6, 4, 1e-5, &w, &b, false).is_err());
assert!(group_norm(&mut x, 2, 6, 6, 0, 1e-5, &w, &b, false).is_err());
let mut short = vec![0.0f32; 5];
assert!(group_norm(&mut short, 2, 6, 6, 3, 1e-5, &w, &b, false).is_err());
let mut x2 = vec![0.0f32; 2 * 6 * 6];
assert!(group_norm(&mut x2, 2, 6, 6, 3, 1e-5, &w[..4], &b, false).is_err());
}
}