use std::{
collections::BTreeSet,
path::{Path, PathBuf},
};
use eredu_checkpoint::{LinearFormat, SourceTensorEncoding};
use eredu_core::{ParallelTopology, QuantizationRequest, SessionCapabilities};
use eredu_nn::{NeuralBackend, NeuralOperatorCapabilities};
use crate::{
ArchitectureGroupTransport, CacheResidencyPolicy, ExecutionGraph, ExecutionUnitLayout,
LayerWeightResidency, LayeredArchitecture, 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>;
}
#[derive(Debug, Clone, Copy, Eq, PartialEq)]
#[non_exhaustive]
pub enum WeightLoweringKind {
Direct,
Transform,
}
#[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.is_empty()
|| physical_shape.contains(&0)
|| logical_shape.is_empty()
|| logical_shape.contains(&0)
|| physical_shape.len() != logical_shape.len()
{
return Err(ReplicatedTextContractError::invalid(
"weight lowering requires positive physical and logical shapes of equal rank",
));
}
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, PartialEq)]
#[non_exhaustive]
pub enum StateResidencyMechanism {
Device,
Paged,
}
#[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, Hash, Ord, PartialEq, PartialOrd)]
pub struct ReplicatedTextPhysicalSource {
tensor: String,
shard: PathBuf,
output: String,
}
impl ReplicatedTextPhysicalSource {
pub fn new(
tensor: impl Into<String>,
shard: impl Into<PathBuf>,
output: impl Into<String>,
) -> Result<Self, ReplicatedTextContractError> {
let tensor = tensor.into();
let shard = shard.into();
let output = output.into();
if tensor.trim().is_empty() || shard.as_os_str().is_empty() || output.trim().is_empty() {
return Err(ReplicatedTextContractError::invalid(
"physical source tensor, shard, and output must be non-empty",
));
}
Ok(Self {
tensor,
shard,
output,
})
}
pub fn tensor(&self) -> &str {
&self.tensor
}
pub fn shard(&self) -> &Path {
&self.shard
}
pub fn output(&self) -> &str {
&self.output
}
}
#[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,
}
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();
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 = presence.has_physical_source();
if has_source
!= (!sources.is_empty() && source_encoding.is_some() && physical_shape.is_some())
{
return Err(ReplicatedTextContractError::invalid(format!(
"logical parameter {name:?} has inconsistent source presence"
)));
}
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 has_source
&& physical_sources
.iter()
.any(|source| !sources.iter().any(|name| name == source.tensor()))
{
return Err(ReplicatedTextContractError::invalid(format!(
"logical parameter {name:?} has provenance outside its selected sources"
)));
}
if physical_shape
.as_ref()
.is_some_and(|shape| shape.is_empty() || shape.contains(&0))
{
return Err(ReplicatedTextContractError::invalid(format!(
"logical parameter {name:?} has an invalid physical shape"
)));
}
if logical_shape.is_empty() || 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:?}"
)));
}
}
Ok(Self {
name,
sources,
physical_sources,
aliases,
source_encoding,
physical_shape,
logical_shape,
role,
owner,
presence,
native_executable,
transform,
})
}
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 const fn transform_constraint(&self) -> ParameterTransformConstraint {
self.transform
}
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() - 1)
});
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
))
})?,
self.logical_shape.clone(),
packed_axis,
)
}
}
#[derive(Debug, Clone, Eq, PartialEq, thiserror::Error)]
#[error("invalid replicated text contract: {message}")]
pub struct ReplicatedTextContractError {
message: String,
}
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 {
operators: NeuralOperatorCapabilities,
execution_graph: ExecutionGraph,
execution_units: ExecutionUnitLayout,
group_transports: Vec<ArchitectureGroupTransport>,
state_layout: StateLayout,
parameters: Vec<ReplicatedTextParameterRequirement>,
grouped_operations: Vec<GroupedOperationRequirement>,
}
impl ReplicatedTextRequirements {
pub fn new(
operators: NeuralOperatorCapabilities,
execution_graph: ExecutionGraph,
execution_units: ExecutionUnitLayout,
group_transports: Vec<ArchitectureGroupTransport>,
state_layout: StateLayout,
parameters: Vec<ReplicatedTextParameterRequirement>,
) -> Result<Self, ReplicatedTextContractError> {
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()
)));
}
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 {
operators,
execution_graph,
execution_units,
group_transports,
state_layout,
parameters,
grouped_operations: Vec::new(),
})
}
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 fn parameters(&self) -> &[ReplicatedTextParameterRequirement] {
&self.parameters
}
pub fn grouped_operations(&self) -> &[GroupedOperationRequirement] {
&self.grouped_operations
}
}
#[derive(Debug, Clone, Copy, Eq, Hash, Ord, PartialEq, PartialOrd)]
#[non_exhaustive]
pub enum GroupedOperationRequirement {
GatedProduct,
GatedProductTensorParallelPartial,
Relu2,
Relu2TensorParallelPartial,
}
#[derive(Debug, Clone, Eq, PartialEq)]
pub struct BackendMechanismCapabilities {
operators: NeuralOperatorCapabilities,
weight_lowerings: Vec<WeightLoweringCapability>,
weight_residencies: Vec<WeightResidencyMechanism>,
state_residencies: Vec<StateResidencyMechanism>,
session: SessionCapabilities,
prompt_cache: bool,
exact_completion: bool,
grouped_operations: Vec<GroupedOperationRequirement>,
}
impl BackendMechanismCapabilities {
pub fn new(
operators: NeuralOperatorCapabilities,
weight_lowerings: Vec<WeightLoweringCapability>,
weight_residencies: Vec<WeightResidencyMechanism>,
state_residencies: Vec<StateResidencyMechanism>,
) -> Self {
Self {
operators,
weight_lowerings,
weight_residencies,
state_residencies,
session: SessionCapabilities::default(),
prompt_cache: false,
exact_completion: false,
grouped_operations: Vec::new(),
}
}
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 operators(&self) -> NeuralOperatorCapabilities {
self.operators
}
pub fn weight_lowerings(&self) -> &[WeightLoweringCapability] {
&self.weight_lowerings
}
pub fn weight_residencies(&self) -> &[WeightResidencyMechanism] {
&self.weight_residencies
}
pub fn state_residencies(&self) -> &[StateResidencyMechanism] {
&self.state_residencies
}
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)]
pub struct ReplicatedTextSelectionRequest {
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 {
topology: None,
residency,
state,
quantization: None,
session: SessionCapabilities::default(),
prompt_cache: false,
exact_completion: false,
}
}
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,
}
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 SelectedReplicatedTextRealization {
topology: ParallelTopology,
residency: LayerWeightResidency,
state: CacheResidencyPolicy,
parameters: Vec<SelectedParameterRealization>,
session: SessionCapabilities,
prompt_cache: bool,
exact_completion: bool,
grouped_operations: Vec<GroupedOperationRequirement>,
}
impl SelectedReplicatedTextRealization {
pub const fn topology(&self) -> ParallelTopology {
self.topology
}
pub const fn residency(&self) -> LayerWeightResidency {
self.residency
}
pub const fn state(&self) -> &CacheResidencyPolicy {
&self.state
}
pub fn parameters(&self) -> &[SelectedParameterRealization] {
&self.parameters
}
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 state_residency = match &request.state {
CacheResidencyPolicy::Device => StateResidencyMechanism::Device,
CacheResidencyPolicy::Paged(_) => StateResidencyMechanism::Paged,
};
if !capabilities.state_residencies.contains(&state_residency) {
issues.push(format!("state residency {state_residency:?}"));
}
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 names = BTreeSet::new();
for parameter in &requirements.parameters {
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.presence.has_physical_source() {
continue;
}
let candidate = match request.quantization {
Some(request) => match parameter.transform_target(request) {
Ok(Some(target)) => Some((target.executable(), target.descriptor().clone())),
Ok(None) => Some((
parameter.native_executable,
parameter
.lowering_descriptor(parameter.native_executable)
.expect("validated parameter forms native descriptor"),
)),
Err(error) => {
issues.push(error.to_string());
None
}
},
None => Some((
parameter.native_executable,
parameter
.lowering_descriptor(parameter.native_executable)
.expect("validated parameter forms native descriptor"),
)),
};
let Some((executable, descriptor)) = candidate else {
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 {:?}",
parameter.source_encoding, executable, parameter.name
));
continue;
};
parameters.push(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: lowering.kind,
});
}
if !issues.is_empty() {
return Err(ReplicatedTextSelectionError { issues });
}
Ok(SelectedReplicatedTextRealization {
topology: request
.topology
.unwrap_or_else(|| ParallelTopology::new(1, 1, 1, 1).expect("replicated topology")),
residency: request.residency,
state: request.state.clone(),
parameters,
session: request.session,
prompt_cache: request.prompt_cache,
exact_completion: request.exact_completion,
grouped_operations: requirements.grouped_operations.clone(),
})
}
#[cfg(test)]
mod tests {
use super::*;
use crate::{
ArchitectureGroupKind, ArchitectureGroupPlacement, ArchitectureGroupTransport,
ArchitectureMergeDestination, DenseDiskStreamLoadOptions, ExecutionGroupSpec,
ExecutionUnitLayout, LayerwiseLoadOptions, StateLayout,
};
use eredu_checkpoint::{AffineQuantization, StoredDtype};
use eredu_core::{cache::LayerCachePolicy, 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, "/checkpoint/model.safetensors", name).unwrap()
}
fn requirements() -> ReplicatedTextRequirements {
let graph =
ExecutionGraph::new(vec![ExecutionGroupSpec::root("decoder")], "decoder").unwrap();
let execution_units = ExecutionUnitLayout::new(&graph, [1]).unwrap();
ReplicatedTextRequirements::new(
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(),
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 },
)
.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()
}
#[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 physical_provenance_distinguishes_outputs_from_one_sharded_tensor() {
let shard = "/checkpoint/model-00002-of-00003.gguf";
let weight = ReplicatedTextPhysicalSource::new(
"blk.0.ffn_gate.weight",
shard,
"blk.0.ffn_gate.weight",
)
.unwrap();
let scales = ReplicatedTextPhysicalSource::new(
"blk.0.ffn_gate.weight",
shard,
"blk.0.ffn_gate.scales",
)
.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);
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,
],
vec![
StateResidencyMechanism::Device,
StateResidencyMechanism::Paged,
],
)
.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 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(), &paged_state());
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 selection_reports_all_missing_mechanisms_together() {
let capabilities = BackendMechanismCapabilities::new(
NeuralOperatorCapabilities::NONE,
Vec::new(),
Vec::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 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") }));
}
}