use super::{ExpertLinear,ExpertLoRABase,PackedExpertLoRA,FrozenPackedExpertProjection};
use alloc::{format,string::String};
use ruda_model::{module::{Module,ModuleDTypeRecord},record::{Record,PrecisionSettings,Recorder,RecorderError},
serde::{Serialize,Deserialize},tensor::{DType,MoeExpertStrategy,backend::Backend}};
#[derive(Clone,Copy,Debug,PartialEq,Eq,Serialize,Deserialize)]
#[serde(crate="ruda_model::serde")]
pub enum ExpertExecutionSchema {Scalar,Auto,TensorCore}
impl From<MoeExpertStrategy> for ExpertExecutionSchema {
fn from(value:MoeExpertStrategy) -> Self {match value {
MoeExpertStrategy::Scalar=>Self::Scalar,MoeExpertStrategy::Auto=>Self::Auto,MoeExpertStrategy::TensorCore=>Self::TensorCore}}
}
#[derive(Clone,Debug,PartialEq,Serialize,Deserialize)]
#[serde(crate="ruda_model::serde")]
pub enum ExpertAdapterBaseSchema {
Floating {dtype:DType,execution:[ExpertExecutionSchema;2]},
Nf4 {block_size:usize,tile_rows:usize,use_tensor_core:bool},
Awq {group_size:usize,scale_dtype:DType,bias_dtype:Option<DType>},
}
pub trait ExpertAdapterRecordBase<B:Backend>:ExpertLoRABase<B> {
fn adapter_base_schema(&self) -> Result<ExpertAdapterBaseSchema,RecorderError>;
}
fn invalid(reason:&str) -> RecorderError {RecorderError::Unknown(format!("Invalid native expert adapter record: {reason}"))}
impl<B:Backend> ExpertAdapterRecordBase<B> for ExpertLinear<B> {
fn adapter_base_schema(&self) -> Result<ExpertAdapterBaseSchema,RecorderError> {
if self.weight.val().is_require_grad() {return Err(invalid("original floating expert base is not frozen"));}
self.validate();Ok(ExpertAdapterBaseSchema::Floating {dtype:self.weight.val().dtype(),
execution:[self.forward_strategy.into(),self.backward_strategy.into()]})
}
}
impl<B:Backend> ExpertAdapterRecordBase<B> for FrozenPackedExpertProjection<B> {
fn adapter_base_schema(&self) -> Result<ExpertAdapterBaseSchema,RecorderError> {
match self {
Self::Nf4(value)=>{
if value.payload.scales.val().is_require_grad() || value.payload.codebook.val().is_require_grad() {
return Err(invalid("original NF4 scale/codebook is not frozen"));}
value.validate();Ok(ExpertAdapterBaseSchema::Nf4 {block_size:value.payload.block_size,
tile_rows:value.payload.tile_rows,use_tensor_core:value.payload.use_tensor_core})
},
Self::Awq(value)=>{
if value.scales.val().is_require_grad() || value.bias.as_ref().is_some_and(|bias|bias.val().is_require_grad()) {
return Err(invalid("original AWQ scale/bias is not frozen"));}
value.validate();Ok(ExpertAdapterBaseSchema::Awq {group_size:value.group_size,scale_dtype:value.scales.val().dtype(),
bias_dtype:value.bias.as_ref().map(|bias|bias.val().dtype())})
},
}
}
}
#[derive(Clone,Debug,PartialEq,Serialize,Deserialize)]
#[serde(crate="ruda_model::serde")]
pub struct ExpertLoRAAdapterSchema {
pub version:u32,
pub base_id:String,
pub base_shape:[usize;3],
pub base:ExpertAdapterBaseSchema,
pub rank:usize,
pub scale:f64,
pub dropout:f64,
pub adapter_dtypes:[DType;2],
pub adapter_execution:[[ExpertExecutionSchema;2];2],
pub trainable:[bool;2],
}
impl<B:Backend> Record<B> for ExpertLoRAAdapterSchema {
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 ExpertLoRAAdapterSchema {
pub fn capture<B:Backend,Base:ExpertAdapterRecordBase<B>>(layer:&PackedExpertLoRA<B,Base>,base_id:&str) -> Result<Self,RecorderError> {
if base_id.is_empty() {return Err(invalid("exact original base identity must be supplied"));}
let base=layer.base.adapter_base_schema()?;let [e,k,n]=layer.base.dimensions();
let a=layer.adapter_a.weight.val();let b=layer.adapter_b.weight.val();let [ae,rank,ak]=a.dims();
if k==0 || n==0 || rank==0 || [ae,ak]!=[e,k] || b.dims()!=[e,n,rank] {
return Err(invalid("original base and actual expert A/B geometry differs"));}
if !matches!(a.dtype(),DType::F16|DType::BF16|DType::F32) || !matches!(b.dtype(),DType::F16|DType::BF16|DType::F32)
|| a.device()!=layer.base.device() || b.device()!=layer.base.device() {
return Err(invalid("actual A/B storage or resident device differs"));}
if !layer.scale.is_finite() || !layer.dropout.prob.is_finite() || !(0.0..1.0).contains(&layer.dropout.prob) {
return Err(invalid("original expert adapter scale/dropout is invalid"));}
Ok(Self {version:1,base_id:base_id.into(),base_shape:[e,k,n],base,rank,scale:layer.scale,dropout:layer.dropout.prob,
adapter_dtypes:[a.dtype(),b.dtype()],trainable:[a.is_require_grad(),b.is_require_grad()],
adapter_execution:[[layer.adapter_a.forward_strategy.into(),layer.adapter_a.backward_strategy.into()],
[layer.adapter_b.forward_strategy.into(),layer.adapter_b.backward_strategy.into()]]})
}
pub fn validate_for<B:Backend,Base:ExpertAdapterRecordBase<B>>(&self,layer:&PackedExpertLoRA<B,Base>,base_id:&str) -> Result<(),RecorderError> {
if self.version!=1 {return Err(invalid("unsupported expert adapter format version"));}
let actual=Self::capture(layer,base_id)?;
if self.base_id!=actual.base_id || self.base_shape!=actual.base_shape || self.base!=actual.base || self.rank!=actual.rank
|| self.adapter_dtypes!=actual.adapter_dtypes || self.adapter_execution!=actual.adapter_execution || self.trainable!=actual.trainable
|| self.scale.to_bits()!=actual.scale.to_bits() || self.dropout.to_bits()!=actual.dropout.to_bits() {
return Err(invalid("original base, rank, storage, execution, training flags or forward configuration differs"));}Ok(())
}
}
pub struct ExpertLoRAAdapterRecord<B:Backend> {
pub schema:ExpertLoRAAdapterSchema,
adapter_a:<ExpertLinear<B> as Module<B>>::Record,
adapter_b:<ExpertLinear<B> as Module<B>>::Record,
dtypes:ModuleDTypeRecord,
}
impl<B:Backend> Record<B> for ExpertLoRAAdapterRecord<B> {
type Item<S:PrecisionSettings> = (
<ExpertLoRAAdapterSchema as Record<B>>::Item<S>,
<<ExpertLinear<B> as Module<B>>::Record as Record<B>>::Item<S>,
<<ExpertLinear<B> as Module<B>>::Record as Record<B>>::Item<S>,
<ModuleDTypeRecord as Record<B>>::Item<S>,
);
fn into_item<S:PrecisionSettings>(self) -> Self::Item<S> {
(<ExpertLoRAAdapterSchema as Record<B>>::into_item::<S>(self.schema),self.adapter_a.into_item::<S>(),self.adapter_b.into_item::<S>(),
<ModuleDTypeRecord as Record<B>>::into_item::<S>(self.dtypes))
}
fn from_item<S:PrecisionSettings>(item:Self::Item<S>,device:&B::Device) -> Self {
Self {schema:<ExpertLoRAAdapterSchema as Record<B>>::from_item::<S>(item.0,device),
adapter_a:<<ExpertLinear<B> as Module<B>>::Record as Record<B>>::from_item::<S>(item.1,device),
adapter_b:<<ExpertLinear<B> as Module<B>>::Record as Record<B>>::from_item::<S>(item.2,device),
dtypes:<ModuleDTypeRecord as Record<B>>::from_item::<S>(item.3,device)}
}
}
impl<B:Backend> ExpertLoRAAdapterRecord<B> {
pub fn capture<Base:ExpertAdapterRecordBase<B>>(layer:&PackedExpertLoRA<B,Base>,base_id:&str) -> Result<Self,RecorderError> {
let schema=ExpertLoRAAdapterSchema::capture(layer,base_id)?;
let snapshot=|mut adapter:ExpertLinear<B>| {
adapter.weight=adapter.weight.map(|value| {let trainable=value.is_require_grad();value.detach().set_require_grad(trainable)});adapter
};
let adapters=(snapshot(layer.adapter_a.clone()),snapshot(layer.adapter_b.clone()));let dtypes=ModuleDTypeRecord::capture(&adapters)?;
Ok(Self {schema,adapter_a:adapters.0.into_record(),adapter_b:adapters.1.into_record(),dtypes})
}
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<Base:ExpertAdapterRecordBase<B>>(self,mut layer:PackedExpertLoRA<B,Base>,base_id:&str)
-> Result<PackedExpertLoRA<B,Base>,RecorderError> {
self.schema.validate_for(&layer,base_id)?;let schema=self.schema;let device=layer.base.device();
let a=layer.adapter_a.load_record(self.adapter_a).fork(&device);let b=layer.adapter_b.load_record(self.adapter_b).fork(&device);
let (a,b)=self.dtypes.apply((a,b))?;layer.adapter_a=a;layer.adapter_b=b;
schema.validate_for(&layer,base_id)?;Ok(layer)
}
}
impl<B:Backend,Base:ExpertAdapterRecordBase<B>> PackedExpertLoRA<B,Base> {
pub fn adapter_record(&self,base_id:&str) -> Result<ExpertLoRAAdapterRecord<B>,RecorderError> {ExpertLoRAAdapterRecord::capture(self,base_id)}
}