use crate::scalar::Scalar;
#[inline]
pub(crate) unsafe fn pack_panels<T: Scalar>(
dst: *mut T,
src: *const T,
lead: isize,
depth: isize,
n_lead: usize,
depth_len: usize,
width: usize,
) {
unsafe {
let tile = crate::tuning::pack_transpose_tile();
let mut d = dst;
let mut base = 0usize;
while base < n_lead {
let live = core::cmp::min(width, n_lead - base);
if lead == 1 {
for p in 0..depth_len {
let s = src.offset(base as isize + p as isize * depth);
core::ptr::copy_nonoverlapping(s, d, live);
for i in live..width {
*d.add(i) = T::ZERO;
}
d = d.add(width);
}
} else {
let panel = d;
let mut p0 = 0;
while p0 < depth_len {
let pe = core::cmp::min(p0 + tile, depth_len);
for i in 0..width {
if i < live {
let row = src.offset((base + i) as isize * lead);
for p in p0..pe {
*panel.add(p * width + i) = *row.offset(p as isize * depth);
}
} else {
for p in p0..pe {
*panel.add(p * width + i) = T::ZERO;
}
}
}
p0 = pe;
}
d = panel.add(depth_len * width);
}
base += width;
}
}
}
#[cfg(any(feature = "int8", feature = "half"))]
#[allow(clippy::too_many_arguments)]
#[inline]
pub(crate) unsafe fn pack_kgroup_panels<T: Scalar, const Q: usize, F: Fn(T) -> T>(
dst: *mut T,
src: *const T,
lead: isize,
depth: isize,
n_lead: usize,
depth_len: usize,
width: usize,
xform: F,
) {
unsafe {
let kc_pad = depth_len.next_multiple_of(Q);
let full_groups = depth_len / Q;
let has_tail = !depth_len.is_multiple_of(Q);
let pad = xform(T::ZERO);
let mut d = dst;
let mut base = 0usize;
while base < n_lead {
let live = core::cmp::min(width, n_lead - base);
for g in 0..full_groups {
let gbase = g * width * Q;
let dp0 = (g * Q) as isize;
if depth == 1 {
for i in 0..live {
let s = src.offset((base + i) as isize * lead + dp0);
let dd = d.add(gbase + i * Q);
for t in 0..Q {
*dd.add(t) = xform(*s.add(t));
}
}
} else {
for i in 0..live {
let row = src.offset((base + i) as isize * lead + dp0 * depth);
let dd = d.add(gbase + i * Q);
for t in 0..Q {
*dd.add(t) = xform(*row.offset(t as isize * depth));
}
}
}
for i in live..width {
let dd = d.add(gbase + i * Q);
for t in 0..Q {
*dd.add(t) = pad;
}
}
}
if has_tail {
let g = full_groups;
let gbase = g * width * Q;
for i in 0..width {
let live_lead = base + i < n_lead;
let dd = d.add(gbase + i * Q);
for t in 0..Q {
let dp = g * Q + t;
let v = if live_lead && dp < depth_len {
xform(*src.offset((base + i) as isize * lead + dp as isize * depth))
} else {
pad
};
*dd.add(t) = v;
}
}
}
d = d.add(width * kc_pad);
base += width;
}
}
}
#[cfg(test)]
mod panels_tests {
use super::*;
fn reference<T: Scalar>(
base: *const T,
lead: isize,
depth: isize,
n_lead: usize,
depth_len: usize,
width: usize,
) -> Vec<T> {
let panels = n_lead.div_ceil(width);
let mut out = vec![T::ZERO; panels * width * depth_len];
let mut d = 0usize;
let mut b = 0usize;
while b < n_lead {
for p in 0..depth_len {
for i in 0..width {
let lead_pos = b + i;
let v = if lead_pos < n_lead {
unsafe { *base.offset(lead_pos as isize * lead + p as isize * depth) }
} else {
T::ZERO
};
out[d + p * width + i] = v;
}
}
d += width * depth_len;
b += width;
}
out
}
fn same_bytes<T>(a: &[T], b: &[T]) -> bool {
let (pa, la) = (a.as_ptr() as *const u8, core::mem::size_of_val(a));
let (pb, lb) = (b.as_ptr() as *const u8, core::mem::size_of_val(b));
unsafe { core::slice::from_raw_parts(pa, la) == core::slice::from_raw_parts(pb, lb) }
}
fn check_case(n_lead: usize, depth_len: usize, width: usize, depth_stride: isize) {
let depth = depth_stride;
let lead = depth_len as isize * depth + 1;
let max_off = if n_lead == 0 || depth_len == 0 {
0
} else {
((n_lead - 1) as isize * lead + (depth_len - 1) as isize * depth) as usize
};
let backing: Vec<f32> = (0..=max_off)
.map(|i| f32::from_bits((i as u32).wrapping_mul(2_654_435_761)))
.collect();
let base = backing.as_ptr();
let expected = reference::<f32>(base, lead, depth, n_lead, depth_len, width);
let panels = n_lead.div_ceil(width);
let mut actual = vec![0.0f32; panels * width * depth_len];
unsafe {
pack_panels::<f32>(
actual.as_mut_ptr(),
base,
lead,
depth,
n_lead,
depth_len,
width,
);
}
assert!(
same_bytes(&actual, &expected),
"n_lead={n_lead} depth_len={depth_len} width={width} depth={depth}"
);
}
fn transpose_tail<T: Scalar>(
dst: *mut T,
src: *const T,
lead: isize,
depth: isize,
live: usize,
depth_len: usize,
width: usize,
) {
let tile = crate::tuning::pack_transpose_tile();
unsafe {
let panel = dst;
let mut p0 = 0;
while p0 < depth_len {
let pe = core::cmp::min(p0 + tile, depth_len);
for i in 0..width {
if i < live {
let row = src.offset(i as isize * lead);
for p in p0..pe {
*panel.add(p * width + i) = *row.offset(p as isize * depth);
}
} else {
for p in p0..pe {
*panel.add(p * width + i) = T::ZERO;
}
}
}
p0 = pe;
}
}
}
#[test]
#[ignore = "microbench; run with --release --ignored --nocapture"]
fn bench_tail_panel_pack() {
use std::time::Instant;
let (width, live, depth_len, m) = (32usize, 8usize, 4096usize, 520isize);
let (lead, depth) = (1isize, m); let max_off = (live as isize - 1) * lead + (depth_len as isize - 1) * depth;
let backing: Vec<f32> = (0..=max_off as usize).map(|i| i as f32 * 0.5).collect();
let base = backing.as_ptr();
let mut dst = vec![0.0f32; width * depth_len];
let bench = |reps: usize, mut f: Box<dyn FnMut()>| -> f64 {
for _ in 0..20 {
f();
}
let mut s: Vec<f64> = Vec::with_capacity(reps);
for _ in 0..reps {
let t = Instant::now();
for _ in 0..200 {
f();
}
s.push(t.elapsed().as_secs_f64() * 1e9 / 200.0);
}
s.sort_by(f64::total_cmp);
s[reps / 2]
};
let (p, b) = (dst.as_mut_ptr(), base);
let t_new = bench(
25,
Box::new(move || unsafe {
pack_panels::<f32>(p, b, lead, depth, live, depth_len, width);
core::hint::black_box(p);
}),
);
let (p, b) = (dst.as_mut_ptr(), base);
let t_old = bench(
25,
Box::new(move || {
transpose_tail::<f32>(p, b, lead, depth, live, depth_len, width);
core::hint::black_box(p);
}),
);
println!(
"\ntail-panel pack (live={live}/{width}, depth={depth_len}, stride={m}): straight-copy {t_new:7.1} ns transpose {t_old:7.1} ns ({:.2}x)",
t_old / t_new.max(1e-9)
);
}
#[test]
fn panels_bit_identical() {
const N_LEADS: [usize; 8] = [1, 3, 4, 5, 7, 8, 9, 17];
const DEPTHS: [usize; 5] = [1, 2, 3, 5, 8];
const WIDTHS: [usize; 5] = [1, 3, 4, 6, 8];
const STRIDES: [isize; 2] = [1, 4];
for &n_lead in &N_LEADS {
for &depth_len in &DEPTHS {
for &width in &WIDTHS {
for &stride in &STRIDES {
check_case(n_lead, depth_len, width, stride);
}
}
}
}
}
}
#[cfg(all(test, any(feature = "int8", feature = "half")))]
mod tests {
use super::*;
#[allow(clippy::too_many_arguments)]
fn reference<T: Scalar>(
base: *const T,
lead: isize,
depth: isize,
n_lead: usize,
depth_len: usize,
width: usize,
q: usize,
xform: impl Fn(T) -> T,
) -> Vec<T> {
let kc_pad = depth_len.next_multiple_of(q);
let ngroups = kc_pad / q;
let pad = xform(T::ZERO);
let panels = n_lead.div_ceil(width);
let mut out = vec![T::ZERO; panels * width * kc_pad];
let mut d = 0usize;
let mut b = 0usize;
while b < n_lead {
for g in 0..ngroups {
for i in 0..width {
let lead_pos = b + i;
let live_lead = lead_pos < n_lead;
for t in 0..q {
let dp = g * q + t;
let v = if live_lead && dp < depth_len {
unsafe {
xform(*base.offset(lead_pos as isize * lead + dp as isize * depth))
}
} else {
pad
};
out[d + g * width * q + i * q + t] = v;
}
}
}
d += width * kc_pad;
b += width;
}
out
}
fn same_bytes<T>(a: &[T], b: &[T]) -> bool {
let (pa, la) = (a.as_ptr() as *const u8, core::mem::size_of_val(a));
let (pb, lb) = (b.as_ptr() as *const u8, core::mem::size_of_val(b));
unsafe { core::slice::from_raw_parts(pa, la) == core::slice::from_raw_parts(pb, lb) }
}
fn check_case<T: Scalar, const Q: usize>(
n_lead: usize,
depth_len: usize,
width: usize,
depth_stride: isize,
val: impl Fn(usize) -> T,
xform: impl Fn(T) -> T + Copy,
) {
let depth = depth_stride;
let lead = depth_len as isize * depth + 1;
let max_off = if n_lead == 0 || depth_len == 0 {
0
} else {
((n_lead - 1) as isize * lead + (depth_len - 1) as isize * depth) as usize
};
let backing: Vec<T> = (0..=max_off).map(&val).collect();
let base = backing.as_ptr();
let expected = reference::<T>(base, lead, depth, n_lead, depth_len, width, Q, xform);
let kc_pad = depth_len.next_multiple_of(Q);
let panels = n_lead.div_ceil(width);
let mut actual = vec![T::ZERO; panels * width * kc_pad];
unsafe {
pack_kgroup_panels::<T, Q, _>(
actual.as_mut_ptr(),
base,
lead,
depth,
n_lead,
depth_len,
width,
xform,
);
}
assert!(
same_bytes(&actual, &expected),
"Q={Q} n_lead={n_lead} depth_len={depth_len} width={width} depth={depth}"
);
}
const N_LEADS: [usize; 7] = [1, 3, 7, 8, 9, 16, 17];
const DEPTHS: [usize; 8] = [1, 2, 3, 4, 5, 6, 8, 11];
const WIDTHS: [usize; 5] = [1, 3, 4, 5, 8];
const STRIDES: [isize; 2] = [1, 5];
fn i8_val(i: usize) -> i8 {
(i as u32).wrapping_mul(2_654_435_761) as u8 as i8
}
#[cfg(feature = "int8")]
#[test]
fn kgroup_bit_identical_i8() {
let plus128 = |v: i8| ((v as i32 + 128) as u8) as i8;
let ident = |v: i8| v;
for &n_lead in &N_LEADS {
for &depth_len in &DEPTHS {
for &width in &WIDTHS {
for &stride in &STRIDES {
check_case::<i8, 4>(n_lead, depth_len, width, stride, i8_val, plus128);
check_case::<i8, 4>(n_lead, depth_len, width, stride, i8_val, ident);
check_case::<i8, 2>(n_lead, depth_len, width, stride, i8_val, plus128);
check_case::<i8, 2>(n_lead, depth_len, width, stride, i8_val, ident);
}
}
}
}
}
#[cfg(feature = "half")]
#[test]
fn kgroup_bit_identical_bf16() {
use half::bf16;
let val = |i: usize| bf16::from_bits((i as u32).wrapping_mul(40_503) as u16);
let ident = |v: bf16| v;
for &n_lead in &N_LEADS {
for &depth_len in &DEPTHS {
for &width in &WIDTHS {
for &stride in &STRIDES {
check_case::<bf16, 2>(n_lead, depth_len, width, stride, val, ident);
}
}
}
}
}
}