use std::path::PathBuf;
fn main() {
println!("cargo::rustc-check-cfg=cfg(gemmology_simd)");
println!("cargo::rustc-check-cfg=cfg(fast_gemm)");
println!("cargo:rerun-if-env-changed=FXTRANSLATE_REQUIRE_SIMD");
let arch = std::env::var("CARGO_CFG_TARGET_ARCH").unwrap_or_default();
let target_features = std::env::var("CARGO_CFG_TARGET_FEATURE").unwrap_or_default();
let wasm_simd128 = arch == "wasm32"
&& target_features.split(',').any(|f| f == "simd128")
&& std::env::var_os("CARGO_FEATURE_PORTABLE").is_none();
if wasm_simd128 {
println!("cargo::rustc-cfg=fast_gemm");
}
let require_simd = std::env::var_os("FXTRANSLATE_REQUIRE_SIMD").is_some();
let bail = |reason: String| {
if require_simd {
panic!("fxtranslate: FXTRANSLATE_REQUIRE_SIMD is set but {reason}");
}
println!("cargo::warning=fxtranslate: {reason} — using the scalar kernel.");
};
if std::env::var_os("CARGO_FEATURE_GEMMOLOGY").is_none() {
return;
}
if std::env::var_os("CARGO_FEATURE_PORTABLE").is_some() {
bail("`portable` is set, which disables the SIMD kernel (no SIMD, no C++)".into());
return;
}
let (arch_define, arch_flag) = match arch.as_str() {
"aarch64" => ("FXT_GEMM_I8MM", "-march=armv8.4-a+i8mm"),
"x86_64" => ("FXT_GEMM_AVX2", "-mavx2"),
_ => {
bail(format!("no SIMD kernel wired for `{arch}` yet"));
return;
}
};
let vendor = PathBuf::from(std::env::var("CARGO_MANIFEST_DIR").unwrap()).join("vendor");
let gemmology = std::env::var_os("GEMMOLOGY_DIR")
.map(PathBuf::from)
.unwrap_or_else(|| vendor.join("gemmology"));
let xsimd = std::env::var_os("XSIMD_INCLUDE_DIR")
.map(PathBuf::from)
.unwrap_or_else(|| vendor.join("xsimd/include"));
let header = gemmology.join("gemmology.h");
if !header.exists() {
bail(format!(
"gemmology headers not found at {}",
header.display()
));
return;
}
println!("cargo:rerun-if-changed=src/gemmology_shim.cpp");
println!("cargo:rerun-if-changed={}", header.display());
println!("cargo:rerun-if-env-changed=GEMMOLOGY_DIR");
println!("cargo:rerun-if-env-changed=XSIMD_INCLUDE_DIR");
let mut build = cc::Build::new();
build
.cpp(true)
.std("c++17")
.flag(arch_flag)
.define(arch_define, None)
.opt_level(3)
.include(&gemmology)
.include(&xsimd)
.file("src/gemmology_shim.cpp");
if std::env::var_os("CARGO_FEATURE_GEMM_THREADS").is_some() {
build.define("FXT_GEMM_THREADS", None);
}
let result = build.try_compile("gemmology_shim");
match result {
Ok(()) => {
println!("cargo::rustc-cfg=gemmology_simd");
println!("cargo::rustc-cfg=fast_gemm");
}
Err(e) => bail(format!("gemmology C++ build failed ({e})")),
}
}