use alloc::{collections::BTreeMap,vec::Vec};
use core::fmt;
use ruda_model::{module::{Module,ModuleVisitor,Param,ParamId},record::RecorderError,tensor::{Tensor,Int,DType,MoeOptions,MoeExpertStrategy,
MoeReceivedOps,FrozenPackedExpertOps,ExpertProjectionOps,NativeSwiGluOps,backend::Backend}};
use super::{NativeSwiGluExperts,AdaptedFloatingSwiGluExperts,SelectableOwnedExperts,OwnedFloatingExpertAdapters,
SelectablePackedExperts,OwnedPackedExperts,FrozenPackedSwiGluExperts,FrozenNf4SwiGluExperts,AdaptedPackedSwiGluExperts,
MixedAdaptedExperts,LoRALinearConfig,ExpertAdapterTarget,FrozenExpertGeometry,FrozenSelectedExperts,
ExpertAdapterProjections,ExpertAdapterProjectionRef,ExpertAdapterMapper,AdaptedExpertError,FloatingExpertError,
FrozenPackedExpertProjection,AdaptedExpertProjection};
use super::expert_parallel::{ExpertOwnership,ExpertParallelGeometry,ExpertParallelReceived,ExpertParallelMoeLayer};
#[derive(Debug)]
pub enum MixedExpertParallelSource<B:Backend> {
Native(NativeSwiGluExperts<B>),
Floating(AdaptedFloatingSwiGluExperts<B>),
Packed(SelectablePackedExperts<B>),
}
impl<B:Backend> From<NativeSwiGluExperts<B>> for MixedExpertParallelSource<B> {fn from(value:NativeSwiGluExperts<B>) -> Self {Self::Native(value)}}
impl<B:Backend> From<AdaptedFloatingSwiGluExperts<B>> for MixedExpertParallelSource<B> {fn from(value:AdaptedFloatingSwiGluExperts<B>) -> Self {Self::Floating(value)}}
impl<B:Backend> From<SelectablePackedExperts<B>> for MixedExpertParallelSource<B> {fn from(value:SelectablePackedExperts<B>) -> Self {Self::Packed(value)}}
macro_rules! packed_source {
($source:ident)=>{impl<B:Backend> From<$source<B>> for MixedExpertParallelSource<B> {fn from(value:$source<B>) -> Self {Self::Packed(value.into())}}};
}
packed_source!(FrozenPackedSwiGluExperts);
packed_source!(FrozenNf4SwiGluExperts);
packed_source!(AdaptedPackedSwiGluExperts);
impl<B:Backend,E:FrozenExpertGeometry<B>+Into<SelectablePackedExperts<B>>> From<MixedAdaptedExperts<B,E>> for MixedExpertParallelSource<B> {
fn from(value:MixedAdaptedExperts<B,E>) -> Self {match value {
MixedAdaptedExperts::Original(value)=>Self::Packed(value.into()),MixedAdaptedExperts::Packed(value)=>Self::Packed(value.into()),
MixedAdaptedExperts::Floating(value)=>Self::Floating(value)}}
}
impl<B:Backend> MixedExpertParallelSource<B> {
pub fn dimensions(&self) -> [usize;3] {match self {Self::Native(value)=>value.dimensions(),Self::Floating(value)=>value.dimensions(),Self::Packed(value)=>value.dimensions()}}
pub fn validate(&self) {match self {Self::Native(value)=>value.validate(),Self::Floating(value)=>value.validate(),Self::Packed(value)=>value.validate()}}
pub(super) fn validate_adapter_targets(&self,config:&LoRALinearConfig,targets:&[ExpertAdapterTarget],dtype:DType,options:MoeOptions) {
match self {
Self::Native(value)=>AdaptedFloatingSwiGluExperts::from_native(value.clone(),options.forward,options.backward).validate_adapter_targets(config,targets,dtype),
Self::Floating(value)=>value.validate_adapter_targets(config,targets,dtype),
Self::Packed(SelectablePackedExperts::Original(value))=>AdaptedPackedSwiGluExperts::from_frozen(value.clone()).validate_adapter_targets(config,targets,dtype),
Self::Packed(SelectablePackedExperts::Adapted(value))=>value.validate_adapter_targets(config,targets,dtype),
}
}
}
#[derive(Default)]
struct PartitionOccurrences {floating:BTreeMap<(ParamId,bool),usize>,integer:BTreeMap<ParamId,usize>}
impl<B:Backend> ModuleVisitor<B> for PartitionOccurrences {
fn visit_float<const D:usize>(&mut self,parameter:&Param<Tensor<B,D>>) {
*self.floating.entry((parameter.id,parameter.val().is_require_grad())).or_default()+=1;
}
fn visit_int<const D:usize>(&mut self,parameter:&Param<Tensor<B,D,Int>>) {*self.integer.entry(parameter.id).or_default()+=1;}
}
impl PartitionOccurrences {
fn unchanged_book<B:Backend>(&mut self,base:&FrozenPackedExpertProjection<B>) {
let parameter=match base {FrozenPackedExpertProjection::Nf4(value)=>&value.payload.codebook,
FrozenPackedExpertProjection::Nf4Window(value)=>&value.codebook,FrozenPackedExpertProjection::Awq(_)=>return};
let key=(parameter.id,parameter.val().is_require_grad());let count=self.floating.get_mut(&key).expect("actual visited original NF4 codebook");
*count-=1;if *count==0 {self.floating.remove(&key);}
}
}
pub(super) fn validate_ownership_aliases<B:Backend,M:Module<B>>(model:&M,sources:Vec<MixedExpertParallelSource<B>>) {
let mut all=PartitionOccurrences::default();model.visit(&mut all);let mut owned=PartitionOccurrences::default();
for source in sources {
match source {
MixedExpertParallelSource::Native(value)=>value.visit(&mut owned),
MixedExpertParallelSource::Floating(value)=>value.visit(&mut owned),
MixedExpertParallelSource::Packed(value)=>{
let mut selected=PartitionOccurrences::default();value.visit(&mut selected);
match value {
SelectablePackedExperts::Original(value)=>{for base in [&value.gate,&value.up,&value.down] {selected.unchanged_book(base);}},
SelectablePackedExperts::Adapted(value)=>{for projection in [&value.gate,&value.up,&value.down] {
selected.unchanged_book(match projection {AdaptedExpertProjection::Frozen(base)=>base,AdaptedExpertProjection::LoRA(layer)=>&layer.base});}},
}
for (key,count) in selected.floating {*owned.floating.entry(key).or_default()+=count;}
for (key,count) in selected.integer {*owned.integer.entry(key).or_default()+=count;}
},
}
}
for (key,count) in owned.floating {assert_eq!(all.floating.get(&key),Some(&count),"an owned expert parameter is tied to an unchanged local/non-expert parameter; its alias cannot be split");}
for (key,count) in owned.integer {assert_eq!(all.integer.get(&key),Some(&count),"owned packed expert words are tied to unchanged non-expert storage; its alias cannot be split");}
}
#[derive(Module,Debug)]
pub enum MixedOwnedExperts<B:Backend> {
Floating(SelectableOwnedExperts<B>),
Packed(OwnedPackedExperts<B>),
}
impl<B:Backend> MixedOwnedExperts<B> {
pub fn dimensions(&self) -> [usize;3] {match self {Self::Floating(value)=>ExpertParallelGeometry::dimensions(value),Self::Packed(value)=>value.dimensions()}}
pub fn device(&self) -> B::Device {match self {Self::Floating(value)=>ExpertParallelGeometry::device(value),Self::Packed(value)=>value.device()}}
pub fn validate(&self) {match self {Self::Floating(value)=>ExpertParallelGeometry::validate(value),Self::Packed(value)=>value.validate()}}
pub fn with_adapters(self,config:&LoRALinearConfig,targets:&[ExpertAdapterTarget],dtype:DType,use_rslora:bool,
source_options:MoeOptions,forward:MoeExpertStrategy,backward:MoeExpertStrategy) -> Self {
let value=match self {
Self::Floating(SelectableOwnedExperts::Original(value))=>Self::Floating(SelectableOwnedExperts::Adapted(
OwnedFloatingExpertAdapters::from_native(value,source_options.forward,source_options.backward).with_adapters(config,targets,dtype,use_rslora,forward,backward))),
Self::Floating(SelectableOwnedExperts::Adapted(value))=>Self::Floating(SelectableOwnedExperts::Adapted(value.with_adapters(config,targets,dtype,use_rslora,forward,backward))),
Self::Packed(value)=>Self::Packed(value.with_adapters(config,targets,dtype,use_rslora,forward,backward)),
};value.validate();value
}
pub fn adapter_parameter_ids(&self) -> Vec<ParamId> {match self {Self::Floating(value)=>value.adapter_parameter_ids(),Self::Packed(value)=>value.adapter_parameter_ids()}}
}
impl<B:Backend> ExpertParallelGeometry<B> for MixedOwnedExperts<B> {
fn dimensions(&self) -> [usize;3] {self.dimensions()}
fn ownership(&self) -> &ExpertOwnership {match self {Self::Floating(value)=>value.ownership(),Self::Packed(value)=>&value.ownership}}
fn rank(&self) -> usize {match self {Self::Floating(value)=>value.rank(),Self::Packed(value)=>value.rank}}
fn device(&self) -> B::Device {self.device()}
fn validate(&self) {self.validate();}
fn parameter_ids(&self) -> Vec<ParamId> {match self {Self::Floating(value)=>value.parameter_ids(),Self::Packed(value)=>value.parameter_ids()}}
}
#[derive(Debug)]
pub enum MixedOwnedExpertError<M:fmt::Debug,P:fmt::Debug,G:fmt::Debug,S:fmt::Debug> {
Native(M),
Packed(AdaptedExpertError<P,G,S>),
Floating(FloatingExpertError<G,S>),
}
impl<M:fmt::Debug,P:fmt::Debug,G:fmt::Debug,S:fmt::Debug> fmt::Display for MixedOwnedExpertError<M,P,G,S> {
fn fmt(&self,f:&mut fmt::Formatter<'_>) -> fmt::Result {match self {Self::Native(error)=>write!(f,"native owned expert: {error:?}"),Self::Packed(error)=>write!(f,"{error}"),Self::Floating(error)=>write!(f,"{error}")}}
}
impl<M:fmt::Debug,P:fmt::Debug,G:fmt::Debug,S:fmt::Debug> core::error::Error for MixedOwnedExpertError<M,P,G,S> {}
impl<B:MoeReceivedOps+FrozenPackedExpertOps+ExpertProjectionOps+NativeSwiGluOps> ExpertParallelReceived<B> for MixedOwnedExperts<B> {
type Error=MixedOwnedExpertError<B::MoeError,B::PackedExpertError,B::ExpertProjectionError,B::SwiGluError>;
fn routing_error(error:B::MoeError) -> Self::Error {MixedOwnedExpertError::Native(error)}
fn forward_received(&self,input:Tensor<B,2>,ids:Tensor<B,1,Int>,options:MoeOptions) -> Result<Tensor<B,2>,Self::Error> {
self.validate();match self {
Self::Floating(SelectableOwnedExperts::Original(value))=>value.forward_received(input,ids,options).map_err(MixedOwnedExpertError::Native),
Self::Floating(SelectableOwnedExperts::Adapted(value))=>value.experts.forward(input,ids,value.ownership.range(value.rank).start).map_err(MixedOwnedExpertError::Floating),
Self::Packed(value)=>value.experts.forward(input,ids,value.ownership.range(value.rank).start).map_err(MixedOwnedExpertError::Packed),
}
}
}
impl<B:Backend> FrozenExpertGeometry<B> for MixedOwnedExperts<B> {
fn dimensions(&self) -> [usize;3] {self.dimensions()}
fn device(&self) -> B::Device {self.device()}
fn validate(&self) {self.validate();}
}
impl<B:Backend> ExpertAdapterProjections<B> for MixedOwnedExperts<B> {
fn expert_adapter_projections(&self) -> Vec<(ExpertAdapterTarget,ExpertAdapterProjectionRef<'_,B>)> {match self {
Self::Floating(value)=>value.expert_adapter_projections(),Self::Packed(value)=>value.expert_adapter_projections()}}
fn map_expert_adapters<M:ExpertAdapterMapper<B>>(self,mapper:&mut M) -> Result<Self,RecorderError> {
let owned=match self {Self::Floating(value)=>Self::Floating(value.map_expert_adapters(mapper)?),Self::Packed(value)=>Self::Packed(value.map_expert_adapters(mapper)?)};
owned.validate();Ok(owned)
}
}
pub type MixedExpertParallelMoeLayer<B,P> = ExpertParallelMoeLayer<B,P,MixedOwnedExperts<B>>;
pub type MixedExpertParallelTransformerModel<B,P> = crate::transformer::ExpertParallelTransformerModel<B,P,MixedOwnedExperts<B>>;