use std::{marker::PhantomData, ops::Range};
use eredu_core::{
cache::{
LayerCachePolicy, PromptCacheError, PromptCacheModelIdentity, PromptCacheStateSegment,
PromptCacheTopology, StateComponentPolicy, StateTensorRole,
},
LayerSchedule,
};
use eredu_nn::NeuralBackend;
#[derive(Debug, Clone, Copy, Eq, PartialEq)]
#[non_exhaustive]
pub enum ArchitectureStatePlacement {
GroupUnits {
group: usize,
},
OutputOwner,
}
#[derive(Debug, Clone, Eq, PartialEq)]
pub struct ArchitectureStatePartitionRule {
layers: Range<usize>,
placement: ArchitectureStatePlacement,
}
impl ArchitectureStatePartitionRule {
pub fn group_units(group: usize, layers: Range<usize>) -> Self {
Self {
layers,
placement: ArchitectureStatePlacement::GroupUnits { group },
}
}
pub fn output_owner(layers: Range<usize>) -> Self {
Self {
layers,
placement: ArchitectureStatePlacement::OutputOwner,
}
}
pub fn layers(&self) -> Range<usize> {
self.layers.clone()
}
pub const fn placement(&self) -> ArchitectureStatePlacement {
self.placement
}
}
#[derive(Debug, Clone, Eq, PartialEq)]
pub struct ArchitectureStatePartitionPlan {
rules: Vec<ArchitectureStatePartitionRule>,
}
impl ArchitectureStatePartitionPlan {
pub fn new(rules: impl IntoIterator<Item = ArchitectureStatePartitionRule>) -> Self {
Self {
rules: rules.into_iter().collect(),
}
}
pub fn rules(&self) -> &[ArchitectureStatePartitionRule] {
&self.rules
}
}
#[derive(Debug, Clone, Eq, PartialEq, thiserror::Error)]
#[non_exhaustive]
pub enum ArchitectureStatePartitionError {
#[error("architecture state partition plan must contain at least one rule")]
EmptyPlan,
#[error("architecture state partition rule has empty range {start}..{end}")]
EmptyRange {
start: usize,
end: usize,
},
#[error("architecture state partition range {start}..{end} exceeds the {layers}-layer layout")]
RangeOutOfBounds {
start: usize,
end: usize,
layers: usize,
},
#[error(
"architecture state partition range starts at {start}, before prior frontier {frontier}"
)]
OverlappingRange {
start: usize,
frontier: usize,
},
#[error("architecture state layer {layer} is not assigned by the partition plan")]
UnassignedLayer {
layer: usize,
},
#[error("architecture state partition references unknown execution group {group}")]
UnknownGroup {
group: usize,
},
#[error(
"architecture state range {start}..{end} has {} layers but group {group} has {units} units",
end - start
)]
GroupLengthMismatch {
group: usize,
start: usize,
end: usize,
units: usize,
},
#[error(
"architecture state partition selects discontiguous ranges ending at {frontier} and starting at {start}"
)]
DiscontiguousSelection {
frontier: usize,
start: usize,
},
#[error("architecture state partition layout is invalid: {0}")]
InvalidLayout(String),
}
pub const DEFAULT_STATE_SEGMENT_ID: &str = "state";
#[derive(Debug, Clone, Eq, PartialEq, Ord, PartialOrd, Hash)]
pub struct StateSegmentId(String);
impl StateSegmentId {
pub fn new(id: impl Into<String>) -> Result<Self, StateError> {
let id = id.into();
if id.trim().is_empty() {
return Err(StateError::EmptySegmentId);
}
Ok(Self(id))
}
pub fn as_str(&self) -> &str {
&self.0
}
}
impl std::fmt::Display for StateSegmentId {
fn fmt(&self, formatter: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
formatter.write_str(&self.0)
}
}
#[derive(Debug, Clone, Copy, Eq, PartialEq, Ord, PartialOrd, Hash)]
#[non_exhaustive]
pub enum StateSegmentLifetime {
Persistent,
FrameLocal,
}
#[derive(Debug, Clone, Eq, PartialEq)]
pub struct StateSegmentSpec {
id: StateSegmentId,
layers: Range<usize>,
lifetime: StateSegmentLifetime,
processed_token_offset: i32,
}
impl StateSegmentSpec {
pub fn new(
id: impl Into<String>,
layers: Range<usize>,
lifetime: StateSegmentLifetime,
processed_token_offset: i32,
) -> Result<Self, StateError> {
let id = StateSegmentId::new(id)?;
if layers.is_empty() {
return Err(StateError::EmptySegmentRange {
segment: id,
start: layers.start,
end: layers.end,
});
}
if processed_token_offset > 0 {
return Err(StateError::PositiveSegmentOffset {
segment: id,
offset: processed_token_offset,
});
}
Ok(Self {
id,
layers,
lifetime,
processed_token_offset,
})
}
pub const fn id(&self) -> &StateSegmentId {
&self.id
}
pub fn layers(&self) -> Range<usize> {
self.layers.clone()
}
pub const fn lifetime(&self) -> StateSegmentLifetime {
self.lifetime
}
pub const fn processed_token_offset(&self) -> i32 {
self.processed_token_offset
}
}
#[derive(Debug, Clone, Eq, PartialEq)]
pub struct StateLayout {
layers: LayerSchedule<LayerCachePolicy>,
components: Vec<Vec<StateComponentPolicy>>,
segments: Vec<StateSegmentSpec>,
}
impl StateLayout {
pub fn new(layers: LayerSchedule<LayerCachePolicy>) -> Result<Self, StateError> {
if layers.is_empty() {
return Err(StateError::EmptyLayout);
}
let count = layers.len();
Self::segmented(
layers,
[StateSegmentSpec::new(
DEFAULT_STATE_SEGMENT_ID,
0..count,
StateSegmentLifetime::Persistent,
0,
)?],
)
}
pub fn segmented(
layers: LayerSchedule<LayerCachePolicy>,
segments: impl IntoIterator<Item = StateSegmentSpec>,
) -> Result<Self, StateError> {
if layers.is_empty() {
return Err(StateError::EmptyLayout);
}
for (layer, policy) in layers.iter().enumerate() {
policy
.validate()
.map_err(|error| StateError::InvalidLayer {
layer,
reason: error.to_string(),
})?;
}
let components = layers.iter().map(LayerCachePolicy::components).collect();
let mut segments = segments.into_iter().collect::<Vec<_>>();
if segments.is_empty() {
return Err(StateError::EmptySegments);
}
segments.sort_by(|left, right| {
left.layers
.start
.cmp(&right.layers.start)
.then_with(|| left.layers.end.cmp(&right.layers.end))
.then_with(|| left.id.cmp(&right.id))
});
let mut identities = std::collections::BTreeSet::new();
let mut frontier = 0usize;
for segment in &segments {
if !identities.insert(segment.id.clone()) {
return Err(StateError::DuplicateSegment {
segment: segment.id.clone(),
});
}
if segment.layers.end > layers.len() {
return Err(StateError::SegmentOutOfBounds {
segment: segment.id.clone(),
start: segment.layers.start,
end: segment.layers.end,
layers: layers.len(),
});
}
if segment.layers.start < frontier {
return Err(StateError::OverlappingSegment {
segment: segment.id.clone(),
start: segment.layers.start,
frontier,
});
}
if segment.layers.start > frontier {
return Err(StateError::UnassignedStateLayer { layer: frontier });
}
frontier = segment.layers.end;
}
if frontier != layers.len() {
return Err(StateError::UnassignedStateLayer { layer: frontier });
}
Ok(Self {
layers,
components,
segments,
})
}
pub fn len(&self) -> usize {
self.layers.len()
}
pub fn is_empty(&self) -> bool {
self.layers.is_empty()
}
pub fn layer(&self, layer: usize) -> Option<&LayerCachePolicy> {
self.layers.get(layer)
}
pub const fn layers(&self) -> &LayerSchedule<LayerCachePolicy> {
&self.layers
}
pub fn components(&self, layer: usize) -> Option<&[StateComponentPolicy]> {
self.components.get(layer).map(Vec::as_slice)
}
pub fn segments(&self) -> &[StateSegmentSpec] {
&self.segments
}
pub fn segment(&self, id: &StateSegmentId) -> Option<&StateSegmentSpec> {
self.segments.iter().find(|segment| segment.id() == id)
}
pub fn segment_for_layer(&self, layer: usize) -> Option<&StateSegmentSpec> {
self.segments
.iter()
.find(|segment| segment.layers.contains(&layer))
}
pub fn layer_prefix_offsets(&self) -> Vec<i32> {
let mut offsets = Vec::with_capacity(self.len());
for segment in &self.segments {
offsets.extend(std::iter::repeat_n(
segment.processed_token_offset(),
segment.layers.len(),
));
}
offsets
}
pub fn slice(&self, layers: Range<usize>) -> Result<Self, StateError> {
if layers.is_empty() || layers.end > self.len() {
return Err(StateError::InvalidLayoutSlice {
start: layers.start,
end: layers.end,
layers: self.len(),
});
}
let policies = self
.layers
.iter()
.skip(layers.start)
.take(layers.len())
.cloned()
.collect::<Vec<_>>();
let mut segments = Vec::new();
for segment in &self.segments {
let start = segment.layers.start.max(layers.start);
let end = segment.layers.end.min(layers.end);
if start < end {
segments.push(StateSegmentSpec::new(
segment.id.as_str(),
start - layers.start..end - layers.start,
segment.lifetime,
segment.processed_token_offset,
)?);
}
}
Self::segmented(
LayerSchedule::new(policies.len(), policies)
.map_err(|error| StateError::InvalidResidency(error.to_string()))?,
segments,
)
}
}
pub trait RuntimeLayerState<B: NeuralBackend> {
type RetainedValues<'a>: Iterator<Item = &'a B::Tensor>
where
Self: 'a,
B::Tensor: 'a;
fn retained_values(&self) -> Self::RetainedValues<'_>;
}
pub trait ResettableRuntimeLayerState<B: NeuralBackend>: RuntimeLayerState<B> {
fn reset(&mut self) -> Result<(), StateError>;
}
pub trait RuntimeStateComponents<B: NeuralBackend>: RuntimeLayerState<B> {
fn position(&self) -> i32;
fn fixed_component(
&mut self,
role: StateTensorRole,
) -> Result<&mut Option<B::Tensor>, StateError>;
fn advance_fixed(&mut self, tokens: i32) -> Result<(), StateError>;
}
pub trait RuntimeState<B: NeuralBackend> {
type RetainedValues<'a>: Iterator<Item = &'a B::Tensor>
where
Self: 'a,
B::Tensor: 'a;
fn layout(&self) -> &StateLayout;
fn retained_values(
&self,
ordinal: usize,
address: crate::ExecutionUnitAddress,
) -> Result<Self::RetainedValues<'_>, StateError>;
}
pub trait ArchitectureStateFactory<B: NeuralBackend> {
type State: RuntimeState<B>;
type Error;
fn realize(&mut self, layout: &StateLayout) -> Result<Self::State, Self::Error>;
}
#[derive(Debug, thiserror::Error)]
#[non_exhaustive]
pub enum ArchitectureStateRealizationError<ArchitectureError, FactoryError> {
#[error("architecture state layout selection failed")]
Architecture(#[source] ArchitectureError),
#[error("backend state realization failed")]
Factory(#[source] FactoryError),
#[error("backend state realization changed the selected architecture layout")]
LayoutMismatch,
}
pub fn realize_architecture_state<B, M, F>(
architecture: &M,
factory: &mut F,
) -> Result<F::State, ArchitectureStateRealizationError<M::DefinitionError, F::Error>>
where
B: NeuralBackend,
M: crate::ArchitectureParameters<B>,
F: ArchitectureStateFactory<B>,
{
let layout = architecture
.state_layout()
.map_err(ArchitectureStateRealizationError::Architecture)?;
let state = factory
.realize(&layout)
.map_err(ArchitectureStateRealizationError::Factory)?;
if state.layout() != &layout {
return Err(ArchitectureStateRealizationError::LayoutMismatch);
}
Ok(state)
}
pub trait ResettableRuntimeState<B: NeuralBackend>: RuntimeState<B> {
fn reset_segment(&mut self, segment: &StateSegmentId) -> Result<(), StateError>;
}
pub trait LayerRuntimeState<B: NeuralBackend>: RuntimeState<B> {
type LayerState: RuntimeLayerState<B>;
fn layer(&mut self, layer: usize) -> Result<&mut Self::LayerState, StateError>;
}
#[derive(Debug)]
pub struct DeviceState<B: NeuralBackend, L> {
layout: StateLayout,
layers: Vec<L>,
backend: PhantomData<fn() -> B>,
}
impl<B: NeuralBackend, L: Clone> Clone for DeviceState<B, L> {
fn clone(&self) -> Self {
Self {
layout: self.layout.clone(),
layers: self.layers.clone(),
backend: PhantomData,
}
}
fn clone_from(&mut self, source: &Self) {
self.layout.clone_from(&source.layout);
self.layers.clone_from(&source.layers);
}
}
impl<B: NeuralBackend, L> DeviceState<B, L> {
pub fn create<E>(
layout: StateLayout,
mut create: impl FnMut(usize, &LayerCachePolicy) -> Result<L, E>,
) -> Result<Self, E> {
let layers = layout
.layers()
.iter()
.enumerate()
.map(|(layer, policy)| create(layer, policy))
.collect::<Result<Vec<_>, _>>()?;
Ok(Self {
layout,
layers,
backend: PhantomData,
})
}
}
impl<B, L> RuntimeState<B> for DeviceState<B, L>
where
B: NeuralBackend,
L: RuntimeLayerState<B>,
{
type RetainedValues<'a>
= L::RetainedValues<'a>
where
Self: 'a,
B::Tensor: 'a;
fn layout(&self) -> &StateLayout {
&self.layout
}
fn retained_values(
&self,
_ordinal: usize,
address: crate::ExecutionUnitAddress,
) -> Result<Self::RetainedValues<'_>, StateError> {
let layer = address.index();
self.layers
.get(layer)
.map(|layer| layer.retained_values())
.ok_or(StateError::UnknownLayer {
layer,
count: self.layers.len(),
})
}
}
impl<B, L> LayerRuntimeState<B> for DeviceState<B, L>
where
B: NeuralBackend,
L: RuntimeLayerState<B>,
{
type LayerState = L;
fn layer(&mut self, layer: usize) -> Result<&mut Self::LayerState, StateError> {
let count = self.layers.len();
self.layers
.get_mut(layer)
.ok_or(StateError::UnknownLayer { layer, count })
}
}
impl<B, L> ResettableRuntimeState<B> for DeviceState<B, L>
where
B: NeuralBackend,
L: ResettableRuntimeLayerState<B>,
{
fn reset_segment(&mut self, segment: &StateSegmentId) -> Result<(), StateError> {
let range = self
.layout
.segment(segment)
.map(StateSegmentSpec::layers)
.ok_or_else(|| StateError::UnknownSegment {
segment: segment.clone(),
})?;
for layer in &mut self.layers[range] {
layer.reset()?;
}
Ok(())
}
}
impl<B: NeuralBackend, L> AsRef<[L]> for DeviceState<B, L> {
fn as_ref(&self) -> &[L] {
&self.layers
}
}
impl<B: NeuralBackend, L> AsMut<[L]> for DeviceState<B, L> {
fn as_mut(&mut self) -> &mut [L] {
&mut self.layers
}
}
#[derive(Debug, Clone, Eq, PartialEq)]
pub struct ModelStateIdentity {
model_family: String,
effective_model_type: String,
architecture_fingerprint: String,
layer_count: usize,
global_layer_start: usize,
sink_tokens: usize,
topology: PromptCacheTopology,
}
impl ModelStateIdentity {
#[allow(clippy::too_many_arguments)]
pub fn new(
model_family: impl Into<String>,
effective_model_type: impl Into<String>,
architecture_fingerprint: impl Into<String>,
layer_count: usize,
global_layer_start: usize,
sink_tokens: usize,
topology: PromptCacheTopology,
) -> Result<Self, PromptCacheError> {
let model_family = model_family.into();
let effective_model_type = effective_model_type.into();
let architecture_fingerprint = architecture_fingerprint.into();
if model_family.trim().is_empty()
|| effective_model_type.trim().is_empty()
|| architecture_fingerprint.trim().is_empty()
{
return Err(PromptCacheError::Malformed(
"model-state identity strings must be non-empty".into(),
));
}
if layer_count == 0 || global_layer_start > layer_count {
return Err(PromptCacheError::Malformed(format!(
"model-state layer start {global_layer_start} is invalid for {layer_count} layers"
)));
}
topology.validate()?;
Ok(Self {
model_family,
effective_model_type,
architecture_fingerprint,
layer_count,
global_layer_start,
sink_tokens,
topology,
})
}
pub const fn layer_count(&self) -> usize {
self.layer_count
}
pub const fn global_layer_start(&self) -> usize {
self.global_layer_start
}
pub fn prompt_cache_identity(
&self,
layout: &StateLayout,
) -> Result<PromptCacheModelIdentity, PromptCacheError> {
let global_layer_end = self
.global_layer_start
.checked_add(layout.len())
.ok_or_else(|| PromptCacheError::Malformed("owned layer range overflowed".into()))?;
PromptCacheModelIdentity::new(
self.model_family.clone(),
self.effective_model_type.clone(),
self.architecture_fingerprint.clone(),
self.layer_count,
self.global_layer_start,
global_layer_end,
self.sink_tokens,
self.topology.clone(),
layout.layers().clone(),
layout.layer_prefix_offsets(),
layout
.segments()
.iter()
.map(|segment| {
PromptCacheStateSegment::new(segment.id().as_str(), segment.layers())
})
.collect::<Result<Vec<_>, _>>()?,
)
}
}
#[derive(Debug, Clone, Eq, PartialEq, thiserror::Error)]
#[non_exhaustive]
pub enum StateError {
#[error("runtime state layout must contain at least one layer")]
EmptyLayout,
#[error("runtime state layout must contain at least one named segment")]
EmptySegments,
#[error("runtime state segment identity must not be empty")]
EmptySegmentId,
#[error("runtime state segment {segment:?} has empty layer range {start}..{end}")]
EmptySegmentRange {
segment: StateSegmentId,
start: usize,
end: usize,
},
#[error("runtime state segment {segment:?} has positive processed-token offset {offset}")]
PositiveSegmentOffset {
segment: StateSegmentId,
offset: i32,
},
#[error("runtime state layout cannot select range {start}..{end} from {layers} layers")]
InvalidLayoutSlice {
start: usize,
end: usize,
layers: usize,
},
#[error("runtime state segment {segment:?} is declared more than once")]
DuplicateSegment {
segment: StateSegmentId,
},
#[error(
"runtime state segment {segment:?} range {start}..{end} exceeds {layers} layout layers"
)]
SegmentOutOfBounds {
segment: StateSegmentId,
start: usize,
end: usize,
layers: usize,
},
#[error(
"runtime state segment {segment:?} starts at layer {start}, before prior frontier {frontier}"
)]
OverlappingSegment {
segment: StateSegmentId,
start: usize,
frontier: usize,
},
#[error("runtime state layer {layer} is not assigned to a named segment")]
UnassignedStateLayer {
layer: usize,
},
#[error("runtime state layout has no segment {segment:?}")]
UnknownSegment {
segment: StateSegmentId,
},
#[error("runtime state reset failed: {0}")]
ResetFailed(String),
#[error("invalid runtime state policy for layer {layer}: {reason}")]
InvalidLayer {
layer: usize,
reason: String,
},
#[error("invalid runtime state residency plan: {0}")]
InvalidResidency(String),
#[error("runtime state layer {layer} is outside the {count}-layer layout")]
UnknownLayer {
layer: usize,
count: usize,
},
#[error("runtime state layer does not declare fixed component {role:?}")]
UnknownComponent {
role: StateTensorRole,
},
#[error("invalid fixed-state advance: {0}")]
InvalidAdvance(String),
}
#[cfg(test)]
mod tests {
use super::*;
use eredu_core::{AttentionPolicy, LayerSchedule};
fn layout() -> StateLayout {
StateLayout::new(
LayerSchedule::new(
2,
vec![
LayerCachePolicy::key_value(AttentionPolicy::Full, 2, 8).unwrap(),
LayerCachePolicy::key_value(
AttentionPolicy::from_sliding_window(Some(16)).unwrap(),
2,
8,
)
.unwrap(),
],
)
.unwrap(),
)
.unwrap()
}
#[test]
fn prompt_identity_is_derived_from_layout_and_placement() {
let layout = layout();
let identity = ModelStateIdentity {
model_family: "fixture".into(),
effective_model_type: "fixture-v1".into(),
architecture_fingerprint: "geometry-1".into(),
layer_count: 4,
global_layer_start: 1,
sink_tokens: 0,
topology: PromptCacheTopology::default(),
}
.prompt_cache_identity(&layout)
.unwrap();
assert_eq!(identity.global_layer_start(), 1);
assert_eq!(identity.global_layer_end(), 3);
assert_eq!(identity.layer_layout(), layout.layers());
}
#[test]
fn prompt_identity_derives_offsets_from_segments() {
let policy = LayerCachePolicy::key_value(AttentionPolicy::Full, 1, 8).unwrap();
let layout = StateLayout::segmented(
LayerSchedule::new(2, vec![policy.clone(), policy]).unwrap(),
[
StateSegmentSpec::new("target", 0..1, StateSegmentLifetime::Persistent, 0).unwrap(),
StateSegmentSpec::new("prediction", 1..2, StateSegmentLifetime::Persistent, -1)
.unwrap(),
],
)
.unwrap();
let identity = ModelStateIdentity {
model_family: "fixture".into(),
effective_model_type: "fixture-v1".into(),
architecture_fingerprint: "geometry-1".into(),
layer_count: 2,
global_layer_start: 0,
sink_tokens: 0,
topology: PromptCacheTopology::default(),
}
.prompt_cache_identity(&layout)
.unwrap();
assert_eq!(identity.layer_prefix_offsets(), [0, -1]);
assert_eq!(identity.state_segments().len(), 2);
assert_eq!(identity.state_segments()[0].id(), "target");
assert_eq!(identity.state_segments()[0].layers(), 0..1);
assert_eq!(identity.state_segments()[1].id(), "prediction");
assert_eq!(identity.state_segments()[1].layers(), 1..2);
let prediction = identity.select_state_segment("prediction").unwrap();
assert_eq!(prediction.global_layer_start(), 1);
assert_eq!(prediction.global_layer_end(), 2);
assert_eq!(prediction.layer_prefix_offsets(), [-1]);
assert_eq!(prediction.state_segments()[0].id(), "prediction");
assert_eq!(prediction.state_segments()[0].layers(), 0..1);
}
#[test]
fn state_layout_exposes_stable_semantic_component_names() {
let layout = StateLayout::new(
LayerSchedule::new(
1,
vec![
LayerCachePolicy::compressed_latent_rotary(AttentionPolicy::Full, 16, 8)
.unwrap(),
],
)
.unwrap(),
)
.unwrap();
let names = layout
.components(0)
.unwrap()
.iter()
.map(|component| component.role().stable_name())
.collect::<Vec<_>>();
assert_eq!(
names,
["attention.compressed_latent", "attention.rotary_keys"]
);
}
fn four_layer_schedule() -> LayerSchedule<LayerCachePolicy> {
LayerSchedule::new(
4,
(0..4)
.map(|_| LayerCachePolicy::key_value(AttentionPolicy::Full, 1, 8).unwrap())
.collect(),
)
.unwrap()
}
#[test]
fn composite_state_segments_are_canonical_and_cover_every_layer() {
let layout = StateLayout::segmented(
four_layer_schedule(),
[
StateSegmentSpec::new("depth", 2..4, StateSegmentLifetime::FrameLocal, 0).unwrap(),
StateSegmentSpec::new("temporal", 0..2, StateSegmentLifetime::Persistent, 0)
.unwrap(),
],
)
.unwrap();
assert_eq!(
layout
.segments()
.iter()
.map(|segment| (segment.id().as_str(), segment.layers(), segment.lifetime()))
.collect::<Vec<_>>(),
[
("temporal", 0..2, StateSegmentLifetime::Persistent),
("depth", 2..4, StateSegmentLifetime::FrameLocal),
]
);
assert_eq!(
layout.segment_for_layer(0).unwrap().id().as_str(),
"temporal"
);
assert_eq!(layout.segment_for_layer(3).unwrap().id().as_str(), "depth");
assert!(layout.segment_for_layer(4).is_none());
}
#[test]
fn state_layout_slice_preserves_and_rebases_segment_frontiers() {
let layout = StateLayout::segmented(
four_layer_schedule(),
[
StateSegmentSpec::new("target", 0..2, StateSegmentLifetime::Persistent, 0).unwrap(),
StateSegmentSpec::new("prediction", 2..4, StateSegmentLifetime::Persistent, -1)
.unwrap(),
],
)
.unwrap();
let sliced = layout.slice(1..4).unwrap();
assert_eq!(sliced.segments()[0].layers(), 0..1);
assert_eq!(sliced.segments()[1].layers(), 1..3);
assert_eq!(sliced.layer_prefix_offsets(), [0, -1, -1]);
}
#[test]
fn segment_identity_lifetime_and_offset_participate_in_layout_equality() {
let layout = |depth_name, lifetime, offset| {
StateLayout::segmented(
four_layer_schedule(),
[
StateSegmentSpec::new("temporal", 0..2, StateSegmentLifetime::Persistent, 0)
.unwrap(),
StateSegmentSpec::new(depth_name, 2..4, lifetime, offset).unwrap(),
],
)
.unwrap()
};
let canonical = layout("depth", StateSegmentLifetime::FrameLocal, 0);
assert_ne!(
canonical,
layout("predictor", StateSegmentLifetime::FrameLocal, 0)
);
assert_ne!(
canonical,
layout("depth", StateSegmentLifetime::Persistent, 0)
);
assert_ne!(
canonical,
layout("depth", StateSegmentLifetime::FrameLocal, -1)
);
}
#[test]
fn malformed_segment_partitions_fail_closed() {
let duplicate = StateLayout::segmented(
four_layer_schedule(),
[
StateSegmentSpec::new("cache", 0..2, StateSegmentLifetime::Persistent, 0).unwrap(),
StateSegmentSpec::new("cache", 2..4, StateSegmentLifetime::FrameLocal, 0).unwrap(),
],
)
.unwrap_err();
assert!(matches!(duplicate, StateError::DuplicateSegment { .. }));
let overlap = StateLayout::segmented(
four_layer_schedule(),
[
StateSegmentSpec::new("left", 0..3, StateSegmentLifetime::Persistent, 0).unwrap(),
StateSegmentSpec::new("right", 2..4, StateSegmentLifetime::FrameLocal, 0).unwrap(),
],
)
.unwrap_err();
assert!(matches!(overlap, StateError::OverlappingSegment { .. }));
let gap = StateLayout::segmented(
four_layer_schedule(),
[
StateSegmentSpec::new("left", 0..1, StateSegmentLifetime::Persistent, 0).unwrap(),
StateSegmentSpec::new("right", 2..4, StateSegmentLifetime::FrameLocal, 0).unwrap(),
],
)
.unwrap_err();
assert_eq!(gap, StateError::UnassignedStateLayer { layer: 1 });
let outside = StateLayout::segmented(
four_layer_schedule(),
[StateSegmentSpec::new("all", 0..5, StateSegmentLifetime::Persistent, 0).unwrap()],
)
.unwrap_err();
assert!(matches!(outside, StateError::SegmentOutOfBounds { .. }));
}
}