use alloc::{collections::{BTreeMap,BTreeSet},format,string::String,vec::Vec};
use ruda_model::{module::{Module,ModuleVisitor,ModuleMapper,Param,ParamId},record::{Record,PrecisionSettings,Recorder,RecorderError},
serde::{Serialize,Deserialize},tensor::{Tensor,Int,Bool,DType,TensorData,TensorPrimitive,read_sync,backend::Backend}};
use crate::expert_parallel::ExpertParallelGeometry;
use super::{TransformerProjectionShape,ExpertParallelTransformerModel,ExpertParallelTransformerLayer,ExpertAdapterOwnershipEntry};
#[derive(Clone,Copy,Debug,PartialEq,Eq,PartialOrd,Ord,Serialize,Deserialize)]
#[serde(crate="ruda_model::serde")]
pub enum ExpertModelParameterKind {Float,Integer,Boolean}
#[derive(Clone,Debug,PartialEq,Eq,Serialize,Deserialize)]
#[serde(crate="ruda_model::serde")]
pub struct ExpertModelParameterBinding {
pub path:Vec<String>,
pub id:u64,
pub shape:Vec<usize>,
pub dtype:DType,
pub kind:ExpertModelParameterKind,
pub trainable:bool,
}
type BindingKey=(u64,ExpertModelParameterKind,bool);
fn key(binding:&ExpertModelParameterBinding) -> BindingKey {(binding.id,binding.kind,binding.trainable)}
fn invalid(reason:&str) -> RecorderError {RecorderError::Unknown(format!("Invalid expert-owned model state: {reason}"))}
enum StoredValue<B:Backend> {
Float(B::FloatTensorPrimitive),Integer(B::IntTensorPrimitive),Boolean(B::BoolTensorPrimitive),
Archived(ExpertModelParameterKind,TensorData),
}
impl<B:Backend> StoredValue<B> {
fn into_data(self) -> TensorData {
read_sync(self.into_data_async()).expect("native expert model state readback failed")
}
async fn into_data_async(self) -> Result<TensorData,RecorderError> {
(match self {
Self::Float(value)=>B::float_into_data(value).await,
Self::Integer(value)=>B::int_into_data(value).await,
Self::Boolean(value)=>B::bool_into_data(value).await,
Self::Archived(_,data)=>return Ok(data),
}).map_err(|error|invalid(&format!("native parameter readback failed: {error:?}")))
}
fn from_data(kind:ExpertModelParameterKind,data:TensorData,device:&B::Device) -> Self {match kind {
ExpertModelParameterKind::Float=>Self::Float(B::float_from_data(data,device)),
ExpertModelParameterKind::Integer=>Self::Integer(B::int_from_data(data,device)),
ExpertModelParameterKind::Boolean=>Self::Boolean(B::bool_from_data(data,device)),
}}
}
struct Capture<B:Backend> {
path:Vec<String>,bindings:Vec<ExpertModelParameterBinding>,values:Vec<StoredValue<B>>,
aliases:BTreeMap<BindingKey,(Vec<usize>,DType,B::Device)>,paths:BTreeSet<Vec<String>>,error:Option<RecorderError>,
}
struct Validate<'a,B:Backend> {
path:Vec<String>,expected:BTreeMap<Vec<String>,&'a ExpertModelParameterBinding>,seen:BTreeSet<Vec<String>>,
source_to_target:BTreeMap<BindingKey,BindingKey>,target_to_source:BTreeMap<BindingKey,BindingKey>,devices:BTreeMap<BindingKey,B::Device>,error:Option<RecorderError>,
}
impl<'a,B:Backend> Validate<'a,B> {
fn new(bindings:&'a [ExpertModelParameterBinding]) -> Result<Self,RecorderError> {
let mut expected=BTreeMap::new();for binding in bindings {if expected.insert(binding.path.clone(),binding).is_some() {return Err(invalid("duplicate saved native parameter path"));}}
Ok(Self {path:Vec::new(),expected,seen:BTreeSet::new(),source_to_target:BTreeMap::new(),target_to_source:BTreeMap::new(),devices:BTreeMap::new(),error:None})
}
fn parameter(&mut self,id:ParamId,shape:Vec<usize>,dtype:Option<DType>,kind:ExpertModelParameterKind,trainable:bool,device:B::Device) {
let Some(saved)=self.expected.get(&self.path) else {self.error=Some(invalid("unexpected actual native parameter path"));return;};
if !self.seen.insert(self.path.clone()) || saved.shape!=shape || dtype.is_some_and(|storage|storage!=saved.dtype)
|| saved.kind!=kind || saved.trainable!=trainable {self.error=Some(invalid("actual parameter path/shape/storage/training contract differs"));return;}
let source=key(saved);let destination=(id.val(),kind,trainable);
if self.devices.insert(destination,device.clone()).is_some_and(|previous|previous!=device) {self.error=Some(invalid("prepared tied parameter devices differ"));}
if self.source_to_target.insert(source,destination).is_some_and(|previous|previous!=destination)
|| self.target_to_source.insert(destination,source).is_some_and(|previous|previous!=source) {self.error=Some(invalid("original shared/frozen parameter alias topology differs"));}
}
}
impl<B:Backend> ModuleVisitor<B> for Validate<'_,B> {
fn enter_module(&mut self,name:&str,_kind:&str) {self.path.push(name.into());}
fn exit_module(&mut self,_name:&str,_kind:&str) {self.path.pop();}
fn visit_float<const D:usize>(&mut self,parameter:&Param<Tensor<B,D>>) {
let dtype=parameter.is_initialized().then(||parameter.val().dtype());
self.parameter(parameter.id,parameter.lazy_shape().as_slice().to_vec(),dtype,ExpertModelParameterKind::Float,
B::ad_enabled(¶meter.lazy_device()) && parameter.planned_is_require_grad(),parameter.lazy_device());
}
fn visit_int<const D:usize>(&mut self,parameter:&Param<Tensor<B,D,Int>>) {
let dtype=parameter.is_initialized().then(||parameter.val().dtype());
self.parameter(parameter.id,parameter.lazy_shape().as_slice().to_vec(),dtype,ExpertModelParameterKind::Integer,false,parameter.lazy_device());
}
fn visit_bool<const D:usize>(&mut self,parameter:&Param<Tensor<B,D,Bool>>) {
let dtype=parameter.is_initialized().then(||parameter.val().dtype());
self.parameter(parameter.id,parameter.lazy_shape().as_slice().to_vec(),dtype,ExpertModelParameterKind::Boolean,false,parameter.lazy_device());
}
}
impl<B:Backend> Capture<B> {
fn new() -> Self {Self {path:Vec::new(),bindings:Vec::new(),values:Vec::new(),
aliases:BTreeMap::new(),paths:BTreeSet::new(),error:None}}
fn binding(&mut self,id:ParamId,shape:Vec<usize>,dtype:DType,kind:ExpertModelParameterKind,trainable:bool,device:B::Device) {
let binding=ExpertModelParameterBinding {path:self.path.clone(),id:id.val(),shape:shape.clone(),dtype,kind,trainable};
if !self.paths.insert(binding.path.clone()) {self.error=Some(invalid("duplicate native module parameter path"));}
if let Some((previous,storage,resident))=self.aliases.get(&key(&binding)) {
if previous!=&shape || storage!=&dtype || resident!=&device {self.error=Some(invalid("tied source parameter shape/storage/device differs"));}
} else {self.aliases.insert(key(&binding),(shape,dtype,device));}
self.bindings.push(binding);
}
}
impl<B:Backend> ModuleVisitor<B> for Capture<B> {
fn enter_module(&mut self,name:&str,_kind:&str) {self.path.push(name.into());}
fn exit_module(&mut self,_name:&str,_kind:&str) {self.path.pop();}
fn visit_float<const D:usize>(&mut self,parameter:&Param<Tensor<B,D>>) {
let value=parameter.val();self.binding(parameter.id,value.dims().to_vec(),value.dtype(),ExpertModelParameterKind::Float,value.is_require_grad(),value.device());
self.values.push(StoredValue::Float(parameter.transform_for_save().val().detach().into_primitive().tensor()));
}
fn visit_int<const D:usize>(&mut self,parameter:&Param<Tensor<B,D,Int>>) {
let value=parameter.val();self.binding(parameter.id,value.dims().to_vec(),value.dtype(),ExpertModelParameterKind::Integer,false,value.device());
self.values.push(StoredValue::Integer(parameter.transform_for_save().val().into_primitive()));
}
fn visit_bool<const D:usize>(&mut self,parameter:&Param<Tensor<B,D,Bool>>) {
let value=parameter.val();self.binding(parameter.id,value.dims().to_vec(),value.dtype(),ExpertModelParameterKind::Boolean,false,value.device());
self.values.push(StoredValue::Boolean(parameter.transform_for_save().val().into_primitive()));
}
}
fn ownership<B:Backend,P:TransformerProjectionShape<B>,E:ExpertParallelGeometry<B>>(model:&ExpertParallelTransformerModel<B,P,E>) -> Vec<ExpertAdapterOwnershipEntry> {
let mut result=Vec::new();for (index,layer) in model.layers.iter().enumerate() {if let ExpertParallelTransformerLayer::Parallel(block)=layer {
result.push(ExpertAdapterOwnershipEntry {layer:index,prefix:block.routed.experts.ownership().prefix().to_vec(),rank:block.routed.experts.rank()});
}}result
}
fn validate_metadata<B:Backend,P:TransformerProjectionShape<B>,E:ExpertParallelGeometry<B>>(version:u32,saved_contract:&str,layers:usize,
owners:&[ExpertAdapterOwnershipEntry],bindings:&[ExpertModelParameterBinding],values:usize,
model:&ExpertParallelTransformerModel<B,P,E>,contract_id:&str) -> Result<(),RecorderError> {
if version!=1 || contract_id.is_empty() || saved_contract!=contract_id || layers!=model.layers.len() || owners!=ownership(model) || bindings.len()!=values {
return Err(invalid("version, architecture contract, model layers or original expert ownership differs"));}
let mut validation=Validate::<B>::new(bindings)?;model.visit(&mut validation);if let Some(error)=validation.error {return Err(error);}
if validation.seen.len()!=validation.expected.len() {return Err(invalid("actual parameter path set differs"));}Ok(())
}
pub struct ExpertParallelModelStateRecord<B:Backend> {
version:u32,contract_id:String,layers:usize,ownership:Vec<ExpertAdapterOwnershipEntry>,
bindings:Vec<ExpertModelParameterBinding>,values:Vec<StoredValue<B>>,
}
#[derive(Clone,Serialize,Deserialize)]
#[serde(crate="ruda_model::serde")]
pub struct ExpertParallelModelSnapshot {
version:u32,contract_id:String,layers:usize,ownership:Vec<ExpertAdapterOwnershipEntry>,
bindings:Vec<ExpertModelParameterBinding>,values:Vec<(ExpertModelParameterKind,TensorData)>,
}
impl<B:Backend> Record<B> for ExpertParallelModelSnapshot {
type Item<S:PrecisionSettings>=Self;
fn into_item<S:PrecisionSettings>(self) -> Self {self}
fn from_item<S:PrecisionSettings>(item:Self,_device:&B::Device) -> Self {item}
}
impl ExpertParallelModelSnapshot {
pub fn ownership(&self) -> &[ExpertAdapterOwnershipEntry] {&self.ownership}
pub fn bindings(&self) -> &[ExpertModelParameterBinding] {&self.bindings}
pub fn parameter_occurrences(&self) -> usize {self.bindings.len()}
pub fn validate_for<B:Backend,P:TransformerProjectionShape<B>,E:ExpertParallelGeometry<B>>(&self,model:&ExpertParallelTransformerModel<B,P,E>,contract_id:&str) -> Result<(),RecorderError> {
validate_metadata(self.version,&self.contract_id,self.layers,&self.ownership,&self.bindings,self.values.len(),model,contract_id)
}
pub fn restore_into<B:Backend,P:TransformerProjectionShape<B>,E:ExpertParallelGeometry<B>>(self,model:ExpertParallelTransformerModel<B,P,E>,contract_id:&str,device:&B::Device)
-> Result<ExpertParallelTransformerModel<B,P,E>,RecorderError> {
self.validate_for(&model,contract_id)?;
let values=self.values.into_iter().map(|(kind,data)|StoredValue::Archived(kind,data)).collect();
restore_values(model,self.bindings,values,Some(device.clone()))
}
pub fn into_record<B:Backend>(self,device:&B::Device) -> ExpertParallelModelStateRecord<B> {
ExpertParallelModelStateRecord {version:self.version,contract_id:self.contract_id,layers:self.layers,ownership:self.ownership,bindings:self.bindings,
values:self.values.into_iter().map(|(kind,data)|StoredValue::from_data(kind,data,device)).collect()}
}
}
impl<B:Backend> Record<B> for ExpertParallelModelStateRecord<B> {
type Item<S:PrecisionSettings>=(u32,String,usize,Vec<ExpertAdapterOwnershipEntry>,Vec<ExpertModelParameterBinding>,Vec<(ExpertModelParameterKind,TensorData)>);
fn into_item<S:PrecisionSettings>(self) -> Self::Item<S> {
let values=self.bindings.iter().zip(self.values).map(|(binding,value)|(binding.kind,value.into_data())).collect();
(self.version,self.contract_id,self.layers,self.ownership,self.bindings,values)
}
fn from_item<S:PrecisionSettings>(item:Self::Item<S>,device:&B::Device) -> Self {
Self {version:item.0,contract_id:item.1,layers:item.2,ownership:item.3,bindings:item.4,
values:item.5.into_iter().map(|(kind,data)|StoredValue::from_data(kind,data,device)).collect()}
}
}
impl<B:Backend> ExpertParallelModelStateRecord<B> {
pub fn capture<P:TransformerProjectionShape<B>,E:ExpertParallelGeometry<B>>(model:&ExpertParallelTransformerModel<B,P,E>,contract_id:&str) -> Result<Self,RecorderError> {
if contract_id.is_empty() {return Err(invalid("exact prepared architecture/operator contract ID is required"));}
let mut capture=Capture::new();model.visit(&mut capture);if let Some(error)=capture.error {return Err(error);}
Ok(Self {version:1,contract_id:contract_id.into(),layers:model.layers.len(),ownership:ownership(model),bindings:capture.bindings,values:capture.values})
}
pub fn ownership(&self) -> &[ExpertAdapterOwnershipEntry] {&self.ownership}
pub fn bindings(&self) -> &[ExpertModelParameterBinding] {&self.bindings}
pub fn parameter_occurrences(&self) -> usize {self.bindings.len()}
pub async fn into_snapshot(self) -> Result<ExpertParallelModelSnapshot,RecorderError> {
if self.bindings.len()!=self.values.len() {return Err(invalid("native payload count differs from parameter bindings"));}
let mut values=Vec::with_capacity(self.values.len());
for (binding,value) in self.bindings.iter().zip(self.values) {values.push((binding.kind,value.into_data_async().await?));}
Ok(ExpertParallelModelSnapshot {version:self.version,contract_id:self.contract_id,layers:self.layers,ownership:self.ownership,bindings:self.bindings,values})
}
pub fn validate_for<P:TransformerProjectionShape<B>,E:ExpertParallelGeometry<B>>(&self,model:&ExpertParallelTransformerModel<B,P,E>,contract_id:&str) -> Result<(),RecorderError> {
validate_metadata(self.version,&self.contract_id,self.layers,&self.ownership,&self.bindings,self.values.len(),model,contract_id)
}
pub fn save<R:Recorder<B>>(self,recorder:&R,args:R::RecordArgs) -> Result<R::RecordOutput,RecorderError> {recorder.record(self,args)}
pub fn load<R:Recorder<B>>(recorder:&R,args:R::LoadArgs,device:&B::Device) -> Result<Self,RecorderError> {recorder.load(args,device)}
pub fn restore_into<P:TransformerProjectionShape<B>,E:ExpertParallelGeometry<B>>(self,model:ExpertParallelTransformerModel<B,P,E>,contract_id:&str)
-> Result<ExpertParallelTransformerModel<B,P,E>,RecorderError> {
self.validate_for(&model,contract_id)?;
restore_values(model,self.bindings,self.values,None)
}
}
fn restore_values<B:Backend,P:TransformerProjectionShape<B>,E:ExpertParallelGeometry<B>>(model:ExpertParallelTransformerModel<B,P,E>,
bindings:Vec<ExpertModelParameterBinding>,values:Vec<StoredValue<B>>,upload_device:Option<B::Device>) -> Result<ExpertParallelTransformerModel<B,P,E>,RecorderError> {
let entries=bindings.into_iter().zip(values).map(|(binding,value)|(binding.path.clone(),(binding,value))).collect();
let mut restore=Restore {path:Vec::new(),entries,canonical:BTreeMap::new(),error:None,seen:BTreeSet::new(),upload_device};let model=model.map(&mut restore);
if let Some(error)=restore.error {return Err(error);}
if !restore.entries.is_empty() {return Err(invalid("model mapper omitted an actual parameter path"));}Ok(model)
}
struct Restore<B:Backend> {
path:Vec<String>,entries:BTreeMap<Vec<String>,(ExpertModelParameterBinding,StoredValue<B>)>,
canonical:BTreeMap<BindingKey,StoredValue<B>>,error:Option<RecorderError>,seen:BTreeSet<Vec<String>>,upload_device:Option<B::Device>,
}
impl<B:Backend> Restore<B> {
fn binding(&mut self,kind:ExpertModelParameterKind) -> Option<ExpertModelParameterBinding> {
let Some((binding,_))=self.entries.get(&self.path) else {self.error=Some(invalid("missing validated parameter path"));return None;};
if binding.kind!=kind {self.error=Some(invalid("parameter mapper changed the validated kind"));return None;}
self.seen.insert(self.path.clone());Some(binding.clone())
}
fn take_payload(&mut self,kind:ExpertModelParameterKind) -> Option<StoredValue<B>> {
let (_,value)=self.entries.remove(&self.path).expect("validated native parameter path");
if let StoredValue::Archived(saved,data)=value {
if saved!=kind {self.error=Some(invalid("archived parameter payload kind differs"));return None;}
let device=self.upload_device.as_ref().expect("actual archive upload device");
Some(StoredValue::from_data(saved,data,device))
} else {Some(value)}
}
}
impl<B:Backend> ModuleMapper<B> for Restore<B> {
fn enter_module(&mut self,name:&str,_kind:&str) {self.path.push(name.into());}
fn exit_module(&mut self,_name:&str,_kind:&str) {self.path.pop();}
fn map_float<const D:usize>(&mut self,parameter:Param<Tensor<B,D>>) -> Param<Tensor<B,D>> {
let Some(binding)=self.binding(ExpertModelParameterKind::Float) else {return parameter;};
let Some(StoredValue::Float(value))=self.take_payload(binding.kind) else {self.error=Some(invalid("saved floating payload kind differs"));return parameter;};
let loaded=parameter.transform_for_load(Tensor::from_primitive(TensorPrimitive::Float(value)),ParamId::from(binding.id));
let (id,value,mapper)=loaded.consume();let value=value.detach().set_require_grad(binding.trainable);
if value.dims().to_vec()!=binding.shape || value.dtype()!=binding.dtype || value.is_require_grad()!=binding.trainable {
self.error=Some(invalid("floating load mapper changed source shape/storage/flags"));return Param::from_mapped_value(id,value,mapper);}
let primitive=match self.canonical.get(&key(&binding)) {
Some(StoredValue::Float(previous))=>previous.clone(),Some(_)=>unreachable!("validated floating canonical kind"),
None=>{let primitive=value.into_primitive().tensor();self.canonical.insert(key(&binding),StoredValue::Float(primitive.clone()));primitive},
};Param::from_mapped_value(id,Tensor::from_primitive(TensorPrimitive::Float(primitive)),mapper)
}
fn map_int<const D:usize>(&mut self,parameter:Param<Tensor<B,D,Int>>) -> Param<Tensor<B,D,Int>> {
let Some(binding)=self.binding(ExpertModelParameterKind::Integer) else {return parameter;};
let Some(StoredValue::Integer(value))=self.take_payload(binding.kind) else {self.error=Some(invalid("saved integer payload kind differs"));return parameter;};
let loaded=parameter.transform_for_load(Tensor::from_primitive(value),ParamId::from(binding.id));let (id,value,mapper)=loaded.consume();
if value.dims().to_vec()!=binding.shape || value.dtype()!=binding.dtype {self.error=Some(invalid("integer load mapper changed source shape/storage"));return Param::from_mapped_value(id,value,mapper);}
let primitive=match self.canonical.get(&key(&binding)) {
Some(StoredValue::Integer(previous))=>previous.clone(),Some(_)=>unreachable!("validated integer canonical kind"),
None=>{let primitive=value.into_primitive();self.canonical.insert(key(&binding),StoredValue::Integer(primitive.clone()));primitive},
};Param::from_mapped_value(id,Tensor::from_primitive(primitive),mapper)
}
fn map_bool<const D:usize>(&mut self,parameter:Param<Tensor<B,D,Bool>>) -> Param<Tensor<B,D,Bool>> {
let Some(binding)=self.binding(ExpertModelParameterKind::Boolean) else {return parameter;};
let Some(StoredValue::Boolean(value))=self.take_payload(binding.kind) else {self.error=Some(invalid("saved boolean payload kind differs"));return parameter;};
let loaded=parameter.transform_for_load(Tensor::from_primitive(value),ParamId::from(binding.id));let (id,value,mapper)=loaded.consume();
if value.dims().to_vec()!=binding.shape || value.dtype()!=binding.dtype {self.error=Some(invalid("boolean load mapper changed source shape/storage"));return Param::from_mapped_value(id,value,mapper);}
let primitive=match self.canonical.get(&key(&binding)) {
Some(StoredValue::Boolean(previous))=>previous.clone(),Some(_)=>unreachable!("validated boolean canonical kind"),
None=>{let primitive=value.into_primitive();self.canonical.insert(key(&binding),StoredValue::Boolean(primitive.clone()));primitive},
};Param::from_mapped_value(id,Tensor::from_primitive(primitive),mapper)
}
}
impl<B:Backend,P:TransformerProjectionShape<B>,E:ExpertParallelGeometry<B>> ExpertParallelTransformerModel<B,P,E> {
pub fn expert_model_state_record(&self,contract_id:&str) -> Result<ExpertParallelModelStateRecord<B>,RecorderError> {
ExpertParallelModelStateRecord::capture(self,contract_id)
}
}