use core::mem::size_of;
#[derive(Clone, Copy, Debug, PartialEq)]
pub enum CompressionScheme {
None = 0,
Float16 = 1,
Zstd = 2,
}
pub fn get_vector_compression_config() -> CompressionConfig {
CompressionConfig {
vector_compression_enabled: false,
vector_compression_scheme: CompressionScheme::None,
vector_compression_level: 3,
}
}
#[derive(Clone, Copy, Debug)]
pub struct CompressionConfig {
pub vector_compression_enabled: bool,
pub vector_compression_scheme: CompressionScheme,
pub vector_compression_level: u8,
}
#[repr(C)]
pub struct Float16(u16);
impl Float16 {
pub fn from_f32(value: f32) -> Self {
let bits = value.to_bits();
let sign = (bits >> 16) & 0x8000;
let mut exponent = ((bits >> 23) & 0xFF) as i16;
let mantissa = bits & 0x7FFFFF;
exponent = exponent - 127 + 15;
let result = if exponent <= 0 {
if exponent < -10 {
sign
} else {
let mantissa = (mantissa | 0x800000) >> (1 - exponent);
sign | (mantissa >> 13)
}
} else if exponent >= 31 {
if mantissa == 0 {
sign | 0x7C00
} else {
sign | 0x7FFF
}
} else {
sign | ((exponent as u32) << 10) | (mantissa >> 13)
};
Float16(result as u16)
}
pub fn to_f32(self) -> f32 {
let bits = self.0;
let sign = (bits as u32) << 16;
let exponent = ((bits >> 10) & 0x1F) as i16;
let mantissa = bits & 0x3FF;
let result = if exponent == 0 {
if mantissa == 0 {
sign
} else {
let exponent = -14;
let mantissa = mantissa << 13;
sign | ((exponent + 127) as u32) << 23 | (mantissa as u32)
}
} else if exponent == 0x1F {
sign | 0x7F800000 | ((mantissa << 13) as u32)
} else {
let exponent = exponent - 15 + 127;
sign | ((exponent as u32) << 23) | ((mantissa << 13) as u32)
};
f32::from_bits(result)
}
}
pub fn compress_vector(input: *const f32, dimension: usize, output: *mut u8) -> usize {
let is_compressed = false;
let compression_scheme = 0;
if is_compressed {
match compression_scheme {
1 => {
unsafe {
for i in 0..dimension {
let f32_val = *input.add(i);
let f16_val = Float16::from_f32(f32_val);
let output_ptr = output.add(i * size_of::<u16>()) as *mut u16;
*output_ptr = f16_val.0;
}
}
dimension * size_of::<u16>()
}
2 => {
unsafe {
core::ptr::copy_nonoverlapping(
input as *const u8,
output.add(4),
dimension * size_of::<f32>(),
);
let size_ptr = output as *mut u32;
*size_ptr = dimension as u32 * size_of::<f32>() as u32;
}
dimension * size_of::<f32>() + 4
}
_ => {
unsafe {
core::ptr::copy_nonoverlapping(
input as *const u8,
output,
dimension * size_of::<f32>(),
);
}
dimension * size_of::<f32>()
}
}
} else {
unsafe {
core::ptr::copy_nonoverlapping(
input as *const u8,
output,
dimension * size_of::<f32>(),
);
}
dimension * size_of::<f32>()
}
}
pub fn decompress_vector(input: *const u8, dimension: usize, output: *mut f32) {
let is_compressed = false;
let compression_scheme = 0;
if is_compressed {
match compression_scheme {
1 => {
unsafe {
for i in 0..dimension {
let input_ptr = input.add(i * size_of::<u16>()) as *const u16;
let f16_val = Float16(*input_ptr);
*output.add(i) = f16_val.to_f32();
}
}
}
2 => {
unsafe {
core::ptr::copy_nonoverlapping(input.add(4) as *const f32, output, dimension);
}
}
_ => {
unsafe {
core::ptr::copy_nonoverlapping(input as *const f32, output, dimension);
}
}
}
} else {
unsafe {
core::ptr::copy_nonoverlapping(input as *const f32, output, dimension);
}
}
}
pub fn get_compressed_size(dimension: usize) -> usize {
let is_compressed = false;
let compression_scheme = 0;
if is_compressed {
match compression_scheme {
1 => dimension * size_of::<u16>(),
2 => dimension * size_of::<f32>() + 4,
_ => dimension * size_of::<f32>(),
}
} else {
dimension * size_of::<f32>()
}
}
pub fn is_vector_compression_enabled() -> bool {
false
}
pub fn get_current_compression_scheme() -> u8 {
0
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn test_float16_conversion() {
let test_values = [0.0f32, 1.0f32, -1.0f32, 0.5f32, 2.0f32];
for val in test_values {
let f16 = Float16::from_f32(val);
let f32 = f16.to_f32();
assert!((f32 - val).abs() < 0.001 || (val == 0.0 && f32 == 0.0));
}
}
#[test]
fn test_vector_compression() {
let dimension = 4;
let test_vector = [1.0f32, 2.0f32, 3.0f32, 4.0f32];
let mut compressed = [0u8; 16]; let compressed_size =
compress_vector(test_vector.as_ptr(), dimension, compressed.as_mut_ptr());
let mut decompressed = [0.0f32; 4];
decompress_vector(compressed.as_ptr(), dimension, decompressed.as_mut_ptr());
for (i, (orig, decomp)) in test_vector.iter().zip(decompressed.iter()).enumerate() {
assert!(
(orig - decomp).abs() < 0.001,
"Index {}: orig={}, decomp={}",
i,
orig,
decomp
);
}
}
}