Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension


Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
4 changes: 4 additions & 0 deletions Cargo.toml
Original file line number Diff line number Diff line change
Expand Up @@ -45,6 +45,10 @@ test = true
name = "ocr_benchmark"
required-features = ["std"]

[[example]]
name = "bf16_rne_exhaustive"
required-features = ["std"]

[[example]]
name = "splat3d_flex"
required-features = ["splat3d"]
Expand Down
239 changes: 132 additions & 107 deletions README-DE.md

Large diffs are not rendered by default.

165 changes: 75 additions & 90 deletions README.md

Large diffs are not rendered by default.

16 changes: 8 additions & 8 deletions crates/burn/src/ops/matmul.rs
Original file line number Diff line number Diff line change
Expand Up @@ -375,9 +375,9 @@ pub fn build_distance_table_vnni(centroids_u8: &[u8], k: usize, dim: usize) -> V
// Tier 2: avx512vnni VPDPBUSD zmm (512-bit) 64 MACs/instr Cascade Lake+, Zen 4+
// Stable detection: is_x86_feature_detected!("avx512vnni")
//
// Tier 1: avxvnniint8 VPDPBSSD ymm (256-bit) ~32 MACs/instr Sierra Forest+, Arrow Lake+
// VNNI2: signed×signed dot product. Stable detection on Rust 1.94.
// TODO: implement ymm-width kernel when hardware available.
// Tier 1: avxvnni VEX VPDPBUSD ymm (256-bit) ~32 MACs/instr Alder Lake+, Sierra Forest
// u8×i8 dot product (`simd_amx::vnni2_dot_u8_i8`). NOT avxvnniint8,
// which is VPDPBSSD/VPDPBUUD and absent on Alder Lake.
//
// Tier 0: Scalar loop 1 MAC/iter any CPU
//
Expand All @@ -389,8 +389,8 @@ pub fn build_distance_table_vnni(centroids_u8: &[u8], k: usize, dim: usize) -> V
3 // AMX present — use avx512vnni as bridge
} else if is_x86_feature_detected!("avx512vnni") {
2 // AVX-512 VNNI: 64 MACs/instr
} else if is_x86_feature_detected!("avxvnniint8") {
1 // VNNI2: signed i8×i8 (ymm, ~32 MACs) — TODO: needs ymm kernel
} else if is_x86_feature_detected!("avx2") && is_x86_feature_detected!("avxvnni") {
1 // AVX-VNNI: VEX VPDPBUSD ymm, u8×i8 (~32 MACs)
} else {
0
}
Expand All @@ -408,10 +408,10 @@ pub fn build_distance_table_vnni(centroids_u8: &[u8], k: usize, dim: usize) -> V
#[cfg(not(target_arch = "x86_64"))]
ndarray::simd_amx::vnni_dot_u8_i8_scalar(a, b)
},
// Tier 1: avxvnniint8 — ymm-width VPDPBUSD (32 MACs/instr)
// For NUC 14 i9-185H (Arrow Lake) and similar non-AVX-512 CPUs
// Tier 1: avxvnni — ymm-width VEX VPDPBUSD (32 MACs/instr)
// For Alder Lake and later non-AVX-512 client CPUs
1 => |a, b| {
// SAFETY: avxvnniint8 confirmed via is_x86_feature_detected above
// SAFETY: avx2 + avxvnni confirmed via is_x86_feature_detected above
#[cfg(target_arch = "x86_64")]
unsafe { ndarray::simd_amx::vnni2_dot_u8_i8(a, b) }
#[cfg(not(target_arch = "x86_64"))]
Expand Down
96 changes: 93 additions & 3 deletions crates/simd-masking-parity/src/lib.rs
Original file line number Diff line number Diff line change
Expand Up @@ -27,7 +27,8 @@
//! the index-addressed permutation/scatter family); `0xD3x` the
//! index-addressed `masked_group_sum_i32_via` (two-hop zero-fallback); `0xD4x`
//! `eq_u32_via_to_mask` (the same index lane, packed as a predicate rather than
//! folded into a sum). `main.rs` (native / qemu) and
//! folded into a sum); `0xExx` the `F64x8` lane compares (all six relations
//! over every pair of 16 IEEE edge values: NaN, ±0, ±inf, subnormal). `main.rs` (native / qemu) and
//! `selfcheck()` (the wasm cdylib export, driven by `run.mjs`) both call
//! [`run`].

Expand All @@ -43,11 +44,11 @@ use ndarray::simd::{
masked_strided_group_sum, masked_sum_i32, masked_sum_wrapping_add_i32, ne_i32_to_mask, ne_i32_to_mask_under,
ne_u32_to_mask, ne_u32_to_mask_under, ne_u64_to_mask, ne_u8_to_mask, ternary_match_strided_to_mask,
ternary_match_u32_to_mask, ternary_match_u32_to_mask_under, ternary_match_u64_to_mask,
ternary_match_u64_to_mask_under, ternlog, I32x16, KeyRunCarry, MortonDir, U32x16, U64x8,
ternary_match_u64_to_mask_under, ternlog, F64x8, I32x16, KeyRunCarry, MortonDir, U32x16, U64x8,
};

/// Number of check groups [`run`] executes (for the log line only).
pub const CHECKS: usize = 13;
pub const CHECKS: usize = 14;

/// The wasm export: identical to [`run`], `extern "C"` so `run.mjs` can call it.
#[no_mangle]
Expand All @@ -61,6 +62,7 @@ pub fn run() -> u32 {
check_ternlog_all_tables, check_u64x8_algebra, check_i32x16_compare, check_predicates_to_mask,
check_mask_algebra, check_care_match, check_masked_reductions, check_blend, check_morton_shift,
check_predicates_under, check_set_range, check_unsigned_compare_to_mask, check_gather_scatter_group,
check_f64x8_compare,
];
for g in groups {
if let Err(code) = g() {
Expand Down Expand Up @@ -393,6 +395,94 @@ fn check_i32x16_compare() -> Result<(), u32> {
Ok(())
}

// ── 0xExx: F64x8 lane compares — every pair of IEEE edge values ─────────────
//
// All six relations against plain scalar `f64` operators. The edge set covers
// NaN (positive, negative, payload-carrying), ±0, ±inf, the extremes and a
// subnormal. Expected semantics are IEEE-754 / Rust's own: `==`,`<`,`>`,`<=`,
// `>=` are ORDERED (false if either operand is NaN), `!=` is UNORDERED (true
// if either is NaN), and `+0.0 == -0.0`. The mask is read back through
// `select`, the only accessor every realization of `F64Mask8` shares.

fn f64_mask_bits(m: impl FnOnce(F64x8, F64x8) -> F64x8) -> u8 {
let lanes = m(F64x8::splat(1.0), F64x8::splat(0.0)).to_array();
let mut bits = 0u8;
for (i, &v) in lanes.iter().enumerate() {
if v == 1.0 {
bits |= 1 << i;
}
}
bits
}

fn check_f64x8_compare() -> Result<(), u32> {
let edges: [f64; 16] = [
f64::NAN,
-f64::NAN,
f64::from_bits(0x7FF0_0000_0000_0001), // signalling-pattern NaN with payload
0.0,
-0.0,
f64::INFINITY,
f64::NEG_INFINITY,
f64::MAX,
f64::MIN,
f64::MIN_POSITIVE,
f64::from_bits(1), // smallest subnormal
1.0,
-1.0,
1.0 + f64::EPSILON,
0.5,
1.0,
];
let mut pairs = 0u32;
for chunk in 0..(16 * 16 / 8) {
let mut a_arr = [0.0f64; 8];
let mut b_arr = [0.0f64; 8];
for lane in 0..8 {
let k = chunk * 8 + lane;
a_arr[lane] = edges[k / 16];
b_arr[lane] = edges[k % 16];
}
let (a, b) = (F64x8::from_array(a_arr), F64x8::from_array(b_arr));
let expect = |rel: fn(f64, f64) -> bool| -> u8 {
let mut bits = 0u8;
for lane in 0..8 {
if rel(a_arr[lane], b_arr[lane]) {
bits |= 1 << lane;
}
}
bits
};
let got = [
f64_mask_bits(|t, f| a.simd_eq(b).select(t, f)),
f64_mask_bits(|t, f| a.simd_ne(b).select(t, f)),
f64_mask_bits(|t, f| a.simd_lt(b).select(t, f)),
f64_mask_bits(|t, f| a.simd_le(b).select(t, f)),
f64_mask_bits(|t, f| a.simd_gt(b).select(t, f)),
f64_mask_bits(|t, f| a.simd_ge(b).select(t, f)),
];
let want = [
expect(|x, y| x == y),
expect(|x, y| x != y),
expect(|x, y| x < y),
expect(|x, y| x <= y),
expect(|x, y| x > y),
expect(|x, y| x >= y),
];
for rel in 0..6 {
if got[rel] != want[rel] {
return Err(0xE00 | rel as u32);
}
}
pairs += 8;
}
// Anti-vacuity: every one of the 16 x 16 edge pairs was checked.
if pairs != 256 {
return Err(0xE0F);
}
Ok(())
}

// ── 0x5xx: predicate → mask, every tail shape, full-overwrite + zero tail ────

fn check_predicates_to_mask() -> Result<(), u32> {
Expand Down
139 changes: 139 additions & 0 deletions examples/bf16_rne_exhaustive.rs
Original file line number Diff line number Diff line change
@@ -0,0 +1,139 @@
//! Exhaustive F32 -> BF16 round-to-nearest-even parity: ALL 2^32 f32 bit patterns.
//!
//! Reproduces the README claim "4,294,967,296 inputs, 0 mismatches". For every
//! `u32` bit pattern `b`, `simd::f32_to_bf16_batch_rne` is compared bit-for-bit
//! (u16 equality, so NaN payloads and signs count) against
//!
//! 1. `simd::f32_to_bf16_scalar_rne`, the repo's scalar reference, and
//! 2. an INDEPENDENT oracle defined here: pick whichever of the two adjacent
//! BF16 values is nearer in f64, ties to the even mantissa; overflow past
//! BF16::MAX rounds to infinity at the 2^128 midpoint.
//!
//! The intended semantics are Intel SDM `VCVTNEPS2BF16`:
//! - NaN keeps its sign and top 7 payload bits, and the quiet bit is forced;
//! - subnormal input flushes to a signed zero (DAZ);
//! - infinities and zeros pass through.
//!
//! The oracle encodes exactly that and nothing else.
//!
//! Method: the 2^32 space is split into `threads` equal contiguous ranges. Each
//! range is enumerated in chunks of 65,536 inputs; each chunk is converted in
//! one batch call. An order-independent checksum of every converted output
//! (Σ out(b)·(b|1) mod 2^64) is printed so the work cannot be optimized away
//! and runs with different thread counts can be compared.
//!
//! Run (x86_64; the batch takes the AVX-512F path when the host has it):
//! ```text
//! cargo run --release --example bf16_rne_exhaustive [threads]
//! ```

#[cfg(target_arch = "x86_64")]
fn oracle_bf16(bits: u32) -> u16 {
let exp = bits & 0x7F80_0000;
let mant = bits & 0x007F_FFFF;
if exp == 0x7F80_0000 && mant != 0 {
return ((bits >> 16) as u16) | 0x0040; // NaN: keep sign + top payload, force quiet
}
if exp == 0 {
return ((bits >> 16) as u16) & 0x8000; // zero or subnormal -> signed zero
}
if exp == 0x7F80_0000 {
return (bits >> 16) as u16; // infinity
}
let x = f32::from_bits(bits) as f64;
let lo = (bits >> 16) as u16; // truncation: toward zero
let hi = lo + 1; // next magnitude up, same sign
let f_lo = f32::from_bits((lo as u32) << 16) as f64;
let hi_bits = (hi as u32) << 16;
let d_lo = (x - f_lo).abs();
let d_hi = if (hi_bits & 0x7F80_0000) == 0x7F80_0000 {
// `hi` is infinity: the rounding midpoint to infinity is 2^128 - 2^119.
2f64.powi(128) - x.abs()
} else {
(f32::from_bits(hi_bits) as f64 - x).abs()
};
if d_lo < d_hi || (d_lo == d_hi && lo & 1 == 0) {
lo
} else {
hi
}
}

#[cfg(target_arch = "x86_64")]
fn main() {
use ndarray::simd::{f32_to_bf16_batch_rne, f32_to_bf16_scalar_rne};
use std::time::Instant;

let threads: u64 = std::env::args()
.nth(1)
.and_then(|s| s.parse().ok())
.unwrap_or(4);
assert!(threads.is_power_of_two() && threads <= 64, "threads must be a power of two <= 64");
let path = if is_x86_feature_detected!("avx512f") {
"AVX-512F"
} else {
"scalar fallback"
};
println!("bf16_rne_exhaustive: batch path = {path}, threads = {threads}, chunk = 65536");

let t0 = Instant::now();
let span = (1u64 << 32) / threads;
let results: Vec<(u64, u64, u64, u64, Option<u32>)> = std::thread::scope(|s| {
let handles: Vec<_> = (0..threads)
.map(|t| {
s.spawn(move || {
const CHUNK: usize = 1 << 16;
let mut inp = vec![0f32; CHUNK];
let mut out = vec![0u16; CHUNK];
let (mut n, mut vs_scalar, mut vs_oracle, mut sum) = (0u64, 0u64, 0u64, 0u64);
let mut first = None;
let mut base = t * span;
while base < (t + 1) * span {
for (i, x) in inp.iter_mut().enumerate() {
*x = f32::from_bits((base + i as u64) as u32);
}
f32_to_bf16_batch_rne(&inp, &mut out);
for i in 0..CHUNK {
let b = (base + i as u64) as u32;
let got = out[i];
// Order-independent, so the value is the same for any thread count.
sum = sum.wrapping_add((got as u64).wrapping_mul(b as u64 | 1));
if got != f32_to_bf16_scalar_rne(inp[i]) {
vs_scalar += 1;
}
if got != oracle_bf16(b) {
vs_oracle += 1;
first.get_or_insert(b);
}
}
n += CHUNK as u64;
base += CHUNK as u64;
}
(n, vs_scalar, vs_oracle, sum, first)
})
})
.collect();
handles.into_iter().map(|h| h.join().unwrap()).collect()
});
let n: u64 = results.iter().map(|r| r.0).sum();
let vs_scalar: u64 = results.iter().map(|r| r.1).sum();
let vs_oracle: u64 = results.iter().map(|r| r.2).sum();
let checksum = results.iter().fold(0u64, |a, r| a.wrapping_add(r.3));
println!("inputs checked : {n}");
println!("mismatch vs scalar : {vs_scalar}");
match results.iter().find_map(|r| r.4) {
Some(b) => println!("mismatch vs oracle : {vs_oracle} (first input 0x{b:08x})"),
None => println!("mismatch vs oracle : {vs_oracle}"),
}
println!("output checksum : 0x{checksum:016x}");
println!("elapsed : {:.1} s", t0.elapsed().as_secs_f64());
assert_eq!(n, 1u64 << 32, "did not cover the full input space");
if vs_scalar != 0 || vs_oracle != 0 {
std::process::exit(1);
}
}

#[cfg(not(target_arch = "x86_64"))]
fn main() {
eprintln!("bf16_rne_exhaustive: x86_64 only (the batch path under test is AVX-512F)");
}
Loading
Loading