#[inline]
const fn div_ceil(n: usize, d: usize) -> usize {
n.div_ceil(d)
}
pub fn pack(weights: &[i8], bits: u8) -> Vec<i8> {
match bits {
8 => weights.to_vec(),
4 => {
let n = div_ceil(weights.len(), 2);
let mut out = vec![0i8; n];
for (i, &w) in weights.iter().enumerate() {
let nib = (w as u8) & 0x0F;
out[i / 2] = ((out[i / 2] as u8) | (nib << ((i & 1) * 4))) as i8;
}
out
}
2 => {
let n = div_ceil(weights.len(), 4);
let mut out = vec![0i8; n];
for (i, &w) in weights.iter().enumerate() {
let crumb = (w as u8) & 0x03;
out[i / 4] = ((out[i / 4] as u8) | (crumb << ((i & 3) * 2))) as i8;
}
out
}
_ => panic!("pack: unsupported bits {bits} (must be 2, 4 or 8)"),
}
}
#[inline]
pub const fn weights_per_byte(bits: u8) -> usize {
(8 / bits) as usize
}
#[inline]
pub const fn packed_byte_len(n_weights: usize, bits: u8) -> usize {
div_ceil(n_weights, weights_per_byte(bits))
}
pub fn split_weights_by_oc_lane(
packed: &[i8],
c_out: usize,
inner: usize,
bits: u8,
p: usize,
) -> Vec<Vec<i8>> {
use rlx_cortexm::quant::read_weight;
assert!(p > 0);
assert_eq!(
c_out % p,
0,
"split_weights_by_oc_lane: c_out ({c_out}) not divisible by p ({p})"
);
let mut lanes_logical: Vec<Vec<i8>> = vec![Vec::with_capacity((c_out / p) * inner); p];
for oc in 0..c_out {
let lane = oc % p;
for k in 0..inner {
let logical_idx = oc * inner + k;
let v = read_weight(packed, logical_idx, bits);
lanes_logical[lane].push(v as i8);
}
}
lanes_logical.into_iter().map(|w| pack(&w, bits)).collect()
}
pub fn split_table_by_oc_lane<T: Copy>(table: &[T], p: usize) -> Vec<Vec<T>> {
assert!(p > 0);
assert_eq!(
table.len() % p,
0,
"split_table_by_oc_lane: table.len ({}) not divisible by p ({p})",
table.len()
);
let mut lanes: Vec<Vec<T>> = vec![Vec::with_capacity(table.len() / p); p];
for (oc, &v) in table.iter().enumerate() {
lanes[oc % p].push(v);
}
lanes
}
#[cfg(test)]
mod tests {
use super::*;
use rlx_cortexm::quant::read_weight;
#[test]
fn pack_roundtrip_8bit() {
let w: Vec<i8> = (-128..=127)
.collect::<Vec<_>>()
.iter()
.map(|&x| x as i8)
.collect();
let packed = pack(&w, 8);
for (i, &expected) in w.iter().enumerate() {
assert_eq!(read_weight(&packed, i, 8), expected as i32);
}
}
#[test]
fn pack_roundtrip_4bit() {
let w: Vec<i8> = (-7..=7).collect();
let packed = pack(&w, 4);
assert_eq!(packed.len(), packed_byte_len(w.len(), 4));
for (i, &expected) in w.iter().enumerate() {
assert_eq!(read_weight(&packed, i, 4), expected as i32);
}
}
#[test]
fn pack_roundtrip_2bit_ternary() {
let w: Vec<i8> = vec![1, -1, 0, 1, 0, 1, -1, 0, -1, 1];
let packed = pack(&w, 2);
assert_eq!(packed.len(), packed_byte_len(w.len(), 2));
for (i, &expected) in w.iter().enumerate() {
assert_eq!(read_weight(&packed, i, 2), expected as i32);
}
}
#[test]
fn split_lanes_8bit_roundtrip() {
let w: Vec<i8> = (1i8..=8).collect(); let packed = pack(&w, 8);
let lanes = split_weights_by_oc_lane(&packed, 4, 2, 8, 2);
assert_eq!(lanes.len(), 2);
assert_eq!(read_weight(&lanes[0], 0, 8), 1);
assert_eq!(read_weight(&lanes[0], 1, 8), 2);
assert_eq!(read_weight(&lanes[0], 2, 8), 5);
assert_eq!(read_weight(&lanes[0], 3, 8), 6);
assert_eq!(read_weight(&lanes[1], 0, 8), 3);
assert_eq!(read_weight(&lanes[1], 1, 8), 4);
assert_eq!(read_weight(&lanes[1], 2, 8), 7);
assert_eq!(read_weight(&lanes[1], 3, 8), 8);
}
#[test]
fn split_lanes_2bit_roundtrip() {
let w: Vec<i8> = vec![
1, -1, 0, 1, 0, 1, -1, 0, -1, 1, 0, -1, 1, 0, 1, -1, ];
let packed = pack(&w, 2);
let lanes = split_weights_by_oc_lane(&packed, 4, 4, 2, 2);
assert_eq!(lanes.len(), 2);
for (i, &expected) in [1, -1, 0, 1, -1, 1, 0, -1].iter().enumerate() {
assert_eq!(read_weight(&lanes[0], i, 2), expected);
}
for (i, &expected) in [0, 1, -1, 0, 1, 0, 1, -1].iter().enumerate() {
assert_eq!(read_weight(&lanes[1], i, 2), expected);
}
}
#[test]
fn split_table_partitions_by_oc_mod_p() {
let bias: Vec<i32> = (10..=17).collect(); let lanes = split_table_by_oc_lane(&bias, 4);
assert_eq!(lanes.len(), 4);
assert_eq!(lanes[0], vec![10, 14]); assert_eq!(lanes[1], vec![11, 15]); assert_eq!(lanes[2], vec![12, 16]); assert_eq!(lanes[3], vec![13, 17]); }
}