urna_format/encoding/
int4.rs1use crate::bytes::le_u32;
22use crate::error::UrnaError;
23
24pub const INT4_PAYLOAD_VERSION: u32 = 1;
25pub const INT4_SCALE_KIND_PER_GROUP: u32 = 1;
26pub const INT4_PREFIX_SIZE: usize = 8;
27pub const INT4_BLOCK: usize = 64;
29
30#[inline]
32pub fn int4_blocks_per_row(dim: usize) -> usize {
33 dim / INT4_BLOCK
34}
35
36pub fn quantize_f32_to_i4(values: &[f32], dim: usize) -> (Vec<half::f16>, Vec<i8>) {
42 let blocks = int4_blocks_per_row(dim);
43 let mut scales: Vec<half::f16> = Vec::with_capacity(blocks);
44 let mut codes: Vec<i8> = Vec::with_capacity(dim);
45 for g in 0..blocks {
46 let blk = &values[g * INT4_BLOCK..(g + 1) * INT4_BLOCK];
47 let max_abs = blk.iter().fold(0.0f32, |acc, &v| acc.max(v.abs()));
48 let scale_f16 = half::f16::from_f32(if max_abs == 0.0 { 1.0 } else { max_abs / 7.0 });
51 let scale = scale_f16.to_f32();
52 let inv = if scale > 0.0 { 1.0 / scale } else { 0.0 };
53 for &v in blk {
54 let q = (v * inv).round().clamp(-7.0, 7.0);
55 codes.push(q as i8);
56 }
57 scales.push(scale_f16);
58 }
59 (scales, codes)
60}
61
62#[inline]
66pub fn pack_nibbles(codes: &[i8]) -> Vec<u8> {
67 let mut out = Vec::with_capacity(codes.len().div_ceil(2));
68 for pair in codes.chunks(2) {
69 let lo = (pair[0] as u8) & 0x0F;
70 let hi = pair.get(1).map(|&c| (c as u8) & 0x0F).unwrap_or(0);
71 out.push(lo | (hi << 4));
72 }
73 out
74}
75
76#[inline]
78pub fn nibble_to_i4(b: u8) -> i8 {
79 let n = b & 0x0F;
80 if n & 0x08 != 0 {
81 (n | 0xF0) as i8
82 } else {
83 n as i8
84 }
85}
86
87pub fn encode_int4_embeddings(embeddings: &[f32], n: usize, dim: usize) -> crate::Result<Vec<u8>> {
90 if dim == 0 || dim % INT4_BLOCK != 0 {
91 return Err(UrnaError::InvalidInput(format!(
92 "encode_int4_embeddings: dim={dim} must be a nonzero multiple of {INT4_BLOCK}"
93 )));
94 }
95 if embeddings.len() != n * dim {
96 return Err(UrnaError::InvalidInput(format!(
97 "encode_int4_embeddings: got {} f32 values for n={n} dim={dim}",
98 embeddings.len()
99 )));
100 }
101 let blocks = int4_blocks_per_row(dim);
102 let mut out = Vec::with_capacity(INT4_PREFIX_SIZE + n * blocks * 2 + n * dim / 2);
103 out.extend_from_slice(&INT4_PAYLOAD_VERSION.to_le_bytes());
104 out.extend_from_slice(&INT4_SCALE_KIND_PER_GROUP.to_le_bytes());
105 let mut scale_bytes: Vec<u8> = Vec::with_capacity(n * blocks * 2);
106 let mut body: Vec<u8> = Vec::with_capacity(n * dim / 2);
107 for i in 0..n {
108 let row = &embeddings[i * dim..(i + 1) * dim];
109 let (scales, codes) = quantize_f32_to_i4(row, dim);
110 for s in &scales {
111 scale_bytes.extend_from_slice(&s.to_le_bytes());
112 }
113 body.extend_from_slice(&pack_nibbles(&codes));
114 }
115 out.extend_from_slice(&scale_bytes);
116 out.extend_from_slice(&body);
117 Ok(out)
118}
119
120pub struct Int4EmbeddingsView<'a> {
123 pub scales: &'a [u8],
125 pub codes: &'a [u8],
127 pub n: usize,
128 pub dim: usize,
129 pub blocks: usize,
130}
131
132impl<'a> Int4EmbeddingsView<'a> {
133 pub fn parse(bytes: &'a [u8], n: usize, dim: usize) -> crate::Result<Self> {
134 if dim == 0 || dim % INT4_BLOCK != 0 {
135 return Err(UrnaError::MalformedSectionPayload {
136 section_id: crate::layout::SECTION_EMBEDDINGS,
137 reason: format!("int4 dim={dim} must be a nonzero multiple of {INT4_BLOCK}"),
138 });
139 }
140 let blocks = int4_blocks_per_row(dim);
141 let want = super::expected_embeddings_size("int4", n, dim).unwrap_or(usize::MAX);
144 if bytes.len() != want {
145 return Err(UrnaError::EmbeddingSizeMismatch {
146 expected: want,
147 got: bytes.len(),
148 });
149 }
150 let version = le_u32(&bytes[0..4])?;
151 if version != INT4_PAYLOAD_VERSION {
152 return Err(UrnaError::UnsupportedSectionVersion {
153 section_id: crate::layout::SECTION_EMBEDDINGS,
154 version,
155 });
156 }
157 let kind = le_u32(&bytes[4..8])?;
158 if kind != INT4_SCALE_KIND_PER_GROUP {
159 return Err(UrnaError::MalformedSectionPayload {
160 section_id: crate::layout::SECTION_EMBEDDINGS,
161 reason: format!("int4 scale_kind {kind} not supported"),
162 });
163 }
164 let scales_end = INT4_PREFIX_SIZE + n * blocks * 2;
165 Ok(Self {
166 scales: &bytes[INT4_PREFIX_SIZE..scales_end],
167 codes: &bytes[scales_end..],
168 n,
169 dim,
170 blocks,
171 })
172 }
173
174 #[inline]
176 pub fn group_scale(&self, i: usize, g: usize) -> f32 {
177 let off = (i * self.blocks + g) * 2;
178 half::f16::from_le_bytes([self.scales[off], self.scales[off + 1]]).to_f32()
179 }
180
181 #[inline]
183 pub fn row_codes(&self, i: usize) -> &'a [u8] {
184 let rs = self.dim / 2;
185 let start = i * rs;
186 &self.codes[start..start + rs]
187 }
188
189 #[inline]
193 pub fn row_scales_f32(&self, i: usize) -> Vec<f32> {
194 (0..self.blocks).map(|g| self.group_scale(i, g)).collect()
195 }
196
197 #[inline]
205 pub fn row_scales_into(&self, i: usize, out: &mut [f32]) {
206 assert_eq!(out.len(), self.blocks, "int4: one scale slot per block");
207 for (g, slot) in out.iter_mut().enumerate() {
208 *slot = self.group_scale(i, g);
209 }
210 }
211}
212
213