use std::{
collections::{BTreeMap, BTreeSet},
path::{Path, PathBuf},
};
use eredu_checkpoint::{LinearFormat, SourceTensorEncoding, StoredDtype};
use eredu_core::{
cache::{StateComponentPolicy, StateTensorDtype},
checkpoint::TensorDtype,
ParallelTopology, QuantizationRequest, SessionCapabilities,
};
use eredu_nn::{NeuralBackend, NeuralOperatorCapabilities};
use crate::{
ArchitectureGroupTransport, ArchitectureParameterDescription, ArchitecturePartition,
CacheResidencyPolicy, ExecutionGraph, ExecutionGroupId, ExecutionUnitLayout,
LayerWeightResidency, LayeredArchitecture, ParameterGroupOwner, ParameterGroupSpec,
RuntimeState, StateLayout,
};
pub trait ReplicatedTextArchitecture<B, S>: LayeredArchitecture<B, S>
where
B: NeuralBackend,
S: RuntimeState<B>,
{
fn text_input<'a>(tokens: &'a B::Tensor, mask: Option<&'a B::Tensor>) -> Self::Input<'a>;
fn text_output_selection(&self) -> ReplicatedTextOutputSelection {
ReplicatedTextOutputSelection::LastSequencePosition
}
}
#[derive(Debug, Clone, Copy, Eq, PartialEq)]
#[non_exhaustive]
pub enum ReplicatedTextOutputSelection {
LastSequencePosition,
}
impl ReplicatedTextOutputSelection {
pub const fn sequence_index(self) -> i32 {
match self {
Self::LastSequencePosition => -1,
}
}
}
#[derive(Debug, Clone, Copy, Eq, PartialEq)]
#[non_exhaustive]
pub enum WeightLoweringKind {
Direct,
Derived,
Transform,
DerivedTransform,
}
#[derive(Debug, Clone, Eq, PartialEq)]
pub struct WeightLoweringCapability {
descriptor: WeightLoweringDescriptor,
kind: WeightLoweringKind,
}
impl WeightLoweringCapability {
pub fn new(descriptor: WeightLoweringDescriptor, kind: WeightLoweringKind) -> Self {
Self { descriptor, kind }
}
pub const fn source(&self) -> &SourceTensorEncoding {
self.descriptor.source()
}
pub const fn executable(&self) -> LinearFormat {
self.descriptor.executable()
}
pub const fn kind(&self) -> WeightLoweringKind {
self.kind
}
pub const fn descriptor(&self) -> &WeightLoweringDescriptor {
&self.descriptor
}
}
#[derive(Debug, Clone, Eq, PartialEq)]
pub struct WeightLoweringDescriptor {
source: SourceTensorEncoding,
executable: LinearFormat,
physical_shape: Vec<usize>,
logical_shape: Vec<usize>,
packed_axis: Option<usize>,
}
impl WeightLoweringDescriptor {
pub fn new(
source: SourceTensorEncoding,
executable: LinearFormat,
physical_shape: Vec<usize>,
logical_shape: Vec<usize>,
packed_axis: Option<usize>,
) -> Result<Self, ReplicatedTextContractError> {
if physical_shape.contains(&0)
|| logical_shape.contains(&0)
|| physical_shape.len() != logical_shape.len()
{
return Err(ReplicatedTextContractError::invalid(
"weight lowering requires positive extents and equal physical and logical ranks",
));
}
if packed_axis.is_some_and(|axis| axis >= logical_shape.len()) {
return Err(ReplicatedTextContractError::invalid(
"weight lowering packed axis is outside the logical shape",
));
}
Ok(Self {
source,
executable,
physical_shape,
logical_shape,
packed_axis,
})
}
pub const fn source(&self) -> &SourceTensorEncoding {
&self.source
}
pub const fn executable(&self) -> LinearFormat {
self.executable
}
pub fn physical_shape(&self) -> &[usize] {
&self.physical_shape
}
pub fn logical_shape(&self) -> &[usize] {
&self.logical_shape
}
pub const fn packed_axis(&self) -> Option<usize> {
self.packed_axis
}
pub fn packed_extent(&self) -> Option<usize> {
self.packed_axis.map(|axis| self.logical_shape[axis])
}
}
#[derive(Debug, Clone, Copy, Eq, PartialEq)]
#[non_exhaustive]
pub enum WeightResidencyMechanism {
Resident,
Windowed,
DiskStreamed,
}
#[derive(Debug, Clone, Copy, Eq, Hash, Ord, PartialEq, PartialOrd)]
#[non_exhaustive]
pub enum StateComponentPlacement {
Device,
Paged,
}
#[derive(Debug, Clone, Copy, Eq, PartialEq)]
#[non_exhaustive]
pub enum StateStorageDtype {
F16,
Bf16,
F32,
F64,
Complex64,
I32,
U32,
}
impl StateStorageDtype {
pub const fn bytes(self) -> std::num::NonZeroU8 {
let bytes = match self {
Self::F16 | Self::Bf16 => 2,
Self::F32 | Self::I32 | Self::U32 => 4,
Self::F64 | Self::Complex64 => 8,
};
std::num::NonZeroU8::new(bytes).unwrap()
}
pub const fn is_floating(self) -> bool {
!matches!(self, Self::I32 | Self::U32)
}
pub const fn resolve(policy: StateTensorDtype, floating: Option<Self>) -> Option<Self> {
match policy {
StateTensorDtype::Floating => match floating {
Some(dtype) if dtype.is_floating() => Some(dtype),
_ => None,
},
StateTensorDtype::Float32 => Some(Self::F32),
StateTensorDtype::Int32 => Some(Self::I32),
StateTensorDtype::Uint32 => Some(Self::U32),
}
}
}
#[derive(Debug, Clone, Eq, PartialEq)]
pub struct StateComponentMechanism {
layer: usize,
component: StateComponentPolicy,
device_placement: Option<StateComponentPlacement>,
paged_placement: Option<StateComponentPlacement>,
}
impl StateComponentMechanism {
pub fn new(
layer: usize,
component: StateComponentPolicy,
device_placement: Option<StateComponentPlacement>,
paged_placement: Option<StateComponentPlacement>,
) -> Self {
Self {
layer,
component,
device_placement,
paged_placement,
}
}
pub const fn layer(&self) -> usize {
self.layer
}
pub const fn component(&self) -> &StateComponentPolicy {
&self.component
}
pub const fn placement(
&self,
policy: &CacheResidencyPolicy,
) -> Option<StateComponentPlacement> {
match policy {
CacheResidencyPolicy::Device => self.device_placement,
CacheResidencyPolicy::Paged(_) => self.paged_placement,
}
}
}
#[derive(Debug, Clone, Eq, PartialEq)]
pub struct StateMechanismCapabilities {
floating_state: Option<(TensorDtype, StateStorageDtype)>,
components: Vec<StateComponentMechanism>,
checkpoint: bool,
rollback: bool,
reset: bool,
prompt_cache: bool,
observation_retention: bool,
}
impl StateMechanismCapabilities {
pub fn new(components: impl IntoIterator<Item = StateComponentMechanism>) -> Self {
Self {
floating_state: None,
components: components.into_iter().collect(),
checkpoint: false,
rollback: false,
reset: false,
prompt_cache: false,
observation_retention: false,
}
}
pub fn with_floating_state_dtype(
mut self,
source: TensorDtype,
dtype: StateStorageDtype,
) -> Self {
self.floating_state = Some((source, dtype));
self
}
pub fn floating_state_dtype(&self) -> Option<(&TensorDtype, StateStorageDtype)> {
self.floating_state
.as_ref()
.map(|(source, dtype)| (source, *dtype))
}
pub const fn with_transactions(mut self, checkpoint: bool, rollback: bool) -> Self {
self.checkpoint = checkpoint;
self.rollback = rollback;
self
}
pub const fn with_reset(mut self, supported: bool) -> Self {
self.reset = supported;
self
}
pub const fn with_prompt_cache(mut self, supported: bool) -> Self {
self.prompt_cache = supported;
self
}
pub const fn with_observation_retention(mut self, supported: bool) -> Self {
self.observation_retention = supported;
self
}
pub fn components(&self) -> &[StateComponentMechanism] {
&self.components
}
pub const fn checkpoint(&self) -> bool {
self.checkpoint
}
pub const fn rollback(&self) -> bool {
self.rollback
}
pub const fn reset(&self) -> bool {
self.reset
}
pub const fn prompt_cache(&self) -> bool {
self.prompt_cache
}
pub const fn observation_retention(&self) -> bool {
self.observation_retention
}
}
#[derive(Debug, Clone, Eq, PartialEq)]
pub struct ParameterTransformTarget {
request: QuantizationRequest,
executable: LinearFormat,
descriptor: WeightLoweringDescriptor,
}
impl ParameterTransformTarget {
fn new(
request: QuantizationRequest,
executable: LinearFormat,
descriptor: WeightLoweringDescriptor,
) -> Self {
Self {
request,
executable,
descriptor,
}
}
pub const fn request(&self) -> QuantizationRequest {
self.request
}
pub const fn executable(&self) -> LinearFormat {
self.executable
}
pub const fn descriptor(&self) -> &WeightLoweringDescriptor {
&self.descriptor
}
}
#[derive(Debug, Clone, Copy, Eq, PartialEq)]
#[non_exhaustive]
pub enum ParameterTransformConstraint {
None,
Linear {
packed_axis: usize,
},
}
#[derive(Debug, Clone, Copy, Eq, PartialEq)]
#[non_exhaustive]
pub enum ReplicatedTextParameterRole {
Embedding,
LinearWeight,
LinearBias,
Normalization,
FormatCompanion,
Other,
}
#[derive(Debug, Clone, Eq, PartialEq)]
#[non_exhaustive]
pub enum ReplicatedTextParameterOwner {
StaticRole(String),
ExecutionUnit {
group: String,
unit: usize,
},
}
#[derive(Debug, Clone, Eq, PartialEq)]
#[non_exhaustive]
pub enum ReplicatedTextParameterPresence {
Required,
OptionalPresent,
OptionalAbsent,
Tied {
target: String,
},
Derived {
recipe: String,
},
}
impl ReplicatedTextParameterPresence {
pub fn has_physical_source(&self) -> bool {
matches!(self, Self::Required | Self::OptionalPresent)
}
}
#[derive(Debug, Clone, Eq, PartialEq)]
pub struct ReplicatedTextPhysicalSource {
catalog_key: String,
tensor: String,
shard: PathBuf,
output: String,
source_encoding: SourceTensorEncoding,
encoded_byte_len: u64,
}
impl ReplicatedTextPhysicalSource {
pub fn new(
catalog_key: impl Into<String>,
tensor: impl Into<String>,
shard: impl Into<PathBuf>,
output: impl Into<String>,
source_encoding: SourceTensorEncoding,
encoded_byte_len: u64,
) -> Result<Self, ReplicatedTextContractError> {
let catalog_key = catalog_key.into();
let tensor = tensor.into();
let shard = shard.into();
let output = output.into();
if catalog_key.trim().is_empty()
|| tensor.trim().is_empty()
|| shard.as_os_str().is_empty()
|| output.trim().is_empty()
|| encoded_byte_len == 0
{
return Err(ReplicatedTextContractError::invalid(
"physical source key, tensor, shard, output, and byte length must be valid",
));
}
Ok(Self {
catalog_key,
tensor,
shard,
output,
source_encoding,
encoded_byte_len,
})
}
pub fn catalog_key(&self) -> &str {
&self.catalog_key
}
pub fn tensor(&self) -> &str {
&self.tensor
}
pub fn shard(&self) -> &Path {
&self.shard
}
pub fn output(&self) -> &str {
&self.output
}
pub const fn source_encoding(&self) -> &SourceTensorEncoding {
&self.source_encoding
}
pub const fn encoded_byte_len(&self) -> u64 {
self.encoded_byte_len
}
}
#[derive(Debug, Clone, Eq, PartialEq)]
pub struct ReplicatedTextParameterRequirement {
name: String,
sources: Vec<String>,
physical_sources: Vec<ReplicatedTextPhysicalSource>,
aliases: Vec<String>,
source_encoding: Option<SourceTensorEncoding>,
physical_shape: Option<Vec<usize>>,
logical_shape: Vec<usize>,
role: ReplicatedTextParameterRole,
owner: ReplicatedTextParameterOwner,
presence: ReplicatedTextParameterPresence,
native_executable: LinearFormat,
transform: ParameterTransformConstraint,
linear_companion: Option<(eredu_nn::LinearCompanionRole, String)>,
transform_companions: Option<(String, String)>,
permitted_native_source_dtypes: Vec<eredu_checkpoint::recipe::RecipeDtype>,
}
impl ReplicatedTextParameterRequirement {
#[allow(
clippy::too_many_arguments,
reason = "the constructor validates one complete immutable catalog record"
)]
pub fn new(
name: impl Into<String>,
sources: Vec<String>,
physical_sources: Vec<ReplicatedTextPhysicalSource>,
aliases: Vec<String>,
source_encoding: Option<SourceTensorEncoding>,
physical_shape: Option<Vec<usize>>,
logical_shape: Vec<usize>,
native_executable: LinearFormat,
role: ReplicatedTextParameterRole,
owner: ReplicatedTextParameterOwner,
presence: ReplicatedTextParameterPresence,
transform: ParameterTransformConstraint,
) -> Result<Self, ReplicatedTextContractError> {
let name = name.into();
native_executable
.validate()
.map_err(|error| ReplicatedTextContractError::invalid(error.to_string()))?;
if name.trim().is_empty() {
return Err(ReplicatedTextContractError::invalid(
"logical parameter identity is empty",
));
}
if sources.iter().any(|source| source.trim().is_empty())
|| aliases.iter().any(|alias| alias.trim().is_empty())
{
return Err(ReplicatedTextContractError::invalid(format!(
"logical parameter {name:?} has an empty physical identity"
)));
}
let has_source = !sources.is_empty();
let has_physical_facts = source_encoding.is_some() && physical_shape.is_some();
if source_encoding.is_some() != physical_shape.is_some()
|| (has_source && !has_physical_facts)
{
return Err(ReplicatedTextContractError::invalid(format!(
"logical parameter {name:?} has inconsistent source presence"
)));
}
match presence {
ReplicatedTextParameterPresence::Required
| ReplicatedTextParameterPresence::OptionalPresent
if !has_source =>
{
return Err(ReplicatedTextContractError::invalid(format!(
"physical logical parameter {name:?} has no lowering source"
)));
}
ReplicatedTextParameterPresence::OptionalAbsent
| ReplicatedTextParameterPresence::Tied { .. }
if has_source =>
{
return Err(ReplicatedTextContractError::invalid(format!(
"source-free logical parameter {name:?} has a lowering source"
)));
}
_ => {}
}
let provenance_required =
has_source || matches!(presence, ReplicatedTextParameterPresence::Derived { .. });
if provenance_required != !physical_sources.is_empty() {
return Err(ReplicatedTextContractError::invalid(format!(
"logical parameter {name:?} has inconsistent physical provenance"
)));
}
if physical_sources.is_empty() && has_physical_facts {
return Err(ReplicatedTextContractError::invalid(format!(
"logical parameter {name:?} has physical facts without provenance"
)));
}
if physical_shape
.as_ref()
.is_some_and(|shape| shape.contains(&0))
{
return Err(ReplicatedTextContractError::invalid(format!(
"logical parameter {name:?} has an invalid physical shape"
)));
}
if logical_shape.contains(&0) {
return Err(ReplicatedTextContractError::invalid(format!(
"logical parameter {name:?} has an invalid shape {logical_shape:?}"
)));
}
if let ParameterTransformConstraint::Linear { packed_axis } = transform {
if packed_axis >= logical_shape.len() {
return Err(ReplicatedTextContractError::invalid(format!(
"logical parameter {name:?} has packing axis {packed_axis} outside shape {logical_shape:?}"
)));
}
}
let requirement = Self {
name,
sources,
physical_sources,
aliases,
source_encoding,
physical_shape,
logical_shape,
role,
owner,
presence,
native_executable,
transform,
linear_companion: None,
transform_companions: None,
permitted_native_source_dtypes: Vec::new(),
};
Ok(requirement)
}
pub fn with_permitted_native_source_dtypes(
mut self,
dtypes: Vec<eredu_checkpoint::recipe::RecipeDtype>,
) -> Self {
self.permitted_native_source_dtypes =
dtypes.into_iter().fold(Vec::new(), |mut out, dtype| {
if !out.contains(&dtype) {
out.push(dtype);
}
out
});
self
}
pub fn with_linear_companion(
mut self,
role: eredu_nn::LinearCompanionRole,
primary: impl Into<String>,
) -> Result<Self, ReplicatedTextContractError> {
let primary = primary.into();
if self.role != ReplicatedTextParameterRole::FormatCompanion
|| primary.trim().is_empty()
|| primary == self.name
{
return Err(ReplicatedTextContractError::invalid(format!(
"parameter {:?} has an invalid encoded-linear companion relationship",
self.name
)));
}
self.linear_companion = Some((role, primary));
Ok(self)
}
pub fn with_transform_companions(
mut self,
scale: impl Into<String>,
affine_bias: impl Into<String>,
) -> Result<Self, ReplicatedTextContractError> {
let scale = scale.into();
let affine_bias = affine_bias.into();
if !matches!(self.transform, ParameterTransformConstraint::Linear { .. })
|| scale.trim().is_empty()
|| affine_bias.trim().is_empty()
|| scale == affine_bias
|| scale == self.name
|| affine_bias == self.name
{
return Err(ReplicatedTextContractError::invalid(format!(
"parameter {:?} has invalid transform companion identities",
self.name
)));
}
self.transform_companions = Some((scale, affine_bias));
Ok(self)
}
pub fn name(&self) -> &str {
&self.name
}
pub fn sources(&self) -> &[String] {
&self.sources
}
pub fn physical_sources(&self) -> &[ReplicatedTextPhysicalSource] {
&self.physical_sources
}
pub fn aliases(&self) -> &[String] {
&self.aliases
}
pub const fn source_encoding(&self) -> Option<&SourceTensorEncoding> {
self.source_encoding.as_ref()
}
pub fn physical_shape(&self) -> Option<&[usize]> {
self.physical_shape.as_deref()
}
pub fn logical_shape(&self) -> &[usize] {
&self.logical_shape
}
pub const fn role(&self) -> ReplicatedTextParameterRole {
self.role
}
pub const fn owner(&self) -> &ReplicatedTextParameterOwner {
&self.owner
}
pub const fn presence(&self) -> &ReplicatedTextParameterPresence {
&self.presence
}
pub fn has_lowering_source(&self) -> bool {
!self.sources.is_empty() || !self.physical_sources.is_empty()
}
pub const fn transform_constraint(&self) -> ParameterTransformConstraint {
self.transform
}
pub fn linear_companion(&self) -> Option<(eredu_nn::LinearCompanionRole, &str)> {
self.linear_companion
.as_ref()
.map(|(role, primary)| (*role, primary.as_str()))
}
pub fn transform_companions(&self) -> Option<(&str, &str)> {
self.transform_companions
.as_ref()
.map(|(scale, bias)| (scale.as_str(), bias.as_str()))
}
pub fn permitted_native_source_dtypes(&self) -> &[eredu_checkpoint::recipe::RecipeDtype] {
&self.permitted_native_source_dtypes
}
pub const fn native_executable(&self) -> LinearFormat {
self.native_executable
}
pub fn transform_target(
&self,
request: QuantizationRequest,
) -> Result<Option<ParameterTransformTarget>, ReplicatedTextContractError> {
let packed_axis = match self.transform {
ParameterTransformConstraint::None => return Ok(None),
ParameterTransformConstraint::Linear { packed_axis } => packed_axis,
};
let extent = self.logical_shape[packed_axis];
let executable = match request {
QuantizationRequest::Affine { group_size, bits } => {
let group_size = i32::try_from(group_size).map_err(|_| {
ReplicatedTextContractError::invalid("affine group size exceeds i32")
})?;
let format = eredu_checkpoint::AffineQuantization::new(group_size, i32::from(bits))
.map_err(|error| ReplicatedTextContractError::invalid(error.to_string()))?;
let group_size = usize::try_from(format.group_size).map_err(|_| {
ReplicatedTextContractError::invalid("affine group size is negative")
})?;
if group_size > extent || !extent.is_multiple_of(group_size) {
return Err(ReplicatedTextContractError::invalid(format!(
"affine group size {group_size} does not divide packed extent {extent}"
)));
}
LinearFormat::Affine(format)
}
QuantizationRequest::MxFp4 => {
const MXFP4_BLOCK_SIZE: usize = 32;
if !extent.is_multiple_of(MXFP4_BLOCK_SIZE) {
return Err(ReplicatedTextContractError::invalid(format!(
"MXFP4 packed extent {extent} is not divisible by block size {MXFP4_BLOCK_SIZE}"
)));
}
LinearFormat::MxFp4
}
_ => {
return Err(ReplicatedTextContractError::invalid(
"unknown load-time transform request",
))
}
};
let descriptor = self.lowering_descriptor(executable)?;
Ok(Some(ParameterTransformTarget::new(
request, executable, descriptor,
)))
}
pub fn lowering_descriptor(
&self,
executable: LinearFormat,
) -> Result<WeightLoweringDescriptor, ReplicatedTextContractError> {
let packed_axis = match self.transform {
ParameterTransformConstraint::None => None,
ParameterTransformConstraint::Linear { packed_axis } => Some(packed_axis),
};
let packed_axis = packed_axis
.or_else(|| {
(self.role == ReplicatedTextParameterRole::Embedding
&& executable != LinearFormat::Dense)
.then(|| self.logical_shape.len().checked_sub(1))
.flatten()
})
.or_else(|| {
(matches!(
self.presence,
ReplicatedTextParameterPresence::Derived { .. }
) && executable != LinearFormat::Dense)
.then(|| {
self.physical_shape
.as_ref()
.and_then(|shape| shape.len().checked_sub(1))
})
.flatten()
});
let alias_backed_packed_output = matches!(
self.source_encoding,
Some(
SourceTensorEncoding::Safetensors(StoredDtype::U32)
| SourceTensorEncoding::RecipeOutput(StoredDtype::U32)
)
);
let lowering_shape = if matches!(
self.presence,
ReplicatedTextParameterPresence::Derived { .. }
) && !alias_backed_packed_output
{
self.physical_shape.as_ref().unwrap_or(&self.logical_shape)
} else {
&self.logical_shape
};
WeightLoweringDescriptor::new(
self.source_encoding.clone().ok_or_else(|| {
ReplicatedTextContractError::invalid(format!(
"logical parameter {:?} has no physical lowering source",
self.name
))
})?,
executable,
self.physical_shape.clone().ok_or_else(|| {
ReplicatedTextContractError::invalid(format!(
"logical parameter {:?} has no physical source shape",
self.name
))
})?,
lowering_shape.clone(),
packed_axis,
)
}
}
#[derive(Debug, Clone, Eq, PartialEq, thiserror::Error)]
#[error("invalid replicated text contract: {message}")]
pub struct ReplicatedTextContractError {
message: String,
}
#[derive(Debug, Clone, Copy, Eq, Hash, Ord, PartialEq, PartialOrd)]
#[non_exhaustive]
pub enum ReplicatedTextStateAccess {
Stateless,
KeyValue,
Fixed,
AttentionWithFixed,
CompressedAttention,
CompressedAttentionWithFixed,
}
impl ReplicatedTextContractError {
fn invalid(message: impl Into<String>) -> Self {
Self {
message: message.into(),
}
}
pub fn message(&self) -> &str {
&self.message
}
}
#[derive(Debug, Clone, Eq, PartialEq)]
pub struct ReplicatedTextRequirements {
floating_state_source: Option<TensorDtype>,
architecture_identity: String,
operators: NeuralOperatorCapabilities,
execution_graph: ExecutionGraph,
execution_units: ExecutionUnitLayout,
group_transports: Vec<ArchitectureGroupTransport>,
state_layout: StateLayout,
state_access: ReplicatedTextStateAccess,
parameters: Vec<ReplicatedTextParameterRequirement>,
auxiliary_parameters: Vec<ReplicatedTextParameterRequirement>,
derived_recipes: BTreeMap<String, eredu_checkpoint::recipe::DerivedWeightRecipe>,
derived_recipe_outputs: BTreeMap<String, eredu_checkpoint::recipe::RecipeMetadata>,
shared_source_keys: BTreeSet<String>,
grouped_operations: Vec<GroupedOperationRequirement>,
}
impl ReplicatedTextRequirements {
#[allow(
clippy::too_many_arguments,
reason = "the constructor validates one complete immutable architecture contract"
)]
pub fn new(
architecture_identity: impl Into<String>,
operators: NeuralOperatorCapabilities,
execution_graph: ExecutionGraph,
execution_units: ExecutionUnitLayout,
group_transports: Vec<ArchitectureGroupTransport>,
state_layout: StateLayout,
state_access: ReplicatedTextStateAccess,
parameters: Vec<ReplicatedTextParameterRequirement>,
) -> Result<Self, ReplicatedTextContractError> {
let architecture_identity = architecture_identity.into();
if architecture_identity.trim().is_empty() {
return Err(ReplicatedTextContractError::invalid(
"architecture identity is empty",
));
}
if group_transports.len() != execution_graph.groups().len() {
return Err(ReplicatedTextContractError::invalid(format!(
"{} group transports do not match {} execution groups",
group_transports.len(),
execution_graph.groups().len()
)));
}
if execution_units.group_count() != execution_graph.groups().len()
|| execution_graph
.groups()
.iter()
.enumerate()
.any(|(index, group)| {
execution_units
.group_id(index)
.is_none_or(|id| id.as_str() != group.id())
})
{
return Err(ReplicatedTextContractError::invalid(
"execution-unit layout group identities differ from the execution graph",
));
}
validate_state_access_profile(&state_layout, state_access)?;
let mut names = BTreeSet::new();
if parameters
.iter()
.any(|parameter| !names.insert(parameter.name()))
{
return Err(ReplicatedTextContractError::invalid(
"logical parameter identities are not unique",
));
}
Ok(Self {
floating_state_source: None,
architecture_identity,
operators,
execution_graph,
execution_units,
group_transports,
state_layout,
state_access,
parameters,
auxiliary_parameters: Vec::new(),
derived_recipes: BTreeMap::new(),
derived_recipe_outputs: BTreeMap::new(),
shared_source_keys: BTreeSet::new(),
grouped_operations: Vec::new(),
})
}
pub fn with_floating_state_source(mut self, dtype: TensorDtype) -> Self {
self.floating_state_source = Some(dtype);
self
}
pub fn floating_state_source(&self) -> Option<&TensorDtype> {
self.floating_state_source.as_ref()
}
pub fn with_auxiliary_parameters(
mut self,
parameters: Vec<ReplicatedTextParameterRequirement>,
recipes: BTreeMap<String, eredu_checkpoint::recipe::DerivedWeightRecipe>,
outputs: BTreeMap<String, eredu_checkpoint::recipe::RecipeMetadata>,
) -> Result<Self, ReplicatedTextContractError> {
let primary = self
.parameters
.iter()
.map(|parameter| parameter.name())
.collect::<BTreeSet<_>>();
let mut names = BTreeSet::new();
if parameters
.iter()
.any(|parameter| primary.contains(parameter.name()) || !names.insert(parameter.name()))
{
return Err(ReplicatedTextContractError::invalid(
"auxiliary parameter identities overlap or are not unique",
));
}
if recipes.keys().ne(outputs.keys())
|| recipes
.keys()
.any(|target| !names.contains(target.as_str()))
|| recipes
.keys()
.any(|target| self.derived_recipes.contains_key(target))
{
return Err(ReplicatedTextContractError::invalid(
"auxiliary derivations do not match auxiliary parameters",
));
}
self.derived_recipes.extend(recipes);
self.derived_recipe_outputs.extend(outputs);
self.auxiliary_parameters = parameters;
Ok(self)
}
pub fn with_derived_recipes(
mut self,
recipes: BTreeMap<String, eredu_checkpoint::recipe::DerivedWeightRecipe>,
outputs: BTreeMap<String, eredu_checkpoint::recipe::RecipeMetadata>,
) -> Result<Self, ReplicatedTextContractError> {
self.set_derived_recipes(recipes, outputs, BTreeSet::new())?;
Ok(self)
}
pub fn with_derived_recipes_and_shared_sources(
mut self,
recipes: BTreeMap<String, eredu_checkpoint::recipe::DerivedWeightRecipe>,
outputs: BTreeMap<String, eredu_checkpoint::recipe::RecipeMetadata>,
shared_source_keys: BTreeSet<String>,
) -> Result<Self, ReplicatedTextContractError> {
self.set_derived_recipes(recipes, outputs, shared_source_keys)?;
Ok(self)
}
fn set_derived_recipes(
&mut self,
recipes: BTreeMap<String, eredu_checkpoint::recipe::DerivedWeightRecipe>,
outputs: BTreeMap<String, eredu_checkpoint::recipe::RecipeMetadata>,
shared_source_keys: BTreeSet<String>,
) -> Result<(), ReplicatedTextContractError> {
if recipes.keys().ne(outputs.keys()) {
return Err(ReplicatedTextContractError::invalid(
"derived recipe targets and inferred outputs differ",
));
}
for source in &shared_source_keys {
let claims = recipes
.values()
.filter(|recipe| recipe.source_keys().contains(&source.as_str()))
.count();
if claims < 2 {
return Err(ReplicatedTextContractError::invalid(format!(
"declared shared source {source:?} is not claimed by multiple derived targets"
)));
}
}
for target in recipes.keys() {
let recipe = recipes
.get(target)
.expect("recipe target came from the same map");
let parameter = self
.parameters
.iter_mut()
.find(|parameter| parameter.name == *target)
.ok_or_else(|| {
ReplicatedTextContractError::invalid(format!(
"derived recipe target {target:?} is not a declared parameter"
))
})?;
if matches!(
parameter.presence,
ReplicatedTextParameterPresence::OptionalAbsent
| ReplicatedTextParameterPresence::Tied { .. }
) {
return Err(ReplicatedTextContractError::invalid(format!(
"derived recipe target {target:?} has no independent artifact value"
)));
}
parameter.presence = ReplicatedTextParameterPresence::Derived {
recipe: "architecture.recipe".into(),
};
parameter.sources = recipe
.source_keys()
.into_iter()
.map(str::to_owned)
.collect();
}
self.derived_recipes = recipes;
self.derived_recipe_outputs = outputs;
self.shared_source_keys = shared_source_keys;
Ok(())
}
pub fn architecture_identity(&self) -> &str {
&self.architecture_identity
}
pub fn with_extension_architecture_identity(
mut self,
architecture_identity: impl Into<String>,
) -> Result<Self, ReplicatedTextContractError> {
let architecture_identity = architecture_identity.into();
if architecture_identity.trim().is_empty() {
return Err(ReplicatedTextContractError::invalid(
"extension architecture identity is empty",
));
}
self.architecture_identity = architecture_identity;
Ok(self)
}
pub fn with_grouped_operations(
mut self,
operations: impl IntoIterator<Item = GroupedOperationRequirement>,
) -> Self {
self.grouped_operations = operations.into_iter().collect();
self
}
pub const fn operators(&self) -> NeuralOperatorCapabilities {
self.operators
}
pub const fn execution_graph(&self) -> &ExecutionGraph {
&self.execution_graph
}
pub const fn execution_units(&self) -> &ExecutionUnitLayout {
&self.execution_units
}
pub fn group_transports(&self) -> &[ArchitectureGroupTransport] {
&self.group_transports
}
pub const fn state_layout(&self) -> &StateLayout {
&self.state_layout
}
pub const fn state_access(&self) -> ReplicatedTextStateAccess {
self.state_access
}
pub fn parameters(&self) -> &[ReplicatedTextParameterRequirement] {
&self.parameters
}
pub fn auxiliary_parameters(&self) -> &[ReplicatedTextParameterRequirement] {
&self.auxiliary_parameters
}
pub fn derived_recipes(
&self,
) -> &BTreeMap<String, eredu_checkpoint::recipe::DerivedWeightRecipe> {
&self.derived_recipes
}
pub fn derived_recipe_outputs(
&self,
) -> &BTreeMap<String, eredu_checkpoint::recipe::RecipeMetadata> {
&self.derived_recipe_outputs
}
pub fn shared_source_keys(&self) -> &BTreeSet<String> {
&self.shared_source_keys
}
pub fn grouped_operations(&self) -> &[GroupedOperationRequirement] {
&self.grouped_operations
}
}
fn validate_state_access_profile(
layout: &StateLayout,
access: ReplicatedTextStateAccess,
) -> Result<(), ReplicatedTextContractError> {
use eredu_core::cache::StateComponentRole;
let roles = (0..layout.len())
.flat_map(|layer| {
layout
.components(layer)
.expect("validated state layout exposes every layer")
})
.map(StateComponentPolicy::role)
.collect::<Vec<_>>();
let ordinary = |role| {
matches!(
role,
StateComponentRole::AttentionKeys | StateComponentRole::AttentionValues
)
};
let compressed = |role| {
matches!(
role,
StateComponentRole::CompressedLatent | StateComponentRole::RotaryKeys
)
};
let fixed = |role| matches!(role, StateComponentRole::Fixed(_));
let has_ordinary = roles.iter().copied().any(ordinary);
let has_compressed = roles.iter().copied().any(compressed);
let has_fixed = roles.iter().copied().any(fixed);
let coherent = match access {
ReplicatedTextStateAccess::Stateless => roles.is_empty(),
ReplicatedTextStateAccess::KeyValue => roles.iter().copied().all(ordinary) && has_ordinary,
ReplicatedTextStateAccess::Fixed => roles.iter().copied().all(fixed) && has_fixed,
ReplicatedTextStateAccess::AttentionWithFixed => {
roles
.iter()
.copied()
.all(|role| ordinary(role) || fixed(role))
&& has_ordinary
&& has_fixed
}
ReplicatedTextStateAccess::CompressedAttention => {
roles.iter().copied().all(compressed) && has_compressed
}
ReplicatedTextStateAccess::CompressedAttentionWithFixed => {
roles
.iter()
.copied()
.all(|role| compressed(role) || fixed(role))
&& has_compressed
&& has_fixed
}
};
if !coherent {
return Err(ReplicatedTextContractError::invalid(format!(
"state access profile {access:?} does not match component roles {roles:?}"
)));
}
Ok(())
}
#[derive(Debug, Clone, Copy, Eq, Hash, Ord, PartialEq, PartialOrd)]
#[non_exhaustive]
pub enum GroupedOperationRequirement {
GatedProduct,
GatedProductTensorParallelPartial,
Relu2,
Relu2TensorParallelPartial,
}
#[derive(Debug, Clone, Copy, Eq, PartialEq)]
pub struct AddressableStorageCapabilities {
bulk_access: bool,
incremental_access: bool,
lease_completion: bool,
maximum_compact_bytes: u64,
tiers: AddressableStorageTiers,
}
#[derive(Debug, Clone, Copy, Eq, PartialEq)]
pub struct AddressableStorageTiers {
device: bool,
host: bool,
disk: bool,
}
impl AddressableStorageTiers {
pub const fn new(device: bool, host: bool, disk: bool) -> Self {
Self { device, host, disk }
}
pub const fn device(self) -> bool {
self.device
}
pub const fn host(self) -> bool {
self.host
}
pub const fn disk(self) -> bool {
self.disk
}
}
impl AddressableStorageCapabilities {
pub const fn new(
bulk_access: bool,
incremental_access: bool,
lease_completion: bool,
maximum_compact_bytes: u64,
) -> Self {
Self {
bulk_access,
incremental_access,
lease_completion,
maximum_compact_bytes,
tiers: AddressableStorageTiers::new(true, true, true),
}
}
pub const fn with_tiers(mut self, tiers: AddressableStorageTiers) -> Self {
self.tiers = tiers;
self
}
pub const fn bulk_access(self) -> bool {
self.bulk_access
}
pub const fn incremental_access(self) -> bool {
self.incremental_access
}
pub const fn lease_completion(self) -> bool {
self.lease_completion
}
pub const fn maximum_compact_bytes(self) -> u64 {
self.maximum_compact_bytes
}
pub const fn tiers(self) -> AddressableStorageTiers {
self.tiers
}
}
#[derive(Debug, Clone, Eq, PartialEq)]
pub struct BackendMechanismCapabilities {
operators: NeuralOperatorCapabilities,
weight_lowerings: Vec<WeightLoweringCapability>,
weight_residencies: Vec<WeightResidencyMechanism>,
state: StateMechanismCapabilities,
session: SessionCapabilities,
prompt_cache: bool,
exact_completion: bool,
grouped_operations: Vec<GroupedOperationRequirement>,
indexed_movement: bool,
addressable_storage: Option<AddressableStorageCapabilities>,
}
impl BackendMechanismCapabilities {
pub fn new(
operators: NeuralOperatorCapabilities,
weight_lowerings: Vec<WeightLoweringCapability>,
weight_residencies: Vec<WeightResidencyMechanism>,
state: StateMechanismCapabilities,
) -> Self {
Self {
operators,
weight_lowerings,
weight_residencies,
state,
session: SessionCapabilities::default(),
prompt_cache: false,
exact_completion: false,
grouped_operations: Vec::new(),
indexed_movement: false,
addressable_storage: None,
}
}
pub const fn with_session(mut self, session: SessionCapabilities) -> Self {
self.session = session;
self
}
pub const fn with_prompt_cache(mut self, supported: bool) -> Self {
self.prompt_cache = supported;
self
}
pub const fn with_exact_completion(mut self, supported: bool) -> Self {
self.exact_completion = supported;
self
}
pub fn with_grouped_operations(
mut self,
operations: impl IntoIterator<Item = GroupedOperationRequirement>,
) -> Self {
self.grouped_operations = operations.into_iter().collect();
self
}
pub const fn with_indexed_movement(mut self, supported: bool) -> Self {
self.indexed_movement = supported;
self
}
pub const fn with_addressable_storage(
mut self,
capabilities: AddressableStorageCapabilities,
) -> Self {
self.addressable_storage = Some(capabilities);
self
}
pub const fn operators(&self) -> NeuralOperatorCapabilities {
self.operators
}
pub fn weight_lowerings(&self) -> &[WeightLoweringCapability] {
&self.weight_lowerings
}
pub fn weight_residencies(&self) -> &[WeightResidencyMechanism] {
&self.weight_residencies
}
pub const fn state(&self) -> &StateMechanismCapabilities {
&self.state
}
pub const fn session(&self) -> SessionCapabilities {
self.session
}
pub const fn prompt_cache(&self) -> bool {
self.prompt_cache
}
pub const fn exact_completion(&self) -> bool {
self.exact_completion
}
pub fn grouped_operations(&self) -> &[GroupedOperationRequirement] {
&self.grouped_operations
}
pub const fn indexed_movement(&self) -> bool {
self.indexed_movement
}
pub const fn addressable_storage(&self) -> Option<AddressableStorageCapabilities> {
self.addressable_storage
}
}
#[derive(Debug, Clone, Eq, PartialEq)]
pub struct ReplicatedTextSelectionRequest {
max_cached_shards: usize,
topology: Option<ParallelTopology>,
residency: LayerWeightResidency,
state: CacheResidencyPolicy,
quantization: Option<QuantizationRequest>,
session: SessionCapabilities,
prompt_cache: bool,
exact_completion: bool,
}
impl ReplicatedTextSelectionRequest {
pub fn new(residency: LayerWeightResidency, state: CacheResidencyPolicy) -> Self {
Self {
max_cached_shards: residency.max_cached_shards(),
topology: None,
residency,
state,
quantization: None,
session: SessionCapabilities::default(),
prompt_cache: false,
exact_completion: false,
}
}
pub const fn with_max_cached_shards(mut self, maximum: std::num::NonZeroUsize) -> Self {
self.max_cached_shards = maximum.get();
self
}
pub const fn max_cached_shards(&self) -> usize {
self.max_cached_shards
}
pub const fn with_topology(mut self, topology: ParallelTopology) -> Self {
self.topology = Some(topology);
self
}
pub const fn with_quantization(mut self, quantization: QuantizationRequest) -> Self {
self.quantization = Some(quantization);
self
}
pub const fn with_session(mut self, session: SessionCapabilities) -> Self {
self.session = session;
self
}
pub const fn with_prompt_cache(mut self, required: bool) -> Self {
self.prompt_cache = required;
self
}
pub const fn with_exact_completion(mut self, required: bool) -> Self {
self.exact_completion = required;
self
}
pub const fn topology(&self) -> Option<ParallelTopology> {
self.topology
}
pub const fn residency(&self) -> LayerWeightResidency {
self.residency
}
pub const fn state(&self) -> &CacheResidencyPolicy {
&self.state
}
pub const fn quantization(&self) -> Option<QuantizationRequest> {
self.quantization
}
pub const fn session(&self) -> SessionCapabilities {
self.session
}
pub const fn prompt_cache(&self) -> bool {
self.prompt_cache
}
pub const fn exact_completion(&self) -> bool {
self.exact_completion
}
}
#[derive(Debug, Clone, Eq, PartialEq)]
pub struct SelectedParameterRealization {
name: String,
sources: Vec<String>,
physical_sources: Vec<ReplicatedTextPhysicalSource>,
source_encoding: SourceTensorEncoding,
executable: LinearFormat,
lowering: WeightLoweringKind,
}
#[derive(Debug, Clone, Eq, PartialEq)]
pub struct ReplicatedTextMaterializationTask {
name: String,
sources: Vec<String>,
physical_sources: Vec<ReplicatedTextPhysicalSource>,
aliases: Vec<String>,
source_encoding: SourceTensorEncoding,
physical_shape: Vec<usize>,
logical_shape: Vec<usize>,
role: ReplicatedTextParameterRole,
owner: ReplicatedTextParameterOwner,
presence: ReplicatedTextParameterPresence,
executable: LinearFormat,
lowering: WeightLoweringKind,
lowering_descriptor: WeightLoweringDescriptor,
derived_recipe: Option<eredu_checkpoint::recipe::DerivedWeightRecipe>,
derived_output: Option<eredu_checkpoint::recipe::RecipeMetadata>,
shared_source_keys: BTreeSet<String>,
permitted_native_source_dtypes: Vec<eredu_checkpoint::recipe::RecipeDtype>,
output_companions: Vec<ReplicatedTextOutputCompanion>,
}
#[derive(Debug, Clone, Eq, PartialEq)]
pub struct ReplicatedTextMaterializationPartitionPlan {
task_count: usize,
static_tasks: Vec<usize>,
unit_tasks: Vec<Vec<usize>>,
}
impl ReplicatedTextMaterializationPartitionPlan {
pub const fn task_count(&self) -> usize {
self.task_count
}
pub fn static_task_indices(&self) -> &[usize] {
&self.static_tasks
}
pub fn unit_task_indices(&self) -> &[Vec<usize>] {
&self.unit_tasks
}
pub fn static_tasks<'a>(
&self,
tasks: &'a [ReplicatedTextMaterializationTask],
) -> Result<Vec<&'a ReplicatedTextMaterializationTask>, ReplicatedTextContractError> {
self.validate_task_slice(tasks)?;
Ok(self
.static_tasks
.iter()
.map(|index| &tasks[*index])
.collect())
}
pub fn unit_tasks<'a>(
&self,
tasks: &'a [ReplicatedTextMaterializationTask],
) -> Result<Vec<Vec<&'a ReplicatedTextMaterializationTask>>, ReplicatedTextContractError> {
self.validate_task_slice(tasks)?;
Ok(self
.unit_tasks
.iter()
.map(|indices| indices.iter().map(|index| &tasks[*index]).collect())
.collect())
}
fn validate_task_slice(
&self,
tasks: &[ReplicatedTextMaterializationTask],
) -> Result<(), ReplicatedTextContractError> {
if tasks.len() != self.task_count {
return Err(ReplicatedTextContractError::invalid(format!(
"materialization partition plan expects {} tasks, got {}",
self.task_count,
tasks.len()
)));
}
Ok(())
}
}
#[derive(Debug, Clone, Eq, PartialEq)]
pub struct ReplicatedTextTransformGroup {
quantization: eredu_checkpoint::WeightQuantization,
task_indices: Vec<usize>,
}
impl ReplicatedTextTransformGroup {
pub const fn quantization(&self) -> eredu_checkpoint::WeightQuantization {
self.quantization
}
pub fn task_indices(&self) -> &[usize] {
&self.task_indices
}
pub fn tasks<'a>(
&self,
tasks: &'a [ReplicatedTextMaterializationTask],
) -> Result<Vec<&'a ReplicatedTextMaterializationTask>, ReplicatedTextContractError> {
self.task_indices
.iter()
.map(|index| {
tasks.get(*index).ok_or_else(|| {
ReplicatedTextContractError::invalid(
"transform group was applied to a different materialization task slice",
)
})
})
.collect()
}
}
#[derive(Debug, Clone, Eq, PartialEq)]
pub struct ReplicatedTextOutputCompanion {
name: String,
role: eredu_nn::LinearCompanionRole,
logical_shape: Vec<usize>,
owner: ParameterGroupOwner,
materialization_task: Option<Box<ReplicatedTextMaterializationTask>>,
catalog_source: Option<ReplicatedTextPhysicalSource>,
derived_recipe: Option<eredu_checkpoint::recipe::DerivedWeightRecipe>,
derived_output: Option<eredu_checkpoint::recipe::RecipeMetadata>,
}
impl ReplicatedTextOutputCompanion {
pub fn new(
name: impl Into<String>,
role: eredu_nn::LinearCompanionRole,
logical_shape: Vec<usize>,
owner: ParameterGroupOwner,
) -> Result<Self, ReplicatedTextContractError> {
let name = name.into();
if name.trim().is_empty() || logical_shape.is_empty() || logical_shape.contains(&0) {
return Err(ReplicatedTextContractError::invalid(
"materialization output companion identity or geometry is invalid",
));
}
Ok(Self {
name,
role,
logical_shape,
owner,
materialization_task: None,
catalog_source: None,
derived_recipe: None,
derived_output: None,
})
}
pub(crate) fn with_derived_recipe(
mut self,
recipe: eredu_checkpoint::recipe::DerivedWeightRecipe,
output: eredu_checkpoint::recipe::RecipeMetadata,
) -> Self {
self.derived_recipe = Some(recipe);
self.derived_output = Some(output);
self
}
pub fn name(&self) -> &str {
&self.name
}
pub const fn role(&self) -> eredu_nn::LinearCompanionRole {
self.role
}
pub fn logical_shape(&self) -> &[usize] {
&self.logical_shape
}
pub const fn owner(&self) -> &ParameterGroupOwner {
&self.owner
}
pub(crate) fn with_materialization_task(
mut self,
task: ReplicatedTextMaterializationTask,
) -> Result<Self, ReplicatedTextContractError> {
let names_output =
task.name() == self.name || task.aliases().iter().any(|alias| alias == &self.name);
if !names_output || !task.output_companions().is_empty() {
return Err(ReplicatedTextContractError::invalid(format!(
"companion {:?} has an inconsistent standalone materialization task",
self.name
)));
}
self.materialization_task = Some(Box::new(task));
Ok(self)
}
pub fn with_catalog_source(mut self, source: ReplicatedTextPhysicalSource) -> Self {
self.catalog_source = Some(source);
self
}
pub fn materialization_task(&self) -> Option<&ReplicatedTextMaterializationTask> {
self.materialization_task.as_deref()
}
pub const fn catalog_source(&self) -> Option<&ReplicatedTextPhysicalSource> {
self.catalog_source.as_ref()
}
pub const fn derived_recipe(&self) -> Option<&eredu_checkpoint::recipe::DerivedWeightRecipe> {
self.derived_recipe.as_ref()
}
pub const fn derived_output(&self) -> Option<&eredu_checkpoint::recipe::RecipeMetadata> {
self.derived_output.as_ref()
}
}
impl ReplicatedTextMaterializationTask {
#[allow(clippy::too_many_arguments)]
pub fn from_exact_source(
name: impl Into<String>,
physical_source: ReplicatedTextPhysicalSource,
aliases: Vec<String>,
physical_shape: Vec<usize>,
logical_shape: Vec<usize>,
role: ReplicatedTextParameterRole,
owner: ReplicatedTextParameterOwner,
executable: LinearFormat,
lowering: WeightLoweringKind,
lowering_descriptor: WeightLoweringDescriptor,
) -> Result<Self, ReplicatedTextContractError> {
let name = name.into();
if name.trim().is_empty()
|| physical_shape.is_empty()
|| logical_shape.is_empty()
|| physical_shape.contains(&0)
|| logical_shape.contains(&0)
|| lowering_descriptor.source() != physical_source.source_encoding()
|| lowering_descriptor.executable() != executable
|| lowering_descriptor.physical_shape() != physical_shape
|| lowering_descriptor.logical_shape() != logical_shape
{
return Err(ReplicatedTextContractError::invalid(
"exact auxiliary materialization task is internally inconsistent",
));
}
let source = physical_source.catalog_key().to_owned();
Ok(Self {
name,
sources: vec![source],
physical_sources: vec![physical_source],
aliases,
source_encoding: lowering_descriptor.source().clone(),
physical_shape,
logical_shape,
role,
owner,
presence: ReplicatedTextParameterPresence::Required,
executable,
lowering,
lowering_descriptor,
derived_recipe: None,
derived_output: None,
shared_source_keys: BTreeSet::new(),
permitted_native_source_dtypes: Vec::new(),
output_companions: Vec::new(),
})
}
pub(crate) fn set_output_companions(
&mut self,
mut companions: Vec<ReplicatedTextOutputCompanion>,
) -> Result<(), ReplicatedTextContractError> {
companions.sort_by(|left, right| {
left.role
.cmp(&right.role)
.then_with(|| left.name.cmp(&right.name))
});
if companions
.windows(2)
.any(|pair| pair[0].name == pair[1].name || pair[0].role == pair[1].role)
{
return Err(ReplicatedTextContractError::invalid(format!(
"materialization task {:?} has duplicate output companions",
self.name
)));
}
let roles = companions
.iter()
.map(|companion| companion.role)
.collect::<Vec<_>>();
let expected = if companions.is_empty()
&& matches!(
self.lowering,
WeightLoweringKind::Direct | WeightLoweringKind::Derived
) {
Vec::new()
} else {
match self.executable {
LinearFormat::Dense | LinearFormat::GgufIQuant { .. } => Vec::new(),
LinearFormat::MxFp4 | LinearFormat::E4M3BlockFp8(_) => {
vec![eredu_nn::LinearCompanionRole::Scale]
}
LinearFormat::Affine(_) => vec![
eredu_nn::LinearCompanionRole::Scale,
eredu_nn::LinearCompanionRole::AffineBias,
],
}
};
let mut expected = expected;
expected.sort();
if roles != expected {
return Err(ReplicatedTextContractError::invalid(format!(
"materialization task {:?} executable {:?} requires companion roles {:?}, got {:?}",
self.name, self.executable, expected, roles
)));
}
self.output_companions = companions;
Ok(())
}
pub fn with_output_companions(
mut self,
companions: Vec<ReplicatedTextOutputCompanion>,
) -> Result<Self, ReplicatedTextContractError> {
self.set_output_companions(companions)?;
Ok(self)
}
pub fn name(&self) -> &str {
&self.name
}
pub fn sources(&self) -> &[String] {
&self.sources
}
pub fn physical_sources(&self) -> &[ReplicatedTextPhysicalSource] {
&self.physical_sources
}
pub fn aliases(&self) -> &[String] {
&self.aliases
}
pub fn shared_source_keys(&self) -> &BTreeSet<String> {
&self.shared_source_keys
}
pub fn permitted_native_source_dtypes(&self) -> &[eredu_checkpoint::recipe::RecipeDtype] {
&self.permitted_native_source_dtypes
}
pub const fn source_encoding(&self) -> &SourceTensorEncoding {
&self.source_encoding
}
pub fn physical_shape(&self) -> &[usize] {
&self.physical_shape
}
pub fn logical_shape(&self) -> &[usize] {
&self.logical_shape
}
pub const fn role(&self) -> ReplicatedTextParameterRole {
self.role
}
pub const fn owner(&self) -> &ReplicatedTextParameterOwner {
&self.owner
}
pub const fn presence(&self) -> &ReplicatedTextParameterPresence {
&self.presence
}
pub const fn executable(&self) -> LinearFormat {
self.executable
}
pub const fn lowering(&self) -> WeightLoweringKind {
self.lowering
}
pub const fn lowering_descriptor(&self) -> &WeightLoweringDescriptor {
&self.lowering_descriptor
}
pub const fn derived_recipe(&self) -> Option<&eredu_checkpoint::recipe::DerivedWeightRecipe> {
self.derived_recipe.as_ref()
}
pub const fn derived_output(&self) -> Option<&eredu_checkpoint::recipe::RecipeMetadata> {
self.derived_output.as_ref()
}
pub fn output_companions(&self) -> &[ReplicatedTextOutputCompanion] {
&self.output_companions
}
pub fn source_recipe(
&self,
) -> Result<eredu_checkpoint::recipe::DerivedWeightRecipe, ReplicatedTextContractError> {
let expects_recipe = matches!(
self.lowering,
WeightLoweringKind::Derived | WeightLoweringKind::DerivedTransform
);
match (expects_recipe, self.derived_recipe.as_ref()) {
(true, Some(recipe)) => {
let declared = self
.sources
.iter()
.map(String::as_str)
.collect::<BTreeSet<_>>();
let consumed = recipe.source_keys().into_iter().collect::<BTreeSet<_>>();
if declared != consumed {
return Err(ReplicatedTextContractError::invalid(format!(
"materialization task {:?} recipe sources differ from its exact source catalog",
self.name
)));
}
Ok(recipe.clone())
}
(false, None) => {
let [source] = self.sources.as_slice() else {
return Err(ReplicatedTextContractError::invalid(format!(
"direct materialization task {:?} must name exactly one source",
self.name
)));
};
Ok(eredu_checkpoint::recipe::DerivedWeightRecipe::source(
source.clone(),
eredu_checkpoint::store::TensorSelection::Full,
))
}
(true, None) => Err(ReplicatedTextContractError::invalid(format!(
"derived materialization task {:?} has no exact recipe",
self.name
))),
(false, Some(_)) => Err(ReplicatedTextContractError::invalid(format!(
"direct materialization task {:?} unexpectedly carries a recipe",
self.name
))),
}
}
}
pub fn plan_replicated_text_materialization_tasks(
tasks: &[ReplicatedTextMaterializationTask],
layout: &ExecutionUnitLayout,
) -> Result<ReplicatedTextMaterializationPartitionPlan, ReplicatedTextContractError> {
let mut static_tasks = Vec::new();
let mut unit_tasks = vec![Vec::new(); layout.len()];
for (task_index, task) in tasks.iter().enumerate() {
match task.owner() {
ReplicatedTextParameterOwner::StaticRole(_) => static_tasks.push(task_index),
ReplicatedTextParameterOwner::ExecutionUnit { group, unit } => {
let group_index = (0..layout.group_count())
.find(|index| {
layout
.group_id(*index)
.is_some_and(|id| id.as_str() == group)
})
.ok_or_else(|| {
ReplicatedTextContractError::invalid(format!(
"exact task {:?} names unknown execution group {group:?}",
task.name()
))
})?;
let ordinal = layout.ordinal(group_index, *unit).ok_or_else(|| {
ReplicatedTextContractError::invalid(format!(
"exact task {:?} names unknown unit {unit} in group {group:?}",
task.name()
))
})?;
unit_tasks[ordinal].push(task_index);
}
}
}
Ok(ReplicatedTextMaterializationPartitionPlan {
task_count: tasks.len(),
static_tasks,
unit_tasks,
})
}
pub fn plan_local_replicated_text_materialization_tasks(
tasks: &[ReplicatedTextMaterializationTask],
global_layout: &ExecutionUnitLayout,
addresses: &[crate::ExecutionUnitAddress],
) -> Result<ReplicatedTextMaterializationPartitionPlan, ReplicatedTextContractError> {
if addresses.is_empty() {
return Err(ReplicatedTextContractError::invalid(
"local partition has no selected execution units",
));
}
let mut seen = BTreeSet::new();
for address in addresses {
if global_layout.address(
global_layout
.ordinal(address.group(), address.index())
.unwrap_or(usize::MAX),
) != Some(*address)
{
return Err(ReplicatedTextContractError::invalid(format!(
"local partition names unknown global unit {}.{}",
address.group(),
address.index()
)));
}
if !seen.insert((address.group(), address.index())) {
return Err(ReplicatedTextContractError::invalid(format!(
"local partition repeats global unit {}.{}",
address.group(),
address.index()
)));
}
}
let mut static_tasks = Vec::new();
let mut unit_tasks = vec![Vec::new(); addresses.len()];
for (task_index, task) in tasks.iter().enumerate() {
match task.owner() {
ReplicatedTextParameterOwner::StaticRole(_) => static_tasks.push(task_index),
ReplicatedTextParameterOwner::ExecutionUnit { group, unit } => {
let local = addresses
.iter()
.position(|address| {
global_layout
.group_id(address.group())
.is_some_and(|id| id.as_str() == group)
&& address.index() == *unit
})
.ok_or_else(|| {
ReplicatedTextContractError::invalid(format!(
"local task {:?} has no owned global unit {group}.{unit}",
task.name()
))
})?;
unit_tasks[local].push(task_index);
}
}
}
Ok(ReplicatedTextMaterializationPartitionPlan {
task_count: tasks.len(),
static_tasks,
unit_tasks,
})
}
pub fn locally_materialized_replicated_text_outputs(
tasks: &[ReplicatedTextMaterializationTask],
) -> BTreeSet<String> {
tasks
.iter()
.filter(|task| {
matches!(
task.lowering(),
WeightLoweringKind::Transform | WeightLoweringKind::DerivedTransform
)
})
.flat_map(|task| {
std::iter::once(task.name().to_owned()).chain(
task.output_companions()
.iter()
.map(|companion| companion.name().to_owned()),
)
})
.collect()
}
pub fn group_replicated_text_transform_tasks(
tasks: &[ReplicatedTextMaterializationTask],
) -> Result<Vec<ReplicatedTextTransformGroup>, ReplicatedTextContractError> {
let mut groups = Vec::<ReplicatedTextTransformGroup>::new();
for (task_index, task) in tasks.iter().enumerate().filter(|(_, task)| {
matches!(
task.lowering(),
WeightLoweringKind::Transform | WeightLoweringKind::DerivedTransform
)
}) {
let quantization = task.executable().weight_quantization().ok_or_else(|| {
ReplicatedTextContractError::invalid(format!(
"selected materialization task {:?} has no packed output format",
task.name()
))
})?;
if let Some(group) = groups
.iter_mut()
.find(|group| group.quantization == quantization)
{
group.task_indices.push(task_index);
} else {
groups.push(ReplicatedTextTransformGroup {
quantization,
task_indices: vec![task_index],
});
}
}
Ok(groups)
}
pub fn selected_materialization_task_bytes(
task: &ReplicatedTextMaterializationTask,
) -> Result<u64, ReplicatedTextContractError> {
let transforms = matches!(
task.lowering(),
WeightLoweringKind::Transform | WeightLoweringKind::DerivedTransform
);
if !transforms {
if let Some(output) = task.derived_output() {
return Ok(output.byte_len());
}
return task
.physical_sources()
.iter()
.try_fold(0u64, |total, source| {
total.checked_add(source.encoded_byte_len()).ok_or_else(|| {
ReplicatedTextContractError::invalid(format!(
"materialization task {:?} physical byte total overflowed",
task.name()
))
})
});
}
let dtype = task
.source_encoding()
.scalar_dtype()
.map(eredu_checkpoint::recipe::RecipeDtype::from)
.ok_or_else(|| {
ReplicatedTextContractError::invalid(format!(
"materialization task {:?} transforms a non-scalar source",
task.name()
))
})?;
let source_bytes = task
.derived_output()
.map(|output| output.byte_len())
.or_else(|| {
task.physical_sources()
.first()
.map(|source| source.encoded_byte_len())
})
.ok_or_else(|| {
ReplicatedTextContractError::invalid(format!(
"materialization task {:?} has no source byte extent",
task.name()
))
})?;
let metadata = eredu_checkpoint::recipe::RecipeMetadata {
shape: task.logical_shape().to_vec(),
dtype,
byte_len: source_bytes,
};
crate::selected_addressable_parameter_bytes(task, &metadata)
.map_err(|error| ReplicatedTextContractError::invalid(error.to_string()))
}
pub fn replicated_text_materialization_tasks(
selected: &SelectedReplicatedTextRealization,
) -> Result<Vec<ReplicatedTextMaterializationTask>, ReplicatedTextContractError> {
if selected.materialization_tasks.is_empty() && !selected.parameters.is_empty() {
return Err(ReplicatedTextContractError::invalid(
"selected realization omitted its authoritative materialization tasks",
));
}
Ok(selected.materialization_tasks.clone())
}
fn build_replicated_text_materialization_tasks(
selected: &SelectedReplicatedTextRealization,
) -> Result<Vec<ReplicatedTextMaterializationTask>, ReplicatedTextContractError> {
build_materialization_tasks(
selected.requirements(),
selected.requirements().parameters(),
selected.parameters(),
)
}
fn build_materialization_tasks(
requirements: &ReplicatedTextRequirements,
parameter_requirements: &[ReplicatedTextParameterRequirement],
selected_parameters: &[SelectedParameterRealization],
) -> Result<Vec<ReplicatedTextMaterializationTask>, ReplicatedTextContractError> {
let mut tasks = selected_parameters
.iter()
.map(|realization| {
let requirement = parameter_requirements
.iter()
.find(|requirement| requirement.name() == realization.name())
.ok_or_else(|| {
ReplicatedTextContractError::invalid(format!(
"selected parameter {:?} has no architecture requirement",
realization.name()
))
})?;
if requirement.sources() != realization.sources()
|| requirement.physical_sources() != realization.physical_sources()
|| requirement.source_encoding() != Some(realization.source_encoding())
{
return Err(ReplicatedTextContractError::invalid(format!(
"selected parameter {:?} changed admitted source provenance",
realization.name()
)));
}
let physical_shape = requirement.physical_shape().ok_or_else(|| {
ReplicatedTextContractError::invalid(format!(
"selected parameter {:?} has no physical geometry",
realization.name()
))
})?;
let lowering_descriptor = requirement.lowering_descriptor(realization.executable())?;
if lowering_descriptor.source() != realization.source_encoding() {
return Err(ReplicatedTextContractError::invalid(format!(
"selected parameter {:?} changed its lowering source encoding",
realization.name()
)));
}
let mut derived_recipe = requirements
.derived_recipes()
.get(realization.name())
.cloned();
let derived_output = requirements
.derived_recipe_outputs()
.get(realization.name())
.cloned();
if derived_recipe.is_some() != derived_output.is_some() {
return Err(ReplicatedTextContractError::invalid(format!(
"selected parameter {:?} has incomplete derived metadata",
realization.name()
)));
}
if derived_recipe.is_none()
&& matches!(
realization.lowering(),
WeightLoweringKind::Derived | WeightLoweringKind::DerivedTransform
)
{
let [source] = realization.sources() else {
return Err(ReplicatedTextContractError::invalid(format!(
"derived selected parameter {:?} has no exact recipe and does not name one source",
realization.name()
)));
};
derived_recipe = Some(
eredu_checkpoint::recipe::DerivedWeightRecipe::source(
source.clone(),
eredu_checkpoint::store::TensorSelection::Full,
),
);
}
Ok(ReplicatedTextMaterializationTask {
name: realization.name().to_owned(),
sources: realization.sources().to_vec(),
physical_sources: realization.physical_sources().to_vec(),
aliases: requirement.aliases().to_vec(),
source_encoding: realization.source_encoding().clone(),
physical_shape: physical_shape.to_vec(),
logical_shape: requirement.logical_shape().to_vec(),
role: requirement.role(),
owner: requirement.owner().clone(),
presence: requirement.presence().clone(),
executable: realization.executable(),
lowering: realization.lowering(),
lowering_descriptor,
derived_recipe,
derived_output,
shared_source_keys: requirements.shared_source_keys().clone(),
permitted_native_source_dtypes: requirement
.permitted_native_source_dtypes()
.to_vec(),
output_companions: Vec::new(),
})
})
.collect::<Result<Vec<_>, _>>()?;
let task_by_name = tasks
.iter()
.map(|task| (task.name().to_owned(), task.clone()))
.collect::<BTreeMap<_, _>>();
let requirement_by_name = parameter_requirements
.iter()
.map(|requirement| (requirement.name(), requirement))
.collect::<BTreeMap<_, _>>();
let mut declared = BTreeMap::<String, Vec<ReplicatedTextOutputCompanion>>::new();
let mut companion_names = BTreeSet::new();
for requirement in parameter_requirements {
let Some((role, primary)) = requirement.linear_companion() else {
continue;
};
let owner = parameter_group_owner(requirement.owner())?;
let mut companion = ReplicatedTextOutputCompanion::new(
requirement.name(),
role,
requirement.logical_shape().to_vec(),
owner,
)?;
if let Some(task) = task_by_name.get(requirement.name()) {
companion = companion.with_materialization_task(task.clone())?;
}
declared
.entry(primary.to_owned())
.or_default()
.push(companion);
companion_names.insert(requirement.name().to_owned());
}
for task in &mut tasks {
let transforms = matches!(
task.lowering(),
WeightLoweringKind::Transform | WeightLoweringKind::DerivedTransform
);
let outputs = if transforms {
let requirement = requirement_by_name.get(task.name()).ok_or_else(|| {
ReplicatedTextContractError::invalid(format!(
"selected task {:?} has no retained requirement",
task.name()
))
})?;
let (scale, affine_bias) = requirement.transform_companions().ok_or_else(|| {
ReplicatedTextContractError::invalid(format!(
"transformed task {:?} has no architecture-selected companion identities",
task.name()
))
})?;
let quantization = task.executable().weight_quantization().ok_or_else(|| {
ReplicatedTextContractError::invalid(format!(
"transformed task {:?} selected a non-quantized executable",
task.name()
))
})?;
let mut shape = task.logical_shape().to_vec();
let input = shape.last_mut().ok_or_else(|| {
ReplicatedTextContractError::invalid("transformed scalar parameter")
})?;
let group = usize::try_from(quantization.group_size()).map_err(|_| {
ReplicatedTextContractError::invalid("transform group size exceeds usize")
})?;
if group == 0 || !input.is_multiple_of(group) {
return Err(ReplicatedTextContractError::invalid(format!(
"transformed task {:?} has incompatible companion geometry",
task.name()
)));
}
*input /= group;
let owner = parameter_group_owner(task.owner())?;
let mut outputs = vec![ReplicatedTextOutputCompanion::new(
scale,
eredu_nn::LinearCompanionRole::Scale,
shape.clone(),
owner.clone(),
)?];
if quantization.has_biases() {
outputs.push(ReplicatedTextOutputCompanion::new(
affine_bias,
eredu_nn::LinearCompanionRole::AffineBias,
shape,
owner,
)?);
}
outputs
} else {
declared.remove(task.name()).unwrap_or_default()
};
task.set_output_companions(outputs)?;
}
if !declared.is_empty() {
return Err(ReplicatedTextContractError::invalid(format!(
"selected companion primaries have no materialization task: {:?}",
declared.keys().collect::<Vec<_>>()
)));
}
tasks.retain(|task| !companion_names.contains(task.name()));
Ok(tasks)
}
fn parameter_group_owner(
owner: &ReplicatedTextParameterOwner,
) -> Result<ParameterGroupOwner, ReplicatedTextContractError> {
match owner {
ReplicatedTextParameterOwner::StaticRole(role) => {
Ok(ParameterGroupOwner::static_role(role.clone()))
}
ReplicatedTextParameterOwner::ExecutionUnit { group, unit } => {
let group = ExecutionGroupId::new(group.clone())
.map_err(|error| ReplicatedTextContractError::invalid(error.to_string()))?;
Ok(ParameterGroupOwner::execution_unit(group, *unit))
}
}
}
pub fn partitioned_replicated_text_materialization_tasks<G, A>(
selected: &SelectedReplicatedTextRealization,
parameters: &ArchitectureParameterDescription,
partition: &ArchitecturePartition<G, A>,
) -> Result<Vec<ReplicatedTextMaterializationTask>, ReplicatedTextContractError> {
let tasks = replicated_text_materialization_tasks(selected)?;
partition_selected_replicated_text_materialization_tasks(&tasks, parameters, partition)
}
pub fn partition_selected_replicated_text_materialization_tasks<G, A>(
tasks: &[ReplicatedTextMaterializationTask],
parameters: &ArchitectureParameterDescription,
partition: &ArchitecturePartition<G, A>,
) -> Result<Vec<ReplicatedTextMaterializationTask>, ReplicatedTextContractError> {
let mut tasks = tasks.to_vec();
let mut companions = BTreeMap::<String, Vec<ReplicatedTextOutputCompanion>>::new();
let mut all_targets = BTreeSet::new();
let mut owned_targets = BTreeSet::new();
for tagged in parameters.groups() {
let local = partition.parameter_bindings().iter().any(|binding| {
binding.owner() == tagged.owner()
&& parameter_groups_have_same_members(binding.group(), tagged.group())
});
let group_targets = tagged
.members()
.iter()
.map(|member| member.target())
.collect::<BTreeSet<_>>();
for member in tagged.members() {
if !all_targets.insert(member.target().to_owned()) {
return Err(ReplicatedTextContractError::invalid(format!(
"architecture parameter target {:?} appears more than once",
member.target()
)));
}
if local {
owned_targets.insert(member.target().to_owned());
}
match (member.linear_companion(), member.linear_companion_of()) {
(None, None) => {}
(Some(role), Some(primary)) if group_targets.contains(primary) && local => {
companions.entry(primary.to_owned()).or_default().push(
ReplicatedTextOutputCompanion::new(
member.target(),
role,
member.global_shape().to_vec(),
tagged.owner().clone(),
)?,
);
}
(Some(_), Some(primary)) if group_targets.contains(primary) => {}
(Some(_), Some(primary)) => {
return Err(ReplicatedTextContractError::invalid(format!(
"physical companion {:?} names primary {primary:?} outside its atomic parameter group",
member.target()
)));
}
_ => {
return Err(ReplicatedTextContractError::invalid(format!(
"physical parameter {:?} has incomplete companion metadata",
member.target()
)));
}
}
}
}
let mut topology_targets = BTreeMap::<String, String>::new();
let mut target_claims = BTreeMap::<String, String>::new();
for task in &tasks {
let matches = std::iter::once(task.name())
.chain(task.aliases().iter().map(String::as_str))
.filter(|candidate| all_targets.contains(*candidate))
.collect::<BTreeSet<_>>();
if matches.len() != 1 {
return Err(ReplicatedTextContractError::invalid(format!(
"selected materialization output {:?} resolves to {} architecture topology targets through its canonical identity and admitted aliases: {:?}",
task.name(),
matches.len(),
matches
)));
}
let target = matches.first().expect("one topology target was validated");
if let Some(previous) = target_claims.insert((*target).to_owned(), task.name().to_owned()) {
return Err(ReplicatedTextContractError::invalid(format!(
"selected materialization outputs {previous:?} and {:?} ambiguously resolve to architecture target {target:?}",
task.name()
)));
}
topology_targets.insert(task.name().to_owned(), (*target).to_owned());
}
for task in &mut tasks {
let topology_target = topology_targets
.get(task.name())
.expect("every task has one validated topology target");
if !owned_targets.contains(topology_target) {
continue;
}
let mut declared = companions.remove(topology_target).unwrap_or_default();
declared.sort_by(|left, right| {
left.role()
.cmp(&right.role())
.then_with(|| left.name().cmp(right.name()))
});
let selected = task.output_companions();
if declared.len() != selected.len()
|| declared.iter().zip(selected).any(|(declared, selected)| {
declared.name() != selected.name()
|| declared.role() != selected.role()
|| declared.logical_shape() != selected.logical_shape()
|| !(declared.owner() == selected.owner()
|| matches!(
(declared.owner(), selected.owner()),
(
ParameterGroupOwner::StaticAnyOf(declared_roles),
ParameterGroupOwner::StaticRole(selected_role)
) if declared_roles.iter().any(|role| role == selected_role)
))
})
{
return Err(ReplicatedTextContractError::invalid(format!(
"partition parameter companions for {:?} differ from authoritative selection: constructed={:?}, selected={:?}",
task.name(), declared, selected,
)));
}
}
if !companions.is_empty() {
return Err(ReplicatedTextContractError::invalid(format!(
"architecture companions name missing primary tasks: {:?}",
companions.keys().collect::<Vec<_>>()
)));
}
let mut projected = Vec::new();
for task in tasks {
let topology_target = topology_targets
.get(task.name())
.expect("every task has one validated topology target");
let emitted = std::iter::once(topology_target.as_str())
.chain(
task.output_companions()
.iter()
.map(ReplicatedTextOutputCompanion::name),
)
.collect::<Vec<_>>();
let local = emitted
.iter()
.filter(|target| owned_targets.contains(**target))
.count();
match local {
0 => {}
count if count == emitted.len() => projected.push(task),
count => {
return Err(ReplicatedTextContractError::invalid(format!(
"materialization task {:?} would emit {count} of {} outputs into this partition",
task.name(),
emitted.len()
)));
}
}
}
Ok(projected)
}
fn parameter_groups_have_same_members(
left: &ParameterGroupSpec,
right: &ParameterGroupSpec,
) -> bool {
left.logical_name() == right.logical_name()
&& left.role() == right.role()
&& left.partition_units() == right.partition_units()
&& left.members().len() == right.members().len()
&& left.members().iter().all(|left_member| {
right.members().iter().any(|right_member| {
left_member.target() == right_member.target()
&& left_member.global_shape() == right_member.global_shape()
&& left_member.sharding() == right_member.sharding()
&& left_member.linear_companion() == right_member.linear_companion()
&& left_member.linear_companion_of() == right_member.linear_companion_of()
})
})
}
impl SelectedParameterRealization {
pub fn name(&self) -> &str {
&self.name
}
pub fn sources(&self) -> &[String] {
&self.sources
}
pub fn physical_sources(&self) -> &[ReplicatedTextPhysicalSource] {
&self.physical_sources
}
pub const fn source_encoding(&self) -> &SourceTensorEncoding {
&self.source_encoding
}
pub const fn executable(&self) -> LinearFormat {
self.executable
}
pub const fn lowering(&self) -> WeightLoweringKind {
self.lowering
}
}
#[derive(Debug, Clone, Eq, PartialEq)]
pub struct SelectedStateComponentRealization {
storage_dtype: StateStorageDtype,
layer: usize,
component: StateComponentPolicy,
placement: StateComponentPlacement,
}
impl SelectedStateComponentRealization {
pub const fn storage_dtype(&self) -> StateStorageDtype {
self.storage_dtype
}
pub const fn layer(&self) -> usize {
self.layer
}
pub const fn component(&self) -> &StateComponentPolicy {
&self.component
}
pub const fn placement(&self) -> StateComponentPlacement {
self.placement
}
}
#[derive(Debug, Clone, Eq, PartialEq)]
pub struct SelectedStateRealization {
floating_dtype: Option<StateStorageDtype>,
layout: StateLayout,
access: ReplicatedTextStateAccess,
policy: CacheResidencyPolicy,
components: Vec<SelectedStateComponentRealization>,
checkpoint: bool,
rollback: bool,
reset: bool,
prompt_cache: bool,
observation_retention: bool,
}
impl SelectedStateRealization {
pub const fn floating_dtype(&self) -> Option<StateStorageDtype> {
self.floating_dtype
}
pub const fn layout(&self) -> &StateLayout {
&self.layout
}
pub fn for_partition(
&self,
partition: &crate::PartitionState,
) -> Result<Self, ReplicatedTextContractError> {
let range = partition.global_layers();
let expected = self
.layout
.slice(range.clone())
.map_err(|error| ReplicatedTextContractError::invalid(error.to_string()))?;
if &expected != partition.layout() {
return Err(ReplicatedTextContractError::invalid(
"partition state layout differs from the selected global interval",
));
}
let components = self
.components
.iter()
.filter(|component| range.contains(&component.layer))
.cloned()
.map(|mut component| {
component.layer -= range.start;
component
})
.collect::<Vec<_>>();
let expected_components = (0..partition.layout().len())
.map(|layer| {
partition
.layout()
.components(layer)
.expect("validated local state layout contains every layer")
.len()
})
.sum::<usize>();
if components.len() != expected_components {
return Err(ReplicatedTextContractError::invalid(
"partition state components differ from the selected global interval",
));
}
Ok(Self {
floating_dtype: self.floating_dtype,
layout: partition.layout().clone(),
access: self.access,
policy: self.policy.clone(),
components,
checkpoint: self.checkpoint,
rollback: self.rollback,
reset: self.reset,
prompt_cache: self.prompt_cache,
observation_retention: self.observation_retention,
})
}
pub fn for_partitioned_geometry(
&self,
partition: &crate::PartitionState,
) -> Result<Self, ReplicatedTextContractError> {
let range = partition.global_layers();
let global = self
.layout
.slice(range.clone())
.map_err(|error| ReplicatedTextContractError::invalid(error.to_string()))?;
if global.len() != partition.layout().len() {
return Err(ReplicatedTextContractError::invalid(
"partition state layer count differs from the selected global interval",
));
}
let mut components = Vec::new();
for local_layer in 0..partition.layout().len() {
let global_components = global
.components(local_layer)
.expect("validated selected state contains every local layer");
let local_components = partition
.layout()
.components(local_layer)
.expect("validated partition state contains every local layer");
if global_components.len() != local_components.len() {
return Err(ReplicatedTextContractError::invalid(
"partition state component count differs from selected state",
));
}
let global_layer = range.start + local_layer;
let selected_components = self
.components
.iter()
.filter(|component| component.layer == global_layer)
.collect::<Vec<_>>();
if selected_components.len() != local_components.len() {
return Err(ReplicatedTextContractError::invalid(
"partition state components differ from the selected global interval",
));
}
for ((global_policy, local_policy), selected) in global_components
.iter()
.zip(local_components)
.zip(selected_components)
{
if global_policy.role() != local_policy.role()
|| global_policy.dtype() != local_policy.dtype()
|| global_policy.residency() != local_policy.residency()
|| global_policy.presence() != local_policy.presence()
|| selected.component != *global_policy
{
return Err(ReplicatedTextContractError::invalid(
"partition state component semantics differ from selected state",
));
}
components.push(SelectedStateComponentRealization {
layer: local_layer,
component: local_policy.clone(),
storage_dtype: selected.storage_dtype,
placement: selected.placement,
});
}
}
Ok(Self {
floating_dtype: self.floating_dtype,
layout: partition.layout().clone(),
access: self.access,
policy: self.policy.clone(),
components,
checkpoint: self.checkpoint,
rollback: self.rollback,
reset: self.reset,
prompt_cache: self.prompt_cache,
observation_retention: self.observation_retention,
})
}
pub const fn access(&self) -> ReplicatedTextStateAccess {
self.access
}
pub const fn policy(&self) -> &CacheResidencyPolicy {
&self.policy
}
pub fn components(&self) -> &[SelectedStateComponentRealization] {
&self.components
}
pub const fn checkpoint(&self) -> bool {
self.checkpoint
}
pub const fn rollback(&self) -> bool {
self.rollback
}
pub const fn reset(&self) -> bool {
self.reset
}
pub const fn prompt_cache(&self) -> bool {
self.prompt_cache
}
pub const fn observation_retention(&self) -> bool {
self.observation_retention
}
}
#[derive(Debug, Clone, Eq, PartialEq)]
pub struct SelectedReplicatedTextRealization {
max_cached_shards: usize,
requirements: ReplicatedTextRequirements,
topology: ParallelTopology,
residency: LayerWeightResidency,
state: SelectedStateRealization,
parameters: Vec<SelectedParameterRealization>,
materialization_tasks: Vec<ReplicatedTextMaterializationTask>,
auxiliary_parameters: Vec<SelectedParameterRealization>,
auxiliary_materialization_tasks: Vec<ReplicatedTextMaterializationTask>,
session: SessionCapabilities,
prompt_cache: bool,
exact_completion: bool,
grouped_operations: Vec<GroupedOperationRequirement>,
}
impl SelectedReplicatedTextRealization {
pub const fn max_cached_shards(&self) -> usize {
self.max_cached_shards
}
pub const fn requirements(&self) -> &ReplicatedTextRequirements {
&self.requirements
}
pub const fn topology(&self) -> ParallelTopology {
self.topology
}
pub const fn residency(&self) -> LayerWeightResidency {
self.residency
}
pub const fn state(&self) -> &SelectedStateRealization {
&self.state
}
pub fn parameters(&self) -> &[SelectedParameterRealization] {
&self.parameters
}
pub fn materialization_tasks(&self) -> &[ReplicatedTextMaterializationTask] {
&self.materialization_tasks
}
pub fn auxiliary_parameters(&self) -> &[SelectedParameterRealization] {
&self.auxiliary_parameters
}
pub fn auxiliary_materialization_tasks(&self) -> &[ReplicatedTextMaterializationTask] {
&self.auxiliary_materialization_tasks
}
pub const fn session(&self) -> SessionCapabilities {
self.session
}
pub const fn prompt_cache(&self) -> bool {
self.prompt_cache
}
pub const fn exact_completion(&self) -> bool {
self.exact_completion
}
pub fn grouped_operations(&self) -> &[GroupedOperationRequirement] {
&self.grouped_operations
}
}
#[derive(Debug, Clone, Eq, PartialEq, thiserror::Error)]
#[error("replicated text realization is unsupported: {issues}", issues = .issues.join("; "))]
pub struct ReplicatedTextSelectionError {
issues: Vec<String>,
}
impl ReplicatedTextSelectionError {
pub fn issues(&self) -> &[String] {
&self.issues
}
}
pub fn select_replicated_text_realization(
requirements: &ReplicatedTextRequirements,
request: &ReplicatedTextSelectionRequest,
capabilities: &BackendMechanismCapabilities,
) -> Result<SelectedReplicatedTextRealization, ReplicatedTextSelectionError> {
let mut issues = Vec::new();
if request
.topology
.is_some_and(|topology| !topology.is_replicated())
{
issues.push("replicated execution topology".into());
}
if !capabilities.operators.contains(requirements.operators) {
issues.extend(
capabilities
.operators
.missing_capability_names(requirements.operators)
.into_iter()
.map(|name| format!("neural operation {name}")),
);
}
for operation in &requirements.grouped_operations {
if !capabilities.grouped_operations.contains(operation) {
issues.push(format!("grouped operation {operation:?}"));
}
}
let residency_mechanism = match request.residency {
LayerWeightResidency::FullyResident => WeightResidencyMechanism::Resident,
LayerWeightResidency::LayerwiseHost(_) => WeightResidencyMechanism::Windowed,
LayerWeightResidency::DenseDiskStream(_) => WeightResidencyMechanism::DiskStreamed,
};
if !capabilities
.weight_residencies
.contains(&residency_mechanism)
{
issues.push(format!("weight residency {residency_mechanism:?}"));
}
let floating_dtype = match capabilities.state.floating_state_dtype() {
Some((source, dtype))
if Some(source) == requirements.floating_state_source() && dtype.is_floating() =>
{
Some(dtype)
}
Some(_) => {
issues.push("floating-state dtype support differs from the selected source".into());
None
}
None => None,
};
let mut state_components = Vec::new();
for layer in 0..requirements.state_layout.len() {
for component in requirements
.state_layout
.components(layer)
.expect("state layout exposes every validated layer")
{
let Some(storage_dtype) = StateStorageDtype::resolve(component.dtype(), floating_dtype)
else {
issues.push(format!(
"state component {} at layer {layer} has no selected floating storage dtype",
component.role().stable_name()
));
continue;
};
let matches = capabilities
.state
.components
.iter()
.filter(|mechanism| mechanism.layer == layer && mechanism.component == *component)
.collect::<Vec<_>>();
let role = component.role().stable_name();
match matches.as_slice() {
[mechanism] => match mechanism.placement(&request.state) {
Some(placement) if placement_is_compatible(component, &request.state, placement) => {
state_components.push(SelectedStateComponentRealization {
layer,
component: component.clone(),
storage_dtype,
placement,
});
}
Some(placement) => issues.push(format!(
"state component {role} at layer {layer} has incompatible {placement:?} placement for {:?} and {:?} residency",
request.state,
component.residency()
)),
None => issues.push(format!(
"state component {role} at layer {layer} for {:?}",
request.state
)),
},
[] => issues.push(format!(
"state component {role} at layer {layer} with shape {:?} and dtype {:?}",
component.shape(),
component.dtype()
)),
_ => issues.push(format!(
"unique state component mechanism {role} at layer {layer}"
)),
}
}
}
for (supported, name) in [
(capabilities.state.checkpoint, "state checkpoint"),
(capabilities.state.rollback, "state rollback"),
(capabilities.state.reset, "state reset"),
] {
if !supported {
issues.push(name.into());
}
}
if request.prompt_cache && !capabilities.state.prompt_cache {
issues.push("state prompt-cache persistence".into());
}
if (request.session.output_observation() || request.session.activation_inspection())
&& !capabilities.state.observation_retention
{
issues.push("state observation retention".into());
}
for (required, supported, name) in [
(
request.session.persistent_cache(),
capabilities.session.persistent_cache(),
"persistent_cache",
),
(
request.session.output_observation(),
capabilities.session.output_observation(),
"output_observation",
),
(
request.session.activation_inspection(),
capabilities.session.activation_inspection(),
"activation_inspection",
),
] {
if required && !supported {
issues.push(format!("session capability {name}"));
}
}
if request.prompt_cache && !capabilities.prompt_cache {
issues.push("prompt-cache persistence".into());
}
if request.exact_completion && !capabilities.exact_completion {
issues.push("exact completion ownership".into());
}
let mut parameters = Vec::with_capacity(requirements.parameters.len());
let mut auxiliary_parameters = Vec::with_capacity(requirements.auxiliary_parameters.len());
let mut names = BTreeSet::new();
for (parameter, auxiliary) in requirements
.parameters
.iter()
.map(|parameter| (parameter, false))
.chain(
requirements
.auxiliary_parameters
.iter()
.map(|parameter| (parameter, true)),
)
{
if parameter.name.trim().is_empty() || !names.insert(parameter.name.as_str()) {
issues.push(format!(
"unique nonempty logical parameter identity {:?}",
parameter.name
));
continue;
}
if !parameter.has_lowering_source() {
continue;
}
let native_candidate = || {
parameter
.lowering_descriptor(parameter.native_executable)
.map(|descriptor| (parameter.native_executable, descriptor))
};
let candidate = match request.quantization {
Some(request) => parameter
.transform_target(request)
.and_then(|target| match target {
Some(target) => Ok((target.executable(), target.descriptor().clone())),
None => native_candidate(),
}),
None => native_candidate(),
};
let (executable, descriptor) = match candidate {
Ok(candidate) => candidate,
Err(error) => {
issues.push(error.to_string());
issues.push(format!(
"architecture transform {:?} for {:?}",
request.quantization, parameter.name
));
continue;
}
};
let Some(lowering) = capabilities
.weight_lowerings
.iter()
.find(|lowering| lowering.descriptor == descriptor)
else {
issues.push(format!(
"weight lowering {:?} -> {:?} for {:?} with descriptor {:?}",
parameter.source_encoding, executable, parameter.name, descriptor
));
continue;
};
let selected_parameter = SelectedParameterRealization {
name: parameter.name.clone(),
sources: parameter.sources.clone(),
physical_sources: parameter.physical_sources.clone(),
source_encoding: parameter
.source_encoding
.clone()
.expect("physical parameter has a source encoding"),
executable,
lowering: match (¶meter.presence, lowering.kind) {
(
ReplicatedTextParameterPresence::Derived { .. },
WeightLoweringKind::Transform | WeightLoweringKind::DerivedTransform,
) => WeightLoweringKind::DerivedTransform,
(ReplicatedTextParameterPresence::Derived { .. }, _) => WeightLoweringKind::Derived,
(_, kind) => kind,
},
};
if auxiliary {
auxiliary_parameters.push(selected_parameter);
} else {
parameters.push(selected_parameter);
}
}
if !issues.is_empty() {
return Err(ReplicatedTextSelectionError { issues });
}
let mut selected = SelectedReplicatedTextRealization {
max_cached_shards: request.max_cached_shards,
requirements: requirements.clone(),
topology: request
.topology
.unwrap_or_else(|| ParallelTopology::new(1, 1, 1, 1).expect("replicated topology")),
residency: request.residency,
state: SelectedStateRealization {
floating_dtype,
layout: requirements.state_layout.clone(),
access: requirements.state_access,
policy: request.state.clone(),
components: state_components,
checkpoint: true,
rollback: true,
reset: true,
prompt_cache: request.prompt_cache,
observation_retention: request.session.output_observation()
|| request.session.activation_inspection(),
},
parameters,
materialization_tasks: Vec::new(),
auxiliary_parameters,
auxiliary_materialization_tasks: Vec::new(),
session: request.session,
prompt_cache: request.prompt_cache,
exact_completion: request.exact_completion,
grouped_operations: requirements.grouped_operations.clone(),
};
selected.materialization_tasks = build_replicated_text_materialization_tasks(&selected)
.map_err(|error| ReplicatedTextSelectionError {
issues: vec![error.to_string()],
})?;
selected.auxiliary_materialization_tasks = build_materialization_tasks(
selected.requirements(),
selected.requirements().auxiliary_parameters(),
&selected.auxiliary_parameters,
)
.map_err(|error| ReplicatedTextSelectionError {
issues: vec![error.to_string()],
})?;
Ok(selected)
}
pub(crate) fn placement_is_compatible(
component: &StateComponentPolicy,
policy: &CacheResidencyPolicy,
placement: StateComponentPlacement,
) -> bool {
use eredu_core::cache::StateResidencyClass;
let expected = match (policy, component.residency()) {
(CacheResidencyPolicy::Device, _) => StateComponentPlacement::Device,
(CacheResidencyPolicy::Paged(_), StateResidencyClass::SealablePaged) => {
StateComponentPlacement::Paged
}
(
CacheResidencyPolicy::Paged(_),
StateResidencyClass::AlwaysDeviceMutable | StateResidencyClass::LayerScopedOffloadable,
) => StateComponentPlacement::Device,
};
placement == expected
}
#[cfg(test)]
mod tests {
use super::*;
use crate::{
ArchitectureGroupKind, ArchitectureGroupPlacement, ArchitectureGroupTransport,
ArchitectureMergeDestination, ArchitectureParameterDescription, ArchitecturePartition,
ArchitectureStatePartitionPlan, ArchitectureStatePartitionRule, DenseDiskStreamLoadOptions,
ExecutionGroupSpec, ExecutionUnitLayout, LayerwiseLoadOptions, MemberSharding,
NoAuxiliaryBoundarySchema, OwnedParameterGroupSpec, ParameterGroupSpec,
ParameterMemberSpec, ParameterRole, PartitionOwnership, StateLayout,
};
use eredu_checkpoint::{AffineQuantization, StoredDtype};
use eredu_core::{
cache::{
LayerCachePolicy, MutableStateResidency, StateTensorDimension, StateTensorDtype,
StateTensorPolicy, StateTensorRole,
},
AttentionPolicy, LayerSchedule,
};
fn paged_state() -> CacheResidencyPolicy {
CacheResidencyPolicy::Paged(
crate::PagedCacheOptions::new(4, 1 << 20, 1 << 20, 1)
.unwrap()
.with_full_attention(true),
)
}
fn physical_source(name: &str) -> ReplicatedTextPhysicalSource {
ReplicatedTextPhysicalSource::new(
name,
name,
"/checkpoint/model.safetensors",
name,
SourceTensorEncoding::Safetensors(StoredDtype::F16),
2,
)
.unwrap()
}
fn requirements() -> ReplicatedTextRequirements {
let graph =
ExecutionGraph::new(vec![ExecutionGroupSpec::root("decoder")], "decoder").unwrap();
let execution_units = ExecutionUnitLayout::new(&graph, [1]).unwrap();
ReplicatedTextRequirements::new(
"test.replicated-text",
NeuralOperatorCapabilities::EXP,
graph,
execution_units,
vec![ArchitectureGroupTransport {
placement: ArchitectureGroupPlacement::Pipeline,
kind: ArchitectureGroupKind::Decoder,
first_owner_static_roles: vec!["embedding".into()],
last_owner_static_roles: vec!["output".into()],
merge_destination: ArchitectureMergeDestination::LastOwner,
parallel_subgroup: None,
request_optional: false,
}],
StateLayout::new(
LayerSchedule::new(
1,
vec![LayerCachePolicy::key_value(AttentionPolicy::Full, 1, 8).unwrap()],
)
.unwrap(),
)
.unwrap(),
ReplicatedTextStateAccess::KeyValue,
vec![
ReplicatedTextParameterRequirement::new(
"model.layers.0.mlp.weight",
vec!["blk.0.ffn.weight".into()],
vec![physical_source("blk.0.ffn.weight")],
Vec::new(),
Some(SourceTensorEncoding::Safetensors(StoredDtype::F16)),
Some(vec![64, 64]),
vec![64, 64],
LinearFormat::Dense,
ReplicatedTextParameterRole::LinearWeight,
ReplicatedTextParameterOwner::ExecutionUnit {
group: "decoder".into(),
unit: 0,
},
ReplicatedTextParameterPresence::Required,
ParameterTransformConstraint::Linear { packed_axis: 1 },
)
.and_then(|requirement| {
requirement.with_transform_companions(
"model.layers.0.mlp.scales",
"model.layers.0.mlp.biases",
)
})
.unwrap(),
ReplicatedTextParameterRequirement::new(
"model.layers.0.mlp.bias",
Vec::new(),
Vec::new(),
Vec::new(),
None,
None,
vec![64],
LinearFormat::Dense,
ReplicatedTextParameterRole::LinearBias,
ReplicatedTextParameterOwner::ExecutionUnit {
group: "decoder".into(),
unit: 0,
},
ReplicatedTextParameterPresence::OptionalAbsent,
ParameterTransformConstraint::None,
)
.unwrap(),
ReplicatedTextParameterRequirement::new(
"model.layers.0.norm.weight",
vec!["blk.0.norm.weight".into()],
vec![physical_source("blk.0.norm.weight")],
Vec::new(),
Some(SourceTensorEncoding::Safetensors(StoredDtype::F16)),
Some(vec![64]),
vec![64],
LinearFormat::Dense,
ReplicatedTextParameterRole::Normalization,
ReplicatedTextParameterOwner::ExecutionUnit {
group: "decoder".into(),
unit: 0,
},
ReplicatedTextParameterPresence::Required,
ParameterTransformConstraint::None,
)
.unwrap(),
],
)
.unwrap()
.with_floating_state_source(TensorDtype::F16)
}
#[test]
fn requirements_reject_unit_layout_from_an_equally_sized_different_graph() {
let baseline = requirements();
let other_graph =
ExecutionGraph::new(vec![ExecutionGroupSpec::root("mutated")], "mutated").unwrap();
let other_layout = ExecutionUnitLayout::new(&other_graph, [1]).unwrap();
let error = ReplicatedTextRequirements::new(
baseline.architecture_identity.clone(),
baseline.operators,
baseline.execution_graph.clone(),
other_layout,
baseline.group_transports.clone(),
baseline.state_layout.clone(),
baseline.state_access,
baseline.parameters.clone(),
)
.unwrap_err();
assert!(error.to_string().contains("layout group identities differ"));
}
#[test]
fn parameter_requirement_preserves_every_admitted_alias() {
let requirement = ReplicatedTextParameterRequirement::new(
"model.layers.0.mlp.weight",
vec!["released.layers.0.mlp.weight".into()],
vec![physical_source("released.layers.0.mlp.weight")],
vec![
"legacy.layers.0.mlp.weight".into(),
"vendor.layers.0.mlp.weight".into(),
],
Some(SourceTensorEncoding::Safetensors(StoredDtype::F16)),
Some(vec![64, 64]),
vec![64, 64],
LinearFormat::Dense,
ReplicatedTextParameterRole::LinearWeight,
ReplicatedTextParameterOwner::ExecutionUnit {
group: "decoder".into(),
unit: 0,
},
ReplicatedTextParameterPresence::Required,
ParameterTransformConstraint::Linear { packed_axis: 1 },
)
.unwrap();
assert_eq!(
requirement.aliases(),
["legacy.layers.0.mlp.weight", "vendor.layers.0.mlp.weight"]
);
assert_eq!(requirement.sources(), ["released.layers.0.mlp.weight"]);
let absent_bias = ReplicatedTextParameterRequirement::new(
"model.layers.0.mlp.bias",
Vec::new(),
Vec::new(),
vec!["released.layers.0.mlp.bias".into()],
None,
None,
vec![64],
LinearFormat::Dense,
ReplicatedTextParameterRole::LinearBias,
ReplicatedTextParameterOwner::ExecutionUnit {
group: "decoder".into(),
unit: 0,
},
ReplicatedTextParameterPresence::OptionalAbsent,
ParameterTransformConstraint::None,
)
.unwrap();
assert_eq!(
absent_bias.presence(),
&ReplicatedTextParameterPresence::OptionalAbsent
);
assert!(absent_bias.sources().is_empty());
assert_eq!(
absent_bias.transform_constraint(),
ParameterTransformConstraint::None
);
}
#[test]
fn scalar_parameter_requirement_preserves_rank_zero_geometry() {
let requirement = ReplicatedTextParameterRequirement::new(
"model.audio_tower.input_max",
vec!["model.audio_tower.input_max".into()],
vec![physical_source("model.audio_tower.input_max")],
Vec::new(),
Some(SourceTensorEncoding::Safetensors(StoredDtype::F32)),
Some(Vec::new()),
Vec::new(),
LinearFormat::Dense,
ReplicatedTextParameterRole::Other,
ReplicatedTextParameterOwner::StaticRole("audio".into()),
ReplicatedTextParameterPresence::Required,
ParameterTransformConstraint::None,
)
.unwrap();
let descriptor = requirement
.lowering_descriptor(LinearFormat::Dense)
.unwrap();
assert!(descriptor.physical_shape().is_empty());
assert!(descriptor.logical_shape().is_empty());
assert_eq!(descriptor.packed_axis(), None);
}
#[test]
fn physical_provenance_distinguishes_outputs_from_one_sharded_tensor() {
let shard = "/checkpoint/model-00002-of-00003.gguf";
let weight = ReplicatedTextPhysicalSource::new(
"model.layers.0.gate_proj.weight",
"blk.0.ffn_gate.weight",
shard,
"blk.0.ffn_gate.weight",
SourceTensorEncoding::Safetensors(StoredDtype::F16),
2,
)
.unwrap();
let scales = ReplicatedTextPhysicalSource::new(
"model.layers.0.gate_proj.scales",
"blk.0.ffn_gate.weight",
shard,
"blk.0.ffn_gate.scales",
SourceTensorEncoding::Safetensors(StoredDtype::F16),
2,
)
.unwrap();
assert_eq!(weight.tensor(), scales.tensor());
assert_eq!(weight.shard(), scales.shard());
assert_ne!(weight.output(), scales.output());
}
fn capabilities() -> BackendMechanismCapabilities {
let source = SourceTensorEncoding::Safetensors(StoredDtype::F16);
let requirements = requirements();
let state = StateMechanismCapabilities::new(
(0..requirements.state_layout().len()).flat_map(|layer| {
requirements
.state_layout()
.components(layer)
.unwrap()
.iter()
.cloned()
.map(move |component| {
let paged = match component.role() {
eredu_core::cache::StateComponentRole::AttentionKeys
| eredu_core::cache::StateComponentRole::AttentionValues
| eredu_core::cache::StateComponentRole::CompressedLatent
| eredu_core::cache::StateComponentRole::RotaryKeys => {
StateComponentPlacement::Paged
}
eredu_core::cache::StateComponentRole::Fixed(_) => {
StateComponentPlacement::Device
}
};
StateComponentMechanism::new(
layer,
component,
Some(StateComponentPlacement::Device),
Some(paged),
)
})
}),
)
.with_floating_state_dtype(TensorDtype::F16, StateStorageDtype::F16)
.with_transactions(true, true)
.with_reset(true)
.with_prompt_cache(true)
.with_observation_retention(true);
BackendMechanismCapabilities::new(
NeuralOperatorCapabilities::EXP,
vec![
WeightLoweringCapability::new(
WeightLoweringDescriptor::new(
source.clone(),
LinearFormat::Dense,
vec![64, 64],
vec![64, 64],
Some(1),
)
.unwrap(),
WeightLoweringKind::Direct,
),
WeightLoweringCapability::new(
WeightLoweringDescriptor::new(
source,
LinearFormat::Affine(AffineQuantization::new(64, 4).unwrap()),
vec![64, 64],
vec![64, 64],
Some(1),
)
.unwrap(),
WeightLoweringKind::Transform,
),
WeightLoweringCapability::new(
WeightLoweringDescriptor::new(
SourceTensorEncoding::Safetensors(StoredDtype::F16),
LinearFormat::Dense,
vec![64],
vec![64],
None,
)
.unwrap(),
WeightLoweringKind::Direct,
),
],
vec![
WeightResidencyMechanism::Resident,
WeightResidencyMechanism::Windowed,
WeightResidencyMechanism::DiskStreamed,
],
state,
)
.with_session(SessionCapabilities::new(true, true, true))
.with_prompt_cache(true)
.with_exact_completion(true)
}
fn request(residency: LayerWeightResidency) -> ReplicatedTextSelectionRequest {
ReplicatedTextSelectionRequest::new(residency, paged_state())
.with_session(SessionCapabilities::new(true, true, true))
.with_prompt_cache(true)
.with_exact_completion(true)
}
#[test]
fn complete_requirements_are_invariant_across_all_caller_policy_dimensions() {
let baseline = requirements();
let disk = DenseDiskStreamLoadOptions::new(4096, 8192, 2, 1).unwrap();
let requests = [
ReplicatedTextSelectionRequest::new(
LayerWeightResidency::FullyResident,
CacheResidencyPolicy::Device,
),
ReplicatedTextSelectionRequest::new(
LayerWeightResidency::LayerwiseHost(LayerwiseLoadOptions::default()),
paged_state(),
)
.with_topology(ParallelTopology::new(2, 1, 1, 1).unwrap())
.with_quantization(QuantizationRequest::Affine {
group_size: 64,
bits: 4,
})
.with_session(SessionCapabilities::new(true, true, true))
.with_prompt_cache(true)
.with_exact_completion(true),
ReplicatedTextSelectionRequest::new(
LayerWeightResidency::DenseDiskStream(disk),
CacheResidencyPolicy::Device,
)
.with_quantization(QuantizationRequest::MxFp4),
];
for _request in &requests {
assert_eq!(requirements(), baseline);
}
assert_eq!(requests[0].state(), &CacheResidencyPolicy::Device);
assert!(matches!(
requests[1].residency(),
LayerWeightResidency::LayerwiseHost(_)
));
assert_eq!(requests[1].topology().unwrap().tensor(), 2);
assert_eq!(
requests[1].quantization(),
Some(QuantizationRequest::Affine {
group_size: 64,
bits: 4,
})
);
assert!(requests[1].prompt_cache());
assert!(requests[1].exact_completion());
assert!(requests[1].session().activation_inspection());
assert_eq!(
requests[2].residency(),
LayerWeightResidency::DenseDiskStream(disk)
);
assert_eq!(requests[2].quantization(), Some(QuantizationRequest::MxFp4));
}
#[test]
fn partitioned_tasks_keep_encoded_companions_atomic_and_reject_split_groups() {
let request = request(LayerWeightResidency::FullyResident).with_quantization(
QuantizationRequest::Affine {
group_size: 64,
bits: 4,
},
);
let selected =
select_replicated_text_realization(&requirements(), &request, &capabilities()).unwrap();
let graph = selected.requirements().execution_graph().clone();
let layout = selected.requirements().execution_units().clone();
let format = eredu_nn::LinearFormatSpec::affine(
LinearFormat::Affine(AffineQuantization::new(64, 4).unwrap()),
eredu_nn::ParameterSpec::trainable("model.layers.0.mlp.scales").unwrap(),
eredu_nn::ParameterSpec::trainable("model.layers.0.mlp.biases").unwrap(),
)
.unwrap();
let [physical] = crate::expand_linear_format_parameter_groups(
vec![ParameterGroupSpec::new(
"mlp",
ParameterRole::FeedForwardIntermediate,
[ParameterMemberSpec::new(
"model.layers.0.mlp.weight",
vec![64, 64],
MemberSharding::Replicated,
)],
)
.unwrap()],
|_| Ok(Some(format.clone())),
)
.unwrap()
.try_into()
.unwrap();
let norm = ParameterGroupSpec::new(
"norm",
ParameterRole::Replicated,
[ParameterMemberSpec::new(
"model.layers.0.norm.weight",
vec![64],
MemberSharding::Replicated,
)],
)
.unwrap();
let owner = ParameterGroupOwner::execution_unit(layout.group_id(0).unwrap().clone(), 0);
let description = ArchitectureParameterDescription::new(
&graph,
&layout,
[physical.clone(), norm.clone()],
[
OwnedParameterGroupSpec::new(owner.clone(), physical.clone()),
OwnedParameterGroupSpec::new(owner.clone(), norm.clone()),
],
)
.unwrap();
let ownership =
PartitionOwnership::new(false, false, std::iter::empty::<String>()).unwrap();
let state = selected.requirements().state_layout();
let state_plan =
ArchitectureStatePartitionPlan::new([ArchitectureStatePartitionRule::group_units(
0,
0..state.len(),
)]);
let partition = ArchitecturePartition::from_description(
&description,
[(layout.group_id(0).unwrap().as_str(), 0..1)],
ownership.clone(),
state,
&state_plan,
(),
NoAuxiliaryBoundarySchema::new(64),
)
.unwrap();
let tasks =
partitioned_replicated_text_materialization_tasks(&selected, &description, &partition)
.unwrap();
let task = tasks
.iter()
.find(|task| task.name() == "model.layers.0.mlp.weight")
.unwrap();
assert_eq!(task.output_companions().len(), 2);
let members = physical.members();
let primary = ParameterGroupSpec::new(
"primary",
ParameterRole::FeedForwardIntermediate,
[members[0].clone()],
)
.unwrap();
let companions = ParameterGroupSpec::new(
"companions",
ParameterRole::FeedForwardIntermediate,
members[1..].to_vec(),
)
.unwrap();
let malformed = ArchitectureParameterDescription::new(
&graph,
&layout,
[primary.clone(), companions.clone(), norm.clone()],
[
OwnedParameterGroupSpec::new(owner.clone(), primary),
OwnedParameterGroupSpec::new(owner.clone(), companions),
OwnedParameterGroupSpec::new(owner, norm),
],
)
.unwrap();
let malformed_partition = ArchitecturePartition::from_description(
&malformed,
[(layout.group_id(0).unwrap().as_str(), 0..1)],
ownership,
state,
&state_plan,
(),
NoAuxiliaryBoundarySchema::new(64),
)
.unwrap();
let error = partitioned_replicated_text_materialization_tasks(
&selected,
&malformed,
&malformed_partition,
)
.unwrap_err();
assert!(error
.to_string()
.contains("outside its atomic parameter group"));
}
#[test]
fn partitioned_tasks_require_one_canonical_or_admitted_alias_topology_target() {
let mut requirements = requirements();
requirements.parameters[0].aliases = vec!["architecture.mlp.weight".into()];
let selected = select_replicated_text_realization(
&requirements,
&request(LayerWeightResidency::FullyResident),
&capabilities(),
)
.unwrap();
let graph = selected.requirements().execution_graph().clone();
let layout = selected.requirements().execution_units().clone();
let owner = ParameterGroupOwner::execution_unit(layout.group_id(0).unwrap().clone(), 0);
let ownership =
PartitionOwnership::new(false, false, std::iter::empty::<String>()).unwrap();
let state = selected.requirements().state_layout();
let state_plan =
ArchitectureStatePartitionPlan::new([ArchitectureStatePartitionRule::group_units(
0,
0..state.len(),
)]);
let project = |primary_targets: &[&str]| {
let mut groups = primary_targets
.iter()
.enumerate()
.map(|(index, target)| {
ParameterGroupSpec::new(
format!("mlp-{index}"),
ParameterRole::FeedForwardIntermediate,
[ParameterMemberSpec::new(
*target,
vec![64, 64],
MemberSharding::Replicated,
)],
)
.unwrap()
})
.collect::<Vec<_>>();
groups.push(
ParameterGroupSpec::new(
"norm",
ParameterRole::Replicated,
[ParameterMemberSpec::new(
"model.layers.0.norm.weight",
vec![64],
MemberSharding::Replicated,
)],
)
.unwrap(),
);
let description = ArchitectureParameterDescription::new(
&graph,
&layout,
groups.clone(),
groups
.into_iter()
.map(|group| OwnedParameterGroupSpec::new(owner.clone(), group)),
)
.unwrap();
let partition = ArchitecturePartition::from_description(
&description,
[(layout.group_id(0).unwrap().as_str(), 0..1)],
ownership.clone(),
state,
&state_plan,
(),
NoAuxiliaryBoundarySchema::new(64),
)
.unwrap();
partitioned_replicated_text_materialization_tasks(&selected, &description, &partition)
};
let canonical = project(&["model.layers.0.mlp.weight"]).unwrap();
assert!(canonical
.iter()
.any(|task| task.name() == "model.layers.0.mlp.weight"));
let aliased = project(&["architecture.mlp.weight"]).unwrap();
let task = aliased
.iter()
.find(|task| task.name() == "model.layers.0.mlp.weight")
.unwrap();
assert_eq!(task.aliases(), ["architecture.mlp.weight"]);
let error = project(&["model.layers.0.mlp.weight", "architecture.mlp.weight"]).unwrap_err();
assert!(error
.to_string()
.contains("resolves to 2 architecture topology targets"));
}
#[test]
fn selection_is_deterministic_and_keeps_source_format_distinct() {
let disk = DenseDiskStreamLoadOptions::new(1234, 5678, 3, 2).unwrap();
let request = request(LayerWeightResidency::DenseDiskStream(disk)).with_quantization(
QuantizationRequest::Affine {
group_size: 64,
bits: 4,
},
);
let left =
select_replicated_text_realization(&requirements(), &request, &capabilities()).unwrap();
let right =
select_replicated_text_realization(&requirements(), &request, &capabilities()).unwrap();
assert_eq!(left, right);
assert_eq!(
left.residency(),
LayerWeightResidency::DenseDiskStream(disk)
);
assert_eq!(left.state().policy(), &paged_state());
assert_eq!(left.state().layout(), requirements().state_layout());
assert_eq!(left.parameters().len(), 2);
assert_eq!(requirements().parameters().len(), 3);
assert!(matches!(
requirements().parameters()[1].presence(),
ReplicatedTextParameterPresence::OptionalAbsent
));
assert!(matches!(
requirements().parameters()[2].role(),
ReplicatedTextParameterRole::Normalization
));
assert_eq!(requirements().parameters()[2].logical_shape(), [64]);
assert_eq!(
requirements().parameters()[2].transform_constraint(),
ParameterTransformConstraint::None
);
assert_eq!(
left.parameters()[0].lowering(),
WeightLoweringKind::Transform
);
assert_ne!(
format!("{:?}", left.parameters()[0].source_encoding()),
format!("{:?}", left.parameters()[0].executable())
);
}
#[test]
fn exact_tasks_are_the_authority_for_direct_derived_and_transform_sources() {
use eredu_checkpoint::recipe::{DerivedWeightRecipe, RecipeDtype, RecipeMetadata};
let direct = select_replicated_text_realization(
&requirements(),
&request(LayerWeightResidency::FullyResident),
&capabilities(),
)
.unwrap();
let direct_tasks = replicated_text_materialization_tasks(&direct).unwrap();
assert_eq!(
direct_tasks[0].source_recipe().unwrap(),
DerivedWeightRecipe::source(
"blk.0.ffn.weight",
eredu_checkpoint::store::TensorSelection::Full,
)
);
let recipe = DerivedWeightRecipe::source(
"blk.0.ffn.weight",
eredu_checkpoint::store::TensorSelection::Full,
);
let outputs = BTreeMap::from([(
"model.layers.0.mlp.weight".into(),
RecipeMetadata {
shape: vec![64, 64],
dtype: RecipeDtype::F16,
byte_len: 64 * 64 * 2,
},
)]);
let derived_requirements = requirements()
.with_derived_recipes(
BTreeMap::from([("model.layers.0.mlp.weight".into(), recipe.clone())]),
outputs,
)
.unwrap();
let derived = select_replicated_text_realization(
&derived_requirements,
&request(LayerWeightResidency::FullyResident),
&capabilities(),
)
.unwrap();
let derived_tasks = replicated_text_materialization_tasks(&derived).unwrap();
assert_eq!(derived_tasks[0].lowering(), WeightLoweringKind::Derived);
assert_eq!(derived_tasks[0].source_recipe().unwrap(), recipe);
let transformed = select_replicated_text_realization(
&derived_requirements,
&request(LayerWeightResidency::FullyResident).with_quantization(
QuantizationRequest::Affine {
group_size: 64,
bits: 4,
},
),
&capabilities(),
)
.unwrap();
let transformed_tasks = replicated_text_materialization_tasks(&transformed).unwrap();
assert_eq!(
transformed_tasks[0].lowering(),
WeightLoweringKind::DerivedTransform
);
assert_eq!(transformed_tasks[0].source_recipe().unwrap(), recipe);
let mut corrupt_direct = direct_tasks[0].clone();
corrupt_direct.sources.push("unselected.weight".into());
assert!(corrupt_direct.source_recipe().is_err());
let mut corrupt_kind = derived_tasks[0].clone();
corrupt_kind.lowering = WeightLoweringKind::Direct;
assert!(corrupt_kind.source_recipe().is_err());
let mut corrupt_recipe = derived_tasks[0].clone();
corrupt_recipe.derived_recipe = Some(DerivedWeightRecipe::source(
"unselected.weight",
eredu_checkpoint::store::TensorSelection::Full,
));
assert!(corrupt_recipe.source_recipe().is_err());
let member_output = RecipeMetadata {
shape: vec![64, 64],
dtype: RecipeDtype::F16,
byte_len: 64 * 64 * 2,
};
let member_recipe = direct_tasks[0].source_recipe().unwrap();
let project = |task, selected_bytes| {
crate::AddressableBankParameter::new(
"weight",
task,
member_recipe.clone(),
member_output.clone(),
selected_bytes,
None,
)
};
assert!(project(direct_tasks[0].clone(), member_output.byte_len()).is_ok());
assert!(matches!(
project(direct_tasks[0].clone(), member_output.byte_len() - 1),
Err(crate::AddressableBankMemberError::SelectedByteMismatch { .. })
));
let mut corrupt_source = direct_tasks[0].clone();
corrupt_source.source_encoding = SourceTensorEncoding::Safetensors(StoredDtype::F32);
assert!(project(corrupt_source, member_output.byte_len()).is_err());
let mut corrupt_executable = direct_tasks[0].clone();
corrupt_executable.executable = LinearFormat::MxFp4;
assert!(project(corrupt_executable, member_output.byte_len()).is_err());
let mut corrupt_lowering = direct_tasks[0].clone();
corrupt_lowering.lowering = WeightLoweringKind::Derived;
assert!(project(corrupt_lowering, member_output.byte_len()).is_err());
let mut exact_transform = transformed_tasks[0].clone();
let companion_owner = crate::ParameterGroupOwner::ExecutionUnit {
group: crate::ExecutionGroupId::new("decoder").unwrap(),
global_unit: 0,
};
exact_transform
.set_output_companions(vec![
ReplicatedTextOutputCompanion::new(
"model.layers.0.mlp.weight_scales",
eredu_nn::LinearCompanionRole::Scale,
vec![64, 1],
companion_owner.clone(),
)
.unwrap(),
ReplicatedTextOutputCompanion::new(
"model.layers.0.mlp.weight_biases",
eredu_nn::LinearCompanionRole::AffineBias,
vec![64, 1],
companion_owner,
)
.unwrap(),
])
.unwrap();
let transformed_bytes =
crate::selected_addressable_parameter_bytes(&exact_transform, &member_output).unwrap();
assert!(crate::AddressableBankParameter::new(
"weight",
exact_transform.clone(),
recipe.clone(),
member_output.clone(),
transformed_bytes,
Some(
crate::QuantizationCompanionBindings::new(
"weight_scales",
Some("weight_biases".into()),
)
.unwrap(),
),
)
.is_ok());
assert!(crate::AddressableBankParameter::new(
"weight",
exact_transform,
recipe.clone(),
member_output.clone(),
transformed_bytes,
Some(
crate::QuantizationCompanionBindings::new(
"drifted_scales",
Some("weight_biases".into()),
)
.unwrap(),
),
)
.is_err());
let mut scale_only = transformed_tasks[0].clone();
scale_only.executable = LinearFormat::MxFp4;
scale_only.lowering_descriptor = WeightLoweringDescriptor::new(
scale_only.source_encoding.clone(),
LinearFormat::MxFp4,
scale_only.physical_shape.clone(),
scale_only.logical_shape.clone(),
scale_only.logical_shape.len().checked_sub(1),
)
.unwrap();
scale_only
.set_output_companions(vec![ReplicatedTextOutputCompanion::new(
"model.layers.0.mlp.weight_scales",
eredu_nn::LinearCompanionRole::Scale,
vec![64, 2],
crate::ParameterGroupOwner::ExecutionUnit {
group: crate::ExecutionGroupId::new("decoder").unwrap(),
global_unit: 0,
},
)
.unwrap()])
.unwrap();
let scale_only_bytes =
crate::selected_addressable_parameter_bytes(&scale_only, &member_output).unwrap();
let scale_companions =
crate::QuantizationCompanionBindings::new("weight_scales", None).unwrap();
assert!(crate::AddressableBankParameter::new(
"weight",
scale_only.clone(),
recipe.clone(),
member_output.clone(),
scale_only_bytes,
Some(scale_companions),
)
.is_ok());
let invented_bias = crate::QuantizationCompanionBindings::new(
"weight_scales",
Some("invented_bias".into()),
)
.unwrap();
assert!(crate::AddressableBankParameter::new(
"weight",
scale_only,
recipe,
member_output,
scale_only_bytes,
Some(invented_bias),
)
.is_err());
}
#[test]
fn selection_reports_all_missing_mechanisms_together() {
let capabilities = BackendMechanismCapabilities::new(
NeuralOperatorCapabilities::NONE,
Vec::new(),
Vec::new(),
StateMechanismCapabilities::new(Vec::new()),
);
let error = select_replicated_text_realization(
&requirements(),
&request(LayerWeightResidency::LayerwiseHost(
LayerwiseLoadOptions::default(),
)),
&capabilities,
)
.unwrap_err();
assert!(error.issues().len() >= 7, "{:?}", error.issues());
assert!(error.issues().iter().any(|issue| issue.contains("exp")));
assert!(error
.issues()
.iter()
.any(|issue| issue.contains("weight lowering")));
}
#[test]
fn selection_rejects_paged_fixed_component_placement_even_when_reported() {
let fixed = StateTensorPolicy::new(
StateTensorRole::Recurrent,
vec![
StateTensorDimension::Batch,
StateTensorDimension::fixed(8).unwrap(),
],
StateTensorDtype::Float32,
MutableStateResidency::LayerScopedOffloadable,
)
.unwrap();
let mut requirements = requirements();
requirements.state_layout = StateLayout::new(
LayerSchedule::new(
1,
vec![LayerCachePolicy::key_value_with_fixed_state(
AttentionPolicy::Full,
1,
8,
vec![fixed],
)
.unwrap()],
)
.unwrap(),
)
.unwrap();
requirements.state_access = ReplicatedTextStateAccess::AttentionWithFixed;
let mut capabilities = capabilities();
capabilities.state.components = (0..requirements.state_layout.len())
.flat_map(|layer| {
requirements
.state_layout
.components(layer)
.unwrap()
.iter()
.cloned()
.map(move |component| {
StateComponentMechanism::new(
layer,
component,
Some(StateComponentPlacement::Device),
Some(StateComponentPlacement::Paged),
)
})
})
.collect();
let error = select_replicated_text_realization(
&requirements,
&request(LayerWeightResidency::FullyResident),
&capabilities,
)
.unwrap_err();
assert!(error
.issues()
.iter()
.any(|issue| issue.contains("incompatible Paged placement")));
}
#[test]
fn requirements_reject_state_layout_and_access_profile_mismatch() {
let fixed = StateTensorPolicy::new(
StateTensorRole::Recurrent,
vec![
StateTensorDimension::Batch,
StateTensorDimension::fixed(8).unwrap(),
],
StateTensorDtype::Float32,
MutableStateResidency::LayerScopedOffloadable,
)
.unwrap();
let layout = StateLayout::new(
LayerSchedule::new(
1,
vec![LayerCachePolicy::key_value_with_fixed_state(
AttentionPolicy::Full,
1,
8,
vec![fixed],
)
.unwrap()],
)
.unwrap(),
)
.unwrap();
let base = requirements();
let error = ReplicatedTextRequirements::new(
base.architecture_identity,
base.operators,
base.execution_graph,
base.execution_units,
base.group_transports,
layout,
ReplicatedTextStateAccess::KeyValue,
base.parameters,
)
.unwrap_err();
assert!(error.message().contains("does not match component roles"));
}
#[test]
fn transform_selection_rejects_incompatible_exact_geometry() {
for quantization in [
QuantizationRequest::Affine {
group_size: 96,
bits: 4,
},
QuantizationRequest::Affine {
group_size: 256,
bits: 4,
},
QuantizationRequest::Affine {
group_size: 0,
bits: 4,
},
QuantizationRequest::Affine {
group_size: u32::MAX,
bits: 4,
},
QuantizationRequest::Affine {
group_size: 32,
bits: 0,
},
QuantizationRequest::Affine {
group_size: 32,
bits: 7,
},
] {
let error = select_replicated_text_realization(
&requirements(),
&request(LayerWeightResidency::FullyResident).with_quantization(quantization),
&capabilities(),
)
.unwrap_err();
assert!(error
.issues()
.iter()
.any(|issue| issue.contains("invalid replicated text contract")));
}
let mut indivisible = requirements();
indivisible.parameters[0].logical_shape = vec![64, 48];
let error = select_replicated_text_realization(
&indivisible,
&request(LayerWeightResidency::FullyResident)
.with_quantization(QuantizationRequest::MxFp4),
&capabilities(),
)
.unwrap_err();
assert!(error
.issues()
.iter()
.any(|issue| issue.contains("MXFP4 packed extent 48")));
}
#[test]
fn exact_source_and_physical_geometry_fail_before_construction_or_payload() {
for mutate in [
|requirement: &mut ReplicatedTextParameterRequirement| {
requirement.source_encoding =
Some(SourceTensorEncoding::Safetensors(StoredDtype::U8));
},
|requirement: &mut ReplicatedTextParameterRequirement| {
requirement.physical_shape = Some(vec![64, 32]);
},
] {
let mut requirements = requirements();
mutate(&mut requirements.parameters[0]);
let selected = select_replicated_text_realization(
&requirements,
&request(LayerWeightResidency::FullyResident),
&capabilities(),
);
let error = selected.unwrap_err();
assert!(error
.issues()
.iter()
.any(|issue| issue.contains("weight lowering")));
}
}
#[test]
fn missing_tensor_parallel_grouped_partial_fails_before_construction_or_forward() {
let requirements = requirements().with_grouped_operations([
GroupedOperationRequirement::GatedProduct,
GroupedOperationRequirement::GatedProductTensorParallelPartial,
]);
let capabilities =
capabilities().with_grouped_operations([GroupedOperationRequirement::GatedProduct]);
let selected = select_replicated_text_realization(
&requirements,
&request(LayerWeightResidency::FullyResident),
&capabilities,
);
let error = selected.unwrap_err();
assert!(error
.issues()
.iter()
.any(|issue| { issue.contains("GatedProductTensorParallelPartial") }));
}
}