1use crate::gguf_sharder::GgufTensorInfo;
8
9pub const GGML_TYPE_F32: u32 = 0;
11pub const GGML_TYPE_F16: u32 = 1;
12pub const GGML_TYPE_Q4_0: u32 = 2;
13pub const GGML_TYPE_Q5_0: u32 = 6;
14pub const GGML_TYPE_Q8_0: u32 = 8;
15pub const GGML_TYPE_Q4_K: u32 = 12;
16pub const GGML_TYPE_Q6_K: u32 = 14;
17pub const GGML_TYPE_BF16: u32 = 30;
20pub const GGML_TYPE_Q4_K_SOA: u32 = 112;
30pub const BLOCK_Q4K_SOA_BYTES: usize = 160;
32pub const BLOCK_Q4K_SOA_ELEMS: usize = 256;
33
34#[repr(C)]
36#[derive(Copy, Clone, bytemuck::Pod, bytemuck::Zeroable)]
37pub struct BlockQ6K {
38 pub ql: [u8; 128],
39 pub qh: [u8; 64],
40 pub scales: [i8; 16],
41 pub d: u16,
42}
43
44pub const BLOCK_Q6K_BYTES: usize = 210;
45pub const BLOCK_Q6K_ELEMS: usize = 256;
46
47#[derive(Debug, Clone, Copy, PartialEq, Eq)]
49pub struct GgmlBlockLayout {
50 pub block_elems: usize,
51 pub block_bytes: usize,
52}
53
54pub fn ggml_block_layout(ggml_type: u32) -> Option<GgmlBlockLayout> {
56 match ggml_type {
57 GGML_TYPE_Q4_0 => Some(GgmlBlockLayout {
58 block_elems: 32,
59 block_bytes: 18,
60 }),
61 GGML_TYPE_Q5_0 => Some(GgmlBlockLayout {
62 block_elems: 32,
63 block_bytes: 22,
64 }),
65 GGML_TYPE_Q8_0 => Some(GgmlBlockLayout {
66 block_elems: 32,
67 block_bytes: 34,
68 }),
69 GGML_TYPE_Q4_K => Some(GgmlBlockLayout {
70 block_elems: 256,
71 block_bytes: 144,
72 }),
73 GGML_TYPE_Q4_K_SOA => Some(GgmlBlockLayout {
74 block_elems: BLOCK_Q4K_SOA_ELEMS,
75 block_bytes: BLOCK_Q4K_SOA_BYTES,
76 }),
77 GGML_TYPE_Q6_K => Some(GgmlBlockLayout {
78 block_elems: 256,
79 block_bytes: 210,
80 }),
81 _ => None,
82 }
83}
84
85pub fn ggml_row_bytes(ggml_type: u32, n_elems: usize) -> Option<usize> {
87 match ggml_type {
88 GGML_TYPE_F32 => Some(n_elems.checked_mul(4)?),
89 GGML_TYPE_F16 | GGML_TYPE_BF16 => Some(n_elems.checked_mul(2)?),
90 _ => {
91 let layout = ggml_block_layout(ggml_type)?;
92 if n_elems == 0 {
93 return Some(0);
94 }
95 Some(n_elems.div_ceil(layout.block_elems) * layout.block_bytes)
96 }
97 }
98}
99
100#[derive(Debug, Clone, Copy, PartialEq, Eq)]
101pub enum GgmlDequantError {
102 UnsupportedType,
103 BufferTooSmall,
104 TruncatedInput,
105}
106
107#[derive(Debug, Clone, Copy, PartialEq, Eq)]
109pub enum ExecutionError {
110 TensorNotFound,
111 TokenOutOfRange,
112 UnsupportedType,
113 MmapBounds,
114}
115
116pub fn fetch_token_embedding<'a>(
118 mmap: &'a [u8],
119 tensor_data_start: u64,
120 tensor: &GgufTensorInfo,
121 token_id: u32,
122) -> Result<&'a [u8], ExecutionError> {
123 let n_embd = tensor.dims[0] as usize;
124 let n_vocab = tensor.dims[1] as usize;
125 if n_embd == 0 {
126 return Err(ExecutionError::TensorNotFound);
127 }
128 if token_id as usize >= n_vocab {
129 return Err(ExecutionError::TokenOutOfRange);
130 }
131 let bytes_per_token =
132 ggml_row_bytes(tensor.ggml_type, n_embd).ok_or(ExecutionError::UnsupportedType)?;
133 let start =
134 (tensor_data_start + tensor.byte_offset) as usize + token_id as usize * bytes_per_token;
135 let end = start + bytes_per_token;
136 if end > mmap.len() {
137 return Err(ExecutionError::MmapBounds);
138 }
139 Ok(&mmap[start..end])
140}
141
142pub fn tensor_byte_len(tensor: &GgufTensorInfo) -> Option<usize> {
144 let n0 = tensor.dims[0] as usize;
145 if n0 == 0 {
146 return None;
147 }
148 if tensor.ggml_type == crate::ternary::GGML_TYPE_TERNARY_158 {
152 let n_elems = if tensor.n_dims > 1 && tensor.dims[1] > 0 {
153 n0.checked_mul(tensor.dims[1] as usize)?
154 } else {
155 n0
156 };
157 return Some(crate::ternary::ternary_blob_len(n_elems));
158 }
159 let row = ggml_row_bytes(tensor.ggml_type, n0)?;
160 if tensor.n_dims <= 1 || tensor.dims[1] == 0 {
161 Some(row)
162 } else {
163 Some(row.checked_mul(tensor.dims[1] as usize)?)
164 }
165}
166
167pub fn fetch_tensor_bytes<'a>(
169 mmap: &'a [u8],
170 tensor_data_start: u64,
171 tensor: &GgufTensorInfo,
172) -> Result<&'a [u8], ExecutionError> {
173 let len = tensor_byte_len(tensor).ok_or(ExecutionError::UnsupportedType)?;
174 let start = (tensor_data_start + tensor.byte_offset) as usize;
175 let end = start + len;
176 if end > mmap.len() {
177 return Err(ExecutionError::MmapBounds);
178 }
179 Ok(&mmap[start..end])
180}
181
182pub fn tensor_row_byte_len(tensor: &GgufTensorInfo) -> Result<usize, ExecutionError> {
184 let n0 = tensor.dims[0] as usize;
185 if n0 == 0 {
186 return Err(ExecutionError::TensorNotFound);
187 }
188 ggml_row_bytes(tensor.ggml_type, n0).ok_or(ExecutionError::UnsupportedType)
189}
190
191pub fn fetch_tensor_row_range_bytes<'a>(
193 mmap: &'a [u8],
194 tensor_data_start: u64,
195 tensor: &GgufTensorInfo,
196 row_start: usize,
197 row_count: usize,
198) -> Result<&'a [u8], ExecutionError> {
199 let row_bytes = tensor_row_byte_len(tensor)?;
200 let n_rows = if tensor.n_dims > 1 && tensor.dims[1] > 0 {
201 tensor.dims[1] as usize
202 } else {
203 1
204 };
205 if row_start >= n_rows || row_count == 0 {
206 return Err(ExecutionError::TokenOutOfRange);
207 }
208 let rows = row_count.min(n_rows - row_start);
209 let start = (tensor_data_start + tensor.byte_offset) as usize + row_start * row_bytes;
210 let end = start + rows * row_bytes;
211 if end > mmap.len() {
212 return Err(ExecutionError::MmapBounds);
213 }
214 Ok(&mmap[start..end])
215}
216
217pub fn dequant_matrix_row_into(
219 raw: &[u8],
220 info: &GgufTensorInfo,
221 row: usize,
222 out: &mut [f32],
223) -> Result<usize, GgmlDequantError> {
224 let n0 = info.dims[0] as usize;
225 let row_bytes = ggml_row_bytes(info.ggml_type, n0).ok_or(GgmlDequantError::UnsupportedType)?;
226 let start = row
227 .checked_mul(row_bytes)
228 .ok_or(GgmlDequantError::TruncatedInput)?;
229 if start + row_bytes > raw.len() {
230 return Err(GgmlDequantError::TruncatedInput);
231 }
232 dequantize_row_into(&raw[start..start + row_bytes], info.ggml_type, n0, out)
233}
234
235pub fn dequantize_row_into(
238 raw: &[u8],
239 ggml_type: u32,
240 n_elems: usize,
241 out: &mut [f32],
242) -> Result<usize, GgmlDequantError> {
243 if out.len() < n_elems {
244 return Err(GgmlDequantError::BufferTooSmall);
245 }
246 match ggml_type {
247 GGML_TYPE_F32 => dequant_f32(raw, n_elems, out),
248 GGML_TYPE_F16 => dequant_f16(raw, n_elems, out),
249 GGML_TYPE_BF16 => dequant_bf16(raw, n_elems, out),
250 GGML_TYPE_Q4_0 => dequant_q4_0(raw, n_elems, out),
251 GGML_TYPE_Q5_0 => dequant_q5_0(raw, n_elems, out),
252 GGML_TYPE_Q8_0 => dequant_q8_0(raw, n_elems, out),
253 GGML_TYPE_Q4_K => dequant_q4_k(raw, n_elems, out),
254 GGML_TYPE_Q4_K_SOA => dequant_q4_k_soa(raw, n_elems, out),
255 GGML_TYPE_Q6_K => dequant_q6_k(raw, n_elems, out),
256 _ => Err(GgmlDequantError::UnsupportedType),
257 }
258}
259
260#[inline]
265pub fn q4k_block_to_soa(src: &[u8], dst: &mut [u8]) -> Result<(), GgmlDequantError> {
266 if src.len() < 144 || dst.len() < BLOCK_Q4K_SOA_BYTES {
267 return Err(GgmlDequantError::TruncatedInput);
268 }
269 dst[..128].copy_from_slice(&src[16..144]);
271 let d = half::f16::from_le_bytes([src[0], src[1]]).to_f32();
272 let dmin = half::f16::from_le_bytes([src[2], src[3]]).to_f32();
273 let scales: [u8; 12] = src[4..16].try_into().unwrap_or([0; 12]);
274 for j in 0..8 {
275 let mut sc = 0u8;
276 let mut m = 0u8;
277 get_scale_min_k4(j, &scales, &mut sc, &mut m);
278 let d_bits = half::f16::from_f32(d * sc as f32).to_le_bytes();
279 let m_bits = half::f16::from_f32(dmin * m as f32).to_le_bytes();
280 let o = 128 + j * 2;
281 dst[o] = d_bits[0];
282 dst[o + 1] = d_bits[1];
283 let om = 144 + j * 2;
284 dst[om] = m_bits[0];
285 dst[om + 1] = m_bits[1];
286 }
287 Ok(())
288}
289
290pub fn expand_q4k_tensor_to_soa(
293 raw: &[u8],
294 n_row_elems: usize,
295 n_rows: usize,
296 out: &mut [u8],
297) -> Result<(), GgmlDequantError> {
298 let src_row =
299 ggml_row_bytes(GGML_TYPE_Q4_K, n_row_elems).ok_or(GgmlDequantError::UnsupportedType)?;
300 let dst_row =
301 ggml_row_bytes(GGML_TYPE_Q4_K_SOA, n_row_elems).ok_or(GgmlDequantError::UnsupportedType)?;
302 let need = dst_row
303 .checked_mul(n_rows)
304 .ok_or(GgmlDequantError::TruncatedInput)?;
305 if out.len() < need || raw.len() < src_row.saturating_mul(n_rows) {
306 return Err(GgmlDequantError::TruncatedInput);
307 }
308 let n_blocks = n_row_elems.div_ceil(256);
309 for r in 0..n_rows {
310 let src_base = r * src_row;
311 let dst_base = r * dst_row;
312 for b in 0..n_blocks {
313 let s = src_base + b * 144;
314 let d = dst_base + b * BLOCK_Q4K_SOA_BYTES;
315 if s + 144 > raw.len() || d + BLOCK_Q4K_SOA_BYTES > out.len() {
316 return Err(GgmlDequantError::TruncatedInput);
317 }
318 q4k_block_to_soa(&raw[s..s + 144], &mut out[d..d + BLOCK_Q4K_SOA_BYTES])?;
319 }
320 }
321 Ok(())
322}
323
324fn dequant_q4_k_soa(
325 raw: &[u8],
326 n_elems: usize,
327 out: &mut [f32],
328) -> Result<usize, GgmlDequantError> {
329 let n_blocks = n_elems.div_ceil(BLOCK_Q4K_SOA_ELEMS);
330 if raw.len() < n_blocks * BLOCK_Q4K_SOA_BYTES {
331 return Err(GgmlDequantError::TruncatedInput);
332 }
333 let mut out_idx = 0usize;
334 for b in 0..n_blocks {
335 let block = &raw[b * BLOCK_Q4K_SOA_BYTES..(b + 1) * BLOCK_Q4K_SOA_BYTES];
336 let qs = &block[0..128];
337 let mut d_sub = [0f32; 8];
339 let mut m_sub = [0f32; 8];
340 for j in 0..8 {
341 let o = 128 + j * 2;
342 d_sub[j] = half::f16::from_le_bytes([block[o], block[o + 1]]).to_f32();
343 let om = 144 + j * 2;
344 m_sub[j] = half::f16::from_le_bytes([block[om], block[om + 1]]).to_f32();
345 }
346 let mut q_off = 0usize;
349 for g in 0..4 {
350 let sub0 = g * 2;
351 let sub1 = g * 2 + 1;
352 for l in 0..32 {
353 if out_idx >= n_elems {
354 return Ok(out_idx);
355 }
356 let nib = (qs[q_off + l] & 0xF) as f32;
357 out[out_idx] = d_sub[sub0] * nib - m_sub[sub0];
358 out_idx += 1;
359 }
360 for l in 0..32 {
361 if out_idx >= n_elems {
362 return Ok(out_idx);
363 }
364 let nib = (qs[q_off + l] >> 4) as f32;
365 out[out_idx] = d_sub[sub1] * nib - m_sub[sub1];
366 out_idx += 1;
367 }
368 q_off += 32;
369 }
370 }
371 Ok(out_idx.min(n_elems))
372}
373
374fn dequant_f32(raw: &[u8], n_elems: usize, out: &mut [f32]) -> Result<usize, GgmlDequantError> {
375 let need = n_elems * 4;
376 if raw.len() < need {
377 return Err(GgmlDequantError::TruncatedInput);
378 }
379 for i in 0..n_elems {
380 out[i] = f32::from_le_bytes(raw[i * 4..i * 4 + 4].try_into().unwrap_or([0; 4]));
381 }
382 Ok(n_elems)
383}
384
385fn dequant_f16(raw: &[u8], n_elems: usize, out: &mut [f32]) -> Result<usize, GgmlDequantError> {
386 let need = n_elems * 2;
387 if raw.len() < need {
388 return Err(GgmlDequantError::TruncatedInput);
389 }
390 for i in 0..n_elems {
391 out[i] =
392 half::f16::from_le_bytes(raw[i * 2..i * 2 + 2].try_into().unwrap_or([0; 2])).to_f32();
393 }
394 Ok(n_elems)
395}
396
397fn dequant_bf16(raw: &[u8], n_elems: usize, out: &mut [f32]) -> Result<usize, GgmlDequantError> {
399 let need = n_elems * 2;
400 if raw.len() < need {
401 return Err(GgmlDequantError::TruncatedInput);
402 }
403 for i in 0..n_elems {
404 let bits = u16::from_le_bytes(raw[i * 2..i * 2 + 2].try_into().unwrap_or([0; 2]));
405 out[i] = f32::from_bits((bits as u32) << 16);
406 }
407 Ok(n_elems)
408}
409
410fn dequant_q4_0(raw: &[u8], n_elems: usize, out: &mut [f32]) -> Result<usize, GgmlDequantError> {
411 const BLOCK_ELEMS: usize = 32;
412 const BLOCK_BYTES: usize = 18;
413 let n_blocks = n_elems.div_ceil(BLOCK_ELEMS);
414 if raw.len() < n_blocks * BLOCK_BYTES {
415 return Err(GgmlDequantError::TruncatedInput);
416 }
417 for b in 0..n_blocks {
418 let bs = b * BLOCK_BYTES;
419 let scale = half::f16::from_le_bytes([raw[bs], raw[bs + 1]]).to_f32();
420 let half = BLOCK_ELEMS / 2;
421 for j in 0..half {
422 if b * BLOCK_ELEMS + j >= n_elems {
423 break;
424 }
425 let byte = raw[bs + 2 + j];
426 let x0 = (byte & 0x0F) as i32 - 8;
427 let x1 = ((byte >> 4) & 0x0F) as i32 - 8;
428 out[b * BLOCK_ELEMS + j] = x0 as f32 * scale;
429 let hi = b * BLOCK_ELEMS + j + half;
430 if hi < n_elems {
431 out[hi] = x1 as f32 * scale;
432 }
433 }
434 }
435 Ok(n_elems)
436}
437
438fn dequant_q5_0(raw: &[u8], n_elems: usize, out: &mut [f32]) -> Result<usize, GgmlDequantError> {
440 const BLOCK_ELEMS: usize = 32;
441 const BLOCK_BYTES: usize = 22;
442 let n_blocks = n_elems.div_ceil(BLOCK_ELEMS);
443 if raw.len() < n_blocks * BLOCK_BYTES {
444 return Err(GgmlDequantError::TruncatedInput);
445 }
446 for b in 0..n_blocks {
447 let bs = b * BLOCK_BYTES;
448 let d = half::f16::from_le_bytes([raw[bs], raw[bs + 1]]).to_f32();
449 let qh = u32::from_le_bytes([raw[bs + 2], raw[bs + 3], raw[bs + 4], raw[bs + 5]]);
450 let qs = &raw[bs + 6..bs + 22];
451 let half = BLOCK_ELEMS / 2;
452 for j in 0..half {
453 let xh_0 = ((qh >> j) << 4) & 0x10;
454 let xh_1 = (qh >> (j + 12)) & 0x10;
455 let x0 = ((qs[j] & 0x0F) as u32 | xh_0) as i32 - 16;
456 let x1 = ((qs[j] >> 4) as u32 | xh_1) as i32 - 16;
457 let lo = b * BLOCK_ELEMS + j;
458 if lo < n_elems {
459 out[lo] = x0 as f32 * d;
460 }
461 let hi = lo + half;
462 if hi < n_elems {
463 out[hi] = x1 as f32 * d;
464 }
465 }
466 }
467 Ok(n_elems)
468}
469
470fn dequant_q8_0(raw: &[u8], n_elems: usize, out: &mut [f32]) -> Result<usize, GgmlDequantError> {
471 const BLOCK_ELEMS: usize = 32;
472 const BLOCK_BYTES: usize = 34;
473 let n_blocks = n_elems.div_ceil(BLOCK_ELEMS);
474 if raw.len() < n_blocks * BLOCK_BYTES {
475 return Err(GgmlDequantError::TruncatedInput);
476 }
477 for b in 0..n_blocks {
478 let bs = b * BLOCK_BYTES;
479 let scale = half::f16::from_le_bytes([raw[bs], raw[bs + 1]]).to_f32();
480 let elems = BLOCK_ELEMS.min(n_elems - b * BLOCK_ELEMS);
481 for j in 0..elems {
482 out[b * BLOCK_ELEMS + j] = raw[bs + 2 + j] as i8 as f32 * scale;
483 }
484 }
485 Ok(n_elems)
486}
487
488#[inline]
490fn get_scale_min_k4(j: usize, scales: &[u8; 12], sc: &mut u8, m: &mut u8) {
491 if j < 4 {
492 *sc = scales[j] & 63;
493 *m = scales[j + 4] & 63;
494 } else {
495 *sc = (scales[j + 4] & 0xF) | ((scales[j - 4] >> 6) << 4);
496 *m = (scales[j + 4] >> 4) | ((scales[j] >> 6) << 4);
497 }
498}
499
500fn dequant_q4_k(raw: &[u8], n_elems: usize, out: &mut [f32]) -> Result<usize, GgmlDequantError> {
501 const BLOCK_ELEMS: usize = 256;
502 const BLOCK_BYTES: usize = 144;
503 let n_blocks = n_elems.div_ceil(BLOCK_ELEMS);
504 if raw.len() < n_blocks * BLOCK_BYTES {
505 return Err(GgmlDequantError::TruncatedInput);
506 }
507
508 let mut out_idx = 0usize;
509 for b in 0..n_blocks {
510 let block = &raw[b * BLOCK_BYTES..b * BLOCK_BYTES + BLOCK_BYTES];
511 let d = half::f16::from_le_bytes([block[0], block[1]]).to_f32();
512 let dmin = half::f16::from_le_bytes([block[2], block[3]]).to_f32();
513 let scales: [u8; 12] = block[4..16].try_into().unwrap_or([0; 12]);
514 let qs = &block[16..144];
515
516 let block_elems = BLOCK_ELEMS.min(n_elems - b * BLOCK_ELEMS);
517 let mut q_off = 0usize;
518 let mut is = 0usize;
519 let mut j = 0usize;
520 while j < block_elems && out_idx < n_elems {
521 let mut sc = 0u8;
522 let mut m = 0u8;
523 get_scale_min_k4(is, &scales, &mut sc, &mut m);
524 let d1 = d * sc as f32;
525 let m1 = dmin * m as f32;
526 get_scale_min_k4(is + 1, &scales, &mut sc, &mut m);
527 let d2 = d * sc as f32;
528 let m2 = dmin * m as f32;
529
530 for l in 0..32 {
531 if out_idx >= n_elems || j >= block_elems {
532 break;
533 }
534 out[out_idx] = d1 * (qs[q_off + l] & 0xF) as f32 - m1;
535 out_idx += 1;
536 j += 1;
537 }
538 for l in 0..32 {
539 if out_idx >= n_elems || j >= block_elems {
540 break;
541 }
542 out[out_idx] = d2 * (qs[q_off + l] >> 4) as f32 - m2;
543 out_idx += 1;
544 j += 1;
545 }
546 q_off += 32;
547 is += 2;
548 }
549 }
550 Ok(out_idx.min(n_elems))
551}
552
553fn dequant_q6_k_block(block: &[u8; 210], out: &mut [f32]) {
554 let blk = bytemuck::from_bytes::<BlockQ6K>(block);
555 let d = half::f16::from_bits(blk.d).to_f32();
556 let mut ql_off = 0usize;
557 let mut qh_off = 0usize;
558 let mut sc_off = 0usize;
559 let mut y_off = 0usize;
560
561 for _ in 0..2 {
562 for l in 0..32 {
563 let is = l / 16;
564 let q1 =
565 ((blk.ql[ql_off + l] & 0xF) | (((blk.qh[qh_off + l] >> 0) & 3) << 4)) as i8 - 32;
566 let q2 = ((blk.ql[ql_off + l + 32] & 0xF) | (((blk.qh[qh_off + l] >> 2) & 3) << 4))
567 as i8
568 - 32;
569 let q3 =
570 ((blk.ql[ql_off + l] >> 4) | (((blk.qh[qh_off + l] >> 4) & 3) << 4)) as i8 - 32;
571 let q4 = ((blk.ql[ql_off + l + 32] >> 4) | (((blk.qh[qh_off + l] >> 6) & 3) << 4))
572 as i8
573 - 32;
574 let sc = &blk.scales[sc_off..sc_off + 8];
575 out[y_off + l] = d * sc[is] as f32 * q1 as f32;
576 out[y_off + l + 32] = d * sc[is + 2] as f32 * q2 as f32;
577 out[y_off + l + 64] = d * sc[is + 4] as f32 * q3 as f32;
578 out[y_off + l + 96] = d * sc[is + 6] as f32 * q4 as f32;
579 }
580 y_off += 128;
581 ql_off += 64;
582 qh_off += 32;
583 sc_off += 8;
584 }
585}
586
587fn dequant_q6_k(raw: &[u8], n_elems: usize, out: &mut [f32]) -> Result<usize, GgmlDequantError> {
588 const BLOCK_ELEMS: usize = 256;
589 const BLOCK_BYTES: usize = 210;
590 let n_blocks = n_elems.div_ceil(BLOCK_ELEMS);
591 if raw.len() < n_blocks * BLOCK_BYTES {
592 return Err(GgmlDequantError::TruncatedInput);
593 }
594
595 let mut written = 0usize;
596 for b in 0..n_blocks {
597 let block: &[u8; 210] = raw[b * BLOCK_BYTES..b * BLOCK_BYTES + BLOCK_BYTES]
598 .try_into()
599 .map_err(|_| GgmlDequantError::TruncatedInput)?;
600 let elems = BLOCK_ELEMS.min(n_elems - written);
601 let mut block_out = [0f32; BLOCK_ELEMS];
602 dequant_q6_k_block(block, &mut block_out);
603 out[written..written + elems].copy_from_slice(&block_out[..elems]);
604 written += elems;
605 }
606 Ok(written)
607}
608
609fn quantize_block_f32_to_q4_k_soa(src: &[f32], out: &mut [u8]) {
620 debug_assert!(out.len() >= BLOCK_Q4K_SOA_BYTES);
621 out[..128].fill(0);
623
624 for j in 0..8 {
625 let sub_start = j * 32;
626 let sub_end = (sub_start + 32).min(src.len());
627
628 let mut min_val = 0.0f32;
630 let mut max_val = 0.0f32;
631 if sub_start < sub_end {
632 min_val = src[sub_start];
633 max_val = src[sub_start];
634 for i in sub_start + 1..sub_end {
635 let w = src[i];
636 if w < min_val {
637 min_val = w;
638 }
639 if w > max_val {
640 max_val = w;
641 }
642 }
643 }
644
645 let scale = if max_val > min_val {
646 (max_val - min_val) / 15.0
647 } else {
648 1.0
649 };
650 let d_sub = half::f16::from_f32(scale);
654 let m_sub = half::f16::from_f32(-min_val);
655 let d_actual = d_sub.to_f32();
656 let m_actual = m_sub.to_f32();
657
658 let d_bytes = d_sub.to_le_bytes();
660 let m_bytes = m_sub.to_le_bytes();
661 out[128 + j * 2] = d_bytes[0];
662 out[128 + j * 2 + 1] = d_bytes[1];
663 out[144 + j * 2] = m_bytes[0];
664 out[144 + j * 2 + 1] = m_bytes[1];
665
666 let g = j / 2;
669 let q_off = g * 32;
670 let is_high = j % 2 == 1;
671 for l in 0..32 {
672 let nibble: u8 = if sub_start + l < sub_end {
673 let w = src[sub_start + l];
674 let q = ((w + m_actual) / d_actual).round();
675 q.clamp(0.0, 15.0) as u8
676 } else {
677 0
678 };
679 let bi = q_off + l;
680 if is_high {
681 out[bi] = (out[bi] & 0x0F) | (nibble << 4);
682 } else {
683 out[bi] = (out[bi] & 0xF0) | (nibble & 0x0F);
684 }
685 }
686 }
687}
688
689pub fn quantize_f32_to_q4_k_soa_tensor(
694 src: &[f32],
695 n_row_elems: usize,
696 n_rows: usize,
697 out: &mut [u8],
698) -> Result<(), GgmlDequantError> {
699 let dst_row =
700 ggml_row_bytes(GGML_TYPE_Q4_K_SOA, n_row_elems).ok_or(GgmlDequantError::UnsupportedType)?;
701 let need = dst_row
702 .checked_mul(n_rows)
703 .ok_or(GgmlDequantError::TruncatedInput)?;
704 if out.len() < need {
705 return Err(GgmlDequantError::BufferTooSmall);
706 }
707 if src.len() < n_row_elems * n_rows {
708 return Err(GgmlDequantError::TruncatedInput);
709 }
710
711 let n_blocks = n_row_elems.div_ceil(BLOCK_Q4K_SOA_ELEMS);
712 for r in 0..n_rows {
713 let src_base = r * n_row_elems;
714 let dst_base = r * dst_row;
715 for b in 0..n_blocks {
716 let block_start = src_base + b * BLOCK_Q4K_SOA_ELEMS;
717 let block_end = (block_start + BLOCK_Q4K_SOA_ELEMS).min(src_base + n_row_elems);
718 let block_src = &src[block_start..block_end];
719 let dst_off = dst_base + b * BLOCK_Q4K_SOA_BYTES;
720 quantize_block_f32_to_q4_k_soa(
721 block_src,
722 &mut out[dst_off..dst_off + BLOCK_Q4K_SOA_BYTES],
723 );
724 }
725 }
726 Ok(())
727}
728
729#[cfg(test)]
730mod tests {
731 use super::*;
732
733 #[test]
734 fn q4k_soa_roundtrip_matches_aos_dequant() {
735 let mut aos = [0u8; 144];
737 aos[0..2].copy_from_slice(&half::f16::from_f32(1.0).to_le_bytes());
739 aos[2..4].copy_from_slice(&half::f16::from_f32(0.1).to_le_bytes());
740 for i in 0..4 {
742 aos[4 + i] = 1;
743 aos[8 + i] = 0;
744 }
745 for i in 0..4 {
746 aos[12 + i] = 0;
747 }
748 for i in 16..144 {
750 aos[i] = 0x53;
751 }
752 let mut soa = [0u8; BLOCK_Q4K_SOA_BYTES];
753 q4k_block_to_soa(&aos, &mut soa).unwrap();
754 let mut out_aos = [0f32; 256];
755 let mut out_soa = [0f32; 256];
756 dequant_q4_k(&aos, 256, &mut out_aos).unwrap();
757 dequant_q4_k_soa(&soa, 256, &mut out_soa).unwrap();
758 for i in 0..256 {
759 let d = (out_aos[i] - out_soa[i]).abs();
760 assert!(
761 d < 1e-3,
762 "elem {i}: aos={} soa={} δ={d}",
763 out_aos[i],
764 out_soa[i]
765 );
766 }
767 assert_eq!(ggml_row_bytes(GGML_TYPE_Q4_K_SOA, 256), Some(160));
768 assert_eq!(ggml_row_bytes(GGML_TYPE_Q4_K_SOA, 512), Some(320));
769 }
770
771 #[test]
772 fn quantize_f32_to_soa_roundtrip() {
773 let mut src = [0f32; 256];
775 for j in 0..8 {
776 let base = j as f32 * 0.1 - 0.4;
777 let amp = (j as f32 + 1.0) * 0.05;
778 for l in 0..32 {
779 src[j * 32 + l] = base + amp * (l as f32 / 31.0);
780 }
781 }
782 let mut soa = [0u8; BLOCK_Q4K_SOA_BYTES];
783 quantize_block_f32_to_q4_k_soa(&src, &mut soa);
784 let mut deq = [0f32; 256];
785 dequant_q4_k_soa(&soa, 256, &mut deq).unwrap();
786 for j in 0..8 {
789 let sub_min = src[j * 32..j * 32 + 32]
790 .iter()
791 .cloned()
792 .fold(f32::MAX, f32::min);
793 let sub_max = src[j * 32..j * 32 + 32]
794 .iter()
795 .cloned()
796 .fold(f32::MIN, f32::max);
797 let step = (sub_max - sub_min) / 15.0;
798 for l in 0..32 {
799 let i = j * 32 + l;
800 let err = (src[i] - deq[i]).abs();
801 assert!(
802 err <= step + 1e-4,
803 "elem {i}: src={} deq={} err={} step={}",
804 src[i],
805 deq[i],
806 err,
807 step
808 );
809 }
810 }
811 }
812
813 #[test]
814 fn quantize_f32_to_soa_tensor_roundtrip() {
815 let n_row_elems = 512;
817 let n_rows = 2;
818 let mut src = vec![0f32; n_row_elems * n_rows];
819 for i in 0..src.len() {
820 src[i] = ((i as f32) * 0.01 - 5.0).sin() * 0.5;
821 }
822 let dst_row = ggml_row_bytes(GGML_TYPE_Q4_K_SOA, n_row_elems).unwrap();
823 let mut out = vec![0u8; dst_row * n_rows];
824 quantize_f32_to_q4_k_soa_tensor(&src, n_row_elems, n_rows, &mut out).unwrap();
825 let mut deq = vec![0f32; n_row_elems * n_rows];
826 for r in 0..n_rows {
827 let off = r * dst_row;
828 dequant_q4_k_soa(
829 &out[off..off + dst_row],
830 n_row_elems,
831 &mut deq[r * n_row_elems..(r + 1) * n_row_elems],
832 )
833 .unwrap();
834 }
835 let mut total_err = 0.0;
837 for i in 0..src.len() {
838 total_err += (src[i] - deq[i]).abs();
839 }
840 let avg_err = total_err / src.len() as f32;
841 assert!(avg_err < 0.05, "avg quantization error too high: {avg_err}");
842 }
843
844 #[test]
845 fn bf16_row_bytes_and_dequant() {
846 assert_eq!(ggml_row_bytes(GGML_TYPE_BF16, 1024), Some(2048));
848 let raw: [u8; 4] = [0x80, 0x3F, 0x00, 0xC0];
850 let mut out = [0.0f32; 2];
851 assert_eq!(
852 dequantize_row_into(&raw, GGML_TYPE_BF16, 2, &mut out),
853 Ok(2)
854 );
855 assert!((out[0] - 1.0).abs() < 1e-6, "got {}", out[0]);
856 assert!((out[1] + 2.0).abs() < 1e-6, "got {}", out[1]);
857 let info = GgufTensorInfo {
858 dims: [128, 1, 0, 0],
859 n_dims: 2,
860 ggml_type: GGML_TYPE_BF16,
861 byte_offset: 0,
862 };
863 assert_eq!(tensor_byte_len(&info), Some(256));
864 }
865
866 #[test]
867 fn q4_0_row_bytes_stride() {
868 assert_eq!(ggml_row_bytes(GGML_TYPE_Q4_0, 4096), Some(2304));
870 }
871
872 #[test]
873 fn q5_0_row_bytes_stride() {
874 assert_eq!(ggml_row_bytes(GGML_TYPE_Q5_0, 960), Some(660));
876 }
877
878 #[test]
879 fn ternary_tensor_byte_len_is_whole_blob_not_row_strided() {
880 let info = GgufTensorInfo {
884 dims: [960, 2560, 0, 0], n_dims: 2,
886 ggml_type: crate::ternary::GGML_TYPE_TERNARY_158,
887 byte_offset: 0,
888 };
889 let n = 960 * 2560;
890 assert_eq!(
891 tensor_byte_len(&info),
892 Some(crate::ternary::ternary_blob_len(n))
893 );
894
895 let weights: Vec<f32> = (0..n).map(|i| (i as f32 * 0.001).sin()).collect();
897 let blob = crate::ternary::ternary_blob(&weights);
898 assert_eq!(blob.len(), crate::ternary::ternary_blob_len(n));
899 let got = fetch_tensor_bytes(&blob, 0, &info).expect("ternary fetch must slice");
900 assert_eq!(got.len(), blob.len());
901 assert_eq!(&got[..4], &blob[..4]); }
903
904 #[test]
905 fn q5_0_dequant_matches_gguf_smollm2_row0() {
906 let path = std::path::Path::new(env!("CARGO_MANIFEST_DIR"))
907 .join("../../docs/models/SmolLM2-360M-Instruct-Q4_K_M.gguf");
908 if !path.exists() {
909 return;
910 }
911 let mmap = std::fs::read(&path).expect("read gguf");
912 let index = crate::gguf_sharder::GgufTensorIndex::from_gguf(&mmap);
913 let info = index.get_layer_tensors(0).attn_k.expect("blk.0.attn_k");
914 let raw = fetch_tensor_bytes(&mmap, index.tensor_data_start, &info).expect("fetch attn_k");
915 let row_bytes = ggml_row_bytes(GGML_TYPE_Q5_0, info.dims[0] as usize).unwrap();
916 assert_eq!(row_bytes, 660);
917 let mut out = [0f32; 960];
918 dequantize_row_into(&raw[..row_bytes], GGML_TYPE_Q5_0, 960, &mut out).unwrap();
919 let expected = [
921 -0.0f32,
922 -0.02502441,
923 -0.10009766,
924 -0.07507324,
925 -0.05004883,
926 -0.3503418,
927 -0.0,
928 0.05004883,
929 -0.07507324,
930 0.1751709,
931 ];
932 for (i, &exp) in expected.iter().enumerate() {
933 assert!(
934 (out[i] - exp).abs() < 1e-5,
935 "elem {i}: got {} expected {}",
936 out[i],
937 exp
938 );
939 }
940 }
941
942 #[test]
943 fn smollm_q4km_tensor_types_and_q4k_row0() {
944 let path = std::path::Path::new(env!("CARGO_MANIFEST_DIR"))
945 .join("../../docs/models/SmolLM2-360M-Instruct-Q4_K_M.gguf");
946 if !path.exists() {
947 return;
948 }
949 let mmap = std::fs::read(&path).expect("read gguf");
950 let index = crate::gguf_sharder::GgufTensorIndex::from_gguf(&mmap);
951 let emb = index.token_embd_info().expect("token_embd");
952 println!("token_embd type={} dims={:?}", emb.ggml_type, emb.dims);
953 let lt = index.get_layer_tensors(0);
954 for (name, info) in [
955 ("attn_q", lt.attn_q),
956 ("attn_k", lt.attn_k),
957 ("attn_v", lt.attn_v),
958 ("ffn_gate", lt.ffn_gate),
959 ("ffn_down", lt.ffn_down),
960 ] {
961 if let Some(i) = info {
962 println!("{name} type={} dims={:?}", i.ggml_type, i.dims);
963 }
964 }
965 let mut out = [0f32; 960];
967 let n = index.dequantize_token_embedding_into(&mmap, 504, &mut out);
968 assert_eq!(n, 960);
969 println!(
970 "token504 first10: {:?}",
971 &out[..10]
972 .iter()
973 .map(|v| format!("{v:.6}"))
974 .collect::<Vec<_>>()
975 );
976 assert!(out.iter().any(|&v| v.is_finite() && v != 0.0));
977 }
978
979 #[test]
980 fn q4_k_row_bytes_stride() {
981 assert_eq!(ggml_row_bytes(GGML_TYPE_Q4_K, 2560), Some(1440));
983 }
984
985 #[test]
986 fn q6_k_row_bytes_stride() {
987 assert_eq!(ggml_row_bytes(GGML_TYPE_Q6_K, 2560), Some(2100));
989 }
990
991 #[test]
992 fn q6_k_dequant_matches_gguf_smollm2_ffn_down_row0() {
993 let path = std::path::Path::new(env!("CARGO_MANIFEST_DIR"))
994 .join("../../docs/models/SmolLM2-360M-Instruct-Q4_K_M.gguf");
995 if !path.exists() {
996 return;
997 }
998 let mmap = std::fs::read(&path).expect("read gguf");
999 let index = crate::gguf_sharder::GgufTensorIndex::from_gguf(&mmap);
1000 let info = index.get_layer_tensors(0).ffn_down.expect("blk.0.ffn_down");
1001 assert_eq!(info.ggml_type, GGML_TYPE_Q6_K);
1002 let raw =
1003 fetch_tensor_bytes(&mmap, index.tensor_data_start, &info).expect("fetch ffn_down");
1004 let row_bytes = ggml_row_bytes(GGML_TYPE_Q6_K, info.dims[0] as usize).unwrap();
1005 assert_eq!(row_bytes, 2100);
1006 let mut out = [0f32; 2560];
1007 dequantize_row_into(&raw[..row_bytes], GGML_TYPE_Q6_K, 2560, &mut out).unwrap();
1008 let expected = [
1010 -0.11712998f32,
1011 -0.16535997,
1012 -0.0,
1013 0.15157998,
1014 -0.08956999,
1015 -0.02067,
1016 0.04133999,
1017 0.03445,
1018 0.17913999,
1019 0.22047997,
1020 ];
1021 for (i, &exp) in expected.iter().enumerate() {
1022 assert!(
1023 (out[i] - exp).abs() < 1e-5,
1024 "elem {i}: got {} expected {}",
1025 out[i],
1026 exp
1027 );
1028 }
1029 assert!(out.iter().any(|&v| v.is_finite() && v != 0.0));
1030 }
1031
1032 #[test]
1033 fn block_q6k_layout_matches_ggml() {
1034 assert_eq!(std::mem::size_of::<BlockQ6K>(), BLOCK_Q6K_BYTES);
1035 assert_eq!(std::mem::align_of::<BlockQ6K>(), 2);
1036 }
1037
1038 #[test]
1039 fn fetch_row_range_bounds() {
1040 use crate::gguf_sharder::GgufTensorInfo;
1041 let info = GgufTensorInfo {
1042 dims: [2560, 100, 0, 0],
1043 n_dims: 2,
1044 ggml_type: GGML_TYPE_Q6_K,
1045 byte_offset: 0,
1046 };
1047 let row = tensor_row_byte_len(&info).unwrap();
1048 let total = tensor_byte_len(&info).unwrap();
1049 assert_eq!(total, row * 100);
1050 let fake = vec![0u8; total];
1051 let chunk = fetch_tensor_row_range_bytes(&fake, 0, &info, 10, 8).unwrap();
1052 assert_eq!(chunk.len(), row * 8);
1053 }
1054
1055 #[test]
1056 fn q8_0_block_roundtrip() {
1057 let mut block = [0u8; 34];
1058 block[0] = 0x00;
1059 block[1] = 0x3C; for i in 0..32 {
1061 block[2 + i] = (i + 1) as u8;
1062 }
1063 let mut out = [0f32; 32];
1064 dequant_q8_0(&block, 32, &mut out).unwrap();
1065 assert!((out[0] - 1.0).abs() < 0.01);
1066 assert!((out[31] - 32.0).abs() < 0.01);
1067 }
1068}