#![allow(clippy::module_name_repetitions)]
#![allow(
clippy::cast_possible_truncation,
clippy::cast_possible_wrap,
clippy::cast_sign_loss
)]
use crate::trit::Trit;
pub const I2_S_TAIL_BYTES: usize = 32;
pub const I2_S_BLOCK: usize = 128;
pub const Q1_0_BLOCK: usize = 128;
pub const Q2_0_BLOCK: usize = 128;
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
pub enum BridgeError {
BadLength,
TooShort,
UnsupportedCode,
}
impl core::fmt::Display for BridgeError {
fn fmt(&self, f: &mut core::fmt::Formatter<'_>) -> core::fmt::Result {
match self {
Self::BadLength => {
write!(f, "tensor length is not a multiple of the block size")
}
Self::TooShort => write!(f, "byte buffer shorter than the layout requires"),
Self::UnsupportedCode => {
write!(f, "2-bit code 3 found (reference quantizer cannot emit it)")
}
}
}
}
#[cfg(feature = "std")]
impl std::error::Error for BridgeError {}
const fn trit_to_code(t: Trit) -> u8 {
match t {
Trit::MinusOne => 0,
Trit::Zero => 1,
Trit::One => 2,
}
}
const fn code_to_trit(code: u8) -> Result<Trit, BridgeError> {
match code {
0 => Ok(Trit::MinusOne),
1 => Ok(Trit::Zero),
2 => Ok(Trit::One),
_ => Err(BridgeError::UnsupportedCode),
}
}
#[must_use]
const fn i2_s_encoded_len(n: usize) -> usize {
n / 4 + I2_S_TAIL_BYTES
}
#[must_use]
const fn i2_s_byte_index(i: usize) -> usize {
(i / I2_S_BLOCK) * 32 + (i % 32)
}
#[must_use]
const fn i2_s_lane_shift(i: usize) -> u32 {
6 - 2 * (((i % I2_S_BLOCK) / 32) as u32)
}
pub fn encode_i2_s(trits: &[Trit], scale_bits: u32, out: &mut [u8]) -> Result<usize, BridgeError> {
let n = trits.len();
if !n.is_multiple_of(I2_S_BLOCK) {
return Err(BridgeError::BadLength);
}
let written = i2_s_encoded_len(n);
if out.len() < written {
return Err(BridgeError::TooShort);
}
out[..written].fill(0);
for (i, t) in trits.iter().enumerate() {
out[i2_s_byte_index(i)] |= trit_to_code(*t) << i2_s_lane_shift(i);
}
out[n / 4..n / 4 + 4].copy_from_slice(&scale_bits.to_le_bytes());
Ok(written)
}
pub fn decode_i2_s(bytes: &[u8], trits: &mut [Trit]) -> Result<u32, BridgeError> {
let n = trits.len();
if !n.is_multiple_of(I2_S_BLOCK) {
return Err(BridgeError::BadLength);
}
if bytes.len() < i2_s_encoded_len(n) {
return Err(BridgeError::TooShort);
}
for (i, slot) in trits.iter_mut().enumerate() {
let byte = bytes[i2_s_byte_index(i)];
let code = (byte >> i2_s_lane_shift(i)) & 0x03;
*slot = code_to_trit(code)?;
}
let mut scale = [0u8; 4];
scale.copy_from_slice(&bytes[n / 4..n / 4 + 4]);
Ok(u32::from_le_bytes(scale))
}
#[must_use]
pub const fn q1_0_encoded_len(n: usize) -> usize {
(n / Q1_0_BLOCK) * 18
}
pub fn decode_q1_0(
bytes: &[u8],
trits: &mut [Trit],
scale_bits_out: &mut [u16],
) -> Result<(), BridgeError> {
let n = trits.len();
if !n.is_multiple_of(Q1_0_BLOCK) {
return Err(BridgeError::BadLength);
}
let blocks = n / Q1_0_BLOCK;
if bytes.len() < q1_0_encoded_len(n) {
return Err(BridgeError::TooShort);
}
if scale_bits_out.len() < blocks {
return Err(BridgeError::TooShort);
}
for b in 0..blocks {
let base = b * 18;
scale_bits_out[b] = u16::from_le_bytes([bytes[base], bytes[base + 1]]);
for j in 0..Q1_0_BLOCK {
let sign = (bytes[base + 2 + j / 8] >> (j % 8)) & 1;
trits[b * Q1_0_BLOCK + j] = if sign == 1 { Trit::One } else { Trit::MinusOne };
}
}
Ok(())
}
#[must_use]
pub const fn q2_0_encoded_len(n: usize) -> usize {
(n / Q2_0_BLOCK) * 34
}
pub fn decode_q2_0(
bytes: &[u8],
trits: &mut [Trit],
scale_bits_out: &mut [u16],
) -> Result<(), BridgeError> {
let n = trits.len();
if !n.is_multiple_of(Q2_0_BLOCK) {
return Err(BridgeError::BadLength);
}
let blocks = n / Q2_0_BLOCK;
if bytes.len() < q2_0_encoded_len(n) {
return Err(BridgeError::TooShort);
}
if scale_bits_out.len() < blocks {
return Err(BridgeError::TooShort);
}
for b in 0..blocks {
let base = b * 34;
scale_bits_out[b] = u16::from_le_bytes([bytes[base], bytes[base + 1]]);
for j in 0..Q2_0_BLOCK {
let code = (bytes[base + 2 + j / 4] >> (2 * (j % 4))) & 0x03;
trits[b * Q2_0_BLOCK + j] = code_to_trit(code)?;
}
}
Ok(())
}
pub fn encode_q2_0(trits: &[Trit], scale_bits: &[u16], out: &mut [u8]) -> Result<(), BridgeError> {
let n = trits.len();
if !n.is_multiple_of(Q2_0_BLOCK) {
return Err(BridgeError::BadLength);
}
let blocks = n / Q2_0_BLOCK;
if scale_bits.len() < blocks {
return Err(BridgeError::TooShort);
}
if out.len() < q2_0_encoded_len(n) {
return Err(BridgeError::TooShort);
}
for b in 0..blocks {
let base = b * 34;
out[base..base + 2].copy_from_slice(&scale_bits[b].to_le_bytes());
out[base + 2..base + 34].fill(0);
for j in 0..Q2_0_BLOCK {
let code = trit_to_code(trits[b * Q2_0_BLOCK + j]);
out[base + 2 + j / 4] |= code << (2 * (j % 4));
}
}
Ok(())
}
#[must_use]
pub const fn half_to_f32_bits(h: u16) -> u32 {
let sign = ((h >> 15) as u32) << 31;
let exp = ((h >> 10) & 0x1F) as u32;
let mant = (h & 0x03FF) as u32;
if exp == 0 {
if mant == 0 {
return sign; }
let mut k = 0_u32;
let m = mant;
while (m >> k) > 1 {
k += 1;
}
let e32 = k + 103;
let m32 = (mant ^ (1 << k)) << (23 - k);
return sign | (e32 << 23) | m32;
}
if exp == 0x1F {
return sign | (0xFF_u32 << 23) | (mant << 13);
}
sign | ((exp + 112) << 23) | (mant << 13)
}
#[must_use]
pub fn half_to_milli(h: u16) -> i32 {
let bits = half_to_f32_bits(h);
let negative = (bits >> 31) == 1;
let exp = ((bits >> 23) & 0xFF) as i32;
if exp == 0 {
return 0; }
let mant = (bits & 0x007F_FFFF) | (0x0080_0000); if exp == 0xFF {
if bits.trailing_zeros() >= 23 {
return if negative { i32::MIN } else { i32::MAX };
}
return 0; }
let shift = 127 + 23 - exp; let scaled = i64::from(mant) * 1000;
let milli = (scaled + (1_i64 << (shift - 1))) >> shift;
if negative {
(-milli) as i32
} else {
milli as i32
}
}
#[must_use]
pub fn wire_gamma_to_substrate(gamma_milli: i32) -> i16 {
let scaled = (i64::from(gamma_milli) * i64::from(crate::synapse::SCALE)
+ if gamma_milli >= 0 { 500 } else { -500 })
/ 1000; scaled.clamp(i64::from(i16::MIN), i64::from(i16::MAX)) as i16
}
pub fn repack_i2s_to_kernel(i2s: &[u8], n: usize, out: &mut [u8]) -> Result<(), BridgeError> {
if !n.is_multiple_of(I2_S_BLOCK) {
return Err(BridgeError::BadLength);
}
if i2s.len() < i2_s_encoded_len(n) || out.len() < n / 4 {
return Err(BridgeError::TooShort);
}
out[..n / 4].fill(0);
for i in 0..n {
let code = (i2s[i2_s_byte_index(i)] >> i2_s_lane_shift(i)) & 0x03;
if code == 3 {
return Err(BridgeError::UnsupportedCode);
}
out[i / 4] |= code << (2 * (i % 4));
}
Ok(())
}
#[cfg(test)]
mod tests {
#![allow(clippy::shadow_unrelated)]
use super::*;
use proptest::prelude::*;
#[test]
fn i2_s_known_vector() {
let pat = [
Trit::One,
Trit::MinusOne,
Trit::Zero,
Trit::One,
Trit::Zero,
Trit::Zero,
Trit::MinusOne,
Trit::One,
];
let trits: Vec<Trit> = (0..I2_S_BLOCK).map(|i| pat[i % 8]).collect();
let scale_bits: u32 = 0x4000_0000; let mut buf = [0xFF_u8; 128]; let written = encode_i2_s(&trits, scale_bits, &mut buf).expect("encode");
assert_eq!(written, 32 + I2_S_TAIL_BYTES);
let mut expected = [0_u8; 64];
expected[..32].copy_from_slice(&[0xAA, 0x00, 0x55, 0xAA, 0x55, 0x55, 0x00, 0xAA].repeat(4));
expected[32..36].copy_from_slice(&scale_bits.to_le_bytes());
assert_eq!(
&buf[..written],
&expected,
"i2_s bytes must match reference layout"
);
let mut back = [Trit::Zero; I2_S_BLOCK];
let got_scale = decode_i2_s(&buf, &mut back).expect("decode");
assert_eq!(got_scale, scale_bits);
assert_eq!(&back, &trits[..]);
}
#[test]
fn i2_s_lane_order_golden_vector() {
let code_of = |i: usize| -> u8 {
let t = match i % 3 {
0 => Trit::MinusOne,
1 => Trit::Zero,
_ => Trit::One,
};
trit_to_code(t)
};
let trits: Vec<Trit> = (0..I2_S_BLOCK)
.map(|i| match i % 3 {
0 => Trit::MinusOne,
1 => Trit::Zero,
_ => Trit::One,
})
.collect();
let mut buf = [0_u8; 64];
let written = encode_i2_s(&trits, 0, &mut buf).expect("encode");
assert_eq!(written, 64);
let mut expected = [0_u8; 64];
for (j, slot) in expected.iter_mut().enumerate().take(32) {
for lane in 0..4_usize {
let e = 32 * lane + j; *slot |= code_of(e) << (6 - 2 * lane);
}
}
assert_eq!(
&buf[..32],
&expected[..32],
"i2_s lane order must match the reference transposition"
);
let mut back = [Trit::Zero; I2_S_BLOCK];
decode_i2_s(&buf, &mut back).expect("decode");
assert_eq!(&back, &trits[..]);
}
#[test]
fn i2_s_rejects_bad_lengths() {
let trits = [Trit::One; 8]; let mut out = [0u8; 64];
assert_eq!(
encode_i2_s(&trits, 0, &mut out),
Err(BridgeError::BadLength)
);
let mut back = [Trit::Zero; 8];
assert_eq!(
decode_i2_s(&[0; 64], &mut back),
Err(BridgeError::BadLength)
);
}
#[test]
fn i2_s_rejects_short_buffers() {
let trits = [Trit::One; I2_S_BLOCK];
let mut out = [0u8; 10]; assert_eq!(encode_i2_s(&trits, 0, &mut out), Err(BridgeError::TooShort));
let mut back = [Trit::Zero; I2_S_BLOCK];
assert_eq!(decode_i2_s(&[0; 12], &mut back), Err(BridgeError::TooShort));
}
#[test]
fn i2_s_rejects_code_three() {
let mut back = [Trit::Zero; I2_S_BLOCK];
assert_eq!(
decode_i2_s(&[0xFF; 64], &mut back),
Err(BridgeError::UnsupportedCode)
);
}
proptest! {
#[test]
fn prop_i2_s_round_trip(
blocks in 1_usize..=4,
seed in any::<u32>(),
scale_bits in any::<u32>(),
) {
let n = blocks * I2_S_BLOCK;
let mut x = seed | 1;
let trits: Vec<Trit> = (0..n)
.map(|_| {
x ^= x << 13; x ^= x >> 17; x ^= x << 5;
match x % 3 { 0 => Trit::MinusOne, 1 => Trit::Zero, _ => Trit::One }
})
.collect();
let mut buf = vec![0_u8; i2_s_encoded_len(n) + 8];
let written = encode_i2_s(&trits, scale_bits, &mut buf).unwrap();
prop_assert_eq!(written, i2_s_encoded_len(n));
let mut back = vec![Trit::Zero; n];
let got = decode_i2_s(&buf[..written], &mut back).unwrap();
prop_assert_eq!(got, scale_bits);
prop_assert_eq!(&back, &trits[..]);
}
}
#[test]
fn q1_0_known_vector() {
let pat = [
Trit::One,
Trit::MinusOne,
Trit::One,
Trit::MinusOne,
Trit::One,
Trit::One,
Trit::MinusOne,
Trit::One,
];
let mut bytes = Vec::with_capacity(18);
bytes.extend_from_slice(&0x3C00_u16.to_le_bytes()); bytes.extend(std::iter::repeat_n(0xB5_u8, 16));
let mut trits = [Trit::Zero; Q1_0_BLOCK];
let mut scales = [0_u16; 1];
decode_q1_0(&bytes, &mut trits, &mut scales).expect("decode");
assert_eq!(scales[0], 0x3C00);
let expected: Vec<Trit> = (0..Q1_0_BLOCK).map(|i| pat[i % 8]).collect();
assert_eq!(&trits, &expected[..], "q1_0 sign bits must map LSB-first");
}
#[test]
fn q1_0_byte_order_golden_vector() {
let mut bytes = Vec::with_capacity(18);
bytes.extend_from_slice(&0x3C00_u16.to_le_bytes());
let mut sign_bytes = [0_u8; 16];
for (j, slot) in sign_bytes.iter_mut().enumerate() {
*slot = 1 << (j % 8);
}
bytes.extend_from_slice(&sign_bytes);
let mut trits = [Trit::Zero; Q1_0_BLOCK];
let mut scales = [0_u16; 1];
decode_q1_0(&bytes, &mut trits, &mut scales).expect("decode");
let positives: std::collections::HashSet<usize> =
(0..16_usize).map(|j| 8 * j + j % 8).collect();
for (i, t) in trits.iter().enumerate() {
let want = if positives.contains(&i) {
Trit::One
} else {
Trit::MinusOne
};
assert_eq!(*t, want, "element {i}");
}
}
#[test]
fn q1_0_rejects_bad_input() {
let mut trits = [Trit::Zero; Q1_0_BLOCK];
let mut scales = [0u16; 1];
let mut short_trits = [Trit::Zero; 64];
assert_eq!(
decode_q1_0(&[0; 18], &mut short_trits, &mut scales),
Err(BridgeError::BadLength)
);
assert_eq!(
decode_q1_0(&[0; 17], &mut trits, &mut scales),
Err(BridgeError::TooShort)
);
let mut no_scales = [0u16; 0];
assert_eq!(
decode_q1_0(&[0; 18], &mut trits, &mut no_scales),
Err(BridgeError::TooShort)
);
}
#[test]
fn q2_0_block_geometry_is_pinned() {
assert_eq!(Q2_0_BLOCK, 128);
assert_eq!(q2_0_encoded_len(128), 34);
assert_eq!(q2_0_encoded_len(2560), 680);
assert_eq!(q2_0_encoded_len(128 * 3), 34 * 3);
}
#[test]
fn q2_0_known_vector() {
let pat = [Trit::MinusOne, Trit::Zero, Trit::One, Trit::One];
let mut bytes = Vec::with_capacity(34);
bytes.extend_from_slice(&0x4400_u16.to_le_bytes());
bytes.extend(std::iter::repeat_n(0xA4_u8, 32));
let mut trits = [Trit::Zero; Q2_0_BLOCK];
let mut scales = [0u16; 1];
decode_q2_0(&bytes, &mut trits, &mut scales).expect("decode");
assert_eq!(scales[0], 0x4400);
let expected: Vec<Trit> = (0..Q2_0_BLOCK).map(|i| pat[i % 4]).collect();
assert_eq!(&trits, &expected[..]);
}
#[test]
fn q2_0_code3_rejected() {
let mut bytes = vec![0_u8; 34];
bytes[0..2].copy_from_slice(&0x4400_u16.to_le_bytes());
bytes[2] = 0x03;
let mut trits = [Trit::Zero; Q2_0_BLOCK];
let mut scales = [0_u16; 1];
assert_eq!(
decode_q2_0(&bytes, &mut trits, &mut scales),
Err(BridgeError::UnsupportedCode)
);
let mut bytes = vec![0_u8; 34];
bytes[0..2].copy_from_slice(&0x4400_u16.to_le_bytes());
bytes[33] = 0xC0;
assert_eq!(
decode_q2_0(&bytes, &mut trits, &mut scales),
Err(BridgeError::UnsupportedCode)
);
}
#[test]
fn q2_0_byte_and_lane_order_golden_vector() {
let code_of = |i: usize| -> u8 { (i % 3) as u8 };
let mut bytes = Vec::with_capacity(34);
bytes.extend_from_slice(&0x4400_u16.to_le_bytes()); for j in 0..32_usize {
let mut b = 0_u8;
for lane in 0..4_usize {
b |= code_of(4 * j + lane) << (2 * lane);
}
bytes.push(b);
}
let mut trits = [Trit::Zero; Q2_0_BLOCK];
let mut scales = [0_u16; 1];
decode_q2_0(&bytes, &mut trits, &mut scales).expect("decode");
for (i, t) in trits.iter().enumerate() {
let want = match i % 3 {
0 => Trit::MinusOne,
1 => Trit::Zero,
_ => Trit::One,
};
assert_eq!(*t, want, "element {i}");
}
}
#[test]
fn q2_0_rejects_bad_input() {
let mut trits = [Trit::Zero; Q2_0_BLOCK];
let mut scales = [0u16; 1];
let mut odd = [Trit::Zero; 100];
assert_eq!(
decode_q2_0(&[0; 34], &mut odd, &mut scales),
Err(BridgeError::BadLength)
);
assert_eq!(
decode_q2_0(&[0; 33], &mut trits, &mut scales),
Err(BridgeError::TooShort)
);
let mut two = [Trit::Zero; Q2_0_BLOCK * 2];
assert_eq!(
decode_q2_0(&[0; 34 * 2], &mut two, &mut scales),
Err(BridgeError::TooShort)
);
}
const REAL_Q2_0_FIRST_BLOCK: [u8; 34] = [
0xC8, 0x24, 0x14, 0x44, 0x45, 0x1A, 0x18, 0x68, 0x68, 0x61, 0x8A, 0xA8, 0x91, 0x66, 0x45,
0x42, 0x91, 0x80, 0x11, 0x62, 0x18, 0x11, 0x29, 0x48, 0x61, 0x00, 0x1A, 0x94, 0x81, 0x24,
0x54, 0x0A, 0x86, 0x84,
];
#[test]
fn q2_0_real_artifact_first_block_decodes_and_round_trips() {
let mut trits = [Trit::Zero; Q2_0_BLOCK];
let mut scales = [0u16; 1];
decode_q2_0(&REAL_Q2_0_FIRST_BLOCK, &mut trits, &mut scales)
.expect("real bytes decode (code 3 would mean the pin is wrong)");
assert_eq!(scales[0], 0x24C8);
let census = trits.iter().fold((0, 0, 0), |(p, z, m), t| match t {
Trit::One => (p + 1, z, m),
Trit::Zero => (p, z + 1, m),
Trit::MinusOne => (p, z, m + 1),
});
assert_eq!(
census,
(37, 43, 48),
"recorded probe census: +37 / 0×43 / −48"
);
let mut back = [0u8; 34];
encode_q2_0(&trits, &scales, &mut back).expect("encode real block");
assert_eq!(&back[..], &REAL_Q2_0_FIRST_BLOCK[..]);
}
#[test]
fn encode_q2_0_reproduces_known_vector_bytes() {
let pat = [Trit::MinusOne, Trit::Zero, Trit::One, Trit::One];
let trits: Vec<Trit> = (0..Q2_0_BLOCK).map(|i| pat[i % 4]).collect();
let mut out = [0_u8; 34];
encode_q2_0(&trits, &[0x4400], &mut out).expect("encode");
let mut expected = Vec::with_capacity(34);
expected.extend_from_slice(&0x4400_u16.to_le_bytes());
expected.extend(std::iter::repeat_n(0xA4_u8, 32));
assert_eq!(&out[..], &expected[..]);
}
#[test]
fn encode_q2_0_reproduces_golden_vector_bytes() {
let code_of = |i: usize| -> u8 { (i % 3) as u8 };
let mut bytes = Vec::with_capacity(34);
bytes.extend_from_slice(&0x4400_u16.to_le_bytes());
for j in 0..32_usize {
let mut b = 0_u8;
for lane in 0..4_usize {
b |= code_of(4 * j + lane) << (2 * lane);
}
bytes.push(b);
}
let mut trits = [Trit::Zero; Q2_0_BLOCK];
let mut scales = [0u16; 1];
decode_q2_0(&bytes, &mut trits, &mut scales).expect("decode");
let mut out = [0_u8; 34];
encode_q2_0(&trits, &scales, &mut out).expect("encode");
assert_eq!(&out[..], &bytes[..]);
}
#[test]
fn encode_q2_0_rejects_bad_input() {
let scales = [0x4400_u16; 2];
let odd = [Trit::Zero; 100];
assert_eq!(
encode_q2_0(&odd, &scales, &mut [0; 68]),
Err(BridgeError::BadLength)
);
let two_blocks = [Trit::Zero; Q2_0_BLOCK * 2];
assert_eq!(
encode_q2_0(&two_blocks, &scales, &mut [0; 34]),
Err(BridgeError::TooShort)
);
assert_eq!(
encode_q2_0(&two_blocks, &scales[..1], &mut [0; 68]),
Err(BridgeError::TooShort)
);
}
proptest! {
#[test]
fn prop_q2_0_round_trip(
blocks in 1_usize..=4,
seed in any::<u32>(),
scales_seed in any::<u64>(),
) {
let n = blocks * Q2_0_BLOCK;
let mut x = seed | 1;
let trits: Vec<Trit> = (0..n)
.map(|_| {
x ^= x << 13; x ^= x >> 17; x ^= x << 5;
match x % 3 { 0 => Trit::MinusOne, 1 => Trit::Zero, _ => Trit::One }
})
.collect();
let mut s = scales_seed | 1;
let scales: Vec<u16> = (0..blocks)
.map(|_| {
s ^= s << 13; s ^= s >> 7; s ^= s << 17;
(s & 0xFFFF) as u16
})
.collect();
let mut buf = vec![0xFF_u8; q2_0_encoded_len(n) + 8]; encode_q2_0(&trits, &scales, &mut buf).unwrap();
let mut back = vec![Trit::Zero; n];
let mut back_scales = vec![0u16; blocks];
decode_q2_0(&buf[..q2_0_encoded_len(n)], &mut back, &mut back_scales).unwrap();
prop_assert_eq!(&back_scales, &scales[..]);
prop_assert_eq!(&back, &trits[..]);
}
}
#[test]
fn half_known_vectors() {
let cases: &[(u16, u32, i32)] = &[
(0x0000, 0x0000_0000, 0), (0x8000, 0x8000_0000, 0), (0x3C00, 0x3F80_0000, 1000), (0x3800, 0x3F00_0000, 500), (0xC000, 0xC000_0000, -2000), (0x4400, 0x4080_0000, 4000), (0x7BFF, 0x477F_E000, 65_504_000), (0x03FF, 0x387F_C000, 0), (0x0400, 0x3880_0000, 0), (0x7C00, 0x7F80_0000, i32::MAX), (0xFC00, 0xFF80_0000, i32::MIN), (0x7E00, 0x7FC0_0000, 0), ];
for &(h, f32_bits, milli) in cases {
assert_eq!(half_to_f32_bits(h), f32_bits, "f32 bits for fp16 {h:#06x}");
assert_eq!(half_to_milli(h), milli, "milli for fp16 {h:#06x}");
}
}
#[cfg(feature = "std")]
#[test]
fn half_to_milli_exhaustive_vs_f64() {
for h in 0..=u16::MAX {
let f = f32::from_bits(half_to_f32_bits(h));
let want: i64 = if f.is_nan() {
0
} else if f.is_infinite() {
i64::from(if f > 0.0 { i32::MAX } else { i32::MIN })
} else {
let r = (f64::from(f) * 1000.0).round();
if r >= f64::from(i32::MAX) {
i64::from(i32::MAX)
} else if r <= f64::from(i32::MIN) {
i64::from(i32::MIN)
} else {
r as i64
}
};
assert_eq!(i64::from(half_to_milli(h)), want, "fp16 bits {h:#06x}");
}
}
#[test]
fn decoded_trits_feed_trit_substrate() {
let mut bytes = Vec::with_capacity(18);
bytes.extend_from_slice(&0x3C00_u16.to_le_bytes());
bytes.extend(std::iter::repeat_n(0xB5_u8, 16));
let mut trits = [Trit::Zero; Q1_0_BLOCK];
let mut scales = [0u16; 1];
decode_q1_0(&bytes, &mut trits, &mut scales).expect("decode");
let gamma = 125_i16;
for t in trits {
let w = t.to_weight(gamma);
assert!(
w == gamma || w == 0 || w == -gamma,
"imported trit produced off-grid weight {w}"
);
assert_eq!(
Trit::from_weight(w, gamma),
t,
"classification must round-trip"
);
}
}
proptest! {
#[cfg(feature = "std")]
#[test]
fn prop_half_widening_matches_f32_reference(h in any::<u16>()) {
let f = f32::from_bits(half_to_f32_bits(h));
let reference = half_to_f32_via_f32(h);
if reference.is_nan() {
prop_assert!(f.is_nan());
} else {
prop_assert_eq!(f.to_bits(), reference.to_bits());
}
}
}
#[test]
fn wire_gamma_known_vectors() {
assert_eq!(wire_gamma_to_substrate(24), 24); assert_eq!(wire_gamma_to_substrate(0), 0);
assert_eq!(wire_gamma_to_substrate(125), 125); assert_eq!(wire_gamma_to_substrate(65_504_000), i16::MAX); assert_eq!(wire_gamma_to_substrate(-40_000), i16::MIN);
assert_eq!(wire_gamma_to_substrate(-40_000_000), i16::MIN);
assert_eq!(wire_gamma_to_substrate(-5), -5);
}
#[test]
fn scale_constant_is_pinned() {
assert_eq!(crate::synapse::SCALE, 1000);
}
#[test]
fn repack_known_vector() {
let pat = [
Trit::One,
Trit::MinusOne,
Trit::Zero,
Trit::One,
Trit::Zero,
Trit::Zero,
Trit::One,
Trit::MinusOne,
];
let trits: Vec<Trit> = (0..I2_S_BLOCK).map(|i| pat[i % 8]).collect();
let mut wire = [0_u8; 64];
encode_i2_s(&trits, 0x4000_0000, &mut wire).expect("encode");
let mut kernel = [0_u8; 32];
repack_i2s_to_kernel(&wire, I2_S_BLOCK, &mut kernel).expect("repack");
for (i, &t) in trits.iter().enumerate() {
let code = (kernel[i / 4] >> (2 * (i % 4))) & 0x03;
let got = code_to_trit(code).expect("no code 3 in encode output");
assert_eq!(got, t, "element {i} wrong after repack");
}
}
#[test]
fn repack_rejects_bad_input() {
let mut out = [0_u8; 32];
assert_eq!(
repack_i2s_to_kernel(&[0; 64], 64, &mut out),
Err(BridgeError::BadLength)
);
assert_eq!(
repack_i2s_to_kernel(&[0; 12], 128, &mut out),
Err(BridgeError::TooShort)
);
let mut tiny = [0_u8; 4];
assert_eq!(
repack_i2s_to_kernel(&[0; 64], 128, &mut tiny),
Err(BridgeError::TooShort)
);
let evil = [0xFF_u8; 64];
let mut sink = [0_u8; 32];
assert_eq!(
repack_i2s_to_kernel(&evil, 128, &mut sink),
Err(BridgeError::UnsupportedCode)
);
}
proptest! {
#[test]
fn prop_repack_round_trip(
blocks in 1_usize..=3,
seed in any::<u32>(),
) {
let n = blocks * I2_S_BLOCK;
let mut x = seed | 1;
let trits: Vec<Trit> = (0..n)
.map(|_| {
x ^= x << 13; x ^= x >> 17; x ^= x << 5;
match x % 3 { 0 => Trit::MinusOne, 1 => Trit::Zero, _ => Trit::One }
})
.collect();
let mut wire = vec![0_u8; i2_s_encoded_len(n)];
encode_i2_s(&trits, 0, &mut wire).unwrap();
let mut kernel = vec![0_u8; n / 4];
repack_i2s_to_kernel(&wire, n, &mut kernel).unwrap();
for (i, &t) in trits.iter().enumerate() {
let code = (kernel[i / 4] >> (2 * (i % 4))) & 0x03;
prop_assert_eq!(code_to_trit(code).unwrap(), t, "element {}", i);
}
}
}
#[cfg(feature = "std")]
fn half_to_f32_via_f32(h: u16) -> f32 {
let sign = if (h >> 15) & 1 == 1 { -1.0_f32 } else { 1.0 };
let exp = i32::from((h >> 10) & 0x1F);
let mant = f32::from(h & 0x03FF);
if exp == 0 {
sign * mant * (2.0_f32).powi(-24)
} else if exp == 0x1F {
if mant == 0.0 {
sign * f32::INFINITY
} else {
f32::NAN
}
} else {
sign * (1.0 + mant / 1024.0) * (2.0_f32).powi(exp - 15)
}
}
}