urna_format/encoding/
int8.rs1use crate::bytes::le_u32;
16use crate::error::UrnaError;
17
18pub const INT8_PAYLOAD_VERSION: u32 = 1;
19pub const INT8_SCALE_KIND_PER_VECTOR: u32 = 0;
20pub const INT8_PREFIX_SIZE: usize = 8;
21
22pub fn quantize_f32_to_i8(values: &[f32]) -> (f32, Vec<i8>) {
29 let max_abs = values.iter().fold(0.0f32, |acc, &v| acc.max(v.abs()));
30 if max_abs == 0.0 {
31 return (1.0, vec![0i8; values.len()]);
34 }
35 let scale = max_abs / 127.0;
36 let inv_scale = 1.0 / scale;
37 let q: Vec<i8> = values
38 .iter()
39 .map(|&v| {
40 let scaled = (v * inv_scale).round();
41 scaled.clamp(-127.0, 127.0) as i8
42 })
43 .collect();
44 (scale, q)
45}
46
47pub fn encode_int8_embeddings(embeddings: &[f32], n: usize, dim: usize) -> crate::Result<Vec<u8>> {
53 if embeddings.len() != n * dim {
54 return Err(UrnaError::InvalidInput(format!(
55 "encode_int8_embeddings: got {} f32 values for n={} dim={}",
56 embeddings.len(),
57 n,
58 dim
59 )));
60 }
61 let mut out = Vec::with_capacity(INT8_PREFIX_SIZE + n * 4 + n * dim);
62 out.extend_from_slice(&INT8_PAYLOAD_VERSION.to_le_bytes());
63 out.extend_from_slice(&INT8_SCALE_KIND_PER_VECTOR.to_le_bytes());
64 let mut scales: Vec<u8> = Vec::with_capacity(n * 4);
65 let mut bodies: Vec<u8> = Vec::with_capacity(n * dim);
66 for i in 0..n {
67 let row = &embeddings[i * dim..(i + 1) * dim];
68 let (scale, q) = quantize_f32_to_i8(row);
69 scales.extend_from_slice(&scale.to_le_bytes());
70 bodies.extend(q.iter().map(|&v| v as u8));
72 }
73 out.extend_from_slice(&scales);
74 out.extend_from_slice(&bodies);
75 Ok(out)
76}
77
78pub struct Int8EmbeddingsView<'a> {
81 pub scales: &'a [u8], pub bodies: &'a [u8], pub n: usize,
84 pub dim: usize,
85}
86
87impl<'a> Int8EmbeddingsView<'a> {
88 pub fn parse(bytes: &'a [u8], n: usize, dim: usize) -> crate::Result<Self> {
89 let want = super::expected_embeddings_size("int8", n, dim).unwrap_or(usize::MAX);
92 if bytes.len() != want {
93 return Err(UrnaError::EmbeddingSizeMismatch {
94 expected: want,
95 got: bytes.len(),
96 });
97 }
98 let version = le_u32(&bytes[0..4])?;
99 if version != INT8_PAYLOAD_VERSION {
100 return Err(UrnaError::UnsupportedSectionVersion {
101 section_id: crate::layout::SECTION_EMBEDDINGS,
102 version,
103 });
104 }
105 let kind = le_u32(&bytes[4..8])?;
106 if kind != INT8_SCALE_KIND_PER_VECTOR {
107 return Err(UrnaError::MalformedSectionPayload {
108 section_id: crate::layout::SECTION_EMBEDDINGS,
109 reason: format!("int8 scale_kind {} not supported", kind),
110 });
111 }
112 let scales_end = INT8_PREFIX_SIZE + n * 4;
113 Ok(Self {
114 scales: &bytes[INT8_PREFIX_SIZE..scales_end],
115 bodies: &bytes[scales_end..],
116 n,
117 dim,
118 })
119 }
120
121 #[inline]
123 pub fn scale(&self, i: usize) -> f32 {
124 let off = i * 4;
125 f32::from_le_bytes([
126 self.scales[off],
127 self.scales[off + 1],
128 self.scales[off + 2],
129 self.scales[off + 3],
130 ])
131 }
132
133 #[inline]
135 pub fn row(&self, i: usize) -> &'a [i8] {
136 let start = i * self.dim;
137 let end = start + self.dim;
138 bytemuck::cast_slice(&self.bodies[start..end])
141 }
142}