1pub fn kv_cache_formats() -> (&'static str, &'static str) {
15 static F: std::sync::OnceLock<(&'static str, &'static str)> = std::sync::OnceLock::new();
16 *F.get_or_init(|| {
17 let k = match std::env::var("MEMRA_KV_K").as_deref() {
18 Ok("fp8") => "fp8",
19 Ok("q8_0") | Ok("") | Err(_) => "q8_0",
20 Ok(o) => panic!("MEMRA_KV_K={o} unsupported (q8_0 | fp8)"),
21 };
22 let v = match std::env::var("MEMRA_KV_V").as_deref() {
23 Ok("q4_0") => "q4_0",
24 Ok("fp8") => "fp8",
25 Ok("q5_1") | Ok("") | Err(_) => "q5_1",
26 Ok(o) => panic!("MEMRA_KV_V={o} unsupported (q5_1 | q4_0 | fp8)"),
27 };
28 if (k, v) != ("q8_0", "q5_1") {
29 eprintln!("[memra] KV cache format: K={k} V={v} (non-default — new numeric config)");
30 }
31 (k, v)
32 })
33}
34
35pub fn kv_blk_bytes() -> (usize, usize) {
37 let (k, v) = kv_cache_formats();
38 let kb = match k {
39 "fp8" => 32,
40 _ => 34,
41 };
42 let vb = match v {
43 "q4_0" => 18,
44 "fp8" => 32,
45 _ => 24,
46 };
47 (kb, vb)
48}
49
50#[derive(Clone, Copy, Debug, PartialEq, Eq)]
56pub struct TpKvRankAllocationShape {
57 pub kv_dim_k: usize,
58 pub kv_dim_v: usize,
59 pub k_token_bytes: usize,
60 pub v_token_bytes: usize,
61 pub fixed_bytes: usize,
62}
63
64impl TpKvRankAllocationShape {
65 pub fn bytes_per_token(self) -> usize {
66 self.k_token_bytes.saturating_add(self.v_token_bytes)
67 }
68
69 pub fn allocation_bytes(self, capacity: usize) -> usize {
70 self.bytes_per_token()
71 .saturating_mul(capacity)
72 .saturating_add(self.fixed_bytes)
73 }
74}
75
76#[allow(clippy::manual_is_multiple_of)] pub fn tp_kv_rank_allocation_shape(
78 kv_dim_k: usize,
79 kv_dim_v: usize,
80 ranks: usize,
81) -> Result<TpKvRankAllocationShape, String> {
82 if ranks == 0 || kv_dim_k == 0 || kv_dim_v == 0 {
83 return Err(format!(
84 "TP KV dimensions and rank count must be nonzero: k={kv_dim_k} v={kv_dim_v} \
85 ranks={ranks}"
86 ));
87 }
88 if kv_dim_k % ranks != 0 || kv_dim_v % ranks != 0 {
89 return Err(format!(
90 "TP KV dimensions k={kv_dim_k} v={kv_dim_v} are not divisible by TP={ranks}"
91 ));
92 }
93 let local_k = kv_dim_k / ranks;
94 let local_v = kv_dim_v / ranks;
95 if !local_k.is_multiple_of(32) || !local_v.is_multiple_of(32) {
96 return Err(format!(
97 "TP KV local dimensions k={local_k} v={local_v} must be 32-aligned"
98 ));
99 }
100 let (k_block_bytes, v_block_bytes) = kv_blk_bytes();
101 let k_token_bytes = (local_k / 32)
102 .checked_mul(k_block_bytes)
103 .ok_or("TP KV K token-byte overflow")?;
104 let v_token_bytes = (local_v / 32)
105 .checked_mul(v_block_bytes)
106 .ok_or("TP KV V token-byte overflow")?;
107 Ok(TpKvRankAllocationShape {
108 kv_dim_k: local_k,
109 kv_dim_v: local_v,
110 k_token_bytes,
111 v_token_bytes,
112 fixed_bytes: 8 + 8 + std::mem::size_of::<i32>(),
115 })
116}
117
118pub fn gkv_on() -> bool {
121 static ON: std::sync::OnceLock<bool> = std::sync::OnceLock::new();
122 *ON.get_or_init(|| {
123 std::env::var("MEMRA_GEMMA_GKV")
124 .map(|v| v != "0")
125 .unwrap_or(true)
126 })
127}
128
129pub fn wkv_on() -> bool {
133 static ON: std::sync::OnceLock<bool> = std::sync::OnceLock::new();
134 *ON.get_or_init(|| {
135 std::env::var("MEMRA_GEMMA_WKV")
136 .map(|v| v != "0")
137 .unwrap_or_else(|_| std::env::var("MEMRA_DRAFT").is_err())
138 })
139}
140
141pub static KV_FP8_FORCE: std::sync::atomic::AtomicI8 = std::sync::atomic::AtomicI8::new(-1);
147
148pub fn kv_fp8_on() -> bool {
152 static ENV: std::sync::OnceLock<Option<bool>> = std::sync::OnceLock::new();
153 if let Some(v) = *ENV.get_or_init(|| std::env::var("MEMRA_KV_FP8").ok().map(|v| v == "1")) {
154 return v;
155 }
156 matches!(KV_FP8_FORCE.load(std::sync::atomic::Ordering::Relaxed), 1)
157}
158
159static SWA_RING_DEFAULT: std::sync::atomic::AtomicBool = std::sync::atomic::AtomicBool::new(false);
166
167pub fn set_swa_ring_default(on: bool) {
168 SWA_RING_DEFAULT.store(on, std::sync::atomic::Ordering::Relaxed);
169}
170
171pub fn swa_ring_on() -> bool {
172 static ENV: std::sync::OnceLock<Option<bool>> = std::sync::OnceLock::new();
173 match *ENV.get_or_init(|| match std::env::var("MEMRA_SWA_RING").ok().as_deref() {
174 Some("1") => Some(true),
175 Some("0") => Some(false),
176 _ => None,
177 }) {
178 Some(forced) => forced,
179 None => SWA_RING_DEFAULT.load(std::sync::atomic::Ordering::Relaxed),
180 }
181}
182
183pub const PRIME_CHUNK_MAX_TOKENS: usize = 4096;
186const SWA_VIEW_ALIGNMENT_ROWS: usize = 32;
187
188pub fn swa_ring_rows(window: usize, max_ctx: usize) -> usize {
191 max_ctx.min(
196 window + PRIME_CHUNK_MAX_TOKENS + SWA_REWIND_SLACK_ROWS + (SWA_VIEW_ALIGNMENT_ROWS - 1),
197 )
198}
199
200pub const SWA_REWIND_SLACK_ROWS: usize = 512;
210
211pub fn swa_retain_from(first_row: usize, window: usize, base: usize) -> usize {
221 let ideal = first_row
222 .saturating_sub(window.saturating_sub(1))
223 .saturating_sub(SWA_REWIND_SLACK_ROWS)
224 & !(SWA_VIEW_ALIGNMENT_ROWS - 1);
225 ideal.max(base)
226}
227
228pub const INDEX_RING_WORKING_ROWS: usize = 5120;
256
257pub fn index_ring_rows_for(explicit: Option<usize>, max_ctx: usize) -> Option<usize> {
267 let rows = match explicit {
268 Some(0) => return None,
269 Some(n) => n,
270 None => INDEX_RING_WORKING_ROWS,
271 };
272 (rows > 0 && rows < max_ctx).then_some(rows)
273}
274
275pub fn index_ring_default_rows(max_ctx: usize) -> Option<usize> {
279 index_ring_rows_for(None, max_ctx)
280}
281
282#[allow(clippy::manual_is_multiple_of)] pub fn index_ring_take(
303 ring: usize,
304 pool: usize,
305 pools_ready: usize,
306 cur: usize,
307 remaining: usize,
308) -> Option<usize> {
309 if ring == 0 {
310 return Some(remaining);
311 }
312 debug_assert!(
313 ring % pool == 0,
314 "the effective ring is a whole number of pools"
315 );
316 let live = cur.checked_sub(pools_ready.saturating_mul(pool))?;
317 (ring > live).then(|| (ring - live).min(remaining))
318}
319
320pub fn index_ring_rows(max_ctx: usize) -> Option<usize> {
324 let explicit = std::env::var("MEMRA_DSA_INDEX_RING")
325 .ok()
326 .and_then(|v| v.trim().parse::<usize>().ok());
327 index_ring_rows_for(explicit, max_ctx)
328}
329
330#[derive(Debug, Clone, Copy, PartialEq, Eq)]
331pub struct KvRing {
332 rows: usize,
333 window: usize,
334 base: usize,
335}
336
337#[derive(Debug, Clone, Copy, PartialEq, Eq)]
338pub enum KvRingAppend {
339 Contiguous {
340 write_row: usize,
341 },
342 Rebase {
343 src_row: usize,
344 keep_rows: usize,
345 new_base: usize,
346 write_row: usize,
347 },
348}
349
350impl KvRing {
351 pub fn new(rows: usize, window: usize) -> Self {
352 assert!(window > 0 && rows > 0, "invalid SWA ring geometry");
353 Self {
354 rows,
355 window,
356 base: 0,
357 }
358 }
359
360 pub fn rows(&self) -> usize {
361 self.rows
362 }
363 pub fn base(&self) -> usize {
364 self.base
365 }
366 pub fn window(&self) -> usize {
367 self.window
368 }
369
370 pub fn append_plan(
373 &self,
374 len: usize,
375 retain_from: usize,
376 append_rows: usize,
377 ) -> Result<KvRingAppend, String> {
378 if len < self.base || retain_from < self.base || retain_from > len {
379 return Err(format!(
380 "SWA ring lapped required rows (base {}, retain {retain_from}, len {len})",
381 self.base
382 ));
383 }
384 let used = len - self.base;
385 if used > self.rows {
386 return Err(format!(
387 "SWA ring state exceeds capacity ({used} > {})",
388 self.rows
389 ));
390 }
391 if used.saturating_add(append_rows) <= self.rows {
392 return Ok(KvRingAppend::Contiguous {
393 write_row: used % self.rows,
394 });
395 }
396
397 let keep_rows = len - retain_from;
398 if keep_rows.saturating_add(append_rows) > self.rows {
399 return Err(format!(
400 "SWA ring append does not fit (keep {keep_rows} + append {append_rows} > {})",
401 self.rows
402 ));
403 }
404 Ok(KvRingAppend::Rebase {
405 src_row: retain_from - self.base,
406 keep_rows,
407 new_base: retain_from,
408 write_row: keep_rows,
409 })
410 }
411
412 pub fn apply_rebase(&mut self, new_base: usize) {
413 debug_assert!(new_base >= self.base);
414 self.base = new_base;
415 }
416
417 pub fn physical_range(
418 &self,
419 start: usize,
420 end: usize,
421 ) -> Result<std::ops::Range<usize>, String> {
422 if start < self.base || end < start || end - self.base > self.rows {
423 return Err(format!(
424 "SWA ring view [{start},{end}) is outside resident [{},{})",
425 self.base,
426 self.base + self.rows
427 ));
428 }
429 let start_row = (start - self.base) % self.rows;
430 let len = end - start;
431 debug_assert!(
432 start_row + len <= self.rows,
433 "ring view must be contiguous after rebase"
434 );
435 Ok(start_row..start_row + len)
436 }
437
438 pub fn can_rewind_to(&self, len: usize) -> bool {
440 let raw = len.saturating_sub(self.window - 1);
441 let view_start = raw & !(SWA_VIEW_ALIGNMENT_ROWS - 1);
442 view_start >= self.base
443 }
444
445 pub fn restore_plan(&self, len: usize) -> Result<(usize, std::ops::Range<usize>), String> {
452 let raw = len.saturating_sub(self.window.saturating_sub(1));
453 let new_base = raw & !(SWA_VIEW_ALIGNMENT_ROWS - 1);
454 let physical = self.physical_range(new_base, len)?;
455 Ok((new_base, physical))
456 }
457}
458
459pub trait KvDev {
464 fn zeros(&self, n: usize) -> Result<CudaSlice<f32>, Box<dyn std::error::Error>>;
465 fn uninit(&self, n: usize) -> Result<CudaSlice<f32>, Box<dyn std::error::Error>>;
466 fn alloc_u8(&self, n: usize) -> Result<CudaSlice<u8>, Box<dyn std::error::Error>>;
467 fn htod_i32(&self, v: &[i32]) -> Result<CudaSlice<i32>, Box<dyn std::error::Error>>;
468 fn clone_dtod(
469 &self,
470 src: &CudaSlice<f32>,
471 ) -> Result<CudaSlice<f32>, Box<dyn std::error::Error>>;
472 fn copy_into(
473 &self,
474 dst: &mut CudaSlice<f32>,
475 off: usize,
476 src: &CudaSlice<f32>,
477 len: usize,
478 ) -> Result<(), Box<dyn std::error::Error>>;
479 fn copy_range_into(
483 &self,
484 dst: &mut CudaSlice<f32>,
485 dst_off: usize,
486 src: &CudaSlice<f32>,
487 src_off: usize,
488 len: usize,
489 ) -> Result<(), Box<dyn std::error::Error>>;
490 fn set_i32_one(&self, d: &mut CudaSlice<i32>, v: i32)
491 -> Result<(), Box<dyn std::error::Error>>;
492}
493
494use cudarc::driver::CudaSlice;
495use memra_gguf::config::{LayerKind, ModelConfig};
496use memra_gguf::model_plan::{ModelPlan, ResidualTopology, StatePlan};
497
498pub struct KvLayer {
503 pub k: CudaSlice<u8>, pub v: CudaSlice<u8>, pub kv_dim_k: usize, pub kv_dim_v: usize, pub k_tok_bytes: usize, pub v_tok_bytes: usize, pub len: usize,
510 pub ring: Option<KvRing>,
513 pub len_d: CudaSlice<i32>,
517 pub base_d: Option<CudaSlice<i32>>,
526}
527
528impl KvLayer {
529 pub fn physical_rows(
530 &self,
531 start: usize,
532 end: usize,
533 ) -> Result<std::ops::Range<usize>, String> {
534 match &self.ring {
535 Some(ring) => ring.physical_range(start, end),
536 None => Ok(start..end),
537 }
538 }
539}
540
541pub struct LatentKvLayer {
554 pub rows: CudaSlice<f32>,
556 pub width: usize,
557 pub len: usize,
558 pub len_d: CudaSlice<i32>,
560 pub index_rows: Option<CudaSlice<f32>>,
610 pub index_width: usize,
611 pub index_ring_rows: Option<usize>,
615 pub index_pool_keys: Option<CudaSlice<f32>>,
640 pub index_pools_ready: usize,
642 pub index_pool: usize,
650}
651
652impl LatentKvLayer {
653 pub fn truncate_index_pool_keys(&mut self, pool: usize) {
657 if pool == 0 {
658 return;
659 }
660 self.index_pools_ready = self.index_pools_ready.min(self.len / pool);
661 }
662}
663
664pub fn index_plane_physical_row(ring_rows: usize, pool: usize, abs: usize) -> usize {
669 if ring_rows == 0 {
670 return abs;
671 }
672 debug_assert!(pool > 0, "index plane addressing requires a known pool");
673 let effective = ring_rows / pool * pool;
674 debug_assert!(effective > 0, "the effective ring holds at least one pool");
675 abs % effective
676}
677
678pub struct LatentPlaneSnapshot {
695 pub rows: CudaSlice<f32>,
697 pub width: usize,
698 pub len: usize,
699 pub index_width: usize,
701 pub index_pool: usize,
703 pub index_tail: Option<CudaSlice<f32>>,
706 pub index_pool_keys: Option<CudaSlice<f32>>,
709 pub index_pools_ready: usize,
710}
711
712impl LatentPlaneSnapshot {
713 pub fn bytes(&self) -> usize {
716 let tail = self.index_tail.as_ref().map_or(0, CudaSlice::len);
717 let keys = self.index_pool_keys.as_ref().map_or(0, CudaSlice::len);
718 (self.rows.len() + tail + keys) * std::mem::size_of::<f32>()
719 }
720}
721
722pub struct LatentTailCapture {
733 pub len: usize,
735 pub width: usize,
737 pub index_width: usize,
738 pub index_pool: usize,
740 pub index_pools_ready: usize,
743 pub index_tail: Option<CudaSlice<f32>>,
747}
748
749impl LatentTailCapture {
750 pub fn bytes(&self) -> usize {
752 self.index_tail.as_ref().map_or(0, CudaSlice::len) * std::mem::size_of::<f32>()
753 }
754}
755
756impl LatentKvLayer {
757 pub fn snapshot_plane(
765 &self,
766 e: &impl KvDev,
767 ) -> Result<LatentPlaneSnapshot, Box<dyn std::error::Error>> {
768 let (len, width) = (self.len, self.width);
769 if len == 0 {
770 return Err("latent snapshot at len 0 (record the layer as absent instead)".into());
771 }
772 if self.rows.len() < len * width {
773 return Err(format!(
774 "latent plane holds {} f32 but len {len} x width {width} requires {}",
775 self.rows.len(),
776 len * width,
777 )
778 .into());
779 }
780 let mut rows = e.uninit(len * width)?;
781 e.copy_range_into(&mut rows, 0, &self.rows, 0, len * width)?;
782 if self.index_width == 0 {
783 return Ok(LatentPlaneSnapshot {
784 rows,
785 width,
786 len,
787 index_width: 0,
788 index_pool: 0,
789 index_tail: None,
790 index_pool_keys: None,
791 index_pools_ready: 0,
792 });
793 }
794 let pool = self.index_pool;
795 if pool == 0 {
796 return Err(format!(
797 "latent snapshot: index plane (width {}) has an unresolved pool — no indexer \
798 call ran against this layer, so its derived state cannot be validated",
799 self.index_width,
800 )
801 .into());
802 }
803 let d = self.index_width / 2;
804 let pools_ready = self.index_pools_ready;
805 if pools_ready != len / pool {
806 return Err(format!(
807 "latent snapshot: index_pools_ready {pools_ready} != len/pool {} (len {len}, \
808 pool {pool}); a capture must sit at a drained call boundary or its keys \
809 violate the append-only finality invariant",
810 len / pool,
811 )
812 .into());
813 }
814 let index_pool_keys = if pools_ready > 0 {
815 let src = self
816 .index_pool_keys
817 .as_ref()
818 .ok_or("latent snapshot: pools are ready but the resident key plane is gone")?;
819 if src.len() < pools_ready * d {
820 return Err(format!(
821 "latent snapshot: resident key plane holds {} f32 but {pools_ready} pools \
822 x d {d} require {}",
823 src.len(),
824 pools_ready * d,
825 )
826 .into());
827 }
828 let mut keys = e.uninit(pools_ready * d)?;
829 e.copy_range_into(&mut keys, 0, src, 0, pools_ready * d)?;
830 Some(keys)
831 } else {
832 None
833 };
834 let tail_rows = len - pools_ready * pool;
835 let index_tail = if tail_rows > 0 {
836 let src = self
837 .index_rows
838 .as_ref()
839 .ok_or("latent snapshot: index_width > 0 but the state plane is gone")?;
840 let ring = self.index_ring_rows.unwrap_or(0);
841 let phys = index_plane_physical_row(ring, pool, pools_ready * pool);
844 let want = (phys + tail_rows) * self.index_width;
845 if src.len() < want {
846 return Err(format!(
847 "latent snapshot: index plane holds {} f32 but the live tail window \
848 requires {want}",
849 src.len(),
850 )
851 .into());
852 }
853 let mut tail = e.uninit(tail_rows * self.index_width)?;
854 e.copy_range_into(
855 &mut tail,
856 0,
857 src,
858 phys * self.index_width,
859 tail_rows * self.index_width,
860 )?;
861 Some(tail)
862 } else {
863 None
864 };
865 Ok(LatentPlaneSnapshot {
866 rows,
867 width,
868 len,
869 index_width: self.index_width,
870 index_pool: pool,
871 index_tail,
872 index_pool_keys,
873 index_pools_ready: pools_ready,
874 })
875 }
876
877 pub fn snapshot_tail(
883 &self,
884 e: &impl KvDev,
885 ) -> Result<LatentTailCapture, Box<dyn std::error::Error>> {
886 let (len, width) = (self.len, self.width);
887 if len == 0 {
888 return Err("latent tail capture at len 0 (record the layer as absent instead)".into());
889 }
890 if self.index_width == 0 {
891 return Ok(LatentTailCapture {
892 len,
893 width,
894 index_width: 0,
895 index_pool: 0,
896 index_pools_ready: 0,
897 index_tail: None,
898 });
899 }
900 let pool = self.index_pool;
901 if pool == 0 {
902 return Err(format!(
903 "latent tail capture: index plane (width {}) has an unresolved pool — no \
904 indexer call ran against this layer, so its derived state cannot be validated",
905 self.index_width,
906 )
907 .into());
908 }
909 let pools_ready = self.index_pools_ready;
910 if pools_ready != len / pool {
911 return Err(format!(
912 "latent tail capture: index_pools_ready {pools_ready} != len/pool {} (len \
913 {len}, pool {pool}); a capture must sit at a drained call boundary",
914 len / pool,
915 )
916 .into());
917 }
918 let tail_rows = len - pools_ready * pool;
919 let index_tail = if tail_rows > 0 {
920 let src = self
921 .index_rows
922 .as_ref()
923 .ok_or("latent tail capture: index_width > 0 but the state plane is gone")?;
924 let ring = self.index_ring_rows.unwrap_or(0);
925 let phys = index_plane_physical_row(ring, pool, pools_ready * pool);
926 let want = (phys + tail_rows) * self.index_width;
927 if src.len() < want {
928 return Err(format!(
929 "latent tail capture: index plane holds {} f32 but the live tail window \
930 requires {want}",
931 src.len(),
932 )
933 .into());
934 }
935 let mut tail = e.uninit(tail_rows * self.index_width)?;
936 e.copy_range_into(
937 &mut tail,
938 0,
939 src,
940 phys * self.index_width,
941 tail_rows * self.index_width,
942 )?;
943 Some(tail)
944 } else {
945 None
946 };
947 Ok(LatentTailCapture {
948 len,
949 width,
950 index_width: self.index_width,
951 index_pool: pool,
952 index_pools_ready: pools_ready,
953 index_tail,
954 })
955 }
956
957 pub fn snapshot_plane_at(
966 &self,
967 e: &impl KvDev,
968 cap: LatentTailCapture,
969 ) -> Result<LatentPlaneSnapshot, Box<dyn std::error::Error>> {
970 let (len, width) = (cap.len, cap.width);
971 if len == 0 {
972 return Err("latent boundary publish at len 0".into());
973 }
974 if width != self.width {
975 return Err(format!(
976 "latent boundary publish: captured width {width} != live width {}",
977 self.width,
978 )
979 .into());
980 }
981 if self.len < len {
982 return Err(format!(
983 "latent boundary publish: live len {} < boundary {len} — the plane was \
984 truncated below the capture boundary",
985 self.len,
986 )
987 .into());
988 }
989 if self.rows.len() < len * width {
990 return Err(format!(
991 "latent boundary publish: live plane holds {} f32 but boundary {len} x width \
992 {width} requires {}",
993 self.rows.len(),
994 len * width,
995 )
996 .into());
997 }
998 let mut rows = e.uninit(len * width)?;
999 e.copy_range_into(&mut rows, 0, &self.rows, 0, len * width)?;
1000 if cap.index_width != self.index_width {
1001 return Err(format!(
1002 "latent boundary publish: captured index_width {} != live {}",
1003 cap.index_width, self.index_width,
1004 )
1005 .into());
1006 }
1007 if cap.index_width == 0 {
1008 return Ok(LatentPlaneSnapshot {
1009 rows,
1010 width,
1011 len,
1012 index_width: 0,
1013 index_pool: 0,
1014 index_tail: None,
1015 index_pool_keys: None,
1016 index_pools_ready: 0,
1017 });
1018 }
1019 if cap.index_pool != self.index_pool {
1020 return Err(format!(
1021 "latent boundary publish: captured pool {} != live pool {}",
1022 cap.index_pool, self.index_pool,
1023 )
1024 .into());
1025 }
1026 let d = cap.index_width / 2;
1027 let pools_ready = cap.index_pools_ready;
1028 if self.index_pools_ready < pools_ready {
1029 return Err(format!(
1030 "latent boundary publish: live index_pools_ready {} < boundary {pools_ready} \
1031 — the key plane was clamped below the capture boundary",
1032 self.index_pools_ready,
1033 )
1034 .into());
1035 }
1036 let index_pool_keys = if pools_ready > 0 {
1037 let src = self
1038 .index_pool_keys
1039 .as_ref()
1040 .ok_or("latent boundary publish: pools are ready but the key plane is gone")?;
1041 if src.len() < pools_ready * d {
1042 return Err(format!(
1043 "latent boundary publish: key plane holds {} f32 but {pools_ready} pools \
1044 x d {d} require {}",
1045 src.len(),
1046 pools_ready * d,
1047 )
1048 .into());
1049 }
1050 let mut keys = e.uninit(pools_ready * d)?;
1051 e.copy_range_into(&mut keys, 0, src, 0, pools_ready * d)?;
1052 Some(keys)
1053 } else {
1054 None
1055 };
1056 Ok(LatentPlaneSnapshot {
1057 rows,
1058 width,
1059 len,
1060 index_width: cap.index_width,
1061 index_pool: cap.index_pool,
1062 index_tail: cap.index_tail,
1063 index_pool_keys,
1064 index_pools_ready: pools_ready,
1065 })
1066 }
1067
1068 pub fn validate_restore(
1072 &self,
1073 snap: &LatentPlaneSnapshot,
1074 max_ctx: usize,
1075 ) -> Result<(), String> {
1076 if self.len != 0 {
1077 return Err("restore destination latent plane is not fresh".into());
1078 }
1079 if self.width != snap.width {
1080 return Err(format!(
1081 "snapshot width {} != destination width {}",
1082 snap.width, self.width,
1083 ));
1084 }
1085 if snap.len == 0 || snap.len > max_ctx {
1086 return Err(format!("snapshot len {} outside [1,{max_ctx}]", snap.len));
1087 }
1088 if snap.rows.len() < snap.len * snap.width {
1089 return Err(format!(
1090 "snapshot rows plane holds {} f32 but len {} x width {} requires {} \
1091 (truncated capture)",
1092 snap.rows.len(),
1093 snap.len,
1094 snap.width,
1095 snap.len * snap.width,
1096 ));
1097 }
1098 if self.rows.len() < snap.len * self.width {
1099 return Err(format!(
1100 "destination latent plane holds {} f32 but the restore requires {}",
1101 self.rows.len(),
1102 snap.len * self.width,
1103 ));
1104 }
1105 if self.index_width != snap.index_width {
1106 return Err(format!(
1107 "snapshot index width {} != destination {}",
1108 snap.index_width, self.index_width,
1109 ));
1110 }
1111 if snap.index_width == 0 {
1112 return Ok(());
1113 }
1114 let pool = snap.index_pool;
1115 if pool == 0 {
1116 return Err("snapshot carries an index plane with an unresolved pool".into());
1117 }
1118 if self.index_pool != 0 && self.index_pool != pool {
1119 return Err(format!(
1120 "snapshot pool {pool} != destination resident pool {}",
1121 self.index_pool,
1122 ));
1123 }
1124 let d = snap.index_width / 2;
1125 if snap.index_pools_ready != snap.len / pool {
1126 return Err(format!(
1127 "snapshot index_pools_ready {} != len/pool {} (len {}, pool {pool}): the \
1128 append-only finality invariant does not hold, so its keys are stale",
1129 snap.index_pools_ready,
1130 snap.len / pool,
1131 snap.len,
1132 ));
1133 }
1134 match (&snap.index_pool_keys, snap.index_pools_ready) {
1135 (Some(keys), ready @ 1..) => {
1136 if keys.len() < ready * d {
1137 return Err(format!(
1138 "snapshot key plane holds {} f32 but {ready} pools x d {d} require {}",
1139 keys.len(),
1140 ready * d,
1141 ));
1142 }
1143 }
1144 (None, 0) => {}
1145 (Some(_), 0) => return Err("snapshot carries keys for zero ready pools".into()),
1146 (None, ready) => {
1147 return Err(format!(
1148 "snapshot claims {ready} ready pools but carries no keys"
1149 ));
1150 }
1151 }
1152 let tail_rows = snap.len - snap.index_pools_ready * pool;
1153 match (&snap.index_tail, tail_rows) {
1154 (Some(tail), rows @ 1..) => {
1155 if tail.len() < rows * snap.index_width {
1156 return Err(format!(
1157 "snapshot tail holds {} f32 but {rows} rows x index width {} require {}",
1158 tail.len(),
1159 snap.index_width,
1160 rows * snap.index_width,
1161 ));
1162 }
1163 }
1164 (None, 0) => {}
1165 (Some(_), 0) => return Err("snapshot carries a tail at a pool-aligned boundary".into()),
1166 (None, rows) => {
1167 return Err(format!(
1168 "snapshot owes {rows} live tail rows but carries none"
1169 ));
1170 }
1171 }
1172 if self.index_rows.is_none() {
1173 return Err("destination declares an index plane but allocated none".into());
1174 }
1175 if tail_rows > 0 {
1176 let ring = self.index_ring_rows.unwrap_or(0);
1177 let phys = index_plane_physical_row(ring, pool, snap.index_pools_ready * pool);
1178 let want = (phys + tail_rows) * self.index_width;
1179 let have = self.index_rows.as_ref().map_or(0, CudaSlice::len);
1180 if have < want {
1181 return Err(format!(
1182 "destination index plane holds {have} f32 but the tail window requires \
1183 {want}",
1184 ));
1185 }
1186 }
1187 Ok(())
1188 }
1189
1190 pub fn restore_plane(
1198 &mut self,
1199 e: &impl KvDev,
1200 snap: &LatentPlaneSnapshot,
1201 max_ctx: usize,
1202 ) -> Result<(), Box<dyn std::error::Error>> {
1203 self.validate_restore(snap, max_ctx)?;
1204 e.copy_range_into(&mut self.rows, 0, &snap.rows, 0, snap.len * snap.width)?;
1205 if snap.index_width > 0 {
1206 let pool = snap.index_pool;
1207 let d = snap.index_width / 2;
1208 let mut keys = e.zeros(((max_ctx / pool) * d).max(1))?;
1211 if let Some(src) = &snap.index_pool_keys {
1212 e.copy_range_into(&mut keys, 0, src, 0, snap.index_pools_ready * d)?;
1213 }
1214 self.index_pool_keys = Some(keys);
1215 self.index_pools_ready = snap.index_pools_ready;
1216 self.index_pool = pool;
1217 if let Some(tail) = &snap.index_tail {
1218 let tail_rows = snap.len - snap.index_pools_ready * pool;
1219 let ring = self.index_ring_rows.unwrap_or(0);
1220 let phys = index_plane_physical_row(ring, pool, snap.index_pools_ready * pool);
1221 let dst = self
1222 .index_rows
1223 .as_mut()
1224 .ok_or("destination index plane vanished after validation")?;
1225 e.copy_range_into(
1226 dst,
1227 phys * self.index_width,
1228 tail,
1229 0,
1230 tail_rows * self.index_width,
1231 )?;
1232 }
1233 }
1234 self.len = snap.len;
1235 let len_i32 = i32::try_from(snap.len).map_err(|_| "latent length exceeds i32 mirror")?;
1236 e.set_i32_one(&mut self.len_d, len_i32)?;
1237 Ok(())
1238 }
1239}
1240
1241pub struct RecurLayer {
1245 pub conv_state: CudaSlice<f32>, pub ssm_state: CudaSlice<f32>, pub ssm_state_alt: CudaSlice<f32>,
1257}
1258
1259pub struct ResidentTpKvCacheRank {
1260 k: CudaSlice<u8>,
1261 v: CudaSlice<u8>,
1262 len_d: CudaSlice<i32>,
1263 base_d: Option<CudaSlice<i32>>,
1268}
1269
1270impl ResidentTpKvCacheRank {
1271 pub fn new(k: CudaSlice<u8>, v: CudaSlice<u8>, len_d: CudaSlice<i32>) -> Self {
1272 Self {
1273 k,
1274 v,
1275 len_d,
1276 base_d: None,
1277 }
1278 }
1279
1280 pub fn base_d(&self) -> Option<&CudaSlice<i32>> {
1281 self.base_d.as_ref()
1282 }
1283
1284 pub fn base_d_mut(&mut self) -> Option<&mut CudaSlice<i32>> {
1285 self.base_d.as_mut()
1286 }
1287
1288 pub fn arm_base_d(&mut self, buf: CudaSlice<i32>) {
1289 self.base_d = Some(buf);
1290 }
1291
1292 pub fn k(&self) -> &CudaSlice<u8> {
1293 &self.k
1294 }
1295
1296 pub fn v(&self) -> &CudaSlice<u8> {
1297 &self.v
1298 }
1299
1300 pub fn len_d(&self) -> &CudaSlice<i32> {
1301 &self.len_d
1302 }
1303
1304 pub fn k_mut(&mut self) -> &mut CudaSlice<u8> {
1305 &mut self.k
1306 }
1307
1308 pub fn v_mut(&mut self) -> &mut CudaSlice<u8> {
1309 &mut self.v
1310 }
1311
1312 pub fn planes_mut(&mut self) -> (&mut CudaSlice<u8>, &mut CudaSlice<u8>) {
1313 (&mut self.k, &mut self.v)
1314 }
1315
1316 #[allow(clippy::type_complexity)]
1318 pub fn planes_and_counters_mut(
1319 &mut self,
1320 ) -> (
1321 &mut CudaSlice<u8>,
1322 &mut CudaSlice<u8>,
1323 &CudaSlice<i32>,
1324 Option<&CudaSlice<i32>>,
1325 ) {
1326 (&mut self.k, &mut self.v, &self.len_d, self.base_d.as_ref())
1327 }
1328
1329 pub fn len_d_mut(&mut self) -> &mut CudaSlice<i32> {
1330 &mut self.len_d
1331 }
1332}
1333
1334#[derive(Clone, Copy, Debug, PartialEq, Eq)]
1335pub struct TpKvTransaction {
1336 generation: u64,
1337 base_len: usize,
1338}
1339
1340impl TpKvTransaction {
1341 pub fn generation(self) -> u64 {
1342 self.generation
1343 }
1344
1345 pub fn base_len(self) -> usize {
1346 self.base_len
1347 }
1348}
1349
1350#[derive(Clone, Copy, Debug, PartialEq, Eq)]
1351pub struct TpKvAppendPlan {
1352 transaction: TpKvTransaction,
1353 target: usize,
1354 write_row: usize,
1355 ring_append: Option<KvRingAppend>,
1356}
1357
1358impl TpKvAppendPlan {
1359 pub fn target(self) -> usize {
1360 self.target
1361 }
1362
1363 pub fn write_row(self) -> usize {
1364 self.write_row
1365 }
1366
1367 pub fn ring_append(self) -> Option<KvRingAppend> {
1368 self.ring_append
1369 }
1370}
1371
1372#[derive(Debug, PartialEq, Eq)]
1373pub struct TpKvGrowPlan {
1374 rows: usize,
1375 source_row: usize,
1376 copy_rows: usize,
1377 target_base: usize,
1378 k_bytes: usize,
1379 v_bytes: usize,
1380 source_capacity: usize,
1381 target_capacity: usize,
1382 ring_window: Option<usize>,
1383 target_physical_rows: usize,
1384 kv_dim_k: usize,
1385 kv_dim_v: usize,
1386 k_tok_bytes: usize,
1387 v_tok_bytes: usize,
1388 ranks: usize,
1389 next_generation: u64,
1390}
1391
1392impl TpKvGrowPlan {
1393 pub fn rows(&self) -> usize {
1394 self.rows
1395 }
1396
1397 pub fn source_row(&self) -> usize {
1398 self.source_row
1399 }
1400
1401 pub fn copy_rows(&self) -> usize {
1402 self.copy_rows
1403 }
1404
1405 pub fn k_bytes(&self) -> usize {
1406 self.k_bytes
1407 }
1408
1409 pub fn v_bytes(&self) -> usize {
1410 self.v_bytes
1411 }
1412}
1413
1414#[derive(Clone, Debug, PartialEq, Eq)]
1415struct TpKvTransactionState {
1416 committed_len: usize,
1417 staged_len: usize,
1418 next_generation: u64,
1419 active: Option<TpKvTransaction>,
1420}
1421
1422impl TpKvTransactionState {
1423 fn new() -> Self {
1424 Self {
1425 committed_len: 0,
1426 staged_len: 0,
1427 next_generation: 1,
1428 active: None,
1429 }
1430 }
1431
1432 fn begin(&mut self) -> Result<TpKvTransaction, String> {
1433 if let Some(active) = self.active {
1434 return Err(format!(
1435 "TP KV transaction generation {} is already active at base {}",
1436 active.generation, active.base_len
1437 ));
1438 }
1439 if self.staged_len != self.committed_len {
1440 return Err(format!(
1441 "TP KV cache is half-committed: staged {} != committed {}",
1442 self.staged_len, self.committed_len
1443 ));
1444 }
1445 let transaction = TpKvTransaction {
1446 generation: self.next_generation,
1447 base_len: self.committed_len,
1448 };
1449 self.next_generation = self
1450 .next_generation
1451 .checked_add(1)
1452 .ok_or("TP KV transaction generation overflow")?;
1453 self.active = Some(transaction);
1454 Ok(transaction)
1455 }
1456
1457 fn validate(&self, transaction: TpKvTransaction) -> Result<(), String> {
1458 if self.active != Some(transaction) {
1459 return Err(format!(
1460 "stale TP KV transaction generation {} at base {}",
1461 transaction.generation, transaction.base_len
1462 ));
1463 }
1464 if transaction.base_len != self.committed_len {
1465 return Err(format!(
1466 "TP KV transaction base {} != committed length {}",
1467 transaction.base_len, self.committed_len
1468 ));
1469 }
1470 Ok(())
1471 }
1472
1473 fn append_target(
1474 &self,
1475 transaction: TpKvTransaction,
1476 rows: usize,
1477 capacity: usize,
1478 ) -> Result<usize, String> {
1479 self.validate(transaction)?;
1480 if rows == 0 {
1481 return Err("TP KV append must contain at least one row".into());
1482 }
1483 let target = self
1484 .staged_len
1485 .checked_add(rows)
1486 .ok_or("TP KV staged length overflow")?;
1487 if target > capacity {
1488 return Err(format!(
1489 "TP KV append exceeds capacity: {target} > {capacity}"
1490 ));
1491 }
1492 Ok(target)
1493 }
1494
1495 fn publish_append(
1496 &mut self,
1497 transaction: TpKvTransaction,
1498 target: usize,
1499 ) -> Result<(), String> {
1500 self.validate(transaction)?;
1501 if target <= self.staged_len {
1502 return Err(format!(
1503 "TP KV append target {target} must exceed staged length {}",
1504 self.staged_len
1505 ));
1506 }
1507 self.staged_len = target;
1508 Ok(())
1509 }
1510
1511 fn commit_target(
1512 &self,
1513 transaction: TpKvTransaction,
1514 accepted_rows: usize,
1515 ) -> Result<usize, String> {
1516 self.validate(transaction)?;
1517 let staged_rows = self
1518 .staged_len
1519 .checked_sub(transaction.base_len)
1520 .ok_or("TP KV staged length precedes its transaction base")?;
1521 if accepted_rows > staged_rows {
1522 return Err(format!(
1523 "TP KV commit accepts {accepted_rows} rows from a {staged_rows}-row transaction"
1524 ));
1525 }
1526 transaction
1527 .base_len
1528 .checked_add(accepted_rows)
1529 .ok_or_else(|| "TP KV committed length overflow".to_string())
1530 }
1531
1532 fn publish_finalize(
1533 &mut self,
1534 transaction: TpKvTransaction,
1535 target: usize,
1536 ) -> Result<(), String> {
1537 self.validate(transaction)?;
1538 if target < transaction.base_len || target > self.staged_len {
1539 return Err(format!(
1540 "TP KV finalize target {target} outside transaction range {}..={}",
1541 transaction.base_len, self.staged_len
1542 ));
1543 }
1544 self.committed_len = target;
1545 self.staged_len = target;
1546 self.active = None;
1547 Ok(())
1548 }
1549
1550 fn rewind(&mut self, target: usize, capacity: usize) -> Result<(), String> {
1551 if target > capacity {
1552 return Err(format!(
1553 "TP KV rewind target {target} exceeds capacity {capacity}"
1554 ));
1555 }
1556 self.committed_len = target;
1557 self.staged_len = target;
1558 self.active = None;
1559 Ok(())
1560 }
1561}
1562
1563pub struct ResidentTpKvCache {
1564 ranks: Vec<ResidentTpKvCacheRank>,
1565 kv_dim_k: usize,
1566 kv_dim_v: usize,
1567 k_tok_bytes: usize,
1568 v_tok_bytes: usize,
1569 capacity: usize,
1570 ring: Option<KvRing>,
1571 state: TpKvTransactionState,
1572}
1573
1574impl ResidentTpKvCache {
1575 #[allow(clippy::too_many_arguments)]
1576 pub fn new(
1577 ranks: Vec<ResidentTpKvCacheRank>,
1578 kv_dim_k: usize,
1579 kv_dim_v: usize,
1580 k_tok_bytes: usize,
1581 v_tok_bytes: usize,
1582 capacity: usize,
1583 ) -> Self {
1584 Self::new_inner(
1585 ranks,
1586 kv_dim_k,
1587 kv_dim_v,
1588 k_tok_bytes,
1589 v_tok_bytes,
1590 capacity,
1591 None,
1592 )
1593 }
1594
1595 #[allow(clippy::too_many_arguments)]
1596 pub fn new_swa(
1597 ranks: Vec<ResidentTpKvCacheRank>,
1598 kv_dim_k: usize,
1599 kv_dim_v: usize,
1600 k_tok_bytes: usize,
1601 v_tok_bytes: usize,
1602 capacity: usize,
1603 window: usize,
1604 ) -> Self {
1605 Self::new_inner(
1606 ranks,
1607 kv_dim_k,
1608 kv_dim_v,
1609 k_tok_bytes,
1610 v_tok_bytes,
1611 capacity,
1612 Some(KvRing::new(swa_ring_rows(window, capacity), window)),
1613 )
1614 }
1615
1616 #[allow(clippy::too_many_arguments)]
1617 fn new_inner(
1618 ranks: Vec<ResidentTpKvCacheRank>,
1619 kv_dim_k: usize,
1620 kv_dim_v: usize,
1621 k_tok_bytes: usize,
1622 v_tok_bytes: usize,
1623 capacity: usize,
1624 ring: Option<KvRing>,
1625 ) -> Self {
1626 Self {
1627 ranks,
1628 kv_dim_k,
1629 kv_dim_v,
1630 k_tok_bytes,
1631 v_tok_bytes,
1632 capacity,
1633 ring,
1634 state: TpKvTransactionState::new(),
1635 }
1636 }
1637
1638 pub fn begin_transaction(&mut self) -> Result<TpKvTransaction, String> {
1639 self.state.begin()
1640 }
1641
1642 pub fn committed_len(&self) -> usize {
1643 self.state.committed_len
1644 }
1645
1646 pub fn staged_len(&self) -> usize {
1647 self.state.staged_len
1648 }
1649
1650 pub fn capacity(&self) -> usize {
1651 self.capacity
1652 }
1653
1654 pub fn physical_capacity(&self) -> usize {
1655 self.ring
1656 .as_ref()
1657 .map(KvRing::rows)
1658 .unwrap_or(self.capacity)
1659 }
1660
1661 pub fn ring_window(&self) -> Option<usize> {
1662 self.ring.as_ref().map(KvRing::window)
1663 }
1664
1665 pub fn ring_base(&self) -> Option<usize> {
1666 self.ring.as_ref().map(KvRing::base)
1667 }
1668
1669 pub fn physical_range(
1670 &self,
1671 start: usize,
1672 end: usize,
1673 ) -> Result<std::ops::Range<usize>, String> {
1674 match &self.ring {
1675 Some(ring) => ring.physical_range(start, end),
1676 None => {
1677 if end < start || end > self.capacity {
1678 return Err(format!(
1679 "TP KV linear view [{start},{end}) exceeds capacity {}",
1680 self.capacity
1681 ));
1682 }
1683 Ok(start..end)
1684 }
1685 }
1686 }
1687
1688 pub fn can_rewind_to(&self, target: usize) -> bool {
1689 target <= self.capacity
1690 && self
1691 .ring
1692 .as_ref()
1693 .is_none_or(|ring| ring.can_rewind_to(target))
1694 }
1695
1696 pub fn kv_dim_k(&self) -> usize {
1697 self.kv_dim_k
1698 }
1699
1700 pub fn kv_dim_v(&self) -> usize {
1701 self.kv_dim_v
1702 }
1703
1704 pub fn k_tok_bytes(&self) -> usize {
1705 self.k_tok_bytes
1706 }
1707
1708 pub fn v_tok_bytes(&self) -> usize {
1709 self.v_tok_bytes
1710 }
1711
1712 pub fn ranks_len(&self) -> usize {
1713 self.ranks.len()
1714 }
1715
1716 pub fn rank(&self, rank: usize) -> Option<&ResidentTpKvCacheRank> {
1717 self.ranks.get(rank)
1718 }
1719
1720 pub fn rank_mut(&mut self, rank: usize) -> Option<&mut ResidentTpKvCacheRank> {
1721 self.ranks.get_mut(rank)
1722 }
1723
1724 pub fn ranks(&self) -> &[ResidentTpKvCacheRank] {
1725 &self.ranks
1726 }
1727
1728 pub fn ranks_mut(&mut self) -> &mut [ResidentTpKvCacheRank] {
1729 &mut self.ranks
1730 }
1731
1732 pub fn prepare_grow(
1733 &self,
1734 target_capacity: usize,
1735 rows: usize,
1736 ) -> Result<TpKvGrowPlan, String> {
1737 if let Some(active) = self.state.active {
1738 return Err(format!(
1739 "TP KV grow refuses active transaction generation {} at base {}",
1740 active.generation, active.base_len
1741 ));
1742 }
1743 if self.state.staged_len != self.state.committed_len {
1744 return Err(format!(
1745 "TP KV grow requires quiescent state, got committed/staged={}/{}",
1746 self.state.committed_len, self.state.staged_len
1747 ));
1748 }
1749 if target_capacity <= self.capacity {
1750 return Err(format!(
1751 "TP KV grow target capacity {target_capacity} must exceed source capacity {}",
1752 self.capacity
1753 ));
1754 }
1755 if target_capacity > i32::MAX as usize {
1756 return Err(format!(
1757 "TP KV grow target capacity {target_capacity} exceeds i32 device mirrors"
1758 ));
1759 }
1760 if rows > self.state.committed_len {
1761 return Err(format!(
1762 "TP KV grow rows {rows} exceed committed length {}",
1763 self.state.committed_len
1764 ));
1765 }
1766 let (source_row, copy_rows, target_base, ring_window, target_physical_rows) =
1767 match &self.ring {
1768 Some(ring) => {
1769 let raw = rows.saturating_sub(ring.window().saturating_sub(1));
1770 let target_base = raw & !(SWA_VIEW_ALIGNMENT_ROWS - 1);
1771 let physical = ring.physical_range(target_base, rows)?;
1772 (
1773 physical.start,
1774 physical.len(),
1775 target_base,
1776 Some(ring.window()),
1777 swa_ring_rows(ring.window(), target_capacity),
1778 )
1779 }
1780 None => (0, rows, 0, None, target_capacity),
1781 };
1782 let k_bytes = copy_rows
1783 .checked_mul(self.k_tok_bytes)
1784 .ok_or("TP KV grow K byte extent overflow")?;
1785 let v_bytes = copy_rows
1786 .checked_mul(self.v_tok_bytes)
1787 .ok_or("TP KV grow V byte extent overflow")?;
1788 Ok(TpKvGrowPlan {
1789 rows,
1790 source_row,
1791 copy_rows,
1792 target_base,
1793 k_bytes,
1794 v_bytes,
1795 source_capacity: self.capacity,
1796 target_capacity,
1797 ring_window,
1798 target_physical_rows,
1799 kv_dim_k: self.kv_dim_k,
1800 kv_dim_v: self.kv_dim_v,
1801 k_tok_bytes: self.k_tok_bytes,
1802 v_tok_bytes: self.v_tok_bytes,
1803 ranks: self.ranks.len(),
1804 next_generation: self.state.next_generation,
1805 })
1806 }
1807
1808 pub fn publish_grow(&mut self, plan: TpKvGrowPlan) -> Result<(), String> {
1809 if self.state != TpKvTransactionState::new() {
1810 return Err(format!(
1811 "TP KV grow target must be fresh, got committed/staged={}/{} active={}",
1812 self.state.committed_len,
1813 self.state.staged_len,
1814 self.state.active.is_some()
1815 ));
1816 }
1817 if self.capacity != plan.target_capacity
1818 || self.capacity <= plan.source_capacity
1819 || self.kv_dim_k != plan.kv_dim_k
1820 || self.kv_dim_v != plan.kv_dim_v
1821 || self.k_tok_bytes != plan.k_tok_bytes
1822 || self.v_tok_bytes != plan.v_tok_bytes
1823 || self.ranks.len() != plan.ranks
1824 || self.ring.as_ref().map(KvRing::window) != plan.ring_window
1825 || self.physical_capacity() != plan.target_physical_rows
1826 {
1827 return Err("TP KV grow target layout does not match its source plan".into());
1828 }
1829 if plan.rows > self.capacity {
1830 return Err(format!(
1831 "TP KV grow rows {} exceed target capacity {}",
1832 plan.rows, self.capacity
1833 ));
1834 }
1835 if let Some(ring) = self.ring.as_mut() {
1836 let mut target_ring = *ring;
1837 target_ring.apply_rebase(plan.target_base);
1838 if !target_ring.can_rewind_to(plan.rows) {
1839 return Err(format!(
1840 "TP KV grow target ring base {} cannot expose committed length {}",
1841 target_ring.base(),
1842 plan.rows
1843 ));
1844 }
1845 *ring = target_ring;
1846 }
1847 self.state.committed_len = plan.rows;
1848 self.state.staged_len = plan.rows;
1849 self.state.next_generation = plan.next_generation;
1850 self.state.active = None;
1851 Ok(())
1852 }
1853
1854 pub fn prepare_append(
1855 &self,
1856 transaction: TpKvTransaction,
1857 rows: usize,
1858 ) -> Result<TpKvAppendPlan, String> {
1859 let target = self.state.append_target(transaction, rows, self.capacity)?;
1860 let ring_append = self
1861 .ring
1862 .as_ref()
1863 .map(|ring| {
1864 let staged_retain =
1865 target.saturating_sub(ring.window()) & !(SWA_VIEW_ALIGNMENT_ROWS - 1);
1866 let rollback_retain = transaction
1867 .base_len
1868 .saturating_sub(ring.window().saturating_sub(1))
1869 & !(SWA_VIEW_ALIGNMENT_ROWS - 1);
1870 ring.append_plan(
1871 self.state.staged_len,
1872 staged_retain.min(rollback_retain),
1873 rows,
1874 )
1875 })
1876 .transpose()?;
1877 let write_row = match ring_append {
1878 Some(KvRingAppend::Contiguous { write_row })
1879 | Some(KvRingAppend::Rebase { write_row, .. }) => write_row,
1880 None => self.state.staged_len,
1881 };
1882 Ok(TpKvAppendPlan {
1883 transaction,
1884 target,
1885 write_row,
1886 ring_append,
1887 })
1888 }
1889
1890 pub fn peek_append_ring(&self, rows: usize) -> Result<(usize, bool), String> {
1894 let target = self.state.staged_len + rows;
1895 let plan = self
1896 .ring
1897 .as_ref()
1898 .map(|ring| {
1899 let staged_retain =
1900 target.saturating_sub(ring.window()) & !(SWA_VIEW_ALIGNMENT_ROWS - 1);
1901 let rollback_retain = self
1902 .state
1903 .staged_len
1904 .saturating_sub(ring.window().saturating_sub(1))
1905 & !(SWA_VIEW_ALIGNMENT_ROWS - 1);
1906 ring.append_plan(
1907 self.state.staged_len,
1908 staged_retain.min(rollback_retain),
1909 rows,
1910 )
1911 })
1912 .transpose()?;
1913 Ok(match plan {
1914 Some(KvRingAppend::Contiguous { write_row }) => (write_row, false),
1915 Some(KvRingAppend::Rebase { write_row, .. }) => (write_row, true),
1916 None => (self.state.staged_len, false),
1917 })
1918 }
1919
1920 pub fn publish_append_rebase(&mut self, plan: TpKvAppendPlan) -> Result<(), String> {
1921 self.state.validate(plan.transaction)?;
1922 match (self.ring.as_mut(), plan.ring_append) {
1923 (
1924 Some(ring),
1925 Some(KvRingAppend::Rebase {
1926 new_base,
1927 keep_rows,
1928 ..
1929 }),
1930 ) => {
1931 if keep_rows > ring.rows() {
1932 return Err(format!(
1933 "TP KV ring rebase keeps {keep_rows} rows in {} physical rows",
1934 ring.rows()
1935 ));
1936 }
1937 let mut target_ring = *ring;
1938 target_ring.apply_rebase(new_base);
1939 if !target_ring.can_rewind_to(plan.transaction.base_len) {
1940 return Err(format!(
1941 "TP KV ring rebase to {new_base} laps transaction base {}",
1942 plan.transaction.base_len
1943 ));
1944 }
1945 *ring = target_ring;
1946 Ok(())
1947 }
1948 (Some(_), Some(KvRingAppend::Contiguous { .. })) | (None, None) => Ok(()),
1949 _ => Err("TP KV append plan does not match cache ring layout".into()),
1950 }
1951 }
1952
1953 pub fn publish_append_plan(&mut self, plan: TpKvAppendPlan) -> Result<(), String> {
1954 if let Some(KvRingAppend::Rebase { new_base, .. }) = plan.ring_append {
1955 if self.ring.as_ref().map(KvRing::base) != Some(new_base) {
1956 return Err(format!(
1957 "TP KV append rebase {new_base} was not published before its state"
1958 ));
1959 }
1960 }
1961 self.state.publish_append(plan.transaction, plan.target)
1962 }
1963
1964 pub fn publish_hydration(
1965 &mut self,
1966 logical_len: usize,
1967 resident_start: usize,
1968 ) -> Result<(), Box<dyn std::error::Error>> {
1969 if self.state != TpKvTransactionState::new() {
1970 return Err("TP KV hydration target must be fresh".into());
1971 }
1972 if resident_start > logical_len || logical_len > self.capacity {
1973 return Err(format!(
1974 "TP KV hydration range [{resident_start},{logical_len}) exceeds capacity {}",
1975 self.capacity
1976 )
1977 .into());
1978 }
1979 match self.ring.as_mut() {
1980 Some(ring) => {
1981 let rows = logical_len - resident_start;
1982 if rows > ring.rows() {
1983 return Err(format!(
1984 "TP KV hydration requires {rows} rows in a {}-row ring",
1985 ring.rows()
1986 )
1987 .into());
1988 }
1989 let mut hydrated_ring = *ring;
1990 hydrated_ring.apply_rebase(resident_start);
1991 if !hydrated_ring.can_rewind_to(logical_len) {
1992 return Err(format!(
1993 "TP KV hydration base {resident_start} cannot expose logical length \
1994 {logical_len}"
1995 )
1996 .into());
1997 }
1998 *ring = hydrated_ring;
1999 }
2000 None if resident_start != 0 => {
2001 return Err("linear TP KV hydration must start at absolute row zero".into());
2002 }
2003 None => {}
2004 }
2005 self.rewind_to(logical_len)
2006 }
2007
2008 pub fn append_target(
2009 &self,
2010 transaction: TpKvTransaction,
2011 rows: usize,
2012 ) -> Result<usize, String> {
2013 self.state.append_target(transaction, rows, self.capacity)
2014 }
2015
2016 pub fn publish_append(
2017 &mut self,
2018 transaction: TpKvTransaction,
2019 target: usize,
2020 ) -> Result<(), String> {
2021 self.state.publish_append(transaction, target)
2022 }
2023
2024 pub fn commit_target(
2025 &self,
2026 transaction: TpKvTransaction,
2027 accepted_rows: usize,
2028 ) -> Result<usize, String> {
2029 self.state.commit_target(transaction, accepted_rows)
2030 }
2031
2032 pub fn validate_transaction(&self, transaction: TpKvTransaction) -> Result<(), String> {
2033 self.state.validate(transaction)
2034 }
2035
2036 pub fn publish_finalize(
2037 &mut self,
2038 transaction: TpKvTransaction,
2039 target: usize,
2040 ) -> Result<(), String> {
2041 if !self.can_rewind_to(target) {
2042 return Err(format!(
2043 "TP KV finalize target {target} is outside the resident cache window/capacity"
2044 ));
2045 }
2046 self.state.publish_finalize(transaction, target)
2047 }
2048
2049 pub fn rewind_to(&mut self, target: usize) -> Result<(), Box<dyn std::error::Error>> {
2050 if !self.can_rewind_to(target) {
2051 return Err(format!(
2052 "TP KV rewind target {target} is outside the resident cache window/capacity"
2053 )
2054 .into());
2055 }
2056 let target_i32 =
2057 i32::try_from(target).map_err(|_| "TP KV length exceeds i32 device mirror")?;
2058 for rank in &mut self.ranks {
2059 let stream = rank.len_d.stream().clone();
2060 stream.memcpy_htod(&[target_i32], &mut rank.len_d)?;
2061 }
2062 self.state.rewind(target, self.capacity)?;
2063 Ok(())
2064 }
2065}
2066
2067pub struct Cache {
2068 pub kv: Vec<Option<KvLayer>>,
2069 pub recur: Vec<Option<RecurLayer>>,
2070 pub latent: Vec<Option<LatentKvLayer>>,
2073 pub tp_kv: Vec<Option<ResidentTpKvCache>>,
2076 pub glm5_tp_recur: Vec<Option<Vec<RecurLayer>>>,
2086 pub glm5_tp_latent_peer: Vec<Option<Vec<LatentKvLayer>>>,
2091 pub pos: usize,
2092 pub max_ctx: usize,
2093 pub tainted: bool,
2096 pub last_logits_dev: Option<CudaSlice<f32>>,
2104 pub dflash_taps: Option<DflashTapSink>,
2109 pub hc_taps: Option<HcTapSink>,
2117}
2118
2119#[derive(Clone, Copy, Debug, PartialEq, Eq)]
2123enum FullAttentionClass {
2124 Ordinary,
2125 GemmaGlobal,
2126 GemmaWindowed,
2127}
2128
2129fn full_attention_class(plan: &ModelPlan, il: u32) -> FullAttentionClass {
2130 let layer = plan
2131 .layers
2132 .iter()
2133 .chain(plan.mtp_blocks.iter().map(|block| &block.layer))
2134 .find(|layer| layer.index == il)
2135 .unwrap_or_else(|| panic!("ModelPlan has no layer {il}"));
2136 if !matches!(layer.residual, ResidualTopology::Gemma { .. }) {
2137 return FullAttentionClass::Ordinary;
2138 }
2139 match layer.state {
2140 StatePlan::SlidingKvCache { .. } => FullAttentionClass::GemmaWindowed,
2141 StatePlan::KvCache { .. } => FullAttentionClass::GemmaGlobal,
2142 _ => panic!("Gemma layer {il} does not declare a KV-cache state"),
2143 }
2144}
2145
2146fn full_attention_kv_layout(
2147 cfg: &ModelConfig,
2148 plan: &ModelPlan,
2149 il: u32,
2150) -> (usize, usize, usize, usize) {
2151 debug_assert_eq!(cfg.layer_kind(il), LayerKind::FullAttention);
2152 let class = full_attention_class(plan, il);
2153 let n_head_kv = cfg.n_head_kv as usize;
2154 let (kv_dim_k, kv_dim_v) = match class {
2155 FullAttentionClass::GemmaGlobal | FullAttentionClass::GemmaWindowed => {
2156 let g = cfg
2157 .gemma4
2158 .as_ref()
2159 .expect("Gemma ModelPlan layer requires Gemma cache geometry");
2160 let hd = match class {
2161 FullAttentionClass::GemmaWindowed => g.key_length_swa,
2162 FullAttentionClass::GemmaGlobal => g.key_length_global,
2163 FullAttentionClass::Ordinary => unreachable!(),
2164 } as usize;
2165 let d = match g.head_count_kv.get(il as usize) {
2174 Some(n) => hd * *n as usize,
2175 None => hd * n_head_kv,
2176 };
2177 (d, d)
2178 }
2179 FullAttentionClass::Ordinary => (
2180 cfg.head_dim_k as usize * n_head_kv,
2181 cfg.head_dim_v as usize * n_head_kv,
2182 ),
2183 };
2184 assert!(
2185 kv_dim_k % 32 == 0 && kv_dim_v % 32 == 0,
2186 "KVQUANT requires per-layer kv_dim_k%32==0 && kv_dim_v%32==0 \
2187 (layer {il}: k={kv_dim_k} v={kv_dim_v})"
2188 );
2189 let (kbb, vbb) = kv_blk_bytes();
2190 let g4_global_fp8 = gkv_on() && class == FullAttentionClass::GemmaGlobal;
2191 let g4_windowed_fp8 = wkv_on() && class == FullAttentionClass::GemmaWindowed;
2192 let qwen_fp8 = kv_fp8_on() && class == FullAttentionClass::Ordinary;
2193 let (kbb_l, vbb_l) = if g4_global_fp8 || g4_windowed_fp8 || qwen_fp8 {
2194 (32, 32)
2195 } else {
2196 (kbb, vbb)
2197 };
2198 (kv_dim_k, kv_dim_v, kbb_l, vbb_l)
2199}
2200
2201fn kv_plane_allocation_bytes(rows: usize, token_bytes: usize) -> usize {
2202 rows * token_bytes + 8
2203}
2204
2205pub fn cache_bytes_per_token(cfg: &ModelConfig) -> usize {
2212 cache_bytes_per_token_for_layers(cfg, 0, cfg.n_layer as usize)
2213}
2214
2215pub fn cache_bytes_per_token_for_layers(cfg: &ModelConfig, lo: usize, hi: usize) -> usize {
2219 let plan = ModelPlan::compile(cfg).expect("cache sizing requires a compilable ModelPlan");
2220 cache_bytes_per_token_for_plan(cfg, &plan, lo, hi)
2221}
2222
2223pub fn cache_bytes_per_token_for_plan(
2224 cfg: &ModelConfig,
2225 plan: &ModelPlan,
2226 lo: usize,
2227 hi: usize,
2228) -> usize {
2229 assert!(
2230 lo <= hi && hi <= cfg.n_layer as usize,
2231 "cache layer range out of bounds"
2232 );
2233 let shared = cfg.gemma4.as_ref().map(|g| g.shared_kv_layers).unwrap_or(0);
2234 let full_attn: usize = (lo as u32..hi as u32)
2235 .filter(|&il| cfg.layer_kind(il) == LayerKind::FullAttention)
2236 .filter(|&il| shared == 0 || il < cfg.n_layer - shared)
2237 .map(|il| {
2238 let (kv_dim_k, kv_dim_v, kbb, vbb) = full_attention_kv_layout(cfg, plan, il);
2239 (kv_dim_k / 32) * kbb + (kv_dim_v / 32) * vbb
2240 })
2241 .sum();
2242 full_attn + latent_kv_bytes_per_token_for_plan(cfg, plan, lo, hi)
2243}
2244
2245pub fn latent_kv_bytes_per_token_for_plan(
2275 cfg: &ModelConfig,
2276 plan: &ModelPlan,
2277 lo: usize,
2278 hi: usize,
2279) -> usize {
2280 let ring_disabled = std::env::var("MEMRA_DSA_INDEX_RING")
2281 .ok()
2282 .and_then(|v| v.trim().parse::<usize>().ok())
2283 == Some(0);
2284 plan.layers
2285 .iter()
2286 .filter(|layer| (lo..hi).contains(&(layer.index as usize)))
2287 .map(|layer| match layer.state {
2288 StatePlan::LatentKvCache { width, index_width } => {
2289 let latent = width as usize * std::mem::size_of::<f32>();
2290 let index_width = index_width as usize;
2291 let pool = cfg
2292 .glm5
2293 .as_ref()
2294 .map(|g| g.index_kpool as usize)
2295 .filter(|&p| p > 0)
2296 .unwrap_or(1);
2297 let pool_keys = if index_width > 0 {
2299 (index_width / 2) * std::mem::size_of::<f32>() / pool
2300 } else {
2301 0
2302 };
2303 let flat_plane = if index_width > 0 && ring_disabled {
2304 index_width * std::mem::size_of::<f32>()
2305 } else {
2306 0
2307 };
2308 latent + pool_keys + flat_plane
2309 }
2310 _ => 0,
2311 })
2312 .sum()
2313}
2314
2315pub fn cache_ring_bytes_per_token(cfg: &ModelConfig) -> usize {
2318 cache_ring_bytes_per_token_for_layers(cfg, 0, cfg.n_layer as usize)
2319}
2320
2321pub fn cache_ring_bytes_per_token_for_layers(cfg: &ModelConfig, lo: usize, hi: usize) -> usize {
2323 assert!(
2324 lo <= hi && hi <= cfg.n_layer as usize,
2325 "cache layer range out of bounds"
2326 );
2327 let Ok(plan) = memra_gguf::model_plan::ModelPlan::compile(cfg) else {
2328 return 0;
2329 };
2330 cache_ring_bytes_per_token_for_plan(cfg, &plan, lo, hi)
2331}
2332
2333pub fn cache_ring_bytes_per_token_for_plan(
2334 cfg: &ModelConfig,
2335 plan: &ModelPlan,
2336 lo: usize,
2337 hi: usize,
2338) -> usize {
2339 let total = plan.layers.len() + plan.mtp_blocks.len();
2340 assert!(
2341 lo <= hi && hi <= total,
2342 "cache plan layer range out of bounds"
2343 );
2344 if !swa_ring_on() {
2345 return 0;
2346 }
2347 let shared = cfg.gemma4.as_ref().map(|g| g.shared_kv_layers).unwrap_or(0);
2348 plan.layers
2349 .iter()
2350 .chain(plan.mtp_blocks.iter().map(|block| &block.layer))
2351 .filter(|layer| (lo..hi).contains(&(layer.index as usize)))
2352 .filter(|layer| {
2353 matches!(
2354 layer.state,
2355 memra_gguf::model_plan::StatePlan::SlidingKvCache { .. }
2356 )
2357 })
2358 .filter(|layer| shared == 0 || layer.index < cfg.n_layer - shared)
2359 .map(|layer| {
2360 let (kv_dim_k, kv_dim_v, kbb, vbb) = full_attention_kv_layout(cfg, plan, layer.index);
2361 (kv_dim_k / 32) * kbb + (kv_dim_v / 32) * vbb
2362 })
2363 .sum()
2364}
2365
2366pub fn cache_ring_row_cap(cfg: &ModelConfig) -> usize {
2368 let Ok(plan) = memra_gguf::model_plan::ModelPlan::compile(cfg) else {
2369 return 0;
2370 };
2371 cache_ring_row_cap_for_plan(&plan)
2372}
2373
2374pub fn cache_ring_row_cap_for_plan(plan: &memra_gguf::model_plan::ModelPlan) -> usize {
2375 if !swa_ring_on() {
2376 return 0;
2377 }
2378 plan.layers
2379 .iter()
2380 .chain(plan.mtp_blocks.iter().map(|block| &block.layer))
2381 .filter_map(|layer| match layer.state {
2382 memra_gguf::model_plan::StatePlan::SlidingKvCache { window, .. } => {
2383 Some(window as usize)
2384 }
2385 _ => None,
2386 })
2387 .map(|window| swa_ring_rows(window, usize::MAX))
2388 .max()
2389 .unwrap_or(0)
2390}
2391
2392pub struct DflashTapSink {
2395 pub layer_ids: Vec<usize>,
2396 pub buf: CudaSlice<f32>,
2397 pub hidden: usize,
2398 pub t: usize,
2399 pub base: usize,
2402}
2403
2404pub struct HcTapSink {
2411 pub layer_ids: Vec<usize>,
2413 pub rows: Vec<f32>,
2415 pub hidden: usize,
2416 pub t: usize,
2418 pub base: usize,
2421 pub origin: usize,
2427 pub dev: Vec<Option<CudaSlice<f32>>>,
2435 pub device_stage: bool,
2440}
2441
2442impl HcTapSink {
2443 pub fn new(layer_ids: Vec<usize>, hidden: usize, t: usize) -> Self {
2444 let n_taps = layer_ids.len();
2445 Self {
2446 layer_ids,
2447 rows: vec![0.0; t * n_taps * hidden],
2448 hidden,
2449 t,
2450 base: 0,
2451 origin: 0,
2452 dev: (0..n_taps).map(|_| None).collect(),
2453 device_stage: false,
2454 }
2455 }
2456
2457 pub fn new_at(layer_ids: Vec<usize>, hidden: usize, t: usize, origin: usize) -> Self {
2460 Self {
2461 origin,
2462 ..Self::new(layer_ids, hidden, t)
2463 }
2464 }
2465
2466 pub fn new_device_staged(layer_ids: Vec<usize>, hidden: usize, t: usize) -> Self {
2469 Self {
2470 device_stage: true,
2471 ..Self::new(layer_ids, hidden, t)
2472 }
2473 }
2474}
2475
2476pub struct CacheSnapshot {
2506 pub kv_len: Vec<Option<usize>>, pub tp_kv_len: Vec<Option<usize>>, pub conv: Vec<Option<CudaSlice<f32>>>, pub ssm: Vec<Option<CudaSlice<f32>>>,
2510 pub pos: usize,
2511}
2512
2513impl Cache {
2514 pub fn ensure_usable(&self, path: &str) -> Result<(), Box<dyn std::error::Error>> {
2515 if self.tainted {
2516 return Err(format!(
2517 "{path}: cache was tainted by a failed pipeline wave and cannot be reused"
2518 )
2519 .into());
2520 }
2521 Ok(())
2522 }
2523
2524 pub fn mark_tainted(&mut self) {
2525 self.tainted = true;
2526 self.last_logits_dev = None;
2527 self.dflash_taps = None;
2528 }
2529
2530 pub fn new(
2532 e: &impl KvDev,
2533 cfg: &ModelConfig,
2534 max_ctx: usize,
2535 ) -> Result<Self, Box<dyn std::error::Error>> {
2536 Self::new_inner(&|_| e, cfg, None, max_ctx)
2537 }
2538
2539 pub fn new_planned(
2540 e: &impl KvDev,
2541 cfg: &ModelConfig,
2542 plan: &memra_gguf::model_plan::ModelPlan,
2543 max_ctx: usize,
2544 ) -> Result<Self, Box<dyn std::error::Error>> {
2545 Self::new_inner(&|_| e, cfg, Some(plan), max_ctx)
2546 }
2547
2548 pub fn new_pp2(
2553 dev0: &dyn KvDev,
2554 dev1: &dyn KvDev,
2555 split: usize,
2556 cfg: &ModelConfig,
2557 max_ctx: usize,
2558 ) -> Result<Self, Box<dyn std::error::Error>> {
2559 Self::new_inner(
2560 &|il| if il < split { dev0 } else { dev1 },
2561 cfg,
2562 None,
2563 max_ctx,
2564 )
2565 }
2566
2567 pub fn new_ppn(
2573 devs: &[&dyn KvDev],
2574 fence: &[usize],
2575 cfg: &ModelConfig,
2576 max_ctx: usize,
2577 ) -> Result<Self, Box<dyn std::error::Error>> {
2578 assert_eq!(
2579 devs.len() + 1,
2580 fence.len(),
2581 "ppn cache: devs vs fence mismatch"
2582 );
2583 let pick = |il: usize| -> &dyn KvDev {
2584 let s = match fence[1..fence.len() - 1].binary_search(&il) {
2585 Ok(k) => k + 1,
2586 Err(k) => k,
2587 };
2588 devs[s.min(devs.len() - 1)]
2589 };
2590 Self::new_inner(&pick, cfg, None, max_ctx)
2591 }
2592
2593 pub fn new_ppn_planned(
2594 devs: &[&dyn KvDev],
2595 fence: &[usize],
2596 cfg: &ModelConfig,
2597 plan: &memra_gguf::model_plan::ModelPlan,
2598 max_ctx: usize,
2599 ) -> Result<Self, Box<dyn std::error::Error>> {
2600 assert_eq!(
2601 devs.len() + 1,
2602 fence.len(),
2603 "ppn cache: devs vs fence mismatch"
2604 );
2605 let pick = |il: usize| -> &dyn KvDev {
2606 let stage = match fence[1..fence.len() - 1].binary_search(&il) {
2607 Ok(index) => index + 1,
2608 Err(index) => index,
2609 };
2610 devs[stage.min(devs.len() - 1)]
2611 };
2612 Self::new_inner(&pick, cfg, Some(plan), max_ctx)
2613 }
2614
2615 fn new_inner<'a>(
2618 pick: &dyn Fn(usize) -> &'a dyn KvDev,
2619 cfg: &ModelConfig,
2620 plan: Option<&memra_gguf::model_plan::ModelPlan>,
2621 max_ctx: usize,
2622 ) -> Result<Self, Box<dyn std::error::Error>> {
2623 let fallback_plan = if plan.is_none() {
2624 Some(ModelPlan::compile(cfg)?)
2625 } else {
2626 None
2627 };
2628 let plan = plan
2629 .or(fallback_plan.as_ref())
2630 .expect("cache allocation requires a ModelPlan");
2631 let n = cfg.n_layer as usize;
2632 let mut kv = Vec::with_capacity(n);
2633 let mut recur = Vec::with_capacity(n);
2634 let mut latent = Vec::with_capacity(n);
2635 let head_dim_k = cfg.head_dim_k as usize;
2636 let head_dim_v = cfg.head_dim_v as usize;
2637 for il in 0..cfg.n_layer {
2638 let e = pick(il as usize);
2640 let layer = plan
2641 .layers
2642 .iter()
2643 .chain(plan.mtp_blocks.iter().map(|block| &block.layer))
2644 .find(|layer| layer.index == il)
2645 .ok_or_else(|| format!("cache ModelPlan has no layer {il}"))?;
2646 let g4_shared = cfg.gemma4.as_ref().map(|g| g.shared_kv_layers).unwrap_or(0);
2651 if g4_shared > 0 && il >= cfg.n_layer - g4_shared {
2652 kv.push(None);
2653 recur.push(None);
2654 latent.push(None);
2655 continue;
2656 }
2657 match layer.state {
2658 StatePlan::KvCache { .. } | StatePlan::SlidingKvCache { .. } => {
2659 assert!(
2664 head_dim_k.is_multiple_of(32) && head_dim_v.is_multiple_of(32),
2665 "KVQUANT requires head_dim_k%32==0 && head_dim_v%32==0 \
2666 (layer {il}: k={head_dim_k} v={head_dim_v})"
2667 );
2668 let (kv_dim_k, kv_dim_v, kbb_l, vbb_l) =
2671 full_attention_kv_layout(cfg, plan, il);
2672 let k_tok_bytes = (kv_dim_k / 32) * kbb_l;
2673 let v_tok_bytes = (kv_dim_v / 32) * vbb_l;
2674 let planned_window = plan
2675 .layers
2676 .iter()
2677 .chain(plan.mtp_blocks.iter().map(|block| &block.layer))
2678 .find(|layer| layer.index == il)
2679 .and_then(|layer| match layer.state {
2680 StatePlan::SlidingKvCache { window, .. } => Some(window),
2681 _ => None,
2682 });
2683 let ring = if swa_ring_on() {
2684 planned_window.map(|window| {
2685 let window = window as usize;
2686 KvRing::new(swa_ring_rows(window, max_ctx), window)
2687 })
2688 } else {
2689 None
2690 };
2691 let alloc_rows = ring.as_ref().map(KvRing::rows).unwrap_or(max_ctx);
2692 kv.push(Some(KvLayer {
2693 k: e.alloc_u8(kv_plane_allocation_bytes(alloc_rows, k_tok_bytes))?,
2697 v: e.alloc_u8(kv_plane_allocation_bytes(alloc_rows, v_tok_bytes))?,
2698 kv_dim_k,
2699 kv_dim_v,
2700 k_tok_bytes,
2701 v_tok_bytes,
2702 len: 0,
2703 ring,
2704 len_d: e.htod_i32(&[0])?,
2705 base_d: None,
2706 }));
2707 recur.push(None);
2708 latent.push(None);
2709 }
2710 StatePlan::Recurrent {
2711 conv_width,
2712 conv_kernel,
2713 state_width,
2714 } => {
2715 kv.push(None);
2716 recur.push(Some(RecurLayer {
2717 conv_state: e.zeros(
2718 conv_width as usize * (conv_kernel as usize).saturating_sub(1),
2719 )?,
2720 ssm_state: e.zeros(state_width as usize)?,
2721 ssm_state_alt: e.zeros(state_width as usize)?,
2722 }));
2723 latent.push(None);
2724 }
2725 StatePlan::LatentKvCache { width, index_width } => {
2726 let width = width as usize;
2730 assert!(
2731 width > 0,
2732 "layer {il}: LatentKvCache width must be positive"
2733 );
2734 kv.push(None);
2735 recur.push(None);
2736 let index_width = index_width as usize;
2737 let index_ring = if index_width == 0 {
2741 None
2742 } else {
2743 index_ring_rows(max_ctx)
2744 };
2745 let index_rows = match index_width {
2746 0 => None,
2747 w => Some(e.zeros(index_ring.unwrap_or(max_ctx) * w)?),
2748 };
2749 latent.push(Some(LatentKvLayer {
2750 rows: e.zeros(max_ctx * width)?,
2751 width,
2752 len: 0,
2753 len_d: e.htod_i32(&[0])?,
2754 index_rows,
2755 index_width,
2756 index_ring_rows: index_ring,
2757 index_pool_keys: None,
2760 index_pools_ready: 0,
2761 index_pool: 0,
2762 }));
2763 }
2764 ref state => {
2765 return Err(format!(
2766 "native cache allocator has no implementation for layer {il} state {state:?}"
2767 )
2768 .into());
2769 }
2770 }
2771 }
2772 Ok(Cache {
2773 kv,
2774 recur,
2775 latent,
2776 tp_kv: (0..n).map(|_| None).collect(),
2777 glm5_tp_recur: (0..n).map(|_| None).collect(),
2778 glm5_tp_latent_peer: (0..n).map(|_| None).collect(),
2779 pos: 0,
2780 max_ctx,
2781 tainted: false,
2782 dflash_taps: None,
2783 hc_taps: None,
2784 last_logits_dev: None,
2785 })
2786 }
2787
2788 pub fn has_swa_ring(&self) -> bool {
2789 self.kv.iter().flatten().any(|layer| layer.ring.is_some())
2790 || self
2791 .tp_kv
2792 .iter()
2793 .flatten()
2794 .any(|layer| layer.ring_window().is_some())
2795 }
2796
2797 pub fn can_rollback(&self, snap: &CacheSnapshot, accept_len: usize) -> bool {
2798 let local = self
2799 .kv
2800 .iter()
2801 .zip(&snap.kv_len)
2802 .all(|(layer, saved)| match (layer, saved) {
2803 (Some(layer), Some(saved)) => layer
2804 .ring
2805 .as_ref()
2806 .is_none_or(|ring| ring.can_rewind_to(saved + accept_len)),
2807 _ => true,
2808 });
2809 let tensor = self
2810 .tp_kv
2811 .iter()
2812 .zip(&snap.tp_kv_len)
2813 .all(|(layer, saved)| match (layer, saved) {
2814 (Some(layer), Some(saved)) => saved
2815 .checked_add(accept_len)
2816 .is_some_and(|target| layer.can_rewind_to(target)),
2817 _ => true,
2818 });
2819 local && tensor
2820 }
2821
2822 pub fn snapshot(&self, e: &impl KvDev) -> Result<CacheSnapshot, Box<dyn std::error::Error>> {
2826 if self.glm5_tp_recur.iter().any(Option::is_some)
2830 || self.glm5_tp_latent_peer.iter().any(Option::is_some)
2831 {
2832 return Err(
2833 "cache snapshot is unwired for glm5 TP rank state (MEMRA_GLM5_TP): \
2834 per-rank planes are not carried by CacheSnapshot"
2835 .into(),
2836 );
2837 }
2838 self.ensure_usable("cache snapshot")?;
2839 let n = self.kv.len();
2840 let mut kv_len = Vec::with_capacity(n);
2841 let mut tp_kv_len = Vec::with_capacity(n);
2842 let mut conv = Vec::with_capacity(n);
2843 let mut ssm = Vec::with_capacity(n);
2844 for il in 0..n {
2845 match &self.kv[il] {
2846 Some(kvl) => kv_len.push(Some(kvl.len)),
2847 None => kv_len.push(None),
2848 }
2849 tp_kv_len.push(
2850 self.tp_kv[il]
2851 .as_ref()
2852 .map(ResidentTpKvCache::committed_len),
2853 );
2854 match &self.recur[il] {
2855 Some(rl) => {
2856 conv.push(Some(e.clone_dtod(&rl.conv_state)?));
2857 ssm.push(Some(e.clone_dtod(&rl.ssm_state)?));
2858 }
2859 None => {
2860 conv.push(None);
2861 ssm.push(None);
2862 }
2863 }
2864 }
2865 Ok(CacheSnapshot {
2866 kv_len,
2867 tp_kv_len,
2868 conv,
2869 ssm,
2870 pos: self.pos,
2871 })
2872 }
2873
2874 pub fn snapshot_into(
2879 &self,
2880 e: &impl KvDev,
2881 snap: &mut CacheSnapshot,
2882 ) -> Result<(), Box<dyn std::error::Error>> {
2883 if self.glm5_tp_recur.iter().any(Option::is_some)
2884 || self.glm5_tp_latent_peer.iter().any(Option::is_some)
2885 {
2886 return Err("cache snapshot_into is unwired for glm5 TP rank state \
2887 (MEMRA_GLM5_TP): per-rank planes are not carried by CacheSnapshot"
2888 .into());
2889 }
2890 self.ensure_usable("cache snapshot refresh")?;
2891 let n = self.kv.len();
2892 for il in 0..n {
2893 snap.kv_len[il] = self.kv[il].as_ref().map(|kvl| kvl.len);
2894 snap.tp_kv_len[il] = self.tp_kv[il]
2895 .as_ref()
2896 .map(ResidentTpKvCache::committed_len);
2897 if let Some(rl) = &self.recur[il] {
2898 let dc = snap.conv[il]
2899 .as_mut()
2900 .expect("snapshot_into: shape mismatch (conv)");
2901 let ds = snap.ssm[il]
2902 .as_mut()
2903 .expect("snapshot_into: shape mismatch (ssm)");
2904 let (cn, sn) = (rl.conv_state.len(), rl.ssm_state.len());
2905 e.copy_into(dc, 0, &rl.conv_state, cn)?;
2906 e.copy_into(ds, 0, &rl.ssm_state, sn)?;
2907 }
2908 }
2909 snap.pos = self.pos;
2910 Ok(())
2911 }
2912
2913 pub fn rollback(
2921 &mut self,
2922 e: &impl KvDev,
2923 snap: &CacheSnapshot,
2924 accept_len: usize,
2925 ) -> Result<(), Box<dyn std::error::Error>> {
2926 self.ensure_usable("cache rollback")?;
2927 if !self.can_rollback(snap, accept_len) {
2928 return Err(
2929 "SWA ring rewind checkpoint has been lapped; full re-prime required".into(),
2930 );
2931 }
2932 for il in 0..self.kv.len() {
2933 if let (Some(kvl), Some(saved)) = (self.kv[il].as_mut(), snap.kv_len[il]) {
2934 kvl.len = saved + accept_len;
2935 e.set_i32_one(&mut kvl.len_d, kvl.len as i32)?;
2940 }
2941 if let (Some(kvl), Some(saved)) = (self.tp_kv[il].as_mut(), snap.tp_kv_len[il]) {
2942 kvl.rewind_to(saved + accept_len)?;
2943 }
2944 if let Some(rl) = self.recur[il].as_mut() {
2945 if let Some(c) = &snap.conv[il] {
2946 e.copy_into(&mut rl.conv_state, 0, c, c.len())?;
2947 }
2948 if let Some(s) = &snap.ssm[il] {
2949 e.copy_into(&mut rl.ssm_state, 0, s, s.len())?;
2950 }
2951 }
2952 }
2953 self.pos = snap.pos;
2954 Ok(())
2955 }
2956}
2957
2958#[cfg(test)]
2959mod tp_transaction_tests {
2960 use super::{
2961 Cache, INDEX_RING_WORKING_ROWS, KvRingAppend, ResidentTpKvCache, TpKvTransactionState,
2962 index_ring_default_rows, index_ring_rows_for, index_ring_take, tp_kv_rank_allocation_shape,
2963 };
2964
2965 const GLM5_NEXT_POOL: usize = 4;
2968 const GLM5_NEXT_STATE_ROW_BYTES: usize = 2 * 128 * 4;
2970 const RING_BYTES_PER_LAYER: usize = INDEX_RING_WORKING_ROWS * GLM5_NEXT_STATE_ROW_BYTES;
2973
2974 #[test]
2993 fn the_derived_ring_serves_a_monolithic_prime_of_the_whole_configured_context() {
2994 for max_ctx in [8192usize, 262_144, 1 << 20] {
2997 let rows = index_ring_default_rows(max_ctx).unwrap_or_else(|| {
2998 panic!("the ring must engage at max_ctx {max_ctx}: it is where the saving is")
2999 });
3000 let ring = rows / GLM5_NEXT_POOL * GLM5_NEXT_POOL;
3003
3004 let mut cur = 0usize;
3006 let mut pools_ready = 0usize;
3007 let mut steps = 0usize;
3008 while cur < max_ctx {
3009 let take = index_ring_take(ring, GLM5_NEXT_POOL, pools_ready, cur, max_ctx - cur)
3010 .unwrap_or_else(|| {
3011 panic!(
3012 "MEMRA_CTX={max_ctx}: the {ring}-row ring refused a monolithic prime \
3013 after {cur} of {max_ctx} tokens ({}% of the configured context). \
3014 USABLE CONTEXT MUST BE AT LEAST CONFIGURED CONTEXT. This is the \
3015 4630-of-8192 regression, in arithmetic.",
3016 cur * 100 / max_ctx
3017 )
3018 });
3019 assert!(
3020 take > 0,
3021 "MEMRA_CTX={max_ctx}: the drain made no progress at row {cur}: a zero take \
3022 is an infinite loop in the engine, not a refusal"
3023 );
3024 cur += take;
3025 pools_ready = cur / GLM5_NEXT_POOL;
3026 steps += 1;
3027 assert!(
3028 steps <= max_ctx,
3029 "MEMRA_CTX={max_ctx}: the drain did not terminate"
3030 );
3031 }
3032 assert_eq!(cur, max_ctx, "the whole prompt must be appended");
3033
3034 let ring_bytes = rows * GLM5_NEXT_STATE_ROW_BYTES;
3040 let flat_bytes = max_ctx * GLM5_NEXT_STATE_ROW_BYTES;
3041 assert_eq!(
3042 ring_bytes, RING_BYTES_PER_LAYER,
3043 "MEMRA_CTX={max_ctx}: the ring books {rows} rows, not the context-independent \
3044 {INDEX_RING_WORKING_ROWS}. A plane that tracks max_ctx is the flat plane \
3045 wearing a modulus"
3046 );
3047 assert!(
3048 rows < max_ctx,
3049 "MEMRA_CTX={max_ctx}: a ring of {rows} rows is not shorter than the flat plane \
3050 it replaces, so it would not engage at all"
3051 );
3052 assert!(
3056 ring_bytes <= 16 << 20,
3057 "MEMRA_CTX={max_ctx}: {} MiB per MLA layer, over glm5_next's 12 of them. The ring \
3058 exists to delete 11.94 GiB; a working set this large is not paying for itself",
3059 ring_bytes >> 20
3060 );
3061 println!(
3062 "MEMRA_CTX={max_ctx}: ring {rows} rows (effective {ring}), monolithic prime of \
3063 {max_ctx} tokens admitted in {steps} drain step(s); plane {} MiB/layer vs flat \
3064 {} MiB/layer",
3065 ring_bytes >> 20,
3066 flat_bytes >> 20
3067 );
3068 }
3069 }
3070
3071 fn empty_tp_cache(capacity: usize) -> ResidentTpKvCache {
3072 ResidentTpKvCache::new(Vec::new(), 128, 128, 136, 96, capacity)
3073 }
3074
3075 #[test]
3076 fn step_tp8_rank_allocation_matches_the_official_kv_geometry() {
3077 let shape = tp_kv_rank_allocation_shape(8 * 128, 8 * 128, 8).unwrap();
3078 assert_eq!((shape.kv_dim_k, shape.kv_dim_v), (128, 128));
3079 assert_eq!((shape.k_token_bytes, shape.v_token_bytes), (136, 96));
3080 assert_eq!(shape.bytes_per_token(), 232);
3081 assert_eq!(shape.fixed_bytes, 20);
3082 assert_eq!(shape.allocation_bytes(262_144), 232 * 262_144 + 20);
3083 }
3084
3085 #[test]
3086 fn tp_rank_allocation_refuses_non_divisible_and_non_block_aligned_shards() {
3087 assert!(tp_kv_rank_allocation_shape(1024, 1024, 3).is_err());
3088 assert!(tp_kv_rank_allocation_shape(1024, 1024, 64).is_err());
3089 assert!(tp_kv_rank_allocation_shape(0, 1024, 8).is_err());
3090 }
3091
3092 #[test]
3093 fn partial_commit_publishes_only_the_accepted_prefix() {
3094 let mut state = TpKvTransactionState::new();
3095 let transaction = state.begin().unwrap();
3096 let staged = state.append_target(transaction, 3, 8).unwrap();
3097 state.publish_append(transaction, staged).unwrap();
3098 assert_eq!(state.committed_len, 0);
3099 assert_eq!(state.staged_len, 3);
3100
3101 let committed = state.commit_target(transaction, 2).unwrap();
3102 state.publish_finalize(transaction, committed).unwrap();
3103 assert_eq!(state.committed_len, 2);
3104 assert_eq!(state.staged_len, 2);
3105 assert!(state.active.is_none());
3106 assert!(state.validate(transaction).is_err());
3107 }
3108
3109 #[test]
3110 fn rollback_restores_the_committed_boundary() {
3111 let mut state = TpKvTransactionState::new();
3112 let first = state.begin().unwrap();
3113 let staged = state.append_target(first, 1, 8).unwrap();
3114 state.publish_append(first, staged).unwrap();
3115 let committed = state.commit_target(first, 1).unwrap();
3116 state.publish_finalize(first, committed).unwrap();
3117
3118 let speculative = state.begin().unwrap();
3119 let staged = state.append_target(speculative, 2, 8).unwrap();
3120 state.publish_append(speculative, staged).unwrap();
3121 assert_eq!(state.committed_len, 1);
3122 assert_eq!(state.staged_len, 3);
3123 state
3124 .publish_finalize(speculative, speculative.base_len)
3125 .unwrap();
3126 assert_eq!(state.committed_len, 1);
3127 assert_eq!(state.staged_len, 1);
3128 assert!(state.validate(speculative).is_err());
3129 }
3130
3131 #[test]
3132 fn index_ring_sizing_is_pure_and_carries_no_per_call_t() {
3133 let rows = INDEX_RING_WORKING_ROWS;
3136 assert_eq!(index_ring_rows_for(None, 1 << 20), Some(rows));
3137 assert_eq!(index_ring_default_rows(1 << 20), Some(rows));
3138 assert_eq!(index_ring_rows_for(None, 4096), None);
3141 assert_eq!(index_ring_rows_for(None, rows), None);
3142 assert_eq!(index_ring_rows_for(None, rows + 1), Some(rows));
3143
3144 for max_ctx in [8192usize, 262_144, 1 << 20] {
3149 assert_eq!(
3150 index_ring_rows_for(None, max_ctx),
3151 Some(INDEX_RING_WORKING_ROWS),
3152 "the derived ring must not vary with the configured context"
3153 );
3154 }
3155
3156 assert_eq!(index_ring_rows_for(Some(0), 1 << 20), None);
3159 assert_eq!(index_ring_rows_for(Some(16), 64), Some(16));
3160 assert_eq!(index_ring_rows_for(Some(64), 64), None);
3161 }
3162
3163 #[test]
3165 fn index_ring_take_drains_instead_of_bounding_the_call() {
3166 const POOL: usize = GLM5_NEXT_POOL;
3167 assert_eq!(index_ring_take(0, POOL, 0, 0, 1 << 20), Some(1 << 20));
3169 assert_eq!(index_ring_take(64, POOL, 0, 0, 1024), Some(64));
3172 for cur in 0..64usize {
3175 let ready = cur / POOL;
3176 let take = index_ring_take(64, POOL, ready, cur, 1024).expect("never lapses");
3177 assert!(
3178 (64 - POOL + 1..=64).contains(&take),
3179 "cur {cur}: take {take} outside the guaranteed progress band"
3180 );
3181 }
3182 assert_eq!(index_ring_take(64, POOL, 4, 16, 1), Some(1));
3184 assert_eq!(index_ring_take(16, POOL, 0, 64, 1), None);
3187 assert_eq!(index_ring_take(16, POOL, 0, 16, 1), None);
3188 assert_eq!(index_ring_take(16, POOL, 0, 15, 1), Some(1));
3189 }
3190
3191 #[test]
3192 fn rejects_nested_stale_and_out_of_range_actions() {
3193 let mut state = TpKvTransactionState::new();
3194 let transaction = state.begin().unwrap();
3195 assert!(state.begin().is_err());
3196 assert!(state.append_target(transaction, 0, 2).is_err());
3197 assert!(state.append_target(transaction, 3, 2).is_err());
3198 let staged = state.append_target(transaction, 2, 2).unwrap();
3199 state.publish_append(transaction, staged).unwrap();
3200 assert!(state.commit_target(transaction, 3).is_err());
3201 state.publish_finalize(transaction, 0).unwrap();
3202 assert!(state.publish_append(transaction, 1).is_err());
3203 }
3204
3205 #[test]
3206 fn rewind_resets_visibility_and_invalidates_an_active_transaction() {
3207 let mut state = TpKvTransactionState::new();
3208 let transaction = state.begin().unwrap();
3209 let staged = state.append_target(transaction, 3, 8).unwrap();
3210 state.publish_append(transaction, staged).unwrap();
3211 state.rewind(1, 8).unwrap();
3212 assert_eq!(state.committed_len, 1);
3213 assert_eq!(state.staged_len, 1);
3214 assert!(state.active.is_none());
3215 assert!(state.validate(transaction).is_err());
3216 assert!(state.rewind(9, 8).is_err());
3217 }
3218
3219 #[test]
3220 fn grow_preserves_generation_and_publishes_only_the_checkpoint_prefix() {
3221 let mut source = empty_tp_cache(8);
3222 let first = source.begin_transaction().unwrap();
3223 let staged = source.append_target(first, 5).unwrap();
3224 source.publish_append(first, staged).unwrap();
3225 let committed = source.commit_target(first, 5).unwrap();
3226 source.publish_finalize(first, committed).unwrap();
3227
3228 let rolled_back = source.begin_transaction().unwrap();
3229 source
3230 .publish_finalize(rolled_back, rolled_back.base_len())
3231 .unwrap();
3232 let plan = source.prepare_grow(16, 3).unwrap();
3233 assert_eq!(plan.rows(), 3);
3234 assert_eq!(plan.k_bytes(), 3 * 136);
3235 assert_eq!(plan.v_bytes(), 3 * 96);
3236
3237 let mut target = empty_tp_cache(16);
3238 target.publish_grow(plan).unwrap();
3239 assert_eq!(target.committed_len(), 3);
3240 assert_eq!(target.staged_len(), 3);
3241 assert_eq!(target.capacity(), 16);
3242 let next = target.begin_transaction().unwrap();
3243 assert_eq!(next.generation(), rolled_back.generation() + 1);
3244 assert_eq!(next.base_len(), 3);
3245 }
3246
3247 #[test]
3248 fn grow_refuses_active_source_and_invalid_target_state_or_layout() {
3249 let mut active = empty_tp_cache(8);
3250 active.begin_transaction().unwrap();
3251 assert!(active.prepare_grow(16, 0).is_err());
3252
3253 let mut source = empty_tp_cache(8);
3254 source.rewind_to(5).unwrap();
3255 assert!(source.prepare_grow(8, 5).is_err());
3256 assert!(source.prepare_grow(16, 6).is_err());
3257 let plan = source.prepare_grow(16, 4).unwrap();
3258
3259 let mut wrong_layout = ResidentTpKvCache::new(Vec::new(), 128, 128, 144, 96, 16);
3260 assert!(wrong_layout.publish_grow(plan).is_err());
3261
3262 let plan = source.prepare_grow(16, 4).unwrap();
3263 let mut dirty_target = empty_tp_cache(16);
3264 dirty_target.rewind_to(1).unwrap();
3265 assert!(dirty_target.publish_grow(plan).is_err());
3266 }
3267
3268 #[test]
3269 fn swa_transaction_rebase_preserves_the_rollback_window() {
3270 let mut cache = ResidentTpKvCache::new_swa(Vec::new(), 128, 128, 136, 96, 10_000, 32);
3271 assert_eq!(cache.physical_capacity(), 32 + 4096 + 512 + 31);
3272 cache.publish_hydration(8762, 4096).unwrap();
3276 assert_eq!(cache.ring_base(), Some(4096));
3277
3278 let transaction = cache.begin_transaction().unwrap();
3279 let plan = cache.prepare_append(transaction, 10).unwrap();
3280 assert_eq!(plan.target(), 8772);
3281 assert_eq!(plan.write_row(), 58);
3282 assert_eq!(
3283 plan.ring_append(),
3284 Some(KvRingAppend::Rebase {
3285 src_row: 4608,
3286 keep_rows: 58,
3287 new_base: 8704,
3288 write_row: 58,
3289 })
3290 );
3291 cache.publish_append_rebase(plan).unwrap();
3292 cache.publish_append_plan(plan).unwrap();
3293 assert_eq!(cache.ring_base(), Some(8704));
3294 assert_eq!(cache.physical_range(8740, 8772).unwrap(), 36..68);
3295
3296 let rollback = cache.commit_target(transaction, 0).unwrap();
3297 cache.publish_finalize(transaction, rollback).unwrap();
3298 assert_eq!((cache.committed_len(), cache.staged_len()), (8762, 8762));
3299 assert!(cache.rewind_to(8200).is_err());
3300 }
3301
3302 #[test]
3303 fn swa_grow_normalizes_only_the_live_prefix_and_preserves_generation() {
3304 let mut source = ResidentTpKvCache::new_swa(Vec::new(), 128, 128, 136, 96, 10_000, 32);
3305 source.publish_hydration(8250, 4096).unwrap();
3306 let transaction = source.begin_transaction().unwrap();
3307 source
3308 .publish_finalize(transaction, transaction.base_len())
3309 .unwrap();
3310
3311 let plan = source.prepare_grow(20_000, 8250).unwrap();
3312 assert_eq!(plan.source_row(), 4096);
3313 assert_eq!(plan.copy_rows(), 58);
3314 assert_eq!(plan.k_bytes(), 58 * 136);
3315 assert_eq!(plan.v_bytes(), 58 * 96);
3316
3317 let mut target = ResidentTpKvCache::new_swa(Vec::new(), 128, 128, 136, 96, 20_000, 32);
3318 target.publish_grow(plan).unwrap();
3319 assert_eq!(target.ring_base(), Some(8192));
3320 assert_eq!((target.committed_len(), target.staged_len()), (8250, 8250));
3321 assert_eq!(target.physical_range(8192, 8250).unwrap(), 0..58);
3322 let next = target.begin_transaction().unwrap();
3323 assert_eq!(next.generation(), transaction.generation() + 1);
3324 }
3325
3326 #[test]
3327 fn cache_reports_a_materialized_distributed_swa_ring() {
3328 let mut cache = Cache {
3329 kv: Vec::new(),
3330 recur: Vec::new(),
3331 latent: Vec::new(),
3332 tp_kv: vec![None],
3333 glm5_tp_recur: vec![None],
3334 glm5_tp_latent_peer: vec![None],
3335 pos: 0,
3336 max_ctx: 10_000,
3337 tainted: false,
3338 dflash_taps: None,
3339 hc_taps: None,
3340 last_logits_dev: None,
3341 };
3342 assert!(!cache.has_swa_ring());
3343 cache.tp_kv[0] = Some(ResidentTpKvCache::new_swa(
3344 Vec::new(),
3345 128,
3346 128,
3347 136,
3348 96,
3349 10_000,
3350 512,
3351 ));
3352 assert!(cache.has_swa_ring());
3353 }
3354}
3355
3356#[cfg(test)]
3357mod swa_ring_tests {
3358 use super::{
3359 KvRing, KvRingAppend, PRIME_CHUNK_MAX_TOKENS, SWA_REWIND_SLACK_ROWS,
3360 SWA_VIEW_ALIGNMENT_ROWS, kv_plane_allocation_bytes, swa_retain_from, swa_ring_rows,
3361 };
3362
3363 #[test]
3364 fn allocation_rows_cover_window_max_prime_and_alignment_slack() {
3365 assert_eq!(swa_ring_rows(512, 262_144), 512 + 4096 + 512 + 31);
3366 assert_eq!(swa_ring_rows(512, 4096), 4096);
3367 assert_eq!(
3368 kv_plane_allocation_bytes(5151, 1088),
3369 5151 * 1088 + 8,
3370 "the Step35 session plane allocates ring rows plus the existing tail pad",
3371 );
3372 }
3373
3374 #[test]
3387 fn retain_grants_rewind_slack_but_never_asks_below_base() {
3388 const WINDOW: usize = 512;
3389 let rows = swa_ring_rows(WINDOW, 262_144);
3390 let len = rows;
3391 let aligned = |pos: usize| (pos - (WINDOW - 1)) & !(SWA_VIEW_ALIGNMENT_ROWS - 1);
3392
3393 let retain = swa_retain_from(len, WINDOW, 0);
3396 assert!(retain <= aligned(len) - SWA_REWIND_SLACK_ROWS);
3397 let mut ring = KvRing::new(rows, WINDOW);
3398 ring.apply_rebase(retain);
3399 assert!(
3400 ring.can_rewind_to(len - 1),
3401 "a one-token rewind must survive the rebase"
3402 );
3403 assert!(ring.can_rewind_to(len - SWA_REWIND_SLACK_ROWS));
3404
3405 assert!(len - retain + PRIME_CHUNK_MAX_TOKENS <= rows);
3407
3408 for depth in [1usize, 32, 256, SWA_REWIND_SLACK_ROWS] {
3412 assert!(
3413 ring.can_rewind_to(len - depth),
3414 "a {depth}-row rewind must be resident, not clamped away",
3415 );
3416 }
3417
3418 let base = ring.base();
3425 for depth in [1usize, 32, 256, SWA_REWIND_SLACK_ROWS] {
3426 let window_start = (len - depth - (WINDOW - 1)) & !(SWA_VIEW_ALIGNMENT_ROWS - 1);
3427 assert!(
3428 window_start >= base,
3429 "after a {depth}-row rewind the window starts at {window_start}, below base \
3430 {base} — clamping here would serve rows the ring no longer holds (the pos-8661 \
3431 all-NaN case)",
3432 );
3433 assert!(swa_retain_from(len - depth, WINDOW, base) >= base);
3434 }
3435
3436 let past = len - (SWA_REWIND_SLACK_ROWS + WINDOW);
3439 let past_start = (past - (WINDOW - 1)) & !(SWA_VIEW_ALIGNMENT_ROWS - 1);
3440 assert!(
3441 past_start < base,
3442 "beyond the headroom the window must fall below base"
3443 );
3444 assert!(
3445 !ring.can_rewind_to(past),
3446 "and can_rewind_to must refuse it"
3447 );
3448 }
3449
3450 #[test]
3451 fn ring_matches_flat_bytes_before_wrap() {
3452 let ring = KvRing::new(swa_ring_rows(512, 262_144), 512);
3453 let flat: Vec<u32> = (0..1024).collect();
3454 let mut physical = vec![u32::MAX; ring.rows()];
3455 let KvRingAppend::Contiguous { write_row } = ring.append_plan(0, 0, flat.len()).unwrap()
3456 else {
3457 panic!("first append unexpectedly wrapped")
3458 };
3459 physical[write_row..write_row + flat.len()].copy_from_slice(&flat);
3460 let view = ring.physical_range(0, flat.len()).unwrap();
3461 assert_eq!(&physical[view], flat.as_slice());
3462 }
3463
3464 #[test]
3465 fn wrap_rebases_the_exact_aligned_prime_view() {
3466 let mut ring = KvRing::new(swa_ring_rows(512, 262_144), 512);
3467 let flat: Vec<u32> = (0..8192).collect();
3468 let mut physical = vec![u32::MAX; ring.rows()];
3469 let KvRingAppend::Contiguous { write_row } = ring.append_plan(0, 0, 4096).unwrap() else {
3470 panic!("first prime chunk unexpectedly wrapped")
3471 };
3472 physical[write_row..write_row + 4096].copy_from_slice(&flat[..4096]);
3473
3474 let off = (4096usize - (512 - 1)) & !31usize;
3475 let KvRingAppend::Rebase {
3476 src_row,
3477 keep_rows,
3478 new_base,
3479 write_row,
3480 } = ring.append_plan(4096, off, 4096).unwrap()
3481 else {
3482 panic!("second prime chunk did not wrap")
3483 };
3484 let retained = physical[src_row..src_row + keep_rows].to_vec();
3485 physical[..keep_rows].copy_from_slice(&retained);
3486 ring.apply_rebase(new_base);
3487 physical[write_row..write_row + 4096].copy_from_slice(&flat[4096..8192]);
3488
3489 let view = ring.physical_range(off, 8192).unwrap();
3490 assert_eq!(&physical[view], &flat[off..8192]);
3491 assert_eq!(ring.base(), off);
3492 }
3493
3494 #[test]
3495 fn rewind_declines_once_the_required_window_was_lapped() {
3496 let mut ring = KvRing::new(swa_ring_rows(512, 262_144), 512);
3497 let KvRingAppend::Rebase { new_base, .. } = ring.append_plan(4096, 3584, 4096).unwrap()
3498 else {
3499 panic!("expected wrap")
3500 };
3501 ring.apply_rebase(new_base);
3502 assert!(ring.can_rewind_to(4095));
3503 assert!(!ring.can_rewind_to(4094));
3504 assert!(!ring.can_rewind_to(0));
3505 }
3506
3507 #[test]
3512 fn restore_plan_copies_the_window_not_the_absolute_length() {
3513 let mut ring = KvRing::new(swa_ring_rows(512, 262_144), 512);
3514 let (base, phys) = ring.restore_plan(400).unwrap();
3516 assert_eq!((base, phys), (0, 0..400));
3517
3518 let mut live = 0usize;
3521 while live < 40_960 {
3522 let retain = swa_retain_from(live, 512, ring.base());
3523 if let KvRingAppend::Rebase { new_base, .. } =
3524 ring.append_plan(live, retain, 4096).unwrap()
3525 {
3526 ring.apply_rebase(new_base);
3527 }
3528 live += 4096;
3529 }
3530 assert!(ring.base() > 0, "a 40k walk must have lapped the ring");
3531 let (base, phys) = ring.restore_plan(live).unwrap();
3532 assert_eq!(base, (live - (512 - 1)) & !31usize);
3533 assert!(
3534 base >= ring.base(),
3535 "the plan must stay above the ring floor"
3536 );
3537 assert_eq!(phys.len(), live - base);
3538 assert!(
3539 phys.end <= ring.rows(),
3540 "the copy must fit the physical buffer ({} rows), got {:?}",
3541 ring.rows(),
3542 phys
3543 );
3544
3545 assert!(ring.restore_plan(400).is_err());
3547 }
3548}