1use crate::{FibQuantError, FibQuantizer, Result};
20
21use super::{
22 block::{KvBlockEncodingV1, KvEncodedBlockV1},
23 codec::KvEncodedTensorV1,
24 layout::KvCacheLayoutV1,
25 page::KvEncodedPageV1,
26 profile::{KvAxisPolicyV1, KvCompressionProfileV1, KvFallbackModeV1},
27 receipt::{now_unix_seconds, KvCompressionReceiptV1, KvOperationKindV1, KV_RECEIPT_SCHEMA},
28 shape::{KvRole, KvTensorShapeV1},
29};
30
31#[derive(Debug, Clone, PartialEq)]
33pub struct AppendReceipt {
34 pub token_index: u32,
36 pub key_block_encoding: String,
38 pub value_block_encoding: String,
40 pub key_compressed_bytes: usize,
42 pub value_compressed_bytes: usize,
44 pub fallback_reasons: Vec<String>,
46}
47
48pub struct KvStreamEncoder {
60 shape: KvTensorShapeV1,
62 layout: KvCacheLayoutV1,
64 profile: KvCompressionProfileV1,
66 quantizer: FibQuantizer,
68 current_page_blocks: Vec<KvEncodedBlockV1>,
70 current_page_token_start: u32,
72 completed_pages: Vec<KvEncodedPageV1>,
74 source_digest_state: blake3::Hasher,
76
77 token_index: u32,
80 page_id: u32,
82 compressed_blocks: u32,
84 raw_fallback_blocks: u32,
86 fallback_reasons: Vec<String>,
88 pending_pages: Vec<(u32, u32, u32, Vec<KvEncodedBlockV1>)>,
91}
92
93impl KvStreamEncoder {
94 pub fn new(
103 shape: KvTensorShapeV1,
104 layout: KvCacheLayoutV1,
105 profile: KvCompressionProfileV1,
106 ) -> Result<Self> {
107 shape.validate()?;
108 layout.validate_for_shape(&shape)?;
109 profile.validate_for_shape(&shape)?;
110
111 if shape.batch != 1 || shape.layers != 1 || shape.kv_heads != 1 {
115 return Err(FibQuantError::DependencyUnsupported(
116 "streaming encoder currently supports batch=1, layers=1, kv_heads=1 only".into(),
117 ));
118 }
119
120 let quantizer = build_quantizer(&profile)?;
121
122 let mut source_digest_state = blake3::Hasher::new();
125 source_digest_state.update(b"fib_quant_kv_tensor_f32_v1");
126 source_digest_state.update(&[0]);
127 let total_elements = shape.element_count()? as u64;
128 source_digest_state.update(&total_elements.to_le_bytes());
129
130 Ok(Self {
131 shape,
132 layout,
133 profile,
134 quantizer,
135 current_page_blocks: Vec::new(),
136 current_page_token_start: 0,
137 completed_pages: Vec::new(),
138 source_digest_state,
139 token_index: 0,
140 page_id: 0,
141 compressed_blocks: 0,
142 raw_fallback_blocks: 0,
143 fallback_reasons: Vec::new(),
144 pending_pages: Vec::new(),
145 })
146 }
147
148 pub fn append_token(
158 &mut self,
159 key_vector: &[f32],
160 value_vector: &[f32],
161 ) -> Result<AppendReceipt> {
162 let head_dim = self.shape.head_dim as usize;
163
164 if key_vector.len() != head_dim {
166 return Err(FibQuantError::CorruptPayload(format!(
167 "key_vector has {} elements, expected head_dim={}",
168 key_vector.len(),
169 head_dim
170 )));
171 }
172 if value_vector.len() != head_dim {
173 return Err(FibQuantError::CorruptPayload(format!(
174 "value_vector has {} elements, expected head_dim={}",
175 value_vector.len(),
176 head_dim
177 )));
178 }
179
180 if key_vector.iter().any(|v| !v.is_finite()) {
182 return Err(FibQuantError::CorruptPayload(
183 "key_vector contains non-finite value".into(),
184 ));
185 }
186 if value_vector.iter().any(|v| !v.is_finite()) {
187 return Err(FibQuantError::CorruptPayload(
188 "value_vector contains non-finite value".into(),
189 ));
190 }
191
192 if self.token_index >= self.shape.tokens {
194 return Err(FibQuantError::CorruptPayload(format!(
195 "stream encoder received token {} but shape only has {} tokens",
196 self.token_index, self.shape.tokens
197 )));
198 }
199
200 let token = self.token_index;
201
202 let (encode_vector, other_vector) = match self.shape.role {
205 KvRole::Key => (key_vector, value_vector),
206 KvRole::Value => (value_vector, key_vector),
207 };
208 for v in encode_vector {
209 self.source_digest_state.update(&v.to_le_bytes());
210 }
211 let _ = other_vector; let block_id = self.current_page_blocks.len() as u32;
217 let protected = self
218 .profile
219 .protected_policy
220 .is_protected(&self.shape, 0, 0, token);
221
222 let block = if protected {
223 KvEncodedBlockV1::raw(
224 block_id,
225 0,
226 0,
227 0,
228 token,
229 encode_vector.to_vec(),
230 self.profile.page_geometry.encoded_block_bytes,
231 "protected_region",
232 )
233 } else {
234 self.encode_vector_block(block_id, token, encode_vector)?
235 };
236
237 let mut token_fallback_reasons = Vec::new();
239 if block.raw_fallback {
240 self.raw_fallback_blocks += 1;
241 if !self.fallback_reasons.contains(&block.reason) {
242 self.fallback_reasons.push(block.reason.clone());
243 }
244 token_fallback_reasons.push(block.reason.clone());
245 } else {
246 self.compressed_blocks += 1;
247 }
248
249 let (encoding_type, compressed_bytes) = match &block.encoding {
251 KvBlockEncodingV1::FibQuant { code } => ("fib_quant", code.compact_size()),
252 KvBlockEncodingV1::RawF32 { values } => {
253 ("raw", values.len() * std::mem::size_of::<f32>())
254 }
255 };
256
257 self.current_page_blocks.push(block);
258 self.token_index += 1;
259
260 let tokens_per_page = self.profile.page_geometry.tokens_per_page;
262 let tokens_in_current_page = self.token_index - self.current_page_token_start;
263 if tokens_in_current_page >= tokens_per_page {
264 self.flush_page();
265 }
266
267 let (key_encoding, key_bytes, value_encoding, value_bytes) = match self.shape.role {
269 KvRole::Key => (encoding_type, compressed_bytes, "raw", 0),
270 KvRole::Value => ("raw", 0, encoding_type, compressed_bytes),
271 };
272
273 Ok(AppendReceipt {
274 token_index: token,
275 key_block_encoding: key_encoding.to_string(),
276 value_block_encoding: value_encoding.to_string(),
277 key_compressed_bytes: key_bytes,
278 value_compressed_bytes: value_bytes,
279 fallback_reasons: token_fallback_reasons,
280 })
281 }
282
283 pub fn finish(mut self) -> Result<KvEncodedTensorV1> {
292 if self.token_index == 0 {
294 return Err(FibQuantError::CorruptPayload(
295 "stream encoder finished without any appended tokens".into(),
296 ));
297 }
298 if self.token_index != self.shape.tokens {
299 return Err(FibQuantError::CorruptPayload(format!(
300 "stream encoder finished with {} tokens, expected {}",
301 self.token_index, self.shape.tokens
302 )));
303 }
304
305 if !self.current_page_blocks.is_empty() {
307 self.flush_page();
308 }
309
310 let source_digest = format!("blake3:{}", self.source_digest_state.finalize().to_hex());
312 let profile_digest = self.profile.digest(&self.shape)?;
313
314 for (page_id, token_start, token_count, blocks) in self.pending_pages.drain(..) {
316 let page = KvEncodedPageV1::new(
317 page_id,
318 token_start,
319 token_count,
320 source_digest.clone(),
321 profile_digest.clone(),
322 &self.shape,
323 self.profile.page_geometry.clone(),
324 blocks,
325 )?;
326 self.completed_pages.push(page);
327 }
328
329 let page_digests = self
331 .completed_pages
332 .iter()
333 .map(|p| p.page_digest.clone())
334 .collect();
335
336 let receipt = KvCompressionReceiptV1 {
337 schema_version: KV_RECEIPT_SCHEMA.into(),
338 operation_kind: KvOperationKindV1::Compress,
339 source_digest,
340 profile_digest,
341 shape_digest: self.shape.digest()?,
342 page_digests,
343 codebook_digest: self.profile.codebook_digest.clone(),
344 rotation_digest: self.profile.rotation_digest.clone(),
345 encoded_pages: self.completed_pages.len() as u32,
346 compressed_blocks: self.compressed_blocks,
347 raw_fallback_blocks: self.raw_fallback_blocks,
348 fallback_reasons: std::mem::take(&mut self.fallback_reasons),
349 recorded_unix_seconds: now_unix_seconds(),
350 };
351
352 Ok(KvEncodedTensorV1 {
353 shape: self.shape,
354 layout: self.layout,
355 profile: self.profile,
356 pages: std::mem::take(&mut self.completed_pages),
357 receipt,
358 })
359 }
360
361 fn flush_page(&mut self) {
363 if self.current_page_blocks.is_empty() {
364 return;
365 }
366 let token_start = self.current_page_token_start;
367 let token_count = self.token_index - token_start;
368 let blocks = std::mem::take(&mut self.current_page_blocks);
369 self.pending_pages
370 .push((self.page_id, token_start, token_count, blocks));
371 self.page_id += 1;
372 self.current_page_token_start = self.token_index;
373 }
374
375 fn encode_vector_block(
378 &self,
379 block_id: u32,
380 token: u32,
381 vector: &[f32],
382 ) -> Result<KvEncodedBlockV1> {
383 match self.profile.axis_policy {
384 KvAxisPolicyV1::Raw => Ok(KvEncodedBlockV1::raw(
385 block_id,
386 0,
387 0,
388 0,
389 token,
390 vector.to_vec(),
391 self.profile.page_geometry.encoded_block_bytes,
392 "raw_axis_policy",
393 )),
394 KvAxisPolicyV1::PerToken => match self.quantizer.encode(vector) {
395 Ok(code) => Ok(KvEncodedBlockV1::fib_quant(
396 block_id,
397 0,
398 0,
399 0,
400 token,
401 code,
402 self.profile.page_geometry.encoded_block_bytes,
403 "fib_quant_per_token",
404 )),
405 Err(err) if self.profile.fallback_policy.mode == KvFallbackModeV1::KeepRaw => {
406 Ok(KvEncodedBlockV1::raw(
407 block_id,
408 0,
409 0,
410 0,
411 token,
412 vector.to_vec(),
413 self.profile.page_geometry.encoded_block_bytes,
414 format!("encode_fallback:{err}"),
415 ))
416 }
417 Err(err) => Err(err),
418 },
419 KvAxisPolicyV1::PerChannel | KvAxisPolicyV1::RoleAwareKiviStyle => {
420 if self.profile.fallback_policy.mode == KvFallbackModeV1::KeepRaw {
421 Ok(KvEncodedBlockV1::raw(
422 block_id,
423 0,
424 0,
425 0,
426 token,
427 vector.to_vec(),
428 self.profile.page_geometry.encoded_block_bytes,
429 "unsupported_axis_raw_fallback",
430 ))
431 } else {
432 Err(FibQuantError::DependencyUnsupported(
433 "CPU reference codec supports per-token FibQuant compression only".into(),
434 ))
435 }
436 }
437 }
438 }
439}
440
441fn build_quantizer(profile: &KvCompressionProfileV1) -> Result<FibQuantizer> {
446 let quantizer = FibQuantizer::new(profile.fib_profile.clone())?;
447 if quantizer.codebook().codebook_digest != profile.codebook_digest {
448 return Err(FibQuantError::CodebookDigestMismatch {
449 expected: quantizer.codebook().codebook_digest.clone(),
450 actual: profile.codebook_digest.clone(),
451 });
452 }
453 Ok(quantizer)
454}
455
456#[cfg(test)]
461mod tests {
462 use super::super::codec::encode_kv_tensor;
463 use super::super::layout::KvPageGeometryV1;
464 use super::super::profile::KvAxisPolicyV1;
465 use super::super::shape::{KvAttentionKind, KvDType, KvRopeState};
466 use super::*;
467 use crate::profile::FibQuantProfileV1;
468
469 fn build_test_parts() -> (
475 KvTensorShapeV1,
476 KvCacheLayoutV1,
477 KvCompressionProfileV1,
478 Vec<f32>,
479 ) {
480 let shape = KvTensorShapeV1::new(
481 KvRole::Key,
482 KvAttentionKind::Mha,
483 1, 1, 1, 1, 3, 8, KvDType::F32,
490 KvRopeState::PreRope,
491 );
492 let layout = KvCacheLayoutV1::canonical(&shape).expect("canonical layout");
493 let fib_profile =
494 FibQuantProfileV1::paper_default(8, 4, 32, 42).expect("build fib profile");
495 let quantizer = FibQuantizer::new(fib_profile.clone()).expect("build quantizer");
496 let page_geometry = KvPageGeometryV1::new(2, 8, 64); let profile = KvCompressionProfileV1::from_parts(
498 "test-stream-profile",
499 &shape,
500 fib_profile,
501 quantizer.codebook().codebook_digest.clone(),
502 KvAxisPolicyV1::PerToken,
503 page_geometry,
504 )
505 .expect("build kv profile");
506
507 let total = shape.element_count().expect("element count");
509 let values: Vec<f32> = (0..total).map(|i| (i as f32) * 0.1).collect();
510
511 (shape, layout, profile, values)
512 }
513
514 #[test]
515 fn stream_matches_batch_encode() {
516 let (shape, layout, profile, values) = build_test_parts();
517
518 let batch_result =
520 encode_kv_tensor(shape.clone(), layout.clone(), profile.clone(), &values)
521 .expect("batch encode");
522
523 let mut encoder = KvStreamEncoder::new(shape.clone(), layout.clone(), profile.clone())
525 .expect("build stream encoder");
526 let head_dim = shape.head_dim as usize;
527 for token in 0..shape.tokens {
528 let start = token as usize * head_dim;
529 let key_slice = &values[start..start + head_dim];
530 encoder
533 .append_token(key_slice, key_slice)
534 .expect("append token");
535 }
536 let stream_result = encoder.finish().expect("stream finish");
537
538 assert_eq!(stream_result.pages.len(), batch_result.pages.len());
541 for (stream_page, batch_page) in stream_result.pages.iter().zip(batch_result.pages.iter()) {
542 assert_eq!(stream_page.page_id, batch_page.page_id);
543 assert_eq!(stream_page.token_start, batch_page.token_start);
544 assert_eq!(stream_page.token_count, batch_page.token_count);
545 assert_eq!(
546 stream_page.source_tensor_digest,
547 batch_page.source_tensor_digest
548 );
549 assert_eq!(
550 stream_page.page_digest, batch_page.page_digest,
551 "page digest mismatch for page {}",
552 stream_page.page_id
553 );
554 assert_eq!(
555 stream_page.encoded_blocks.len(),
556 batch_page.encoded_blocks.len()
557 );
558 for (sb, bb) in stream_page
559 .encoded_blocks
560 .iter()
561 .zip(batch_page.encoded_blocks.iter())
562 {
563 assert_eq!(sb.block_id, bb.block_id);
564 assert_eq!(sb.token, bb.token);
565 assert_eq!(sb.raw_fallback, bb.raw_fallback);
566 assert_eq!(
567 sb.encoding, bb.encoding,
568 "block encoding mismatch for block {} (token {})",
569 sb.block_id, sb.token
570 );
571 }
572 }
573
574 assert_eq!(
577 stream_result.receipt.source_digest,
578 batch_result.receipt.source_digest
579 );
580 assert_eq!(
581 stream_result.receipt.profile_digest,
582 batch_result.receipt.profile_digest
583 );
584 assert_eq!(
585 stream_result.receipt.shape_digest,
586 batch_result.receipt.shape_digest
587 );
588 assert_eq!(
589 stream_result.receipt.page_digests,
590 batch_result.receipt.page_digests
591 );
592 assert_eq!(
593 stream_result.receipt.codebook_digest,
594 batch_result.receipt.codebook_digest
595 );
596 assert_eq!(
597 stream_result.receipt.rotation_digest,
598 batch_result.receipt.rotation_digest
599 );
600 assert_eq!(
601 stream_result.receipt.encoded_pages,
602 batch_result.receipt.encoded_pages
603 );
604 assert_eq!(
605 stream_result.receipt.compressed_blocks,
606 batch_result.receipt.compressed_blocks
607 );
608 assert_eq!(
609 stream_result.receipt.raw_fallback_blocks,
610 batch_result.receipt.raw_fallback_blocks
611 );
612 assert_eq!(
613 stream_result.receipt.fallback_reasons,
614 batch_result.receipt.fallback_reasons
615 );
616 }
617
618 #[test]
619 fn append_receipt_fields_correct() {
620 let (shape, layout, profile, values) = build_test_parts();
621 let mut encoder =
622 KvStreamEncoder::new(shape, layout, profile).expect("build stream encoder");
623 let head_dim = 8;
624
625 let r0 = encoder
627 .append_token(&values[0..head_dim], &values[0..head_dim])
628 .expect("append token 0");
629 assert_eq!(r0.token_index, 0);
630 assert!(
633 r0.key_block_encoding == "fib_quant" || r0.key_block_encoding == "raw",
634 "unexpected key_block_encoding: {}",
635 r0.key_block_encoding
636 );
637 assert_eq!(r0.value_block_encoding, "raw");
638 assert_eq!(r0.value_compressed_bytes, 0);
639 if r0.key_block_encoding == "fib_quant" {
640 assert!(r0.key_compressed_bytes > 0);
641 } else {
642 assert_eq!(
643 r0.key_compressed_bytes,
644 head_dim * std::mem::size_of::<f32>()
645 );
646 }
647
648 let r1 = encoder
650 .append_token(
651 &values[head_dim..2 * head_dim],
652 &values[head_dim..2 * head_dim],
653 )
654 .expect("append token 1");
655 assert_eq!(r1.token_index, 1);
656
657 let r2 = encoder
659 .append_token(
660 &values[2 * head_dim..3 * head_dim],
661 &values[2 * head_dim..3 * head_dim],
662 )
663 .expect("append token 2");
664 assert_eq!(r2.token_index, 2);
665
666 if r0.key_block_encoding == "fib_quant"
668 && r1.key_block_encoding == "fib_quant"
669 && r2.key_block_encoding == "fib_quant"
670 {
671 assert!(r0.fallback_reasons.is_empty());
672 assert!(r1.fallback_reasons.is_empty());
673 assert!(r2.fallback_reasons.is_empty());
674 }
675 }
676
677 #[test]
678 fn empty_stream_finish_returns_error() {
679 let (shape, layout, profile, _) = build_test_parts();
680 let encoder = KvStreamEncoder::new(shape, layout, profile).expect("build encoder");
681 let err = encoder.finish().unwrap_err();
682 assert!(
683 matches!(err, FibQuantError::CorruptPayload(ref msg)
684 if msg.contains("without any appended tokens")),
685 "expected empty-stream error, got: {err:?}"
686 );
687 }
688
689 #[test]
690 fn stream_decode_roundtrip() {
691 let (shape, layout, profile, values) = build_test_parts();
692
693 let batch_encoded =
695 encode_kv_tensor(shape.clone(), layout.clone(), profile.clone(), &values)
696 .expect("batch encode");
697 let batch_decoded =
698 super::super::codec::decode_kv_pages(&batch_encoded).expect("batch decode");
699
700 let mut encoder = KvStreamEncoder::new(shape.clone(), layout.clone(), profile.clone())
702 .expect("build stream encoder");
703 let head_dim = shape.head_dim as usize;
704 for token in 0..shape.tokens {
705 let start = token as usize * head_dim;
706 encoder
707 .append_token(
708 &values[start..start + head_dim],
709 &values[start..start + head_dim],
710 )
711 .expect("append token");
712 }
713 let encoded = encoder.finish().expect("stream finish");
714
715 let decoded = super::super::codec::decode_kv_pages(&encoded).expect("decode");
717 assert_eq!(decoded.values.len(), values.len());
718
719 assert_eq!(
722 decoded.values, batch_decoded.values,
723 "stream decode must match batch decode"
724 );
725 }
726
727 #[test]
728 fn too_many_tokens_returns_error() {
729 let (shape, layout, profile, values) = build_test_parts();
730 let mut encoder = KvStreamEncoder::new(shape, layout, profile).expect("build encoder");
731 let head_dim = 8;
732 for token in 0..3 {
734 let start = token as usize * head_dim;
735 encoder
736 .append_token(
737 &values[start..start + head_dim],
738 &values[start..start + head_dim],
739 )
740 .expect("append token");
741 }
742 let extra = vec![0.0f32; head_dim];
744 let err = encoder.append_token(&extra, &extra).unwrap_err();
745 assert!(
746 matches!(err, FibQuantError::CorruptPayload(ref msg)
747 if msg.contains("but shape only has")),
748 "expected too-many-tokens error, got: {err:?}"
749 );
750 }
751
752 #[test]
753 fn partial_stream_returns_error() {
754 let (shape, layout, profile, values) = build_test_parts();
755 let mut encoder = KvStreamEncoder::new(shape, layout, profile).expect("build encoder");
756 let head_dim = 8;
757 for token in 0..2 {
759 let start = token as usize * head_dim;
760 encoder
761 .append_token(
762 &values[start..start + head_dim],
763 &values[start..start + head_dim],
764 )
765 .expect("append token");
766 }
767 let err = encoder.finish().unwrap_err();
768 assert!(
769 matches!(err, FibQuantError::CorruptPayload(ref msg)
770 if msg.contains("finished with 2 tokens, expected 3")),
771 "expected partial-stream error, got: {err:?}"
772 );
773 }
774
775 #[test]
776 fn wrong_vector_length_returns_error() {
777 let (shape, layout, profile, _) = build_test_parts();
778 let mut encoder = KvStreamEncoder::new(shape, layout, profile).expect("build encoder");
779 let short = vec![0.0f32; 4]; let err = encoder.append_token(&short, &short).unwrap_err();
781 assert!(
782 matches!(err, FibQuantError::CorruptPayload(ref msg)
783 if msg.contains("key_vector has 4 elements")),
784 "expected wrong-length error, got: {err:?}"
785 );
786 }
787
788 #[test]
789 fn non_finite_value_returns_error() {
790 let (shape, layout, profile, _) = build_test_parts();
791 let mut encoder = KvStreamEncoder::new(shape, layout, profile).expect("build encoder");
792 let nan_vec = vec![f32::NAN; 8];
793 let err = encoder.append_token(&nan_vec, &nan_vec).unwrap_err();
794 assert!(
795 matches!(err, FibQuantError::CorruptPayload(ref msg)
796 if msg.contains("non-finite")),
797 "expected non-finite error, got: {err:?}"
798 );
799 }
800}