use half::f16;
use crate::distance::dot::dot;
#[inline]
pub fn dot_f16_batch_16(query: &[f16], candidates: &[&[f16]; 16], len: usize) -> [f32; 16] {
assert!(
candidates
.iter()
.all(|candidate| candidate.len() == query.len()),
"all candidate vectors must have the same length as query"
);
assert!(
(1..=16).contains(&len),
"batch length must be in 1..=16, got {len}"
);
#[cfg(all(
kernel_support = "amx_fp16",
target_arch = "x86_64",
target_os = "linux"
))]
{
if query.len() >= 32 && crate::simd::amx_fp16::amx_supported() {
return unsafe { crate::simd::amx_fp16::dot_f16_batch_16_amx(query, candidates, len) };
}
}
dot_f16_batch_16_fallback(query, candidates, len)
}
#[inline]
pub(crate) fn dot_f16_batch_16_fallback(
query: &[f16],
candidates: &[&[f16]; 16],
len: usize,
) -> [f32; 16] {
std::array::from_fn(|i| {
if i < len {
dot(query, candidates[i])
} else {
0.0
}
})
}
pub struct PackedCentroidsF16(Packed);
#[cfg(all(
kernel_support = "amx_fp16",
target_arch = "x86_64",
target_os = "linux"
))]
struct Packed {
centroids: Vec<f16>,
packed: Vec<f16>,
n_padded: usize,
dim: usize,
}
#[cfg(not(all(
kernel_support = "amx_fp16",
target_arch = "x86_64",
target_os = "linux"
)))]
enum Packed {}
pub(crate) fn strided_len(m: usize, stride: usize, row_len: usize) -> Option<usize> {
let Some(last_row) = m.checked_sub(1) else {
return Some(0);
};
last_row.checked_mul(stride)?.checked_add(row_len)
}
impl PackedCentroidsF16 {
pub fn new(centroids: &[f16], n: usize, dim: usize) -> Option<Self> {
let expected = n
.checked_mul(dim)
.unwrap_or_else(|| panic!("centroid shape n = {n} x dim = {dim} overflows usize"));
assert_eq!(
centroids.len(),
expected,
"centroids must hold n*dim = {expected} values, got {}",
centroids.len()
);
if n == 0 || dim == 0 || !amx_fp16_supported() {
return None;
}
#[cfg(all(
kernel_support = "amx_fp16",
target_arch = "x86_64",
target_os = "linux"
))]
{
let n_padded = n.checked_next_multiple_of(32).unwrap_or_else(|| {
panic!("padding n = {n} up to a multiple of 32 overflows usize")
});
let padded_len = n_padded.checked_mul(dim).unwrap_or_else(|| {
panic!("padded centroid shape {n_padded} x dim = {dim} overflows usize")
});
let mut padded = vec![f16::ZERO; padded_len];
padded[..centroids.len()].copy_from_slice(centroids);
let mut packed =
Vec::with_capacity(crate::simd::amx_fp16::packed_centroids_len(n_padded, dim));
crate::simd::amx_fp16::pack_centroids_vnni(&padded, n_padded, dim, &mut packed);
return Some(Self(Packed {
centroids: padded,
packed,
n_padded,
dim,
}));
}
#[allow(unreachable_code)]
None
}
pub fn num_centroids_padded(&self) -> usize {
self.shape().0
}
fn shape(&self) -> (usize, usize) {
#[cfg(all(
kernel_support = "amx_fp16",
target_arch = "x86_64",
target_os = "linux"
))]
{
(self.0.n_padded, self.0.dim)
}
#[cfg(not(all(
kernel_support = "amx_fp16",
target_arch = "x86_64",
target_os = "linux"
)))]
{
match self.0 {}
}
}
pub fn score(
&self,
data: &[f16],
m: usize,
data_stride: usize,
out: &mut [f32],
out_stride: usize,
) {
let (n_padded, dim) = self.shape();
assert_eq!(m % 32, 0, "m ({m}) must be a multiple of 32");
assert!(
data_stride >= dim,
"data_stride ({data_stride}) is below dim ({dim})"
);
assert!(
out_stride >= n_padded,
"out_stride ({out_stride}) is below the padded centroid count ({n_padded})"
);
if m == 0 {
return;
}
let data_needed = strided_len(m, data_stride, dim).unwrap_or_else(|| {
panic!("m = {m} rows of dim {dim} at stride {data_stride} overflow usize")
});
assert!(
data.len() >= data_needed,
"data ({}) holds fewer than m = {m} rows of dim {dim} at stride {data_stride}",
data.len()
);
let out_needed = strided_len(m, out_stride, n_padded).unwrap_or_else(|| {
panic!("m = {m} rows of {n_padded} at stride {out_stride} overflow usize")
});
assert!(
out.len() >= out_needed,
"out ({}) holds fewer than m = {m} rows of {n_padded} at stride {out_stride}",
out.len()
);
#[cfg(all(
kernel_support = "amx_fp16",
target_arch = "x86_64",
target_os = "linux"
))]
unsafe {
crate::simd::amx_fp16::dot_f16_gemm_amx(
data,
m,
data_stride,
&self.0.packed,
&self.0.centroids,
n_padded,
dim,
out,
out_stride,
);
}
}
}
pub fn amx_fp16_supported() -> bool {
#[cfg(all(
kernel_support = "amx_fp16",
target_arch = "x86_64",
target_os = "linux"
))]
{
return crate::simd::amx_fp16::amx_supported();
}
#[allow(unreachable_code)]
false
}
pub fn amx_fp16_available() -> bool {
!amx_fp16_disabled() && amx_fp16_supported()
}
fn amx_fp16_disabled() -> bool {
static DISABLED: std::sync::OnceLock<bool> = std::sync::OnceLock::new();
*DISABLED.get_or_init(|| {
std::env::var("LANCE_DISABLE_AMX").is_ok_and(|value| is_amx_disable_value(&value))
})
}
fn is_amx_disable_value(value: &str) -> bool {
matches!(
value.trim().to_ascii_lowercase().as_str(),
"1" | "true" | "on"
)
}
#[cfg(test)]
mod tests {
use super::*;
use rand::rngs::StdRng;
use rand::{Rng, SeedableRng};
const BATCH_DIMS: &[usize] = &[
1, 7, 31, 32, 33, 47, 64, 96, 100, 127, 128, 200, 256, 384, 768, 1000, 1536,
];
fn make_batch(dim: usize, rng: &mut StdRng) -> (Vec<f16>, Vec<Vec<f16>>) {
let gen_vec = |rng: &mut StdRng| -> Vec<f16> {
(0..dim)
.map(|_| f16::from_f32(rng.random_range(-1.0f32..1.0)))
.collect()
};
let query = gen_vec(rng);
let candidates = (0..16).map(|_| gen_vec(rng)).collect();
(query, candidates)
}
fn ref_dot_f32(query: &[f16], cand: &[f16]) -> f32 {
query
.iter()
.zip(cand.iter())
.map(|(&q, &c)| q.to_f32() * c.to_f32())
.sum()
}
const REL_TOL: f32 = 5e-3;
fn assert_close(got: f32, want: f32, ctx: &str) {
let rel = (got - want).abs() / (want.abs() + 1e-6);
assert!(
rel <= REL_TOL || (got - want).abs() <= 1e-3,
"{ctx}: got {got} want {want} rel_err {rel}"
);
}
#[test]
fn fallback_matches_reference() {
let mut rng = StdRng::seed_from_u64(0xF16);
for &dim in BATCH_DIMS {
let (query, cands) = make_batch(dim, &mut rng);
let candidates: [&[f16]; 16] = std::array::from_fn(|i| cands[i].as_slice());
let got = dot_f16_batch_16_fallback(&query, &candidates, 16);
for i in 0..16 {
assert_close(
got[i],
ref_dot_f32(&query, &cands[i]),
&format!("fb dim={dim} i={i}"),
);
}
}
}
#[test]
fn dispatch_matches_reference() {
let mut rng = StdRng::seed_from_u64(0xBEEF);
for &dim in BATCH_DIMS {
let (query, cands) = make_batch(dim, &mut rng);
let candidates: [&[f16]; 16] = std::array::from_fn(|i| cands[i].as_slice());
let got = dot_f16_batch_16(&query, &candidates, 16);
for i in 0..16 {
assert_close(
got[i],
ref_dot_f32(&query, &cands[i]),
&format!("disp dim={dim} i={i}"),
);
}
}
}
#[test]
#[should_panic(expected = "all candidate vectors must have the same length")]
fn rejects_mismatched_candidate_length() {
let query = vec![f16::from_f32(1.0); 32];
let short = vec![f16::from_f32(1.0); 31];
let candidates: [&[f16]; 16] = std::array::from_fn(|i| {
if i == 0 {
short.as_slice()
} else {
query.as_slice()
}
});
let _ = dot_f16_batch_16(&query, &candidates, 16);
}
#[test]
fn partial_len_matches_full_batch_and_zeroes_the_rest() {
let mut rng = StdRng::seed_from_u64(0x1EE);
for &dim in BATCH_DIMS {
let (query, cands) = make_batch(dim, &mut rng);
let candidates: [&[f16]; 16] = std::array::from_fn(|i| cands[i].as_slice());
let full = dot_f16_batch_16(&query, &candidates, 16);
let full_fb = dot_f16_batch_16_fallback(&query, &candidates, 16);
for len in 1..=16 {
let got = dot_f16_batch_16(&query, &candidates, len);
let got_fb = dot_f16_batch_16_fallback(&query, &candidates, len);
for i in 0..len {
assert_eq!(
got[i].to_bits(),
full[i].to_bits(),
"dim={dim} len={len} i={i}: {} vs {}",
got[i],
full[i]
);
assert_eq!(
got_fb[i].to_bits(),
full_fb[i].to_bits(),
"fb dim={dim} i={i}"
);
}
for i in len..16 {
assert_eq!(got[i], 0.0, "dim={dim} len={len}: lane {i} must be 0");
assert_eq!(got_fb[i], 0.0, "fb dim={dim} len={len}: lane {i} must be 0");
}
}
}
}
#[rstest::rstest]
#[case::zero(0)]
#[case::seventeen(17)]
#[should_panic(expected = "batch length must be in 1..=16")]
fn rejects_out_of_range_len(#[case] len: usize) {
let query = vec![f16::from_f32(1.0); 32];
let candidates: [&[f16]; 16] = std::array::from_fn(|_| query.as_slice());
let _ = dot_f16_batch_16(&query, &candidates, len);
}
#[test]
fn amx_disable_flag_accepts_only_explicit_on() {
for value in ["1", "true", "on", "TRUE", "On", " 1 ", "true\n"] {
assert!(is_amx_disable_value(value), "{value:?} should disable AMX");
}
for value in ["", " ", "0", "false", "off", "no", "yes", "2", "disable"] {
assert!(
!is_amx_disable_value(value),
"{value:?} should not disable AMX"
);
}
}
#[test]
fn amx_is_available_by_default_wherever_it_is_supported() {
if std::env::var_os("LANCE_DISABLE_AMX").is_some() {
return; }
assert_eq!(amx_fp16_available(), amx_fp16_supported());
}
#[test]
fn strided_len_rejects_shapes_it_cannot_represent() {
assert_eq!(strided_len(32, 595_056_260_442_243_601, 32), None);
assert_eq!(strided_len(2, usize::MAX, 1), None);
assert_eq!(strided_len(usize::MAX, 2, 0), None);
assert_eq!(strided_len(32, 768, 768), Some(31 * 768 + 768));
assert_eq!(strided_len(1, usize::MAX, 5), Some(5));
assert_eq!(strided_len(0, usize::MAX, usize::MAX), Some(0));
}
#[cfg(all(
kernel_support = "amx_fp16",
target_arch = "x86_64",
target_os = "linux"
))]
#[test]
fn score_rejects_overflowing_stride_before_ffi() {
use std::panic::{AssertUnwindSafe, catch_unwind};
let centroids = vec![f16::ONE; 32 * 32];
let Some(packed) = PackedCentroidsF16::new(¢roids, 32, 32) else {
return; };
let data = vec![f16::ZERO; 47];
let mut out = vec![0f32; 32 * 32];
let hook = std::panic::take_hook();
std::panic::set_hook(Box::new(|_| {})); let result = catch_unwind(AssertUnwindSafe(|| {
packed.score(&data, 32, 595_056_260_442_243_601, &mut out, 32);
}));
std::panic::set_hook(hook);
assert!(
result.is_err(),
"score accepted a 47-element slice for a stride whose row count overflows usize"
);
}
#[cfg(all(
kernel_support = "amx_fp16",
target_arch = "x86_64",
target_os = "linux"
))]
#[test]
fn amx_path_is_active_and_close() {
if !crate::simd::amx_fp16::amx_supported() {
return;
}
let mut rng = StdRng::seed_from_u64(0xA11);
let mut worst = 0f32;
for &dim in BATCH_DIMS.iter().filter(|&&d| d >= 32) {
let (query, cands) = make_batch(dim, &mut rng);
let candidates: [&[f16]; 16] = std::array::from_fn(|i| cands[i].as_slice());
let amx =
unsafe { crate::simd::amx_fp16::dot_f16_batch_16_amx(&query, &candidates, 16) };
for i in 0..16 {
let want = ref_dot_f32(&query, &cands[i]);
assert_close(amx[i], want, &format!("amx dim={dim} i={i}"));
let rel = (amx[i] - want).abs() / (want.abs() + 1e-6);
worst = worst.max(rel);
}
}
assert!(worst <= REL_TOL, "worst AMX relative error: {worst:.2e}");
}
#[cfg(all(
kernel_support = "amx_fp16",
target_arch = "x86_64",
target_os = "linux"
))]
#[test]
fn kernel_reconfigures_after_foreign_tile_release() {
use crate::simd::amx_fp16::{
amx_supported, clobber_tile_state_for_test, dot_f16_batch_16_amx,
tile_config_is_live_for_test,
};
if !amx_supported() {
return;
}
let mut rng = StdRng::seed_from_u64(0xC10B);
let (query, cands) = make_batch(256, &mut rng);
let candidates: [&[f16]; 16] = std::array::from_fn(|i| cands[i].as_slice());
let before = unsafe { dot_f16_batch_16_amx(&query, &candidates, 16) };
unsafe { clobber_tile_state_for_test() };
let after = unsafe { dot_f16_batch_16_amx(&query, &candidates, 16) };
assert_eq!(
before, after,
"kernel output changed after a foreign TILERELEASE"
);
assert!(
!tile_config_is_live_for_test(),
"a kernel left its tile configuration loaded after returning"
);
}
#[cfg(all(
kernel_support = "amx_fp16",
target_arch = "x86_64",
target_os = "linux"
))]
#[test]
fn search_tile_config_image_is_pinned() {
use crate::simd::amx_fp16::{AMX_CFG_SEARCH, amx_supported, tilecfg_image};
if !amx_supported() {
return;
}
const COLSB: usize = 16;
const ROWS: usize = 48;
let mut want = [0u8; 64];
want[0] = 1; for (tmm, colsb) in [
(0usize, 4u16),
(1, 64),
(2, 4),
(3, 64),
(4, 4),
(5, 64),
(6, 4),
] {
want[COLSB + tmm * 2..COLSB + tmm * 2 + 2].copy_from_slice(&colsb.to_le_bytes());
want[ROWS + tmm] = 16;
}
let got = tilecfg_image(AMX_CFG_SEARCH).expect("search config kind must be known");
assert_eq!(got, want, "batch-16 tile configuration changed");
}
#[cfg(all(
kernel_support = "amx_fp16",
target_arch = "x86_64",
target_os = "linux"
))]
#[test]
fn unknown_tile_config_kind_is_rejected() {
if !crate::simd::amx_fp16::amx_supported() {
return;
}
assert!(crate::simd::amx_fp16::tilecfg_image(-1).is_none());
}
fn make_gemm(m: usize, n: usize, dim: usize, rng: &mut StdRng) -> (Vec<f16>, Vec<f16>) {
let mut sample = |count: usize| -> Vec<f16> {
(0..count)
.map(|_| f16::from_f32(rng.random_range(-1.0f32..1.0)))
.collect()
};
(sample(m * dim), sample(n * dim))
}
fn ref_gemm(
data: &[f16],
m: usize,
data_stride: usize,
centroids: &[f16],
n: usize,
dim: usize,
) -> Vec<f32> {
let mut out = vec![0f32; m * n];
for i in 0..m {
let row = &data[i * data_stride..i * data_stride + dim];
for j in 0..n {
out[i * n + j] = ref_dot_f32(row, ¢roids[j * dim..j * dim + dim]);
}
}
out
}
#[cfg(all(
kernel_support = "amx_fp16",
target_arch = "x86_64",
target_os = "linux"
))]
fn run_gemm(
data: &[f16],
m: usize,
data_stride: usize,
centroids: &[f16],
n: usize,
dim: usize,
out_stride: usize,
) -> Vec<f32> {
use crate::simd::amx_fp16::{dot_f16_gemm_amx, pack_centroids_vnni};
let mut packed = Vec::new();
pack_centroids_vnni(centroids, n, dim, &mut packed);
let mut out = vec![f32::NAN; m * out_stride];
unsafe {
dot_f16_gemm_amx(
data,
m,
data_stride,
&packed,
centroids,
n,
dim,
&mut out,
out_stride,
);
}
out
}
fn assert_gemm_close(
got: &[f32],
want: &[f32],
m: usize,
n: usize,
out_stride: usize,
ctx: &str,
) {
for i in 0..m {
for j in 0..n {
assert_close(
got[i * out_stride + j],
want[i * n + j],
&format!("{ctx} [{i}][{j}]"),
);
}
}
}
#[cfg(all(
kernel_support = "amx_fp16",
target_arch = "x86_64",
target_os = "linux"
))]
#[test]
fn packed_centroid_layout_matches_tile_operand_order() {
use crate::simd::amx_fp16::{pack_centroids_vnni, packed_centroids_len};
let mut packed = Vec::new();
let mut rng = StdRng::seed_from_u64(0x9EC7);
for (n, dim) in [(16usize, 32usize), (32, 64), (48, 100), (32, 31), (32, 768)] {
let centroids: Vec<f16> = (0..n * dim)
.map(|_| f16::from_f32(rng.random_range(-1.0f32..1.0)))
.collect();
pack_centroids_vnni(¢roids, n, dim, &mut packed);
assert_eq!(
packed.len(),
packed_centroids_len(n, dim),
"n={n} dim={dim}"
);
for kb in 0..dim / 32 {
for jb in 0..n / 16 {
for k in 0..16 {
for nn in 0..16 {
for p in 0..2 {
let at = ((kb * (n / 16)) + jb) * 512 + k * 32 + nn * 2 + p;
let from = (jb * 16 + nn) * dim + kb * 32 + 2 * k + p;
assert_eq!(
packed[at], centroids[from],
"n={n} dim={dim} kb={kb} jb={jb} k={k} nn={nn} p={p}"
);
}
}
}
}
}
}
}
#[cfg(all(
kernel_support = "amx_fp16",
target_arch = "x86_64",
target_os = "linux"
))]
#[test]
fn gemm_matches_reference() {
if !crate::simd::amx_fp16::amx_supported() {
return;
}
let mut rng = StdRng::seed_from_u64(0x6E33);
for &m in &[32usize, 64] {
for &n in &[32usize, 64] {
for &dim in &[1usize, 16, 31, 32, 33, 64, 100, 768, 1000, 1536] {
let (data, centroids) = make_gemm(m, n, dim, &mut rng);
let want = ref_gemm(&data, m, dim, ¢roids, n, dim);
let got = run_gemm(&data, m, dim, ¢roids, n, dim, n);
assert_gemm_close(&got, &want, m, n, n, &format!("gemm m={m} n={n} dim={dim}"));
}
}
}
let (m, n, dim) = (64usize, 32usize, 100usize);
let (data_stride, out_stride) = (dim + 7, n + 5);
let mut data: Vec<f16> = vec![f16::from_f32(f32::MAX); m * data_stride];
let (tight, centroids) = make_gemm(m, n, dim, &mut rng);
for i in 0..m {
data[i * data_stride..i * data_stride + dim]
.copy_from_slice(&tight[i * dim..(i + 1) * dim]);
}
let want = ref_gemm(&data, m, data_stride, ¢roids, n, dim);
let got = run_gemm(&data, m, data_stride, ¢roids, n, dim, out_stride);
assert_gemm_close(&got, &want, m, n, out_stride, "gemm padded strides");
}
#[test]
fn packed_centroids_pad_to_the_kernel_block() {
let mut rng = StdRng::seed_from_u64(0x9AD);
for (n, dim) in [(32usize, 64usize), (100, 100), (48, 768)] {
let m = 64;
let (data, centroids) = make_gemm(m, n, dim, &mut rng);
let Some(packed) = PackedCentroidsF16::new(¢roids, n, dim) else {
return; };
let n_padded = packed.num_centroids_padded();
assert_eq!(n_padded, n.next_multiple_of(32), "n={n}");
let out_stride = n_padded + 3;
let mut out = vec![f32::NAN; m * out_stride];
packed.score(&data, m, dim, &mut out, out_stride);
let want = ref_gemm(&data, m, dim, ¢roids, n, dim);
assert_gemm_close(&out, &want, m, n, out_stride, &format!("packed n={n}"));
for i in 0..m {
for j in n..n_padded {
assert_eq!(out[i * out_stride + j], 0.0, "padding n={n} [{i}][{j}]");
}
}
}
}
#[cfg(all(
kernel_support = "amx_fp16",
target_arch = "x86_64",
target_os = "linux"
))]
#[test]
fn gemm_tile_config_image_is_pinned() {
use crate::simd::amx_fp16::{AMX_CFG_GEMM, amx_supported, tilecfg_image};
if !amx_supported() {
return;
}
const COLSB: usize = 16;
const ROWS: usize = 48;
let mut want = [0u8; 64];
want[0] = 1; for tmm in 0..8usize {
want[COLSB + tmm * 2..COLSB + tmm * 2 + 2].copy_from_slice(&64u16.to_le_bytes());
want[ROWS + tmm] = 16;
}
let got = tilecfg_image(AMX_CFG_GEMM).expect("gemm config kind must be known");
assert_eq!(got, want, "gemm tile configuration changed");
}
#[cfg(all(
kernel_support = "amx_fp16",
target_arch = "x86_64",
target_os = "linux"
))]
#[test]
fn interleaved_search_and_gemm_stay_correct() {
if !crate::simd::amx_fp16::amx_supported() {
return;
}
let mut rng = StdRng::seed_from_u64(0x11E12EA5);
let (m, n, dim) = (32usize, 32usize, 96usize);
for round in 0..20 {
let (query, cands) = make_batch(dim, &mut rng);
let candidates: [&[f16]; 16] = std::array::from_fn(|i| cands[i].as_slice());
let batch =
unsafe { crate::simd::amx_fp16::dot_f16_batch_16_amx(&query, &candidates, 16) };
for i in 0..16 {
assert_close(
batch[i],
ref_dot_f32(&query, &cands[i]),
&format!("interleaved batch round={round} i={i}"),
);
}
let (data, centroids) = make_gemm(m, n, dim, &mut rng);
let want = ref_gemm(&data, m, dim, ¢roids, n, dim);
let got = run_gemm(&data, m, dim, ¢roids, n, dim, n);
assert_gemm_close(
&got,
&want,
m,
n,
n,
&format!("interleaved gemm round={round}"),
);
}
}
#[cfg(all(
kernel_support = "amx_fp16",
target_arch = "x86_64",
target_os = "linux"
))]
#[test]
fn concurrent_search_and_gemm_stay_correct() {
if !crate::simd::amx_fp16::amx_supported() {
return;
}
const THREADS: u64 = 8;
std::thread::scope(|scope| {
for t in 0..THREADS {
scope.spawn(move || {
let mut rng = StdRng::seed_from_u64(0xC0FFEE + t);
let dim = 128;
for round in 0..25 {
if t % 2 == 0 {
let (query, cands) = make_batch(dim, &mut rng);
let candidates: [&[f16]; 16] =
std::array::from_fn(|i| cands[i].as_slice());
let batch = unsafe {
crate::simd::amx_fp16::dot_f16_batch_16_amx(&query, &candidates, 16)
};
for i in 0..16 {
assert_close(
batch[i],
ref_dot_f32(&query, &cands[i]),
&format!("concurrent batch t={t} round={round} i={i}"),
);
}
} else {
let (m, n) = (32usize, 32usize);
let (data, centroids) = make_gemm(m, n, dim, &mut rng);
let want = ref_gemm(&data, m, dim, ¢roids, n, dim);
let got = run_gemm(&data, m, dim, ¢roids, n, dim, n);
assert_gemm_close(
&got,
&want,
m,
n,
n,
&format!("concurrent gemm t={t} round={round}"),
);
}
}
});
}
});
}
#[test]
#[ignore]
#[allow(clippy::print_stderr)]
#[allow(unreachable_code, unused_variables)]
fn packed_centroids_gemm_shape_bench() {
use std::time::{Duration, Instant};
const BLOCK_ROWS: &[usize] = &[32, 64, 128, 256, 512, 1024, 2048];
if !amx_fp16_supported() {
eprintln!("[gemm_shape_bench] skipped: amx_fp16_supported=false on this build or host");
return;
}
let env_list = |key: &str, default: &[usize]| -> Vec<usize> {
std::env::var(key)
.ok()
.map(|s| s.split(',').filter_map(|t| t.trim().parse().ok()).collect())
.unwrap_or_else(|| default.to_vec())
};
let dims = env_list("BENCH_DIMS", &[768]);
let ks = env_list("BENCH_KS", &[256, 4096]);
let budget = Duration::from_secs_f64(
std::env::var("BENCH_SECONDS")
.ok()
.and_then(|s| s.parse().ok())
.unwrap_or(3.0),
);
let mut rng = StdRng::seed_from_u64(0x6E33);
let mut random_f16 = |count: usize| -> Vec<f16> {
(0..count)
.map(|_| f16::from_f32(rng.random_range(-1.0f32..1.0)))
.collect()
};
eprintln!(
"[gemm_shape_bench] budget={:.1}s amx_fp16_supported=true",
budget.as_secs_f64()
);
for &dim in &dims {
for &k in &ks {
let centroids = random_f16(k * dim);
let mut packs = 0usize;
let t0 = Instant::now();
while t0.elapsed() < budget {
let packed = PackedCentroidsF16::new(¢roids, k, dim);
std::hint::black_box(&packed);
packs += 1;
}
let pack_us = t0.elapsed().as_secs_f64() * 1e6 / packs as f64;
let packed = PackedCentroidsF16::new(¢roids, k, dim)
.expect("availability was just checked");
let n_padded = packed.num_centroids_padded();
eprintln!(
"[gemm_shape_bench] dim={dim} k={k} n_padded={n_padded} pack_calls={packs} pack_us={pack_us:.1}"
);
let mut best_vec_per_s = 0f64;
for &m in BLOCK_ROWS {
let data = random_f16(m * dim);
let mut out = vec![0f32; m * n_padded];
packed.score(&data, m, dim, &mut out, n_padded);
let t1 = Instant::now();
let mut iters = 0usize;
while t1.elapsed() < budget {
packed.score(&data, m, dim, &mut out, n_padded);
iters += 1;
}
let elapsed = t1.elapsed().as_secs_f64();
std::hint::black_box(&out);
let vec_per_s = (iters * m) as f64 / elapsed;
best_vec_per_s = best_vec_per_s.max(vec_per_s);
eprintln!(
"[gemm_shape_bench] m={m:>5} scratch_kb={:>7} iters={iters:>8} vec_per_s={vec_per_s:>12.0} us_per_vec={:>8.4} Gpair_per_s={:>8.2}",
m * n_padded * 4 / 1024,
1e6 / vec_per_s,
vec_per_s * n_padded as f64 / 1e9,
);
}
eprintln!(
"[gemm_shape_bench] pack_us={pack_us:.1} buys {:.0} vectors of scoring at the best m: packing is amortized above that",
pack_us * 1e-6 * best_vec_per_s,
);
}
}
}
}