use std::sync::OnceLock;
#[derive(Debug, Clone, Copy, PartialEq, Eq, PartialOrd, Ord, Hash)]
pub enum IsaTier {
Scalar = 0,
Avx2 = 1,
AvxVnni = 2,
Avx512Vnni = 3,
Sdot = 4,
Smmla = 5,
WasmSimd128 = 6,
}
impl IsaTier {
#[must_use]
pub fn feature_string(self) -> &'static str {
match self {
IsaTier::Scalar => "scalar",
IsaTier::Avx2 => "x86_64+avx2",
IsaTier::AvxVnni => "x86_64+avx2+avxvnni",
IsaTier::Avx512Vnni => "x86_64+avx512vnni",
IsaTier::Sdot => "aarch64+neon+dotprod",
IsaTier::Smmla => "aarch64+neon+i8mm",
IsaTier::WasmSimd128 => "wasm32+simd128",
}
}
#[must_use]
pub fn tag(self) -> &'static str {
match self {
IsaTier::Scalar => "scalar",
IsaTier::Avx2 => "avx2",
IsaTier::AvxVnni => "avxvnni",
IsaTier::Avx512Vnni => "avx512vnni",
IsaTier::Sdot => "sdot",
IsaTier::Smmla => "smmla",
IsaTier::WasmSimd128 => "wasmsimd128",
}
}
}
#[derive(Debug, Clone, Copy, PartialEq, Eq, PartialOrd, Ord, Hash)]
pub enum EffectiveI8Route {
Autovec,
Scalar,
Avx2,
AvxVnni,
Avx512Vnni,
Sdot,
Smmla,
WasmSimd128,
}
impl EffectiveI8Route {
#[must_use]
pub fn tag(self) -> &'static str {
match self {
Self::Autovec => "autovec",
Self::Scalar => "scalar",
Self::Avx2 => "avx2",
Self::AvxVnni => "avxvnni",
Self::Avx512Vnni => "avx512vnni",
Self::Sdot => "sdot",
Self::Smmla => "smmla",
Self::WasmSimd128 => "wasmsimd128",
}
}
#[must_use]
pub fn feature_string(self) -> &'static str {
match self {
Self::Autovec => "aarch64+llvm-autovec",
Self::Scalar => "scalar",
Self::Avx2 => IsaTier::Avx2.feature_string(),
Self::AvxVnni => IsaTier::AvxVnni.feature_string(),
Self::Avx512Vnni => IsaTier::Avx512Vnni.feature_string(),
Self::Sdot => IsaTier::Sdot.feature_string(),
Self::Smmla => IsaTier::Smmla.feature_string(),
Self::WasmSimd128 => IsaTier::WasmSimd128.feature_string(),
}
}
#[cfg(not(target_arch = "aarch64"))]
fn from_isa(tier: IsaTier) -> Self {
match tier {
IsaTier::Scalar => Self::Scalar,
IsaTier::Avx2 => Self::Avx2,
IsaTier::AvxVnni => Self::AvxVnni,
IsaTier::Avx512Vnni => Self::Avx512Vnni,
IsaTier::Sdot => Self::Sdot,
IsaTier::Smmla => Self::Smmla,
IsaTier::WasmSimd128 => Self::WasmSimd128,
}
}
#[cfg(target_arch = "aarch64")]
fn from_arm(route: super::arm::DenseI8Route) -> Self {
match route {
super::arm::DenseI8Route::Autovec => Self::Autovec,
super::arm::DenseI8Route::Scalar => Self::Scalar,
super::arm::DenseI8Route::Sdot => Self::Sdot,
super::arm::DenseI8Route::Smmla => Self::Smmla,
}
}
}
#[derive(Debug, Clone)]
pub struct Caps {
pub selected: IsaTier,
pub available: Vec<IsaTier>,
}
static CAPS: OnceLock<Caps> = OnceLock::new();
#[must_use]
pub fn caps() -> &'static Caps {
CAPS.get_or_init(detect)
}
#[must_use]
pub fn detected_tier() -> IsaTier {
caps().selected
}
#[must_use]
pub fn available_tiers() -> &'static [IsaTier] {
&caps().available
}
#[must_use]
pub fn tier_string() -> &'static str {
detected_tier().feature_string()
}
#[must_use]
pub fn effective_dense_route() -> EffectiveI8Route {
#[cfg(target_arch = "aarch64")]
{
EffectiveI8Route::from_arm(super::arm::effective_dense_route())
}
#[cfg(not(target_arch = "aarch64"))]
{
EffectiveI8Route::from_isa(detected_tier())
}
}
fn detect() -> Caps {
let mut available: Vec<IsaTier> = Vec::new();
#[cfg(target_arch = "aarch64")]
{
let has_i8mm = std::arch::is_aarch64_feature_detected!("i8mm");
let has_dotprod = std::arch::is_aarch64_feature_detected!("dotprod");
#[cfg(target_vendor = "apple")]
{
if has_dotprod {
available.push(IsaTier::Sdot);
}
if has_i8mm {
available.push(IsaTier::Smmla);
}
}
#[cfg(not(target_vendor = "apple"))]
{
if has_i8mm {
available.push(IsaTier::Smmla);
}
if has_dotprod {
available.push(IsaTier::Sdot);
}
}
}
#[cfg(target_arch = "x86_64")]
{
if std::arch::is_x86_feature_detected!("avx512vnni")
&& std::arch::is_x86_feature_detected!("avx512bw")
&& std::arch::is_x86_feature_detected!("avx512f")
{
available.push(IsaTier::Avx512Vnni);
}
if std::arch::is_x86_feature_detected!("avxvnni") {
available.push(IsaTier::AvxVnni);
}
if std::arch::is_x86_feature_detected!("avx2") {
available.push(IsaTier::Avx2);
}
}
#[cfg(all(target_arch = "wasm32", target_feature = "simd128"))]
{
available.push(IsaTier::WasmSimd128);
}
available.push(IsaTier::Scalar);
let selected = std::env::var("FOCR_FORCE_ARCH")
.ok()
.map(|force| force.trim().to_ascii_lowercase())
.and_then(|want| available.iter().copied().find(|tier| tier.tag() == want))
.unwrap_or(available[0]);
Caps {
selected,
available,
}
}
pub fn igemm_s8s8(a: &[i8], b: &[i8], m: usize, k: usize, n: usize, out: &mut [i32]) {
let _ = igemm_s8s8_with_route(a, b, m, k, n, out);
}
fn igemm_s8s8_with_route(
a: &[i8],
b: &[i8],
m: usize,
k: usize,
n: usize,
out: &mut [i32],
) -> EffectiveI8Route {
#[cfg(target_arch = "aarch64")]
{
EffectiveI8Route::from_arm(super::arm::igemm_s8s8_with_route(a, b, m, k, n, out))
}
#[cfg(target_arch = "x86_64")]
{
super::x86::igemm_s8s8_with_route(a, b, m, k, n, out)
}
#[cfg(all(target_arch = "wasm32", target_feature = "simd128"))]
{
super::wasm128::igemm_s8s8(a, b, m, k, n, out);
EffectiveI8Route::WasmSimd128
}
#[cfg(not(any(
target_arch = "aarch64",
target_arch = "x86_64",
all(target_arch = "wasm32", target_feature = "simd128")
)))]
{
super::scalar::igemm_s8s8(a, b, m, k, n, out);
EffectiveI8Route::Scalar
}
}
pub fn igemm_s8s8_packed_b(
a: &[i8],
b_panels: &[i8],
m: usize,
k: usize,
n: usize,
out: &mut [i32],
) {
#[cfg(target_arch = "aarch64")]
{
super::arm::igemm_s8s8_packed_b(a, b_panels, m, k, n, out);
}
#[cfg(not(target_arch = "aarch64"))]
{
let b = super::pack::smmla_unpack_panels(b_panels, n, k)
.expect("igemm_s8s8_packed_b: panel length contract violated");
out.fill(0);
igemm_s8s8(a, &b, m, k, n, out);
}
}
pub fn igemm_u8s8(a: &[u8], b: &[i8], m: usize, k: usize, n: usize, out: &mut [i32]) {
let _ = igemm_u8s8_with_route(a, b, m, k, n, out);
}
fn igemm_u8s8_with_route(
a: &[u8],
b: &[i8],
m: usize,
k: usize,
n: usize,
out: &mut [i32],
) -> EffectiveI8Route {
#[cfg(target_arch = "aarch64")]
{
EffectiveI8Route::from_arm(super::arm::igemm_u8s8_with_route(a, b, m, k, n, out))
}
#[cfg(target_arch = "x86_64")]
{
super::x86::igemm_u8s8_with_route(a, b, m, k, n, out)
}
#[cfg(all(target_arch = "wasm32", target_feature = "simd128"))]
{
super::wasm128::igemm_u8s8(a, b, m, k, n, out);
EffectiveI8Route::WasmSimd128
}
#[cfg(not(any(
target_arch = "aarch64",
target_arch = "x86_64",
all(target_arch = "wasm32", target_feature = "simd128")
)))]
{
super::scalar::igemm_u8s8(a, b, m, k, n, out);
EffectiveI8Route::Scalar
}
}
#[derive(Debug, Clone)]
pub struct SelftestCase {
pub kind: &'static str,
pub label: &'static str,
pub m: usize,
pub k: usize,
pub n: usize,
pub ok: bool,
pub mismatches: usize,
pub first_bad: Option<(usize, i32, i32)>,
}
#[derive(Debug, Clone)]
pub struct SelftestReport {
pub hardware_selected: IsaTier,
pub effective_route: EffectiveI8Route,
pub executed_routes: Vec<EffectiveI8Route>,
pub route_consistent: bool,
pub available: Vec<IsaTier>,
pub cases: Vec<SelftestCase>,
pub all_ok: bool,
pub models: Vec<(String, bool)>,
}
fn xs32(state: &mut u32) -> u32 {
let mut s = *state;
s ^= s << 13;
s ^= s >> 17;
s ^= s << 5;
*state = s;
s
}
const SELFTEST_SHAPES: &[(&str, usize, usize, usize, u32)] = &[
("edge:1x1x1", 1, 1, 1, 0x1111_1111),
("edge:1x7x3", 1, 7, 3, 0x2222_2222),
("edge:2x3x2", 2, 3, 2, 0x3333_3333),
("ktail:1x15x8", 1, 15, 8, 0x4444_4444),
("ktail:1x16x8", 1, 16, 8, 0x5555_5555),
("ktail:1x17x8", 1, 17, 8, 0x6666_6666),
("ktail:4x33x5", 4, 33, 5, 0x7777_7777),
("model:attn_proj_gemv", 1, 1280, 128, 0x0bad_c0de),
("model:o_proj_gemv", 1, 1280, 1280, 0x1234_5678),
("model:expert_down_gemv", 1, 6848, 256, 0x9abc_def0),
("model:prefill_tile", 4, 1280, 64, 0x0f0f_0f0f),
("overflow:max_mag_k6848", 1, 6848, 4, 0),
("got-ocr2:qkv_fused_gemv", 1, 1024, 3072, 0x6072_0001),
("got-ocr2:o_proj_gemv", 1, 1024, 1024, 0x6072_0002),
("got-ocr2:mlp_down_gemv", 1, 2816, 1024, 0x6072_0003),
("got-ocr2:overflow_k2816", 1, 2816, 4, 0),
("smolvlm2:qkv_fused_gemv", 1, 960, 1600, 0x5601_0001),
("smolvlm2:mlp_down_gemv", 1, 2560, 960, 0x5601_0002),
("smolvlm2:overflow_k2560", 1, 2560, 4, 0),
("onechart:fc1_gemv", 1, 768, 3072, 0x0c4a_0001),
("onechart:fc2_gemv", 1, 3072, 768, 0x0c4a_0002),
("onechart:overflow_k3072", 1, 3072, 4, 0),
];
#[must_use]
pub fn selftest() -> SelftestReport {
use super::scalar;
use std::collections::BTreeSet;
let mut cases = Vec::with_capacity(SELFTEST_SHAPES.len() * 2);
let mut executed_routes = BTreeSet::new();
for &(label, m, k, n, seed) in SELFTEST_SHAPES {
let (a_s, b_s): (Vec<i8>, Vec<i8>) = if seed == 0 {
(vec![i8::MAX; m * k], vec![i8::MIN; n * k])
} else {
let mut st = seed | 1;
(
(0..m * k)
.map(|_| (xs32(&mut st) & 0xff) as u8 as i8)
.collect(),
(0..n * k)
.map(|_| (xs32(&mut st) & 0xff) as u8 as i8)
.collect(),
)
};
let mut got = vec![0i32; m * n];
let mut want = vec![0i32; m * n];
executed_routes.insert(igemm_s8s8_with_route(&a_s, &b_s, m, k, n, &mut got));
scalar::igemm_s8s8(&a_s, &b_s, m, k, n, &mut want);
cases.push(compare_case("s8s8", label, m, k, n, &got, &want));
let (a_u, b_u): (Vec<u8>, Vec<i8>) = if seed == 0 {
(vec![u8::MAX; m * k], vec![i8::MIN; n * k])
} else {
let mut st = seed.rotate_left(7) | 1;
(
(0..m * k).map(|_| (xs32(&mut st) & 0xff) as u8).collect(),
(0..n * k)
.map(|_| (xs32(&mut st) & 0xff) as u8 as i8)
.collect(),
)
};
let mut gotu = vec![0i32; m * n];
let mut wantu = vec![0i32; m * n];
executed_routes.insert(igemm_u8s8_with_route(&a_u, &b_u, m, k, n, &mut gotu));
scalar::igemm_u8s8(&a_u, &b_u, m, k, n, &mut wantu);
cases.push(compare_case("u8s8", label, m, k, n, &gotu, &wantu));
}
let expected_route = effective_dense_route();
let route_consistent = executed_routes.len() == 1 && executed_routes.contains(&expected_route);
let all_ok = cases.iter().all(|c| c.ok) && route_consistent;
let mut models: Vec<(String, bool)> = Vec::new();
for id in ["unlimited-ocr", "got-ocr2", "smolvlm2", "onechart"] {
let ok = cases
.iter()
.filter(|c| match id {
"unlimited-ocr" => {
!c.label.contains(':') || {
let p = c.label.split(':').next().unwrap_or("");
matches!(p, "edge" | "ktail" | "model" | "overflow")
}
}
_ => c.label.starts_with(&format!("{id}:")),
})
.all(|c| c.ok);
models.push((id.to_string(), ok));
}
let snapshot = caps();
SelftestReport {
hardware_selected: snapshot.selected,
effective_route: expected_route,
executed_routes: executed_routes.into_iter().collect(),
route_consistent,
available: snapshot.available.clone(),
cases,
all_ok,
models,
}
}
fn compare_case(
kind: &'static str,
label: &'static str,
m: usize,
k: usize,
n: usize,
got: &[i32],
want: &[i32],
) -> SelftestCase {
let mut mismatches = 0usize;
let mut first_bad = None;
for (i, (&g, &w)) in got.iter().zip(want.iter()).enumerate() {
if g != w {
mismatches += 1;
if first_bad.is_none() {
first_bad = Some((i, g, w));
}
}
}
SelftestCase {
kind,
label,
m,
k,
n,
ok: mismatches == 0,
mismatches,
first_bad,
}
}
#[cfg(test)]
mod tests {
use super::*;
use super::super::scalar;
#[test]
fn detection_does_not_panic_and_has_scalar_floor() {
let c = caps();
assert!(!c.available.is_empty());
assert_eq!(
*c.available.last().expect("non-empty"),
IsaTier::Scalar,
"scalar must always be the floor"
);
assert!(c.available.contains(&c.selected));
}
#[test]
fn caps_is_cached() {
let a = caps();
let b = caps();
assert!(std::ptr::eq(a, b), "caps() must return the cached snapshot");
assert_eq!(detected_tier(), a.selected);
}
#[test]
fn tier_strings_are_stable() {
for t in [
IsaTier::Scalar,
IsaTier::Avx2,
IsaTier::AvxVnni,
IsaTier::Avx512Vnni,
IsaTier::Sdot,
IsaTier::Smmla,
] {
assert!(!t.feature_string().is_empty());
assert!(!t.tag().is_empty());
}
assert_eq!(IsaTier::Scalar.feature_string(), "scalar");
assert_eq!(IsaTier::Sdot.feature_string(), "aarch64+neon+dotprod");
assert_eq!(IsaTier::Smmla.feature_string(), "aarch64+neon+i8mm");
assert_eq!(tier_string(), detected_tier().feature_string());
}
#[test]
fn scalar_is_lowest_rank() {
for t in [
IsaTier::Avx2,
IsaTier::AvxVnni,
IsaTier::Avx512Vnni,
IsaTier::Sdot,
IsaTier::Smmla,
] {
assert!(t > IsaTier::Scalar);
}
}
#[test]
fn dispatch_s8s8_equals_scalar_oracle() {
let a: [i8; 6] = [1, 2, 3, 4, 5, 6];
let b: [i8; 6] = [1, 0, 1, 0, 1, 0]; let mut got = [0i32; 4];
let mut want = [0i32; 4];
igemm_s8s8(&a, &b, 2, 3, 2, &mut got);
scalar::igemm_s8s8(&a, &b, 2, 3, 2, &mut want);
assert_eq!(got, want);
assert_eq!(got, [4, 2, 10, 5]);
}
#[test]
fn dispatch_u8s8_equals_scalar_oracle_randomized() {
let (m, k, n) = (3usize, 19usize, 7usize);
let mut s = 0xc0ffee_u32 | 1;
let mut xs = || {
s ^= s << 13;
s ^= s >> 17;
s ^= s << 5;
s
};
let a: Vec<u8> = (0..m * k).map(|_| (xs() & 0xff) as u8).collect();
let b: Vec<i8> = (0..n * k).map(|_| (xs() & 0xff) as u8 as i8).collect();
let mut got = vec![0i32; m * n];
let mut want = vec![0i32; m * n];
igemm_u8s8(&a, &b, m, k, n, &mut got);
scalar::igemm_u8s8(&a, &b, m, k, n, &mut want);
assert_eq!(got, want);
}
#[test]
fn selftest_passes_on_build_host() {
let report = selftest();
assert!(
!report.cases.is_empty(),
"selftest must exercise at least one shape"
);
assert_eq!(report.cases.len(), SELFTEST_SHAPES.len() * 2);
assert!(report.available.contains(&report.hardware_selected));
assert!(
report.route_consistent,
"every case must execute the predicted route"
);
assert_eq!(report.executed_routes, vec![report.effective_route]);
for case in &report.cases {
assert!(
case.ok,
"tier {:?} diverged from scalar oracle on {} {} ({}x{}x{}): {} lane(s), first {:?}",
report.effective_route,
case.kind,
case.label,
case.m,
case.k,
case.n,
case.mismatches,
case.first_bad,
);
}
assert!(report.all_ok, "headline verdict must reflect all-ok cases");
}
#[test]
fn selftest_reports_a_per_model_verdict_for_every_registered_decoder() {
let report = selftest();
let ids: Vec<&str> = report.models.iter().map(|(id, _)| id.as_str()).collect();
assert_eq!(
ids,
["unlimited-ocr", "got-ocr2", "smolvlm2", "onechart"],
"the per-model rollup must enumerate every registered int8 decoder"
);
for id in ["got-ocr2", "smolvlm2", "onechart"] {
assert!(
report
.cases
.iter()
.any(|c| c.label.starts_with(&format!("{id}:overflow_k"))),
"{id} must carry its own worst-case-K overflow row (doctrine #6 per model)"
);
let model_ok = report.models.iter().find(|(m, _)| m == id).unwrap().1;
let cases_ok = report
.cases
.iter()
.filter(|c| c.label.starts_with(&format!("{id}:")))
.all(|c| c.ok);
assert_eq!(
model_ok, cases_ok,
"{id}: rollup verdict must equal its cases"
);
}
println!(
r#"{{"check":"selftest_per_model_verdicts","models":{},"result":"pass"}}"#,
report.models.len()
);
}
#[test]
fn selftest_overflow_case_is_worst_case_and_exact() {
let (k, n) = (6848usize, 4usize);
let a = vec![u8::MAX; k];
let b = vec![i8::MIN; n * k];
let mut got = vec![0i32; n];
let mut want = vec![0i32; n];
igemm_u8s8(&a, &b, 1, k, n, &mut got);
scalar::igemm_u8s8(&a, &b, 1, k, n, &mut want);
assert_eq!(got, want);
assert!(got.iter().all(|&v| v == -223_518_720));
}
}