use alloc::{collections::{BTreeMap,BTreeSet},format,string::String,vec::Vec};
use ruda_model::{module::{Module,ModuleVisitor,Param},record::{Record,PrecisionSettings,Recorder,RecorderError},
serde::{Serialize,Deserialize},tensor::{Tensor,DType,backend::Backend}};
use crate::{ExpertAdapterTarget,ExpertAdapterProjections,ExpertAdapterMapper,ExpertAdapterRecordBase,
ExpertLoRAAdapterSchema,PackedExpertLoRA};
use super::{TransformerProjectionShape,Nf4MoeTransformerModel,Nf4MoeTransformerLayer};
type ParameterKey=(u64,bool);
fn key<B:Backend>(parameter:&Param<Tensor<B,3>>) -> ParameterKey {(parameter.id.val(),parameter.val().is_require_grad())}
fn invalid(reason:&str) -> RecorderError {RecorderError::Unknown(format!("Invalid native expert model adapter record: {reason}"))}
#[derive(Clone,Copy,Debug,PartialEq,Eq,PartialOrd,Ord,Serialize,Deserialize)]
#[serde(crate="ruda_model::serde")]
pub struct ExpertAdapterPath {
pub layer:usize,
pub role:ExpertAdapterTarget,
}
#[derive(Clone,Serialize,Deserialize)]
#[serde(crate="ruda_model::serde")]
pub struct ExpertAdapterProjectionEntry {
pub path:ExpertAdapterPath,
pub schema:ExpertLoRAAdapterSchema,
pub parameters:[ParameterKey;2],
}
struct ParameterEntry<B:Backend> {
key:ParameterKey,
shape:[usize;3],
dtype:DType,
record:<Param<Tensor<B,3>> as Module<B>>::Record,
}
impl<B:Backend> ParameterEntry<B> {
fn capture(parameter:&Param<Tensor<B,3>>) -> Self {
let value=parameter.val();let key=key(parameter);let shape=value.dims();let dtype=value.dtype();
let record=parameter.clone().map(|value| {
let trainable=value.is_require_grad();value.detach().set_require_grad(trainable)
}).into_record().map(|value|value.detach());
Self {key,shape,dtype,record}
}
}
impl<B:Backend> Record<B> for ParameterEntry<B> {
type Item<S:PrecisionSettings> = (ParameterKey,[usize;3],DType,<<Param<Tensor<B,3>> as Module<B>>::Record as Record<B>>::Item<S>);
fn into_item<S:PrecisionSettings>(self) -> Self::Item<S> {(self.key,self.shape,self.dtype,self.record.into_item::<S>())}
fn from_item<S:PrecisionSettings>(item:Self::Item<S>,device:&B::Device) -> Self {
Self {key:item.0,shape:item.1,dtype:item.2,record:<<Param<Tensor<B,3>> as Module<B>>::Record as Record<B>>::from_item::<S>(item.3,device)}
}
}
pub struct ExpertTransformerAdapterRecord<B:Backend> {
version:u32,
base_id:String,
layers:usize,
projections:Vec<ExpertAdapterProjectionEntry>,
parameters:Vec<ParameterEntry<B>>,
}
impl<B:Backend> Record<B> for ExpertTransformerAdapterRecord<B> {
type Item<S:PrecisionSettings> = (u32,String,usize,Vec<ExpertAdapterProjectionEntry>,
Vec<(ParameterKey,[usize;3],DType,<<Param<Tensor<B,3>> as Module<B>>::Record as Record<B>>::Item<S>)>);
fn into_item<S:PrecisionSettings>(self) -> Self::Item<S> {
(self.version,self.base_id,self.layers,self.projections,self.parameters.into_iter().map(|parameter|parameter.into_item::<S>()).collect())
}
fn from_item<S:PrecisionSettings>(item:Self::Item<S>,device:&B::Device) -> Self {
Self {version:item.0,base_id:item.1,layers:item.2,projections:item.3,
parameters:item.4.into_iter().map(|parameter|ParameterEntry::<B>::from_item::<S>(parameter,device)).collect()}
}
}
struct ParameterOccurrences(BTreeMap<ParameterKey,usize>);
impl<B:Backend> ModuleVisitor<B> for ParameterOccurrences {
fn visit_float<const D:usize>(&mut self,parameter:&Param<Tensor<B,D>>) {
let key=(parameter.id.val(),parameter.val().is_require_grad());*self.0.entry(key).or_default()+=1;
}
}
fn check_isolated<B:Backend,P:TransformerProjectionShape<B>,E:ExpertAdapterProjections<B>>(
model:&Nf4MoeTransformerModel<B,P,E>,expected:&BTreeMap<ParameterKey,usize>) -> Result<(),RecorderError> {
let mut actual=ParameterOccurrences(BTreeMap::new());model.visit(&mut actual);
for (key,count) in expected {if actual.0.get(key)!=Some(count) {
return Err(invalid("expert A/B is tied to an omitted non-expert parameter; use a full model record"));}}
Ok(())
}
impl<B:Backend> ExpertTransformerAdapterRecord<B> {
pub fn capture<P:TransformerProjectionShape<B>,E:ExpertAdapterProjections<B>>(model:&Nf4MoeTransformerModel<B,P,E>,base_id:&str)
-> Result<Self,RecorderError> {
if base_id.is_empty() {return Err(invalid("complete original base identity must be supplied"));}
let mut projections=Vec::new();let mut parameters=BTreeMap::<ParameterKey,ParameterEntry<B>>::new();
let mut occurrences=BTreeMap::new();let mut devices=BTreeMap::new();
for (index,layer) in model.layers.iter().enumerate() {if let Nf4MoeTransformerLayer::Packed(block)=layer {
for (role,projection) in block.routed.experts.expert_adapter_projections() {
let schema=projection.schema(base_id)?;let leaves=projection.parameters();let bindings=[key(leaves[0]),key(leaves[1])];
for parameter in leaves {
let binding=key(parameter);let value=parameter.val();*occurrences.entry(binding).or_default()+=1;
if let Some(previous)=parameters.get(&binding) {
if previous.shape!=value.dims() || previous.dtype!=value.dtype() || devices.get(&binding)!=Some(&value.device()) {
return Err(invalid("shared expert A/B has inconsistent shape, dtype or device"));}
} else {devices.insert(binding,value.device());parameters.insert(binding,ParameterEntry::capture(parameter));}
}
projections.push(ExpertAdapterProjectionEntry {path:ExpertAdapterPath {layer:index,role},schema,parameters:bindings});
}
}}
if projections.is_empty() {return Err(invalid("model has no actual expert adapters"));}
check_isolated(model,&occurrences)?;
Ok(Self {version:1,base_id:base_id.into(),layers:model.layers.len(),projections,parameters:parameters.into_values().collect()})
}
pub fn targets(&self) -> impl Iterator<Item=(&ExpertAdapterPath,&ExpertLoRAAdapterSchema)> {
self.projections.iter().map(|entry|(&entry.path,&entry.schema))
}
pub fn parameter_count(&self) -> usize {self.parameters.len()}
pub fn validate_for<P:TransformerProjectionShape<B>,E:ExpertAdapterProjections<B>>(&self,model:&Nf4MoeTransformerModel<B,P,E>,base_id:&str)
-> Result<(),RecorderError> {
if self.version!=1 || base_id.is_empty() || self.base_id!=base_id || self.layers!=model.layers.len() || self.projections.is_empty() {
return Err(invalid("version, complete original base identity or layer topology differs"));}
let mut entries=BTreeMap::new();for entry in &self.projections {
if entries.insert(entry.path,entry).is_some() {return Err(invalid("duplicate original expert projection path"));}}
let mut parameters=BTreeMap::new();for entry in &self.parameters {
if entry.key.0!=entry.record.id.val() || !matches!(entry.dtype,DType::F16|DType::BF16|DType::F32)
|| parameters.insert(entry.key,entry).is_some() {return Err(invalid("invalid or duplicate canonical A/B parameter"));}}
let mut occurrences=BTreeMap::new();let mut used=BTreeSet::new();let mut count=0;
for (index,layer) in model.layers.iter().enumerate() {if let Nf4MoeTransformerLayer::Packed(block)=layer {
for (role,projection) in block.routed.experts.expert_adapter_projections() {
count+=1;let path=ExpertAdapterPath {layer:index,role};let entry=entries.get(&path).ok_or_else(||invalid("missing original expert projection path"))?;
projection.validate_schema(&entry.schema,base_id)?;
for (leaf,binding) in projection.parameters().into_iter().zip(entry.parameters) {
let value=leaf.val();let old=key(leaf);let saved=parameters.get(&binding).ok_or_else(||invalid("expert A/B references an absent parameter"))?;
if binding.1!=value.is_require_grad() || saved.shape!=value.dims() || saved.dtype!=value.dtype() {
return Err(invalid("canonical expert A/B geometry/storage/flags differ"));}
*occurrences.entry(old).or_default()+=1;used.insert(binding);
}
}
}}
if count!=entries.len() || used.len()!=parameters.len() {return Err(invalid("unexpected expert target or unused A/B parameter"));}
check_isolated(model,&occurrences)?;
let mut actual=ParameterOccurrences(BTreeMap::new());model.visit(&mut actual);
for binding in used {if actual.0.get(&binding).copied().unwrap_or(0)!=occurrences.get(&binding).copied().unwrap_or(0) {
return Err(invalid("saved expert A/B ID collides with an unchanged non-expert parameter"));}}
Ok(())
}
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:ExpertAdapterProjections<B>>(self,mut model:Nf4MoeTransformerModel<B,P,E>,base_id:&str)
-> Result<Nf4MoeTransformerModel<B,P,E>,RecorderError> {
self.validate_for(&model,base_id)?;
let entries=self.projections.into_iter().map(|entry|(entry.path,entry)).collect();
let mut mapper=Restore {layer:0,entries,parameters:self.parameters.into_iter().map(|entry|(entry.key,entry)).collect(),leaves:BTreeMap::new()};
model.layers=model.layers.into_iter().enumerate().map(|(index,layer)| {
mapper.layer=index;match layer {
Nf4MoeTransformerLayer::Packed(mut block)=>{block.routed.experts=block.routed.experts.map_expert_adapters(&mut mapper)?;Ok(Nf4MoeTransformerLayer::Packed(block))},
original=>Ok(original),
}
}).collect::<Result<_,RecorderError>>()?;Ok(model)
}
}
struct Restore<B:Backend> {
layer:usize,
entries:BTreeMap<ExpertAdapterPath,ExpertAdapterProjectionEntry>,
parameters:BTreeMap<ParameterKey,ParameterEntry<B>>,
leaves:BTreeMap<ParameterKey,Tensor<B,3>>,
}
impl<B:Backend> Restore<B> {
fn parameter(&mut self,original:Param<Tensor<B,3>>,binding:ParameterKey) -> Result<Param<Tensor<B,3>>,RecorderError> {
let entry=self.parameters.get(&binding).ok_or_else(||invalid("missing validated A/B parameter"))?;
let device=original.val().device();let loaded=original.load_record(entry.record.clone()).fork(&device)
.map(|value|value.cast(entry.dtype).detach().set_require_grad(binding.1));
let (id,value,parameter_mapper)=loaded.consume();
if value.dims()!=entry.shape || value.dtype()!=entry.dtype || value.device()!=device || value.is_require_grad()!=binding.1 {
return Err(invalid("A/B load mapper changed original geometry/storage/device/flags"));}
let canonical=self.leaves.entry(binding).or_insert(value).clone();Ok(Param::from_mapped_value(id,canonical,parameter_mapper))
}
}
impl<B:Backend> ExpertAdapterMapper<B> for Restore<B> {
fn map<Base:ExpertAdapterRecordBase<B>>(&mut self,role:ExpertAdapterTarget,mut layer:PackedExpertLoRA<B,Base>)
-> Result<PackedExpertLoRA<B,Base>,RecorderError> {
let path=ExpertAdapterPath {layer:self.layer,role};let bindings=self.entries.get(&path).ok_or_else(||invalid("missing validated expert target"))?.parameters;
layer.adapter_a.weight=self.parameter(layer.adapter_a.weight,bindings[0])?;
layer.adapter_b.weight=self.parameter(layer.adapter_b.weight,bindings[1])?;layer.validate();Ok(layer)
}
}
impl<B:Backend,P:TransformerProjectionShape<B>,E:ExpertAdapterProjections<B>> Nf4MoeTransformerModel<B,P,E> {
pub fn expert_adapter_record(&self,base_id:&str) -> Result<ExpertTransformerAdapterRecord<B>,RecorderError> {
ExpertTransformerAdapterRecord::capture(self,base_id)
}
}