use std::cmp::Ordering;
use crate::csr::{matmul_topn_short_circuit, narrow_indptr, CsrMatrix, CsrView};
use crate::index::Index;
use crate::matmul_topn::{SortMode, TopNOptions};
use crate::maxheap::MaxHeap;
use crate::scalar::Scalar;
use crate::tiled::TiledB;
pub(crate) const UNVISITED: u32 = u32::MAX;
pub(crate) const HEAD_NIL: u32 = u32::MAX - 1;
pub(crate) const FALLBACK_L1D_BYTES: usize = 64 * 1024;
pub fn default_chunk_cols<V: Scalar>() -> usize {
const FLOOR: usize = 64;
const CEIL: usize = 1 << 20;
let l1d = detect_l1d_bytes().unwrap_or(FALLBACK_L1D_BYTES);
let per_col_bytes = std::mem::size_of::<V>() + std::mem::size_of::<u32>();
let raw = (l1d / per_col_bytes.max(1)).max(1);
let pow2 = raw.next_power_of_two() / 2;
pow2.clamp(FLOOR, CEIL)
}
pub(crate) fn detect_l1d_bytes() -> Option<usize> {
#[cfg(target_os = "macos")]
if let Some(sz) = macos_sysctl_l1d() {
return Some(sz);
}
cache_size::l1_cache_size()
}
#[cfg(target_os = "macos")]
fn macos_sysctl_l1d() -> Option<usize> {
use std::ffi::CStr;
use std::os::raw::{c_char, c_int, c_void};
extern "C" {
fn sysctlbyname(
name: *const c_char,
oldp: *mut c_void,
oldlenp: *mut usize,
newp: *mut c_void,
newlen: usize,
) -> c_int;
}
fn get(name: &CStr) -> Option<usize> {
let mut val: u64 = 0;
let mut len: usize = std::mem::size_of::<u64>();
let rc = unsafe {
sysctlbyname(
name.as_ptr(),
&mut val as *mut u64 as *mut c_void,
&mut len,
std::ptr::null_mut(),
0,
)
};
(rc == 0 && (len == 4 || len == 8) && val > 0).then_some(val as usize)
}
let perf0 = CStr::from_bytes_with_nul(b"hw.perflevel0.l1dcachesize\0").unwrap();
let plain = CStr::from_bytes_with_nul(b"hw.l1dcachesize\0").unwrap();
get(perf0).or_else(|| get(plain))
}
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
pub enum BProjection {
BinarySearch,
Cursor,
}
#[derive(Debug, Clone, Copy, PartialEq, Eq, Default)]
pub enum AccumMode {
#[default]
Adaptive,
LinkedList,
Dense,
}
pub(crate) const DENSE_MIN_DENSITY: f64 = 0.2;
#[inline]
fn push_row_cursors<V: Scalar, I: Index>(
a: CsrView<'_, V, I>,
b: CsrView<'_, V, I>,
i: usize,
cursors: &mut Vec<usize>,
) {
for jj in a.indptr[i].to_usize()..a.indptr[i + 1].to_usize() {
let j = a.indices[jj].to_usize();
cursors.push(b.indptr[j].to_usize());
}
}
fn row_update_density<V: Scalar, I: Index>(
a: CsrView<'_, V, I>,
b: CsrView<'_, V, I>,
i: usize,
) -> f64 {
let start = a.indptr[i].to_usize();
let end = a.indptr[i + 1].to_usize();
let mut updates: usize = 0;
for jj in start..end {
let j = a.indices[jj].to_usize();
updates += b.indptr[j + 1].to_usize() - b.indptr[j].to_usize();
}
updates as f64 / b.ncols.max(1) as f64
}
pub(crate) fn pick_projection<V: Scalar, I: Index>(
_a: CsrView<'_, V, I>,
b: CsrView<'_, V, I>,
chunk_cols: usize,
) -> BProjection {
let cc = chunk_cols.max(1);
let num_chunks = b.ncols.div_ceil(cc).max(1);
let avg_nnz_b_row = if b.nrows == 0 {
0.0
} else {
b.nnz() as f64 / b.nrows as f64
};
if avg_nnz_b_row >= (num_chunks as f64) * 4.0 {
BProjection::BinarySearch
} else {
BProjection::Cursor
}
}
pub(crate) fn resolve_chunk_cols<V: Scalar>(requested: Option<usize>, ncols: usize) -> usize {
let raw = requested.unwrap_or_else(default_chunk_cols::<V>);
assert!(raw > 0, "chunk_cols must be > 0");
raw.min(ncols.max(1)).min((u32::MAX - 2) as usize)
}
pub fn sp_matmul_topn_chunked<V: Scalar, I: Index>(
a: CsrView<'_, V, I>,
b: CsrView<'_, V, I>,
top_n: usize,
opts: TopNOptions<V>,
) -> CsrMatrix<V, I> {
assert_eq!(
a.ncols, b.nrows,
"sp_matmul_topn_chunked: A.ncols ({}) must equal B.nrows ({})",
a.ncols, b.nrows,
);
if let Some(out) = matmul_topn_short_circuit(a, b, top_n) {
return out;
}
let chunk_cols = resolve_chunk_cols::<V>(opts.chunk_cols, b.ncols);
let projection = opts
.projection
.unwrap_or_else(|| pick_projection(a, b, chunk_cols));
dispatch::<V, I>(a, b, top_n, opts, chunk_cols, projection)
}
pub(crate) const DEFAULT_ROW_BLOCK: usize = 2048;
pub(crate) fn auto_tile<V: Scalar, I: Index>(
a: CsrView<'_, V, I>,
b: CsrView<'_, V, I>,
chunk_cols: usize,
) -> bool {
let n_chunks = b.ncols.div_ceil(chunk_cols.max(1));
n_chunks >= 2
&& TiledB::supports(&b, chunk_cols)
&& a.nnz() >= 2 * b.nrows
&& n_chunks.saturating_mul(b.nrows) <= 4 * b.nnz()
}
pub(crate) fn resolve_blocking<V: Scalar, I: Index>(
a: CsrView<'_, V, I>,
b: CsrView<'_, V, I>,
chunk_cols: usize,
opts: &TopNOptions<V>,
) -> (bool, Option<usize>) {
if opts.row_block == Some(0) {
(false, None)
} else if opts.tile_b || opts.row_block.is_some() {
(
opts.tile_b && TiledB::supports(&b, chunk_cols),
opts.row_block,
)
} else if auto_tile(a, b, chunk_cols) {
(true, Some(DEFAULT_ROW_BLOCK))
} else {
(false, None)
}
}
pub(crate) fn dispatch<V: Scalar, I: Index>(
a: CsrView<'_, V, I>,
b: CsrView<'_, V, I>,
top_n: usize,
opts: TopNOptions<V>,
chunk_cols: usize,
projection: BProjection,
) -> CsrMatrix<V, I> {
let (tile, row_block) = resolve_blocking(a, b, chunk_cols, &opts);
if tile || row_block.is_some() {
let tiled = tile.then(|| TiledB::build(b, chunk_cols));
let opts = TopNOptions { row_block, ..opts };
run_blocked(a, b, top_n, opts, chunk_cols, projection, tiled.as_ref())
} else {
run(a, b, top_n, opts, chunk_cols, projection)
}
}
pub(crate) fn run<V: Scalar, I: Index>(
a: CsrView<'_, V, I>,
b: CsrView<'_, V, I>,
top_n: usize,
opts: TopNOptions<V>,
chunk_cols: usize,
projection: BProjection,
) -> CsrMatrix<V, I> {
let nrows = a.nrows;
let ncols = b.ncols;
let threshold = opts.threshold.unwrap_or_else(V::min_value);
let density = opts.density_hint.unwrap_or(1.0);
let cap_hint = ((nrows as f64) * (top_n as f64) * density).ceil() as usize;
let mut sums: Vec<V> = vec![V::default(); chunk_cols];
let mut next: Vec<u32> = vec![UNVISITED; chunk_cols];
let mut heap = MaxHeap::<V, I>::new(top_n, threshold);
let mut c_indptr: Vec<I> = Vec::with_capacity(nrows + 1);
let mut c_indices: Vec<I> = Vec::with_capacity(cap_hint);
let mut c_data: Vec<V> = Vec::with_capacity(cap_hint);
c_indptr.push(I::zero());
let mut nnz_total: usize = 0;
let mut cursors: Vec<usize> = Vec::new();
for i in 0..nrows {
let n_set = process_row(
a,
b,
i,
opts.sort,
chunk_cols,
projection,
opts.accum_mode,
&mut sums,
&mut next,
&mut heap,
&mut cursors,
&mut c_indices,
&mut c_data,
);
nnz_total += n_set;
c_indptr.push(narrow_indptr(nnz_total));
}
CsrMatrix {
nrows,
ncols,
indptr: c_indptr,
indices: c_indices,
data: c_data,
}
}
#[inline(always)]
fn fmadd<V: Scalar, const FMA: bool>(v: V, x: V, acc: V) -> V {
if FMA {
v.mul_add_fused(x, acc)
} else {
v.mul_add(x, acc)
}
}
#[inline(always)]
fn scatter_dense_unrolled<V: Scalar, I: Index, const FMA: bool>(
seg_idx: &[I],
seg_dat: &[V],
v: V,
c0: usize,
sums: &mut [V],
) -> V {
let n = seg_idx.len();
let mut s = 0;
let mut m0 = V::min_value();
let mut m1 = V::min_value();
let mut m2 = V::min_value();
let mut m3 = V::min_value();
while s + 4 <= n {
let k0 = seg_idx[s].to_usize() - c0;
let k1 = seg_idx[s + 1].to_usize() - c0;
let k2 = seg_idx[s + 2].to_usize() - c0;
let k3 = seg_idx[s + 3].to_usize() - c0;
let n0 = fmadd::<V, FMA>(v, seg_dat[s], sums[k0]);
let n1 = fmadd::<V, FMA>(v, seg_dat[s + 1], sums[k1]);
let n2 = fmadd::<V, FMA>(v, seg_dat[s + 2], sums[k2]);
let n3 = fmadd::<V, FMA>(v, seg_dat[s + 3], sums[k3]);
sums[k0] = n0;
sums[k1] = n1;
sums[k2] = n2;
sums[k3] = n3;
m0 = m0.max(n0);
m1 = m1.max(n1);
m2 = m2.max(n2);
m3 = m3.max(n3);
s += 4;
}
for t in s..n {
let k_local = seg_idx[t].to_usize() - c0;
let nv = fmadd::<V, FMA>(v, seg_dat[t], sums[k_local]);
sums[k_local] = nv;
m0 = m0.max(nv);
}
m0.max(m1).max(m2.max(m3))
}
#[inline(always)]
fn scatter_dense_unrolled_local<V: Scalar, const FMA: bool>(
seg_idx: &[u16],
seg_dat: &[V],
v: V,
sums: &mut [V],
) -> V {
let n = seg_idx.len();
let mut s = 0;
let mut m0 = V::min_value();
let mut m1 = V::min_value();
let mut m2 = V::min_value();
let mut m3 = V::min_value();
while s + 4 <= n {
let k0 = seg_idx[s] as usize;
let k1 = seg_idx[s + 1] as usize;
let k2 = seg_idx[s + 2] as usize;
let k3 = seg_idx[s + 3] as usize;
let n0 = fmadd::<V, FMA>(v, seg_dat[s], sums[k0]);
let n1 = fmadd::<V, FMA>(v, seg_dat[s + 1], sums[k1]);
let n2 = fmadd::<V, FMA>(v, seg_dat[s + 2], sums[k2]);
let n3 = fmadd::<V, FMA>(v, seg_dat[s + 3], sums[k3]);
sums[k0] = n0;
sums[k1] = n1;
sums[k2] = n2;
sums[k3] = n3;
m0 = m0.max(n0);
m1 = m1.max(n1);
m2 = m2.max(n2);
m3 = m3.max(n3);
s += 4;
}
for t in s..n {
let k_local = seg_idx[t] as usize;
let nv = fmadd::<V, FMA>(v, seg_dat[t], sums[k_local]);
sums[k_local] = nv;
m0 = m0.max(nv);
}
m0.max(m1).max(m2.max(m3))
}
#[inline(always)]
fn drain_dense_chunk<V: Scalar, I: Index>(
heap: &mut MaxHeap<V, I>,
sums: &mut [V],
mut min: V,
chunk_max: V,
c0: usize,
chunk_width: usize,
) -> V {
let zero = V::default();
if chunk_max.partial_cmp(&min) == Some(Ordering::Greater) {
let w4 = chunk_width & !3;
let mut k = 0;
while k < w4 {
let s0 = sums[k];
let s1 = sums[k + 1];
let s2 = sums[k + 2];
let s3 = sums[k + 3];
let m = s0.max(s1).max(s2.max(s3));
if m.partial_cmp(&min) == Some(Ordering::Greater) {
for (off, s) in [s0, s1, s2, s3].into_iter().enumerate() {
if s != zero && s.partial_cmp(&min) == Some(Ordering::Greater) {
min = heap.push_pop(I::from_usize(c0 + k + off), s);
}
}
}
k += 4;
}
for (off, &s) in sums[w4..chunk_width].iter().enumerate() {
if s != zero && s.partial_cmp(&min) == Some(Ordering::Greater) {
min = heap.push_pop(I::from_usize(c0 + w4 + off), s);
}
}
}
sums[..chunk_width].fill(zero);
min
}
#[inline(always)]
fn drain_linked_chunk<V: Scalar, I: Index>(
heap: &mut MaxHeap<V, I>,
sums: &mut [V],
next: &mut [u32],
mut head: u32,
length: usize,
mut min: V,
c0: usize,
) -> V {
for _ in 0..length {
let temp = head as usize;
if sums[temp].partial_cmp(&min) == Some(Ordering::Greater) {
min = heap.push_pop(I::from_usize(c0 + temp), sums[temp]);
}
head = next[temp];
next[temp] = UNVISITED;
sums[temp] = V::default();
}
min
}
pub(crate) struct BlockScratch<V: Scalar, I: Index> {
pub sums: Vec<V>,
pub next: Vec<u32>,
pub heaps: Vec<MaxHeap<V, I>>,
pub mins: Vec<V>,
pub use_dense: Vec<bool>,
pub cursors: Vec<usize>,
pub cursor_base: Vec<usize>,
}
impl<V: Scalar, I: Index> BlockScratch<V, I> {
pub fn new(top_n: usize, threshold: V, chunk_cols: usize, block_rows: usize) -> Self {
Self {
sums: vec![V::default(); chunk_cols],
next: vec![UNVISITED; chunk_cols],
heaps: (0..block_rows)
.map(|_| MaxHeap::new(top_n, threshold))
.collect(),
mins: vec![V::default(); block_rows],
use_dense: vec![false; block_rows],
cursors: Vec::new(),
cursor_base: Vec::new(),
}
}
}
#[allow(clippy::too_many_arguments)]
pub(crate) fn process_row_block<V: Scalar, I: Index>(
a: CsrView<'_, V, I>,
b: CsrView<'_, V, I>,
row_lo: usize,
row_hi: usize,
sort: SortMode,
chunk_cols: usize,
projection: BProjection,
accum: AccumMode,
tiled: Option<&TiledB<V>>,
scratch: &mut BlockScratch<V, I>,
out_indices: &mut Vec<I>,
out_data: &mut Vec<V>,
row_nset: &mut Vec<u32>,
) {
#[cfg(all(target_arch = "x86_64", not(target_feature = "fma")))]
if crate::simd::avx2_fma_enabled() {
return unsafe {
process_row_block_avx2(
a,
b,
row_lo,
row_hi,
sort,
chunk_cols,
projection,
accum,
tiled,
scratch,
out_indices,
out_data,
row_nset,
)
};
}
process_row_block_impl::<V, I, false>(
a,
b,
row_lo,
row_hi,
sort,
chunk_cols,
projection,
accum,
tiled,
scratch,
out_indices,
out_data,
row_nset,
)
}
#[cfg(all(target_arch = "x86_64", not(target_feature = "fma")))]
#[target_feature(enable = "avx2,fma")]
#[allow(clippy::too_many_arguments)]
unsafe fn process_row_block_avx2<V: Scalar, I: Index>(
a: CsrView<'_, V, I>,
b: CsrView<'_, V, I>,
row_lo: usize,
row_hi: usize,
sort: SortMode,
chunk_cols: usize,
projection: BProjection,
accum: AccumMode,
tiled: Option<&TiledB<V>>,
scratch: &mut BlockScratch<V, I>,
out_indices: &mut Vec<I>,
out_data: &mut Vec<V>,
row_nset: &mut Vec<u32>,
) {
process_row_block_impl::<V, I, true>(
a,
b,
row_lo,
row_hi,
sort,
chunk_cols,
projection,
accum,
tiled,
scratch,
out_indices,
out_data,
row_nset,
)
}
#[allow(clippy::too_many_arguments)]
#[inline(always)]
fn process_row_block_impl<V: Scalar, I: Index, const FMA: bool>(
a: CsrView<'_, V, I>,
b: CsrView<'_, V, I>,
row_lo: usize,
row_hi: usize,
sort: SortMode,
chunk_cols: usize,
projection: BProjection,
accum: AccumMode,
tiled: Option<&TiledB<V>>,
scratch: &mut BlockScratch<V, I>,
out_indices: &mut Vec<I>,
out_data: &mut Vec<V>,
row_nset: &mut Vec<u32>,
) {
let ncols = b.ncols;
let block_rows = row_hi - row_lo;
debug_assert!(block_rows <= scratch.heaps.len());
for r in 0..block_rows {
let i = row_lo + r;
scratch.mins[r] = scratch.heaps[r].reset();
scratch.use_dense[r] = V::IS_FLOAT
&& match accum {
AccumMode::LinkedList => false,
AccumMode::Dense => true,
AccumMode::Adaptive => row_update_density(a, b, i) >= DENSE_MIN_DENSITY,
};
}
let need_cursors = tiled.is_none() && projection == BProjection::Cursor;
scratch.cursors.clear();
scratch.cursor_base.clear();
if need_cursors {
for i in row_lo..row_hi {
scratch.cursor_base.push(scratch.cursors.len());
push_row_cursors(a, b, i, &mut scratch.cursors);
}
}
let sums = &mut scratch.sums[..];
let next = &mut scratch.next[..];
let mut c0 = 0;
let mut chunk_id = 0;
while c0 < ncols {
let chunk_width = (ncols - c0).min(chunk_cols);
let chunk_end = c0 + chunk_width;
for r in 0..block_rows {
let i = row_lo + r;
let a_row_start = a.indptr[i].to_usize();
let a_row_end = a.indptr[i + 1].to_usize();
let use_dense = scratch.use_dense[r];
let mut min = scratch.mins[r];
let mut head: u32 = HEAD_NIL;
let mut length: usize = 0;
let mut chunk_max = V::min_value();
if let Some(t) = tiled {
if use_dense {
for jj in a_row_start..a_row_end {
let j = a.indices[jj].to_usize();
let (si, sd) = t.segment(chunk_id, j);
let m = scatter_dense_unrolled_local::<V, FMA>(si, sd, a.data[jj], sums);
chunk_max = chunk_max.max(m);
}
} else {
for jj in a_row_start..a_row_end {
let j = a.indices[jj].to_usize();
let v = a.data[jj];
let (si, sd) = t.segment(chunk_id, j);
for (p, &kl) in si.iter().enumerate() {
let k_local = kl as usize;
sums[k_local] = fmadd::<V, FMA>(v, sd[p], sums[k_local]);
if next[k_local] == UNVISITED {
next[k_local] = head;
head = k_local as u32;
length += 1;
}
}
}
}
} else {
match (projection, use_dense) {
(BProjection::BinarySearch, false) => {
for jj in a_row_start..a_row_end {
let j = a.indices[jj].to_usize();
let v = a.data[jj];
let row_start = b.indptr[j].to_usize();
let row_end = b.indptr[j + 1].to_usize();
let row_b_idx = &b.indices[row_start..row_end];
let lo = row_b_idx.partition_point(|x| Index::to_usize(*x) < c0);
let hi = row_b_idx.partition_point(|x| Index::to_usize(*x) < chunk_end);
for (slot, &k_idx) in row_b_idx[lo..hi].iter().enumerate() {
let off = lo + slot;
let k_local = k_idx.to_usize() - c0;
sums[k_local] =
fmadd::<V, FMA>(v, b.data[row_start + off], sums[k_local]);
if next[k_local] == UNVISITED {
next[k_local] = head;
head = k_local as u32;
length += 1;
}
}
}
}
(BProjection::BinarySearch, true) => {
for jj in a_row_start..a_row_end {
let j = a.indices[jj].to_usize();
let v = a.data[jj];
let row_start = b.indptr[j].to_usize();
let row_end = b.indptr[j + 1].to_usize();
let row_b_idx = &b.indices[row_start..row_end];
let lo = row_b_idx.partition_point(|x| Index::to_usize(*x) < c0);
let hi = row_b_idx.partition_point(|x| Index::to_usize(*x) < chunk_end);
let m = scatter_dense_unrolled::<V, I, FMA>(
&row_b_idx[lo..hi],
&b.data[row_start + lo..row_start + hi],
v,
c0,
sums,
);
chunk_max = chunk_max.max(m);
}
}
(BProjection::Cursor, false) => {
let base = scratch.cursor_base[r];
for (idx, jj) in (a_row_start..a_row_end).enumerate() {
let j = a.indices[jj].to_usize();
let v = a.data[jj];
let stop_b = b.indptr[j + 1].to_usize();
let mut cur = scratch.cursors[base + idx];
while cur < stop_b {
let k = b.indices[cur].to_usize();
if k >= chunk_end {
break;
}
debug_assert!(k >= c0);
let k_local = k - c0;
sums[k_local] = fmadd::<V, FMA>(v, b.data[cur], sums[k_local]);
if next[k_local] == UNVISITED {
next[k_local] = head;
head = k_local as u32;
length += 1;
}
cur += 1;
}
scratch.cursors[base + idx] = cur;
}
}
(BProjection::Cursor, true) => {
let base = scratch.cursor_base[r];
for (idx, jj) in (a_row_start..a_row_end).enumerate() {
let j = a.indices[jj].to_usize();
let v = a.data[jj];
let stop_b = b.indptr[j + 1].to_usize();
let cur = scratch.cursors[base + idx];
let seg_len = b.indices[cur..stop_b]
.partition_point(|x| Index::to_usize(*x) < chunk_end);
let m = scatter_dense_unrolled::<V, I, FMA>(
&b.indices[cur..cur + seg_len],
&b.data[cur..cur + seg_len],
v,
c0,
sums,
);
chunk_max = chunk_max.max(m);
scratch.cursors[base + idx] = cur + seg_len;
}
}
}
}
min = if use_dense {
drain_dense_chunk(&mut scratch.heaps[r], sums, min, chunk_max, c0, chunk_width)
} else {
drain_linked_chunk(&mut scratch.heaps[r], sums, next, head, length, min, c0)
};
scratch.mins[r] = min;
}
c0 = chunk_end;
chunk_id += 1;
}
for r in 0..block_rows {
let heap = &mut scratch.heaps[r];
match sort {
SortMode::ByColumn => heap.sort_by_insertion_order(),
SortMode::ByValueDesc => heap.sort_by_value_desc(),
}
let n_set = heap.n_set();
for entry in heap.entries() {
out_indices.push(entry.idx);
out_data.push(entry.val);
}
row_nset.push(n_set as u32);
}
}
pub(crate) fn run_blocked<V: Scalar, I: Index>(
a: CsrView<'_, V, I>,
b: CsrView<'_, V, I>,
top_n: usize,
opts: TopNOptions<V>,
chunk_cols: usize,
projection: BProjection,
tiled: Option<&TiledB<V>>,
) -> CsrMatrix<V, I> {
let nrows = a.nrows;
let ncols = b.ncols;
let threshold = opts.threshold.unwrap_or_else(V::min_value);
let density = opts.density_hint.unwrap_or(1.0);
let cap_hint = ((nrows as f64) * (top_n as f64) * density).ceil() as usize;
let block_rows = opts.row_block.unwrap_or(1).max(1).min(nrows.max(1));
let mut scratch = BlockScratch::<V, I>::new(top_n, threshold, chunk_cols, block_rows);
let mut row_nset: Vec<u32> = Vec::with_capacity(nrows);
let mut c_indices: Vec<I> = Vec::with_capacity(cap_hint);
let mut c_data: Vec<V> = Vec::with_capacity(cap_hint);
let mut lo = 0;
while lo < nrows {
let hi = (lo + block_rows).min(nrows);
process_row_block(
a,
b,
lo,
hi,
opts.sort,
chunk_cols,
projection,
opts.accum_mode,
tiled,
&mut scratch,
&mut c_indices,
&mut c_data,
&mut row_nset,
);
lo = hi;
}
let mut c_indptr: Vec<I> = Vec::with_capacity(nrows + 1);
c_indptr.push(I::zero());
let mut running: usize = 0;
for &n in &row_nset {
running += n as usize;
c_indptr.push(narrow_indptr(running));
}
CsrMatrix {
nrows,
ncols,
indptr: c_indptr,
indices: c_indices,
data: c_data,
}
}
#[allow(clippy::too_many_arguments)]
pub(crate) fn process_row<V: Scalar, I: Index>(
a: CsrView<'_, V, I>,
b: CsrView<'_, V, I>,
i: usize,
sort: SortMode,
chunk_cols: usize,
projection: BProjection,
accum: AccumMode,
sums: &mut [V],
next: &mut [u32],
heap: &mut MaxHeap<V, I>,
cursors: &mut Vec<usize>,
out_indices: &mut Vec<I>,
out_data: &mut Vec<V>,
) -> usize {
#[cfg(all(target_arch = "x86_64", not(target_feature = "fma")))]
if crate::simd::avx2_fma_enabled() {
return unsafe {
process_row_avx2(
a,
b,
i,
sort,
chunk_cols,
projection,
accum,
sums,
next,
heap,
cursors,
out_indices,
out_data,
)
};
}
process_row_impl::<V, I, false>(
a,
b,
i,
sort,
chunk_cols,
projection,
accum,
sums,
next,
heap,
cursors,
out_indices,
out_data,
)
}
#[cfg(all(target_arch = "x86_64", not(target_feature = "fma")))]
#[target_feature(enable = "avx2,fma")]
#[allow(clippy::too_many_arguments)]
unsafe fn process_row_avx2<V: Scalar, I: Index>(
a: CsrView<'_, V, I>,
b: CsrView<'_, V, I>,
i: usize,
sort: SortMode,
chunk_cols: usize,
projection: BProjection,
accum: AccumMode,
sums: &mut [V],
next: &mut [u32],
heap: &mut MaxHeap<V, I>,
cursors: &mut Vec<usize>,
out_indices: &mut Vec<I>,
out_data: &mut Vec<V>,
) -> usize {
process_row_impl::<V, I, true>(
a,
b,
i,
sort,
chunk_cols,
projection,
accum,
sums,
next,
heap,
cursors,
out_indices,
out_data,
)
}
#[allow(clippy::too_many_arguments)]
#[inline(always)]
fn process_row_impl<V: Scalar, I: Index, const FMA: bool>(
a: CsrView<'_, V, I>,
b: CsrView<'_, V, I>,
i: usize,
sort: SortMode,
chunk_cols: usize,
projection: BProjection,
accum: AccumMode,
sums: &mut [V],
next: &mut [u32],
heap: &mut MaxHeap<V, I>,
cursors: &mut Vec<usize>,
out_indices: &mut Vec<I>,
out_data: &mut Vec<V>,
) -> usize {
let ncols = b.ncols;
let mut min = heap.reset();
let a_row_start = a.indptr[i].to_usize();
let a_row_end = a.indptr[i + 1].to_usize();
let use_dense = V::IS_FLOAT
&& match accum {
AccumMode::LinkedList => false,
AccumMode::Dense => true,
AccumMode::Adaptive => row_update_density(a, b, i) >= DENSE_MIN_DENSITY,
};
if projection == BProjection::Cursor {
cursors.clear();
cursors.reserve(a_row_end - a_row_start);
push_row_cursors(a, b, i, cursors);
}
let mut c0 = 0;
while c0 < ncols {
let chunk_width = (ncols - c0).min(chunk_cols);
let chunk_end = c0 + chunk_width;
let mut head: u32 = HEAD_NIL;
let mut length: usize = 0;
let mut chunk_max = V::min_value();
match (projection, use_dense) {
(BProjection::BinarySearch, false) => {
for jj in a_row_start..a_row_end {
let j = a.indices[jj].to_usize();
let v = a.data[jj];
let row_start = b.indptr[j].to_usize();
let row_end = b.indptr[j + 1].to_usize();
let row_b_idx = &b.indices[row_start..row_end];
let lo = row_b_idx.partition_point(|x| Index::to_usize(*x) < c0);
let hi = row_b_idx.partition_point(|x| Index::to_usize(*x) < chunk_end);
for (slot, &k_idx) in row_b_idx[lo..hi].iter().enumerate() {
let off = lo + slot;
let k_local = k_idx.to_usize() - c0;
sums[k_local] = fmadd::<V, FMA>(v, b.data[row_start + off], sums[k_local]);
if next[k_local] == UNVISITED {
next[k_local] = head;
head = k_local as u32;
length += 1;
}
}
}
}
(BProjection::BinarySearch, true) => {
for jj in a_row_start..a_row_end {
let j = a.indices[jj].to_usize();
let v = a.data[jj];
let row_start = b.indptr[j].to_usize();
let row_end = b.indptr[j + 1].to_usize();
let row_b_idx = &b.indices[row_start..row_end];
let lo = row_b_idx.partition_point(|x| Index::to_usize(*x) < c0);
let hi = row_b_idx.partition_point(|x| Index::to_usize(*x) < chunk_end);
let m = scatter_dense_unrolled::<V, I, FMA>(
&row_b_idx[lo..hi],
&b.data[row_start + lo..row_start + hi],
v,
c0,
sums,
);
chunk_max = chunk_max.max(m);
}
}
(BProjection::Cursor, false) => {
for (idx, jj) in (a_row_start..a_row_end).enumerate() {
let j = a.indices[jj].to_usize();
let v = a.data[jj];
let stop_b = b.indptr[j + 1].to_usize();
let mut cur = cursors[idx];
while cur < stop_b {
let k = b.indices[cur].to_usize();
if k >= chunk_end {
break;
}
debug_assert!(k >= c0);
let k_local = k - c0;
sums[k_local] = fmadd::<V, FMA>(v, b.data[cur], sums[k_local]);
if next[k_local] == UNVISITED {
next[k_local] = head;
head = k_local as u32;
length += 1;
}
cur += 1;
}
cursors[idx] = cur;
}
}
(BProjection::Cursor, true) => {
for (idx, jj) in (a_row_start..a_row_end).enumerate() {
let j = a.indices[jj].to_usize();
let v = a.data[jj];
let stop_b = b.indptr[j + 1].to_usize();
let cur = cursors[idx];
let seg_len =
b.indices[cur..stop_b].partition_point(|x| Index::to_usize(*x) < chunk_end);
let m = scatter_dense_unrolled::<V, I, FMA>(
&b.indices[cur..cur + seg_len],
&b.data[cur..cur + seg_len],
v,
c0,
sums,
);
chunk_max = chunk_max.max(m);
cursors[idx] = cur + seg_len;
}
}
}
min = if use_dense {
drain_dense_chunk(heap, sums, min, chunk_max, c0, chunk_width)
} else {
drain_linked_chunk(heap, sums, next, head, length, min, c0)
};
c0 += chunk_width;
}
match sort {
SortMode::ByColumn => heap.sort_by_insertion_order(),
SortMode::ByValueDesc => heap.sort_by_value_desc(),
}
let n_set = heap.n_set();
for entry in heap.entries() {
out_indices.push(entry.idx);
out_data.push(entry.val);
}
n_set
}
#[cfg(test)]
mod tests {
use super::*;
use crate::csr::CsrView;
fn default_is_power_of_two_in_range<V: Scalar>() {
let n = default_chunk_cols::<V>();
assert!(n >= 64, "default {n} < floor");
assert!(n <= (1 << 20), "default {n} > ceil");
assert!(n.is_power_of_two(), "default {n} not a power of two");
}
#[test]
fn default_chunk_cols_f32() {
default_is_power_of_two_in_range::<f32>();
}
#[test]
fn default_chunk_cols_f64() {
default_is_power_of_two_in_range::<f64>();
}
#[test]
fn default_chunk_cols_i32() {
default_is_power_of_two_in_range::<i32>();
}
#[test]
fn default_chunk_cols_i64() {
default_is_power_of_two_in_range::<i64>();
}
type CsrParts = (Vec<i32>, Vec<i32>, Vec<f64>, Vec<i32>, Vec<i32>, Vec<f64>);
fn make_a_b() -> CsrParts {
let a_indptr = vec![0i32, 2, 4];
let a_indices = vec![0i32, 2, 1, 2];
let a_data = vec![1.0f64, 2.0, 3.0, 4.0];
let b_indptr = vec![0i32, 2, 4, 6];
let b_indices = vec![0i32, 3, 1, 3, 2, 3];
let b_data = vec![1.0f64, 5.0, 1.0, 6.0, 1.0, 7.0];
(a_indptr, a_indices, a_data, b_indptr, b_indices, b_data)
}
fn rows_sorted(c: &CsrMatrix<f64, i32>) -> Vec<Vec<(i32, f64)>> {
(0..c.nrows)
.map(|i| {
let s = c.indptr[i] as usize;
let e = c.indptr[i + 1] as usize;
let mut row: Vec<(i32, f64)> = (s..e).map(|k| (c.indices[k], c.data[k])).collect();
row.sort_by_key(|(idx, _)| *idx);
row
})
.collect()
}
#[allow(clippy::type_complexity)]
fn run_both(
a: CsrView<'_, f64, i32>,
b: CsrView<'_, f64, i32>,
top_n: usize,
opts: TopNOptions<f64>,
chunk_cols: usize,
) -> (Vec<Vec<(i32, f64)>>, Vec<Vec<(i32, f64)>>) {
let cc = chunk_cols.min(b.ncols.max(1));
let c_bs = run::<f64, i32>(a, b, top_n, opts, cc, BProjection::BinarySearch);
let c_cu = run::<f64, i32>(a, b, top_n, opts, cc, BProjection::Cursor);
(rows_sorted(&c_bs), rows_sorted(&c_cu))
}
#[test]
fn one_big_chunk_matches_known_product() {
let (a_ip, a_idx, a_d, b_ip, b_idx, b_d) = make_a_b();
let a = CsrView::new(2, 3, &a_ip, &a_idx, &a_d).unwrap();
let b = CsrView::new(3, 4, &b_ip, &b_idx, &b_d).unwrap();
let opts = TopNOptions {
sort: SortMode::ByValueDesc,
..Default::default()
};
let (bs, cu) = run_both(a, b, 2, opts, 4);
let expected = vec![vec![(2, 2.0), (3, 19.0)], vec![(2, 4.0), (3, 46.0)]];
assert_eq!(bs, expected);
assert_eq!(cu, expected);
}
#[test]
fn one_col_chunks_match() {
let (a_ip, a_idx, a_d, b_ip, b_idx, b_d) = make_a_b();
let a = CsrView::new(2, 3, &a_ip, &a_idx, &a_d).unwrap();
let b = CsrView::new(3, 4, &b_ip, &b_idx, &b_d).unwrap();
let opts = TopNOptions {
sort: SortMode::ByColumn,
..Default::default()
};
let (bs, cu) = run_both(a, b, 4, opts, 1);
let expected = vec![
vec![(0, 1.0), (2, 2.0), (3, 19.0)],
vec![(1, 3.0), (2, 4.0), (3, 46.0)],
];
assert_eq!(bs, expected);
assert_eq!(cu, expected);
}
#[test]
fn non_divisor_chunk_width() {
let (a_ip, a_idx, a_d, b_ip, b_idx, b_d) = make_a_b();
let a = CsrView::new(2, 3, &a_ip, &a_idx, &a_d).unwrap();
let b = CsrView::new(3, 4, &b_ip, &b_idx, &b_d).unwrap();
let opts = TopNOptions {
sort: SortMode::ByValueDesc,
..Default::default()
};
let (bs, cu) = run_both(a, b, 4, opts, 3);
let expected = vec![
vec![(0, 1.0), (2, 2.0), (3, 19.0)],
vec![(1, 3.0), (2, 4.0), (3, 46.0)],
];
let resort = |rs: Vec<Vec<(i32, f64)>>| -> Vec<Vec<(i32, f64)>> {
rs.into_iter()
.map(|mut r| {
r.sort_by_key(|(idx, _)| *idx);
r
})
.collect()
};
assert_eq!(resort(bs), expected);
assert_eq!(resort(cu), expected);
}
#[test]
fn most_chunks_empty_for_row() {
let a_indptr = vec![0i32, 1];
let a_indices = vec![3i32];
let a_data = vec![2.0f64];
let b_indptr = vec![0i32, 0, 0, 0, 2, 2, 2, 2, 2, 2, 2];
let b_indices = vec![5i32, 7];
let b_data = vec![3.0f64, 5.0];
let a = CsrView::new(1, 10, &a_indptr, &a_indices, &a_data).unwrap();
let b = CsrView::new(10, 16, &b_indptr, &b_indices, &b_data).unwrap();
let opts = TopNOptions {
sort: SortMode::ByColumn,
..Default::default()
};
let (bs, cu) = run_both(a, b, 4, opts, 4);
let expected = vec![vec![(5, 6.0), (7, 10.0)]];
assert_eq!(bs, expected);
assert_eq!(cu, expected);
}
#[test]
fn later_chunk_displaces_earlier() {
let a_indptr = vec![0i32, 3];
let a_indices = vec![0i32, 1, 2];
let a_data = vec![1.0f64, 1.0, 1.0];
let b_indptr = vec![0i32, 2, 3, 4];
let b_indices = vec![0i32, 5, 2, 4];
let b_data = vec![5.0f64, 1.0, 7.0, 3.0];
let a = CsrView::new(1, 3, &a_indptr, &a_indices, &a_data).unwrap();
let b = CsrView::new(3, 6, &b_indptr, &b_indices, &b_data).unwrap();
let opts = TopNOptions {
sort: SortMode::ByValueDesc,
..Default::default()
};
let (bs, cu) = run_both(a, b, 2, opts, 2);
let expected = vec![vec![(0, 5.0), (2, 7.0)]];
assert_eq!(bs, expected);
assert_eq!(cu, expected);
}
#[test]
fn equal_value_does_not_displace() {
let a_indptr = vec![0i32, 2];
let a_indices = vec![0i32, 1];
let a_data = vec![1.0f64, 1.0];
let b_indptr = vec![0i32, 2, 3];
let b_indices = vec![0i32, 1, 3];
let b_data = vec![3.0f64, 3.0, 3.0];
let a = CsrView::new(1, 2, &a_indptr, &a_indices, &a_data).unwrap();
let b = CsrView::new(2, 4, &b_indptr, &b_indices, &b_data).unwrap();
let opts = TopNOptions {
sort: SortMode::ByColumn,
..Default::default()
};
let (bs, cu) = run_both(a, b, 2, opts, 2);
let expected = vec![vec![(0, 3.0), (1, 3.0)]];
assert_eq!(bs, expected);
assert_eq!(cu, expected);
}
#[test]
fn threshold_filters_across_chunks() {
let (a_ip, a_idx, a_d, b_ip, b_idx, b_d) = make_a_b();
let a = CsrView::new(2, 3, &a_ip, &a_idx, &a_d).unwrap();
let b = CsrView::new(3, 4, &b_ip, &b_idx, &b_d).unwrap();
let opts = TopNOptions {
threshold: Some(10.0),
sort: SortMode::ByValueDesc,
..Default::default()
};
let (bs, cu) = run_both(a, b, 4, opts, 2);
let expected = vec![vec![(3, 19.0)], vec![(3, 46.0)]];
assert_eq!(bs, expected);
assert_eq!(cu, expected);
}
#[test]
fn topn_smaller_than_per_chunk_nnz() {
let (a_ip, a_idx, a_d, b_ip, b_idx, b_d) = make_a_b();
let a = CsrView::new(2, 3, &a_ip, &a_idx, &a_d).unwrap();
let b = CsrView::new(3, 4, &b_ip, &b_idx, &b_d).unwrap();
let opts = TopNOptions {
sort: SortMode::ByValueDesc,
..Default::default()
};
let (bs, cu) = run_both(a, b, 1, opts, 4);
let expected = vec![vec![(3, 19.0)], vec![(3, 46.0)]];
assert_eq!(bs, expected);
assert_eq!(cu, expected);
}
#[test]
fn dense_matches_linked_list() {
let (a_ip, a_idx, a_d, b_ip, b_idx, b_d) = make_a_b();
let a = CsrView::new(2, 3, &a_ip, &a_idx, &a_d).unwrap();
let b = CsrView::new(3, 4, &b_ip, &b_idx, &b_d).unwrap();
for chunk_cols in [1usize, 2, 3, 4, 8] {
for projection in [BProjection::BinarySearch, BProjection::Cursor] {
for top_n in [1usize, 2, 4] {
let mk = |accum_mode| TopNOptions::<f64> {
sort: SortMode::ByColumn,
accum_mode,
..Default::default()
};
let dense =
run::<f64, i32>(a, b, top_n, mk(AccumMode::Dense), chunk_cols, projection);
let linked = run::<f64, i32>(
a,
b,
top_n,
mk(AccumMode::LinkedList),
chunk_cols,
projection,
);
assert_eq!(
rows_sorted(&dense),
rows_sorted(&linked),
"chunk_cols={chunk_cols} projection={projection:?} top_n={top_n}"
);
}
}
}
}
#[test]
fn dense_drops_exact_zero_sum_floats() {
let a_indptr = vec![0i32, 2];
let a_indices = vec![0i32, 1];
let a_data = vec![1.0f64, 1.0];
let b_indptr = vec![0i32, 2, 4];
let b_indices = vec![0i32, 1, 0, 2];
let b_data = vec![1.0f64, 5.0, -1.0, 3.0];
let a = CsrView::new(1, 2, &a_indptr, &a_indices, &a_data).unwrap();
let b = CsrView::new(2, 3, &b_indptr, &b_indices, &b_data).unwrap();
let mk = |accum_mode| TopNOptions::<f64> {
sort: SortMode::ByColumn,
accum_mode,
..Default::default()
};
let linked = run::<f64, i32>(a, b, 3, mk(AccumMode::LinkedList), 4, BProjection::Cursor);
let dense = run::<f64, i32>(a, b, 3, mk(AccumMode::Dense), 4, BProjection::Cursor);
assert_eq!(
rows_sorted(&linked),
vec![vec![(0, 0.0), (1, 5.0), (2, 3.0)]],
);
assert_eq!(rows_sorted(&dense), vec![vec![(1, 5.0), (2, 3.0)]]);
}
#[test]
fn dense_forced_on_ints_falls_back_to_linked_list() {
let a_indptr = vec![0i32, 2];
let a_indices = vec![0i32, 1];
let a_data = vec![1i32, 1];
let b_indptr = vec![0i32, 2, 4];
let b_indices = vec![0i32, 1, 0, 2];
let b_data = vec![2i32, 5, -2, 3];
let a = CsrView::new(1, 2, &a_indptr, &a_indices, &a_data).unwrap();
let b = CsrView::new(2, 3, &b_indptr, &b_indices, &b_data).unwrap();
let opts = TopNOptions::<i32> {
sort: SortMode::ByColumn,
accum_mode: AccumMode::Dense,
..Default::default()
};
let c = run::<i32, i32>(a, b, 3, opts, 4, BProjection::Cursor);
let mut row: Vec<(i32, i32)> = (0..c.nnz()).map(|k| (c.indices[k], c.data[k])).collect();
row.sort_by_key(|(idx, _)| *idx);
assert_eq!(row, vec![(0, 0), (1, 5), (2, 3)]);
}
#[test]
fn row_update_density_exact() {
let (a_ip, a_idx, a_d, b_ip, b_idx, b_d) = make_a_b();
let a = CsrView::new(2, 3, &a_ip, &a_idx, &a_d).unwrap();
let b = CsrView::new(3, 4, &b_ip, &b_idx, &b_d).unwrap();
assert_eq!(row_update_density(a, b, 0), 1.0);
assert_eq!(row_update_density(a, b, 1), 1.0);
}
#[test]
fn pick_projection_extremes() {
let dense_indptr: Vec<i32> = (0..=4).map(|i| i * 8).collect();
let dense_indices: Vec<i32> = (0..32).map(|k| k % 8).collect();
let dense_data = vec![1.0f64; 32];
let dense = CsrView::new(4, 8, &dense_indptr, &dense_indices, &dense_data).unwrap();
let a_ip = vec![0i32, 1];
let a_idx = vec![0i32];
let a_d = vec![1.0f64];
let a = CsrView::new(1, 4, &a_ip, &a_idx, &a_d).unwrap();
assert_eq!(pick_projection(a, dense, 8), BProjection::BinarySearch);
let sparse_indptr = vec![0i32, 1, 1, 1, 1];
let sparse_indices = vec![0i32];
let sparse_data = vec![1.0f64];
let sparse = CsrView::new(4, 1024, &sparse_indptr, &sparse_indices, &sparse_data).unwrap();
let a2 = CsrView::new(1, 4, &a_ip, &a_idx, &a_d).unwrap();
assert_eq!(pick_projection(a2, sparse, 16), BProjection::Cursor);
}
#[test]
fn auto_tile_gates() {
let a_ip: Vec<i32> = (0..=10).map(|i| i * 4).collect();
let a_idx: Vec<i32> = (0..40).map(|k| k % 4).collect();
let a_d = vec![1.0f64; 40];
let a = CsrView::new(10, 4, &a_ip, &a_idx, &a_d).unwrap();
let b_ip: Vec<i32> = (0..=4).map(|i| i * 5).collect();
let b_idx: Vec<i32> = (0..20).map(|k| k * 2).collect();
let b_d = vec![1.0f64; 20];
let b = CsrView::new(4, 40, &b_ip, &b_idx, &b_d).unwrap();
assert!(auto_tile(a, b, 8));
assert!(!auto_tile(a, b, 40));
let a_small = CsrView::new(1, 4, &a_ip[..2], &a_idx[..4], &a_d[..4]).unwrap();
assert!(!auto_tile(a_small, b, 8));
let bs_ip: Vec<i32> = vec![0, 1, 2, 3, 4];
let bs_idx: Vec<i32> = vec![0, 10, 20, 30];
let bs_d = vec![1.0f64; 4];
let b_thin = CsrView::new(4, 40, &bs_ip, &bs_idx, &bs_d).unwrap();
assert!(!auto_tile(a, b_thin, 4));
}
#[test]
fn row_block_zero_forces_classic() {
let a_ip: Vec<i32> = (0..=10).map(|i| i * 4).collect();
let a_idx: Vec<i32> = (0..40).map(|k| k % 4).collect();
let a_d = vec![1.0f64; 40];
let a = CsrView::new(10, 4, &a_ip, &a_idx, &a_d).unwrap();
let b_ip: Vec<i32> = (0..=4).map(|i| i * 5).collect();
let b_idx: Vec<i32> = (0..20).map(|k| k * 2).collect();
let b_d = vec![1.0f64; 20];
let b = CsrView::new(4, 40, &b_ip, &b_idx, &b_d).unwrap();
let auto = TopNOptions::<f64>::default();
assert_eq!(
resolve_blocking(a, b, 8, &auto),
(true, Some(DEFAULT_ROW_BLOCK))
);
let classic = TopNOptions::<f64> {
row_block: Some(0),
..Default::default()
};
assert_eq!(resolve_blocking(a, b, 8, &classic), (false, None));
let c_auto = sp_matmul_topn_chunked(
a,
b,
3,
TopNOptions {
chunk_cols: Some(8),
..auto
},
);
let c_classic = sp_matmul_topn_chunked(
a,
b,
3,
TopNOptions {
chunk_cols: Some(8),
..classic
},
);
assert_eq!(c_auto.indptr, c_classic.indptr);
assert_eq!(c_auto.indices, c_classic.indices);
assert_eq!(c_auto.data, c_classic.data);
}
}