1use std::{marker::PhantomData, ops::Range};
8
9use eredu_core::{
10 cache::{
11 LayerCachePolicy, PromptCacheError, PromptCacheModelIdentity, PromptCacheStateSegment,
12 PromptCacheTopology, StateComponentPolicy, StateTensorRole,
13 },
14 LayerSchedule,
15};
16use eredu_nn::NeuralBackend;
17
18#[derive(Debug, Clone, Copy, Eq, PartialEq)]
20#[non_exhaustive]
21pub enum ArchitectureStatePlacement {
22 GroupUnits {
24 group: usize,
26 },
27 OutputOwner,
29}
30
31#[derive(Debug, Clone, Eq, PartialEq)]
33pub struct ArchitectureStatePartitionRule {
34 layers: Range<usize>,
35 placement: ArchitectureStatePlacement,
36}
37
38impl ArchitectureStatePartitionRule {
39 pub fn group_units(group: usize, layers: Range<usize>) -> Self {
41 Self {
42 layers,
43 placement: ArchitectureStatePlacement::GroupUnits { group },
44 }
45 }
46
47 pub fn output_owner(layers: Range<usize>) -> Self {
49 Self {
50 layers,
51 placement: ArchitectureStatePlacement::OutputOwner,
52 }
53 }
54
55 pub fn layers(&self) -> Range<usize> {
57 self.layers.clone()
58 }
59
60 pub const fn placement(&self) -> ArchitectureStatePlacement {
62 self.placement
63 }
64}
65
66#[derive(Debug, Clone, Eq, PartialEq)]
71pub struct ArchitectureStatePartitionPlan {
72 rules: Vec<ArchitectureStatePartitionRule>,
73}
74
75impl ArchitectureStatePartitionPlan {
76 pub fn new(rules: impl IntoIterator<Item = ArchitectureStatePartitionRule>) -> Self {
78 Self {
79 rules: rules.into_iter().collect(),
80 }
81 }
82
83 pub fn rules(&self) -> &[ArchitectureStatePartitionRule] {
85 &self.rules
86 }
87}
88
89#[derive(Debug, Clone, Eq, PartialEq, thiserror::Error)]
91#[non_exhaustive]
92pub enum ArchitectureStatePartitionError {
93 #[error("architecture state partition plan must contain at least one rule")]
95 EmptyPlan,
96 #[error("architecture state partition rule has empty range {start}..{end}")]
98 EmptyRange {
99 start: usize,
101 end: usize,
103 },
104 #[error("architecture state partition range {start}..{end} exceeds the {layers}-layer layout")]
106 RangeOutOfBounds {
107 start: usize,
109 end: usize,
111 layers: usize,
113 },
114 #[error(
116 "architecture state partition range starts at {start}, before prior frontier {frontier}"
117 )]
118 OverlappingRange {
119 start: usize,
121 frontier: usize,
123 },
124 #[error("architecture state layer {layer} is not assigned by the partition plan")]
126 UnassignedLayer {
127 layer: usize,
129 },
130 #[error("architecture state partition references unknown execution group {group}")]
132 UnknownGroup {
133 group: usize,
135 },
136 #[error(
138 "architecture state range {start}..{end} has {} layers but group {group} has {units} units",
139 end - start
140 )]
141 GroupLengthMismatch {
142 group: usize,
144 start: usize,
146 end: usize,
148 units: usize,
150 },
151 #[error(
153 "architecture state partition selects discontiguous ranges ending at {frontier} and starting at {start}"
154 )]
155 DiscontiguousSelection {
156 frontier: usize,
158 start: usize,
160 },
161 #[error("architecture state partition layout is invalid: {0}")]
163 InvalidLayout(String),
164}
165
166pub const DEFAULT_STATE_SEGMENT_ID: &str = "state";
168
169#[derive(Debug, Clone, Eq, PartialEq, Ord, PartialOrd, Hash)]
171pub struct StateSegmentId(String);
172
173impl StateSegmentId {
174 pub fn new(id: impl Into<String>) -> Result<Self, StateError> {
176 let id = id.into();
177 if id.trim().is_empty() {
178 return Err(StateError::EmptySegmentId);
179 }
180 Ok(Self(id))
181 }
182
183 pub fn as_str(&self) -> &str {
185 &self.0
186 }
187}
188
189impl std::fmt::Display for StateSegmentId {
190 fn fmt(&self, formatter: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
191 formatter.write_str(&self.0)
192 }
193}
194
195#[derive(Debug, Clone, Copy, Eq, PartialEq, Ord, PartialOrd, Hash)]
197#[non_exhaustive]
198pub enum StateSegmentLifetime {
199 Persistent,
201 FrameLocal,
203}
204
205#[derive(Debug, Clone, Eq, PartialEq)]
207pub struct StateSegmentSpec {
208 id: StateSegmentId,
209 layers: Range<usize>,
210 lifetime: StateSegmentLifetime,
211 processed_token_offset: i32,
212}
213
214impl StateSegmentSpec {
215 pub fn new(
217 id: impl Into<String>,
218 layers: Range<usize>,
219 lifetime: StateSegmentLifetime,
220 processed_token_offset: i32,
221 ) -> Result<Self, StateError> {
222 let id = StateSegmentId::new(id)?;
223 if layers.is_empty() {
224 return Err(StateError::EmptySegmentRange {
225 segment: id,
226 start: layers.start,
227 end: layers.end,
228 });
229 }
230 if processed_token_offset > 0 {
231 return Err(StateError::PositiveSegmentOffset {
232 segment: id,
233 offset: processed_token_offset,
234 });
235 }
236 Ok(Self {
237 id,
238 layers,
239 lifetime,
240 processed_token_offset,
241 })
242 }
243
244 pub const fn id(&self) -> &StateSegmentId {
246 &self.id
247 }
248
249 pub fn layers(&self) -> Range<usize> {
251 self.layers.clone()
252 }
253
254 pub const fn lifetime(&self) -> StateSegmentLifetime {
256 self.lifetime
257 }
258
259 pub const fn processed_token_offset(&self) -> i32 {
261 self.processed_token_offset
262 }
263}
264
265#[derive(Debug, Clone, Eq, PartialEq)]
267pub struct StateLayout {
268 layers: LayerSchedule<LayerCachePolicy>,
269 components: Vec<Vec<StateComponentPolicy>>,
270 segments: Vec<StateSegmentSpec>,
271}
272
273impl StateLayout {
274 pub fn new(layers: LayerSchedule<LayerCachePolicy>) -> Result<Self, StateError> {
277 if layers.is_empty() {
278 return Err(StateError::EmptyLayout);
279 }
280 let count = layers.len();
281 Self::segmented(
282 layers,
283 [StateSegmentSpec::new(
284 DEFAULT_STATE_SEGMENT_ID,
285 0..count,
286 StateSegmentLifetime::Persistent,
287 0,
288 )?],
289 )
290 }
291
292 pub fn segmented(
299 layers: LayerSchedule<LayerCachePolicy>,
300 segments: impl IntoIterator<Item = StateSegmentSpec>,
301 ) -> Result<Self, StateError> {
302 if layers.is_empty() {
303 return Err(StateError::EmptyLayout);
304 }
305 for (layer, policy) in layers.iter().enumerate() {
306 policy
307 .validate()
308 .map_err(|error| StateError::InvalidLayer {
309 layer,
310 reason: error.to_string(),
311 })?;
312 }
313 let components = layers.iter().map(LayerCachePolicy::components).collect();
314 let mut segments = segments.into_iter().collect::<Vec<_>>();
315 if segments.is_empty() {
316 return Err(StateError::EmptySegments);
317 }
318 segments.sort_by(|left, right| {
319 left.layers
320 .start
321 .cmp(&right.layers.start)
322 .then_with(|| left.layers.end.cmp(&right.layers.end))
323 .then_with(|| left.id.cmp(&right.id))
324 });
325 let mut identities = std::collections::BTreeSet::new();
326 let mut frontier = 0usize;
327 for segment in &segments {
328 if !identities.insert(segment.id.clone()) {
329 return Err(StateError::DuplicateSegment {
330 segment: segment.id.clone(),
331 });
332 }
333 if segment.layers.end > layers.len() {
334 return Err(StateError::SegmentOutOfBounds {
335 segment: segment.id.clone(),
336 start: segment.layers.start,
337 end: segment.layers.end,
338 layers: layers.len(),
339 });
340 }
341 if segment.layers.start < frontier {
342 return Err(StateError::OverlappingSegment {
343 segment: segment.id.clone(),
344 start: segment.layers.start,
345 frontier,
346 });
347 }
348 if segment.layers.start > frontier {
349 return Err(StateError::UnassignedStateLayer { layer: frontier });
350 }
351 frontier = segment.layers.end;
352 }
353 if frontier != layers.len() {
354 return Err(StateError::UnassignedStateLayer { layer: frontier });
355 }
356 Ok(Self {
357 layers,
358 components,
359 segments,
360 })
361 }
362
363 pub fn len(&self) -> usize {
365 self.layers.len()
366 }
367
368 pub fn is_empty(&self) -> bool {
370 self.layers.is_empty()
371 }
372
373 pub fn layer(&self, layer: usize) -> Option<&LayerCachePolicy> {
375 self.layers.get(layer)
376 }
377
378 pub const fn layers(&self) -> &LayerSchedule<LayerCachePolicy> {
380 &self.layers
381 }
382
383 pub fn components(&self, layer: usize) -> Option<&[StateComponentPolicy]> {
385 self.components.get(layer).map(Vec::as_slice)
386 }
387
388 pub fn segments(&self) -> &[StateSegmentSpec] {
390 &self.segments
391 }
392
393 pub fn segment(&self, id: &StateSegmentId) -> Option<&StateSegmentSpec> {
395 self.segments.iter().find(|segment| segment.id() == id)
396 }
397
398 pub fn segment_for_layer(&self, layer: usize) -> Option<&StateSegmentSpec> {
400 self.segments
401 .iter()
402 .find(|segment| segment.layers.contains(&layer))
403 }
404
405 pub fn layer_prefix_offsets(&self) -> Vec<i32> {
407 let mut offsets = Vec::with_capacity(self.len());
408 for segment in &self.segments {
409 offsets.extend(std::iter::repeat_n(
410 segment.processed_token_offset(),
411 segment.layers.len(),
412 ));
413 }
414 offsets
415 }
416
417 pub fn slice(&self, layers: Range<usize>) -> Result<Self, StateError> {
420 if layers.is_empty() || layers.end > self.len() {
421 return Err(StateError::InvalidLayoutSlice {
422 start: layers.start,
423 end: layers.end,
424 layers: self.len(),
425 });
426 }
427 let policies = self
428 .layers
429 .iter()
430 .skip(layers.start)
431 .take(layers.len())
432 .cloned()
433 .collect::<Vec<_>>();
434 let mut segments = Vec::new();
435 for segment in &self.segments {
436 let start = segment.layers.start.max(layers.start);
437 let end = segment.layers.end.min(layers.end);
438 if start < end {
439 segments.push(StateSegmentSpec::new(
440 segment.id.as_str(),
441 start - layers.start..end - layers.start,
442 segment.lifetime,
443 segment.processed_token_offset,
444 )?);
445 }
446 }
447 Self::segmented(
448 LayerSchedule::new(policies.len(), policies)
449 .map_err(|error| StateError::InvalidResidency(error.to_string()))?,
450 segments,
451 )
452 }
453}
454
455pub trait RuntimeLayerState<B: NeuralBackend> {
457 type RetainedValues<'a>: Iterator<Item = &'a B::Tensor>
459 where
460 Self: 'a,
461 B::Tensor: 'a;
462
463 fn retained_values(&self) -> Self::RetainedValues<'_>;
465}
466
467pub trait ResettableRuntimeLayerState<B: NeuralBackend>: RuntimeLayerState<B> {
472 fn reset(&mut self) -> Result<(), StateError>;
474}
475
476pub trait RuntimeStateComponents<B: NeuralBackend>: RuntimeLayerState<B> {
482 fn position(&self) -> i32;
484
485 fn fixed_component(
487 &mut self,
488 role: StateTensorRole,
489 ) -> Result<&mut Option<B::Tensor>, StateError>;
490
491 fn advance_fixed(&mut self, tokens: i32) -> Result<(), StateError>;
493}
494
495pub trait RuntimeState<B: NeuralBackend> {
497 type RetainedValues<'a>: Iterator<Item = &'a B::Tensor>
499 where
500 Self: 'a,
501 B::Tensor: 'a;
502
503 fn layout(&self) -> &StateLayout;
505
506 fn retained_values(
511 &self,
512 ordinal: usize,
513 address: crate::ExecutionUnitAddress,
514 ) -> Result<Self::RetainedValues<'_>, StateError>;
515}
516
517pub trait ArchitectureStateFactory<B: NeuralBackend> {
523 type State: RuntimeState<B>;
525 type Error;
527
528 fn realize(&mut self, layout: &StateLayout) -> Result<Self::State, Self::Error>;
530}
531
532#[derive(Debug, thiserror::Error)]
534#[non_exhaustive]
535pub enum ArchitectureStateRealizationError<ArchitectureError, FactoryError> {
536 #[error("architecture state layout selection failed")]
538 Architecture(#[source] ArchitectureError),
539 #[error("backend state realization failed")]
541 Factory(#[source] FactoryError),
542 #[error("backend state realization changed the selected architecture layout")]
544 LayoutMismatch,
545}
546
547pub fn realize_architecture_state<B, M, F>(
550 architecture: &M,
551 factory: &mut F,
552) -> Result<F::State, ArchitectureStateRealizationError<M::DefinitionError, F::Error>>
553where
554 B: NeuralBackend,
555 M: crate::ArchitectureParameters<B>,
556 F: ArchitectureStateFactory<B>,
557{
558 let layout = architecture
559 .state_layout()
560 .map_err(ArchitectureStateRealizationError::Architecture)?;
561 let state = factory
562 .realize(&layout)
563 .map_err(ArchitectureStateRealizationError::Factory)?;
564 if state.layout() != &layout {
565 return Err(ArchitectureStateRealizationError::LayoutMismatch);
566 }
567 Ok(state)
568}
569
570pub trait ResettableRuntimeState<B: NeuralBackend>: RuntimeState<B> {
572 fn reset_segment(&mut self, segment: &StateSegmentId) -> Result<(), StateError>;
574}
575
576pub trait LayerRuntimeState<B: NeuralBackend>: RuntimeState<B> {
578 type LayerState: RuntimeLayerState<B>;
580
581 fn layer(&mut self, layer: usize) -> Result<&mut Self::LayerState, StateError>;
583}
584
585#[derive(Debug)]
587pub struct DeviceState<B: NeuralBackend, L> {
588 layout: StateLayout,
589 layers: Vec<L>,
590 backend: PhantomData<fn() -> B>,
591}
592
593impl<B: NeuralBackend, L: Clone> Clone for DeviceState<B, L> {
594 fn clone(&self) -> Self {
595 Self {
596 layout: self.layout.clone(),
597 layers: self.layers.clone(),
598 backend: PhantomData,
599 }
600 }
601
602 fn clone_from(&mut self, source: &Self) {
603 self.layout.clone_from(&source.layout);
604 self.layers.clone_from(&source.layers);
605 }
606}
607
608impl<B: NeuralBackend, L> DeviceState<B, L> {
609 pub fn create<E>(
611 layout: StateLayout,
612 mut create: impl FnMut(usize, &LayerCachePolicy) -> Result<L, E>,
613 ) -> Result<Self, E> {
614 let layers = layout
615 .layers()
616 .iter()
617 .enumerate()
618 .map(|(layer, policy)| create(layer, policy))
619 .collect::<Result<Vec<_>, _>>()?;
620 Ok(Self {
621 layout,
622 layers,
623 backend: PhantomData,
624 })
625 }
626}
627
628impl<B, L> RuntimeState<B> for DeviceState<B, L>
629where
630 B: NeuralBackend,
631 L: RuntimeLayerState<B>,
632{
633 type RetainedValues<'a>
634 = L::RetainedValues<'a>
635 where
636 Self: 'a,
637 B::Tensor: 'a;
638
639 fn layout(&self) -> &StateLayout {
640 &self.layout
641 }
642
643 fn retained_values(
644 &self,
645 _ordinal: usize,
646 address: crate::ExecutionUnitAddress,
647 ) -> Result<Self::RetainedValues<'_>, StateError> {
648 let layer = address.index();
649 self.layers
650 .get(layer)
651 .map(|layer| layer.retained_values())
652 .ok_or(StateError::UnknownLayer {
653 layer,
654 count: self.layers.len(),
655 })
656 }
657}
658
659impl<B, L> LayerRuntimeState<B> for DeviceState<B, L>
660where
661 B: NeuralBackend,
662 L: RuntimeLayerState<B>,
663{
664 type LayerState = L;
665
666 fn layer(&mut self, layer: usize) -> Result<&mut Self::LayerState, StateError> {
667 let count = self.layers.len();
668 self.layers
669 .get_mut(layer)
670 .ok_or(StateError::UnknownLayer { layer, count })
671 }
672}
673
674impl<B, L> ResettableRuntimeState<B> for DeviceState<B, L>
675where
676 B: NeuralBackend,
677 L: ResettableRuntimeLayerState<B>,
678{
679 fn reset_segment(&mut self, segment: &StateSegmentId) -> Result<(), StateError> {
680 let range = self
681 .layout
682 .segment(segment)
683 .map(StateSegmentSpec::layers)
684 .ok_or_else(|| StateError::UnknownSegment {
685 segment: segment.clone(),
686 })?;
687 for layer in &mut self.layers[range] {
688 layer.reset()?;
689 }
690 Ok(())
691 }
692}
693
694impl<B: NeuralBackend, L> AsRef<[L]> for DeviceState<B, L> {
695 fn as_ref(&self) -> &[L] {
696 &self.layers
697 }
698}
699
700impl<B: NeuralBackend, L> AsMut<[L]> for DeviceState<B, L> {
701 fn as_mut(&mut self) -> &mut [L] {
702 &mut self.layers
703 }
704}
705
706#[derive(Debug, Clone, Eq, PartialEq)]
708pub struct ModelStateIdentity {
709 model_family: String,
711 effective_model_type: String,
713 architecture_fingerprint: String,
715 layer_count: usize,
717 global_layer_start: usize,
719 sink_tokens: usize,
721 topology: PromptCacheTopology,
723}
724
725impl ModelStateIdentity {
726 #[allow(clippy::too_many_arguments)]
728 pub fn new(
729 model_family: impl Into<String>,
730 effective_model_type: impl Into<String>,
731 architecture_fingerprint: impl Into<String>,
732 layer_count: usize,
733 global_layer_start: usize,
734 sink_tokens: usize,
735 topology: PromptCacheTopology,
736 ) -> Result<Self, PromptCacheError> {
737 let model_family = model_family.into();
738 let effective_model_type = effective_model_type.into();
739 let architecture_fingerprint = architecture_fingerprint.into();
740 if model_family.trim().is_empty()
741 || effective_model_type.trim().is_empty()
742 || architecture_fingerprint.trim().is_empty()
743 {
744 return Err(PromptCacheError::Malformed(
745 "model-state identity strings must be non-empty".into(),
746 ));
747 }
748 if layer_count == 0 || global_layer_start > layer_count {
749 return Err(PromptCacheError::Malformed(format!(
750 "model-state layer start {global_layer_start} is invalid for {layer_count} layers"
751 )));
752 }
753 topology.validate()?;
754 Ok(Self {
755 model_family,
756 effective_model_type,
757 architecture_fingerprint,
758 layer_count,
759 global_layer_start,
760 sink_tokens,
761 topology,
762 })
763 }
764
765 pub const fn layer_count(&self) -> usize {
767 self.layer_count
768 }
769
770 pub const fn global_layer_start(&self) -> usize {
772 self.global_layer_start
773 }
774
775 pub fn prompt_cache_identity(
777 &self,
778 layout: &StateLayout,
779 ) -> Result<PromptCacheModelIdentity, PromptCacheError> {
780 let global_layer_end = self
781 .global_layer_start
782 .checked_add(layout.len())
783 .ok_or_else(|| PromptCacheError::Malformed("owned layer range overflowed".into()))?;
784 PromptCacheModelIdentity::new(
785 self.model_family.clone(),
786 self.effective_model_type.clone(),
787 self.architecture_fingerprint.clone(),
788 self.layer_count,
789 self.global_layer_start,
790 global_layer_end,
791 self.sink_tokens,
792 self.topology.clone(),
793 layout.layers().clone(),
794 layout.layer_prefix_offsets(),
795 layout
796 .segments()
797 .iter()
798 .map(|segment| {
799 PromptCacheStateSegment::new(segment.id().as_str(), segment.layers())
800 })
801 .collect::<Result<Vec<_>, _>>()?,
802 )
803 }
804}
805
806#[derive(Debug, Clone, Eq, PartialEq, thiserror::Error)]
808#[non_exhaustive]
809pub enum StateError {
810 #[error("runtime state layout must contain at least one layer")]
812 EmptyLayout,
813 #[error("runtime state layout must contain at least one named segment")]
815 EmptySegments,
816 #[error("runtime state segment identity must not be empty")]
818 EmptySegmentId,
819 #[error("runtime state segment {segment:?} has empty layer range {start}..{end}")]
821 EmptySegmentRange {
822 segment: StateSegmentId,
824 start: usize,
826 end: usize,
828 },
829 #[error("runtime state segment {segment:?} has positive processed-token offset {offset}")]
831 PositiveSegmentOffset {
832 segment: StateSegmentId,
834 offset: i32,
836 },
837 #[error("runtime state layout cannot select range {start}..{end} from {layers} layers")]
839 InvalidLayoutSlice {
840 start: usize,
842 end: usize,
844 layers: usize,
846 },
847 #[error("runtime state segment {segment:?} is declared more than once")]
849 DuplicateSegment {
850 segment: StateSegmentId,
852 },
853 #[error(
855 "runtime state segment {segment:?} range {start}..{end} exceeds {layers} layout layers"
856 )]
857 SegmentOutOfBounds {
858 segment: StateSegmentId,
860 start: usize,
862 end: usize,
864 layers: usize,
866 },
867 #[error(
869 "runtime state segment {segment:?} starts at layer {start}, before prior frontier {frontier}"
870 )]
871 OverlappingSegment {
872 segment: StateSegmentId,
874 start: usize,
876 frontier: usize,
878 },
879 #[error("runtime state layer {layer} is not assigned to a named segment")]
881 UnassignedStateLayer {
882 layer: usize,
884 },
885 #[error("runtime state layout has no segment {segment:?}")]
887 UnknownSegment {
888 segment: StateSegmentId,
890 },
891 #[error("runtime state reset failed: {0}")]
893 ResetFailed(String),
894 #[error("invalid runtime state policy for layer {layer}: {reason}")]
896 InvalidLayer {
897 layer: usize,
899 reason: String,
901 },
902 #[error("invalid runtime state residency plan: {0}")]
904 InvalidResidency(String),
905 #[error("runtime state layer {layer} is outside the {count}-layer layout")]
907 UnknownLayer {
908 layer: usize,
910 count: usize,
912 },
913 #[error("runtime state layer does not declare fixed component {role:?}")]
915 UnknownComponent {
916 role: StateTensorRole,
918 },
919 #[error("invalid fixed-state advance: {0}")]
921 InvalidAdvance(String),
922}
923
924#[cfg(test)]
925mod tests {
926 use super::*;
927 use eredu_core::{AttentionPolicy, LayerSchedule};
928
929 fn layout() -> StateLayout {
930 StateLayout::new(
931 LayerSchedule::new(
932 2,
933 vec![
934 LayerCachePolicy::key_value(AttentionPolicy::Full, 2, 8).unwrap(),
935 LayerCachePolicy::key_value(
936 AttentionPolicy::from_sliding_window(Some(16)).unwrap(),
937 2,
938 8,
939 )
940 .unwrap(),
941 ],
942 )
943 .unwrap(),
944 )
945 .unwrap()
946 }
947
948 #[test]
949 fn prompt_identity_is_derived_from_layout_and_placement() {
950 let layout = layout();
951 let identity = ModelStateIdentity {
952 model_family: "fixture".into(),
953 effective_model_type: "fixture-v1".into(),
954 architecture_fingerprint: "geometry-1".into(),
955 layer_count: 4,
956 global_layer_start: 1,
957 sink_tokens: 0,
958 topology: PromptCacheTopology::default(),
959 }
960 .prompt_cache_identity(&layout)
961 .unwrap();
962 assert_eq!(identity.global_layer_start(), 1);
963 assert_eq!(identity.global_layer_end(), 3);
964 assert_eq!(identity.layer_layout(), layout.layers());
965 }
966
967 #[test]
968 fn prompt_identity_derives_offsets_from_segments() {
969 let policy = LayerCachePolicy::key_value(AttentionPolicy::Full, 1, 8).unwrap();
970 let layout = StateLayout::segmented(
971 LayerSchedule::new(2, vec![policy.clone(), policy]).unwrap(),
972 [
973 StateSegmentSpec::new("target", 0..1, StateSegmentLifetime::Persistent, 0).unwrap(),
974 StateSegmentSpec::new("prediction", 1..2, StateSegmentLifetime::Persistent, -1)
975 .unwrap(),
976 ],
977 )
978 .unwrap();
979 let identity = ModelStateIdentity {
980 model_family: "fixture".into(),
981 effective_model_type: "fixture-v1".into(),
982 architecture_fingerprint: "geometry-1".into(),
983 layer_count: 2,
984 global_layer_start: 0,
985 sink_tokens: 0,
986 topology: PromptCacheTopology::default(),
987 }
988 .prompt_cache_identity(&layout)
989 .unwrap();
990 assert_eq!(identity.layer_prefix_offsets(), [0, -1]);
991 assert_eq!(identity.state_segments().len(), 2);
992 assert_eq!(identity.state_segments()[0].id(), "target");
993 assert_eq!(identity.state_segments()[0].layers(), 0..1);
994 assert_eq!(identity.state_segments()[1].id(), "prediction");
995 assert_eq!(identity.state_segments()[1].layers(), 1..2);
996
997 let prediction = identity.select_state_segment("prediction").unwrap();
998 assert_eq!(prediction.global_layer_start(), 1);
999 assert_eq!(prediction.global_layer_end(), 2);
1000 assert_eq!(prediction.layer_prefix_offsets(), [-1]);
1001 assert_eq!(prediction.state_segments()[0].id(), "prediction");
1002 assert_eq!(prediction.state_segments()[0].layers(), 0..1);
1003 }
1004
1005 #[test]
1006 fn state_layout_exposes_stable_semantic_component_names() {
1007 let layout = StateLayout::new(
1008 LayerSchedule::new(
1009 1,
1010 vec![
1011 LayerCachePolicy::compressed_latent_rotary(AttentionPolicy::Full, 16, 8)
1012 .unwrap(),
1013 ],
1014 )
1015 .unwrap(),
1016 )
1017 .unwrap();
1018 let names = layout
1019 .components(0)
1020 .unwrap()
1021 .iter()
1022 .map(|component| component.role().stable_name())
1023 .collect::<Vec<_>>();
1024 assert_eq!(
1025 names,
1026 ["attention.compressed_latent", "attention.rotary_keys"]
1027 );
1028 }
1029
1030 fn four_layer_schedule() -> LayerSchedule<LayerCachePolicy> {
1031 LayerSchedule::new(
1032 4,
1033 (0..4)
1034 .map(|_| LayerCachePolicy::key_value(AttentionPolicy::Full, 1, 8).unwrap())
1035 .collect(),
1036 )
1037 .unwrap()
1038 }
1039
1040 #[test]
1041 fn composite_state_segments_are_canonical_and_cover_every_layer() {
1042 let layout = StateLayout::segmented(
1043 four_layer_schedule(),
1044 [
1045 StateSegmentSpec::new("depth", 2..4, StateSegmentLifetime::FrameLocal, 0).unwrap(),
1046 StateSegmentSpec::new("temporal", 0..2, StateSegmentLifetime::Persistent, 0)
1047 .unwrap(),
1048 ],
1049 )
1050 .unwrap();
1051
1052 assert_eq!(
1053 layout
1054 .segments()
1055 .iter()
1056 .map(|segment| (segment.id().as_str(), segment.layers(), segment.lifetime()))
1057 .collect::<Vec<_>>(),
1058 [
1059 ("temporal", 0..2, StateSegmentLifetime::Persistent),
1060 ("depth", 2..4, StateSegmentLifetime::FrameLocal),
1061 ]
1062 );
1063 assert_eq!(
1064 layout.segment_for_layer(0).unwrap().id().as_str(),
1065 "temporal"
1066 );
1067 assert_eq!(layout.segment_for_layer(3).unwrap().id().as_str(), "depth");
1068 assert!(layout.segment_for_layer(4).is_none());
1069 }
1070
1071 #[test]
1072 fn state_layout_slice_preserves_and_rebases_segment_frontiers() {
1073 let layout = StateLayout::segmented(
1074 four_layer_schedule(),
1075 [
1076 StateSegmentSpec::new("target", 0..2, StateSegmentLifetime::Persistent, 0).unwrap(),
1077 StateSegmentSpec::new("prediction", 2..4, StateSegmentLifetime::Persistent, -1)
1078 .unwrap(),
1079 ],
1080 )
1081 .unwrap();
1082
1083 let sliced = layout.slice(1..4).unwrap();
1084 assert_eq!(sliced.segments()[0].layers(), 0..1);
1085 assert_eq!(sliced.segments()[1].layers(), 1..3);
1086 assert_eq!(sliced.layer_prefix_offsets(), [0, -1, -1]);
1087 }
1088
1089 #[test]
1090 fn segment_identity_lifetime_and_offset_participate_in_layout_equality() {
1091 let layout = |depth_name, lifetime, offset| {
1092 StateLayout::segmented(
1093 four_layer_schedule(),
1094 [
1095 StateSegmentSpec::new("temporal", 0..2, StateSegmentLifetime::Persistent, 0)
1096 .unwrap(),
1097 StateSegmentSpec::new(depth_name, 2..4, lifetime, offset).unwrap(),
1098 ],
1099 )
1100 .unwrap()
1101 };
1102 let canonical = layout("depth", StateSegmentLifetime::FrameLocal, 0);
1103 assert_ne!(
1104 canonical,
1105 layout("predictor", StateSegmentLifetime::FrameLocal, 0)
1106 );
1107 assert_ne!(
1108 canonical,
1109 layout("depth", StateSegmentLifetime::Persistent, 0)
1110 );
1111 assert_ne!(
1112 canonical,
1113 layout("depth", StateSegmentLifetime::FrameLocal, -1)
1114 );
1115 }
1116
1117 #[test]
1118 fn malformed_segment_partitions_fail_closed() {
1119 let duplicate = StateLayout::segmented(
1120 four_layer_schedule(),
1121 [
1122 StateSegmentSpec::new("cache", 0..2, StateSegmentLifetime::Persistent, 0).unwrap(),
1123 StateSegmentSpec::new("cache", 2..4, StateSegmentLifetime::FrameLocal, 0).unwrap(),
1124 ],
1125 )
1126 .unwrap_err();
1127 assert!(matches!(duplicate, StateError::DuplicateSegment { .. }));
1128
1129 let overlap = StateLayout::segmented(
1130 four_layer_schedule(),
1131 [
1132 StateSegmentSpec::new("left", 0..3, StateSegmentLifetime::Persistent, 0).unwrap(),
1133 StateSegmentSpec::new("right", 2..4, StateSegmentLifetime::FrameLocal, 0).unwrap(),
1134 ],
1135 )
1136 .unwrap_err();
1137 assert!(matches!(overlap, StateError::OverlappingSegment { .. }));
1138
1139 let gap = StateLayout::segmented(
1140 four_layer_schedule(),
1141 [
1142 StateSegmentSpec::new("left", 0..1, StateSegmentLifetime::Persistent, 0).unwrap(),
1143 StateSegmentSpec::new("right", 2..4, StateSegmentLifetime::FrameLocal, 0).unwrap(),
1144 ],
1145 )
1146 .unwrap_err();
1147 assert_eq!(gap, StateError::UnassignedStateLayer { layer: 1 });
1148
1149 let outside = StateLayout::segmented(
1150 four_layer_schedule(),
1151 [StateSegmentSpec::new("all", 0..5, StateSegmentLifetime::Persistent, 0).unwrap()],
1152 )
1153 .unwrap_err();
1154 assert!(matches!(outside, StateError::SegmentOutOfBounds { .. }));
1155 }
1156}