use super::vector_types::{VectorQuantType, VectorValueType};
#[derive(Debug, Clone, PartialEq)]
pub struct PreparedVectorData {
pub bytes: Vec<u8>,
pub element_count: usize,
}
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
pub enum PrepareError {
NotU8Range,
NegativeForU8,
NotI8Range,
OverI8Max,
}
impl PrepareError {
pub fn message(self) -> &'static [u8] {
match self {
PrepareError::NotU8Range => {
b"ERR Vector contains element that is < 0 or > 255, operation will lose precision"
}
PrepareError::NegativeForU8 => {
b"ERR Vector contains element that is < 0, operation will lose precision"
}
PrepareError::NotI8Range => {
b"ERR Vector contains element that is < -128 or > 127, operation will lose precision"
}
PrepareError::OverI8Max => {
b"ERR Vector contains element that is > 127, operation will lose precision"
}
}
}
}
pub fn native_format(quant: VectorQuantType) -> VectorValueType {
match quant {
VectorQuantType::NoQuant | VectorQuantType::Q8 | VectorQuantType::Bin => VectorValueType::FP32,
VectorQuantType::XNoQuant_U8 | VectorQuantType::XBin_U8 => VectorValueType::XU8,
VectorQuantType::XNoQuant_I8 | VectorQuantType::XBin_I8 => VectorValueType::XI8,
VectorQuantType::Invalid => VectorValueType::Invalid,
}
}
pub fn prepare_vector_data(
quant: VectorQuantType,
value_type: VectorValueType,
provided: &[u8],
) -> Result<PreparedVectorData, PrepareError> {
match quant {
VectorQuantType::NoQuant | VectorQuantType::Q8 | VectorQuantType::Bin => match value_type {
VectorValueType::FP32 => Ok(convert_f32_for_alignment(provided)),
VectorValueType::XI8 => Ok(convert_i8_to_f32(provided)),
VectorValueType::XU8 => Ok(convert_u8_to_f32(provided)),
VectorValueType::Invalid => Err(PrepareError::NotI8Range),
},
VectorQuantType::XNoQuant_U8 | VectorQuantType::XBin_U8 => match value_type {
VectorValueType::FP32 => convert_f32_to_u8(provided),
VectorValueType::XI8 => convert_i8_to_u8(provided),
VectorValueType::XU8 => Ok(pass_through(provided)),
VectorValueType::Invalid => Err(PrepareError::NotI8Range),
},
VectorQuantType::XNoQuant_I8 | VectorQuantType::XBin_I8 => match value_type {
VectorValueType::FP32 => convert_f32_to_i8(provided),
VectorValueType::XI8 => Ok(pass_through(provided)),
VectorValueType::XU8 => convert_u8_to_i8(provided),
VectorValueType::Invalid => Err(PrepareError::NotI8Range),
},
VectorQuantType::Invalid => Err(PrepareError::NotI8Range),
}
}
fn pass_through(data: &[u8]) -> PreparedVectorData {
PreparedVectorData {
bytes: data.to_vec(),
element_count: data.len(),
}
}
fn convert_f32_for_alignment(data: &[u8]) -> PreparedVectorData {
PreparedVectorData {
bytes: data.to_vec(),
element_count: data.len() / 4,
}
}
fn convert_i8_to_f32(data: &[u8]) -> PreparedVectorData {
let values: Vec<f32> = data
.iter()
.map(|b| f32::from(i8::from_le_bytes([*b])))
.collect();
PreparedVectorData {
bytes: values.iter().flat_map(|v| v.to_le_bytes()).collect(),
element_count: data.len(),
}
}
fn convert_u8_to_f32(data: &[u8]) -> PreparedVectorData {
let values: Vec<f32> = data.iter().map(|b| f32::from(*b)).collect();
PreparedVectorData {
bytes: values.iter().flat_map(|v| v.to_le_bytes()).collect(),
element_count: data.len(),
}
}
fn convert_f32_to_u8(data: &[u8]) -> Result<PreparedVectorData, PrepareError> {
let values = f32_values(data);
if values.iter().any(|v| !(0.0..=255.0).contains(v)) {
return Err(PrepareError::NotU8Range);
}
Ok(PreparedVectorData {
bytes: values.iter().map(|v| *v as u8).collect(),
element_count: values.len(),
})
}
fn convert_i8_to_u8(data: &[u8]) -> Result<PreparedVectorData, PrepareError> {
if data.iter().any(|b| i8::from_le_bytes([*b]) < 0) {
return Err(PrepareError::NegativeForU8);
}
Ok(PreparedVectorData {
bytes: data.to_vec(),
element_count: data.len(),
})
}
fn convert_f32_to_i8(data: &[u8]) -> Result<PreparedVectorData, PrepareError> {
let values = f32_values(data);
if values.iter().any(|v| !(-128.0..=127.0).contains(v)) {
return Err(PrepareError::NotI8Range);
}
Ok(PreparedVectorData {
bytes: values.iter().map(|v| *v as i8 as u8).collect(),
element_count: values.len(),
})
}
fn convert_u8_to_i8(data: &[u8]) -> Result<PreparedVectorData, PrepareError> {
if data.iter().any(|b| *b > 127) {
return Err(PrepareError::OverI8Max);
}
Ok(PreparedVectorData {
bytes: data.to_vec(),
element_count: data.len(),
})
}
fn f32_values(data: &[u8]) -> Vec<f32> {
data
.as_chunks::<4>()
.0
.iter()
.map(|c| f32::from_le_bytes([c[0], c[1], c[2], c[3]]))
.collect()
}
#[cfg(test)]
mod tests {
use super::*;
fn f32_bytes(vals: &[f32]) -> Vec<u8> {
vals.iter().flat_map(|v| v.to_le_bytes()).collect()
}
#[test]
fn redis_quantizers_expect_f32() {
for quant in [
VectorQuantType::NoQuant,
VectorQuantType::Q8,
VectorQuantType::Bin,
] {
assert_eq!(native_format(quant), VectorValueType::FP32);
}
assert_eq!(
native_format(VectorQuantType::XNoQuant_U8),
VectorValueType::XU8
);
assert_eq!(
native_format(VectorQuantType::XBin_I8),
VectorValueType::XI8
);
assert_eq!(
native_format(VectorQuantType::Invalid),
VectorValueType::Invalid
);
let p = prepare_vector_data(
VectorQuantType::NoQuant,
VectorValueType::FP32,
&f32_bytes(&[1.0, 2.0]),
)
.unwrap();
assert_eq!(p.element_count, 2);
let p =
prepare_vector_data(VectorQuantType::NoQuant, VectorValueType::XI8, &[250u8, 3]).unwrap();
assert_eq!(p.element_count, 2);
assert_eq!(f32_values(&p.bytes), vec![-6.0, 3.0]);
let p = prepare_vector_data(VectorQuantType::Q8, VectorValueType::XU8, &[0, 255]).unwrap();
assert_eq!(f32_values(&p.bytes), vec![0.0, 255.0]);
}
#[test]
fn extended_u8_targets() {
let p = prepare_vector_data(
VectorQuantType::XNoQuant_U8,
VectorValueType::FP32,
&f32_bytes(&[0.0, 255.0]),
)
.unwrap();
assert_eq!(p.bytes, vec![0, 255]);
assert_eq!(
prepare_vector_data(
VectorQuantType::XNoQuant_U8,
VectorValueType::FP32,
&f32_bytes(&[-1.0])
)
.unwrap_err(),
PrepareError::NotU8Range
);
assert_eq!(
prepare_vector_data(
VectorQuantType::XNoQuant_U8,
VectorValueType::FP32,
&f32_bytes(&[256.0])
)
.unwrap_err(),
PrepareError::NotU8Range
);
let p = prepare_vector_data(VectorQuantType::XBin_U8, VectorValueType::XI8, &[3, 127]).unwrap();
assert_eq!(p.bytes, vec![3, 127]);
assert_eq!(
prepare_vector_data(VectorQuantType::XNoQuant_U8, VectorValueType::XI8, &[255]).unwrap_err(),
PrepareError::NegativeForU8
);
let p = prepare_vector_data(VectorQuantType::XBin_U8, VectorValueType::XU8, &[9, 8]).unwrap();
assert_eq!(p.bytes, vec![9, 8]);
assert_eq!(p.element_count, 2);
}
#[test]
fn extended_i8_targets() {
let p = prepare_vector_data(
VectorQuantType::XNoQuant_I8,
VectorValueType::FP32,
&f32_bytes(&[-128.0, 127.0]),
)
.unwrap();
assert_eq!(p.bytes, vec![128, 127]);
assert_eq!(
prepare_vector_data(
VectorQuantType::XBin_I8,
VectorValueType::FP32,
&f32_bytes(&[127.5])
)
.unwrap_err(),
PrepareError::NotI8Range
);
assert_eq!(
prepare_vector_data(VectorQuantType::XNoQuant_I8, VectorValueType::XU8, &[128]).unwrap_err(),
PrepareError::OverI8Max
);
let p = prepare_vector_data(
VectorQuantType::XNoQuant_I8,
VectorValueType::XU8,
&[0, 127],
)
.unwrap();
assert_eq!(p.bytes, vec![0, 127]);
let p = prepare_vector_data(VectorQuantType::XBin_I8, VectorValueType::XI8, &[254, 2]).unwrap();
assert_eq!(p.bytes, vec![254, 2]);
assert_eq!(p.element_count, 2);
}
#[test]
fn error_messages_match_csharp() {
assert!(
PrepareError::NotU8Range
.message()
.starts_with(b"ERR Vector contains element that is < 0 or > 255")
);
assert!(
PrepareError::NegativeForU8
.message()
.starts_with(b"ERR Vector contains element that is < 0,")
);
assert!(
PrepareError::NotI8Range
.message()
.starts_with(b"ERR Vector contains element that is < -128 or > 127")
);
assert!(
PrepareError::OverI8Max
.message()
.starts_with(b"ERR Vector contains element that is > 127")
);
}
}