use crate::model::Layer;
pub fn is_ternary(layer: &Layer) -> bool {
matches!(
layer,
Layer::Conv2d {
weight_bits: 2,
weight_encoding: crate::model::WeightEncoding::SignedInt,
w_zp: 0,
..
} | Layer::Dense {
weight_bits: 2,
weight_encoding: crate::model::WeightEncoding::SignedInt,
w_zp: 0,
..
}
)
}
#[cfg(test)]
mod tests {
use super::*;
use crate::model::Layer;
use crate::quant::quantize_multiplier;
#[test]
fn ternary_dense_detected() {
let (m0, sh) = quantize_multiplier(0.5);
let l = Layer::Dense {
name: "d",
in_features: 4,
out_features: 1,
x_zp: 0,
w_zp: 0,
out_zp: 0,
weight_bits: 2,
weight_encoding: crate::model::WeightEncoding::SignedInt,
requant: vec![(m0, sh)],
weights: vec![0x4D],
bias: None,
};
assert!(is_ternary(&l));
}
#[test]
fn eight_bit_dense_rejected() {
let (m0, sh) = quantize_multiplier(0.5);
let l = Layer::Dense {
name: "d",
in_features: 4,
out_features: 1,
x_zp: 0,
w_zp: 0,
out_zp: 0,
weight_bits: 8,
weight_encoding: crate::model::WeightEncoding::SignedInt,
requant: vec![(m0, sh)],
weights: vec![1; 4],
bias: None,
};
assert!(!is_ternary(&l));
}
}