use alloc::{collections::BTreeMap,vec::Vec};
use ruda_model::tensor::backend::Backend;
use crate::NativeMoeLayer;
use super::{AdaptTransformerProjection,TransformerAdapterConfig,AttentionAdapterTarget,FeedForwardAdapterTarget,
NativeMoeFeedForward,NativeMoeTransformerBlock,NativeMoeTransformerLayer,NativeMoeTransformerStack,NativeMoeTransformerModel};
#[derive(Clone,Debug)]
pub enum NativeMoeFeedForwardAdapterTargets {
Dense(Vec<FeedForwardAdapterTarget>),
Routed {
router:bool,
shared:Vec<FeedForwardAdapterTarget>,
},
}
#[derive(Clone,Debug)]
pub struct NativeMoeLayerAdapterConfig {
pub layer:usize,
pub adapter:TransformerAdapterConfig,
pub attention:Vec<AttentionAdapterTarget>,
pub feed_forward:NativeMoeFeedForwardAdapterTargets,
}
impl<B:Backend,P:AdaptTransformerProjection<B>> NativeMoeLayer<B,P> {
pub fn with_router_adapter(mut self,config:&TransformerAdapterConfig) -> Self {
self.router=self.router.with_adapter(config);self
}
}
impl<B:Backend,P:AdaptTransformerProjection<B>> NativeMoeFeedForward<B,P> {
fn validate_adapters(&self,config:&TransformerAdapterConfig,router:bool,shared:&[FeedForwardAdapterTarget]) {
if router {self.routed.router.validate_adapter(config);}
if !shared.is_empty() {self.shared.as_ref().expect("selected original routed FFN has no shared branch").validate_adapters(config,shared);}
}
pub fn with_adapters(mut self,config:&TransformerAdapterConfig,router:bool,shared:&[FeedForwardAdapterTarget]) -> Self {
self.validate_adapters(config,router,shared);if router {self.routed=self.routed.with_router_adapter(config);}
if !shared.is_empty() {self.shared=self.shared.map(|branch|branch.with_adapters(config,shared));}self
}
}
impl<B:Backend,P:AdaptTransformerProjection<B>> NativeMoeTransformerBlock<B,P> {
fn validate_adapters(&self,config:&TransformerAdapterConfig,attention:&[AttentionAdapterTarget],router:bool,shared:&[FeedForwardAdapterTarget]) {
assert!(!attention.is_empty() || router || !shared.is_empty(),"selected routed block requires an actual native projection target");
self.attention.validate_adapters(config,attention);self.feed_forward.validate_adapters(config,router,shared);
}
pub fn with_adapters(mut self,config:&TransformerAdapterConfig,attention:&[AttentionAdapterTarget],router:bool,shared:&[FeedForwardAdapterTarget]) -> Self {
self.validate_adapters(config,attention,router,shared);self.attention=self.attention.with_adapters(config,attention);
self.feed_forward=self.feed_forward.with_adapters(config,router,shared);self
}
}
impl<B:Backend,P:AdaptTransformerProjection<B>> NativeMoeTransformerLayer<B,P> {
fn validate_adapters(&self,config:&NativeMoeLayerAdapterConfig) {
match (self,&config.feed_forward) {
(Self::Dense(block),NativeMoeFeedForwardAdapterTargets::Dense(targets))=>block.validate_adapters(&config.adapter,&config.attention,targets),
(Self::Routed(block),NativeMoeFeedForwardAdapterTargets::Routed {router,shared})=>block.validate_adapters(&config.adapter,&config.attention,*router,shared),
_=>panic!("native adapter target branch differs from the actual loaded dense/routed layer"),
}
}
pub fn with_adapters(self,config:&NativeMoeLayerAdapterConfig) -> Self {
self.validate_adapters(config);match (self,&config.feed_forward) {
(Self::Dense(block),NativeMoeFeedForwardAdapterTargets::Dense(targets))=>Self::Dense(block.with_adapters(&config.adapter,&config.attention,targets)),
(Self::Routed(block),NativeMoeFeedForwardAdapterTargets::Routed {router,shared})=>Self::Routed(block.with_adapters(&config.adapter,&config.attention,*router,shared)),
_=>unreachable!("validated original native layer variant"),
}
}
}
impl<B:Backend,P:AdaptTransformerProjection<B>> NativeMoeTransformerStack<B,P> {
fn selected_adapters<'a>(&self,targets:&'a [NativeMoeLayerAdapterConfig]) -> BTreeMap<usize,&'a NativeMoeLayerAdapterConfig> {
let mut selected=BTreeMap::new();for config in targets {
assert!(config.layer<self.layers.len(),"native MoE adapter layer index exceeds actual loaded layers");
assert!(selected.insert(config.layer,config).is_none(),"duplicate native MoE adapter layer index");self.layers[config.layer].validate_adapters(config);
}selected
}
pub fn with_adapters(self,targets:&[NativeMoeLayerAdapterConfig]) -> Self {
let selected=self.selected_adapters(targets);Self {layers:self.layers.into_iter().enumerate().map(|(index,layer)| {
if let Some(config)=selected.get(&index) {layer.with_adapters(config)} else {layer}
}).collect()}
}
}
impl<B:Backend,P:AdaptTransformerProjection<B>> NativeMoeTransformerModel<B,P> {
pub fn with_adapters(mut self,targets:&[NativeMoeLayerAdapterConfig],head:Option<&TransformerAdapterConfig>) -> Self {
let _=self.backbone.selected_adapters(targets);if let Some(config)=head {self.head.projection.validate_adapter(config);}
self.backbone=self.backbone.with_adapters(targets);if let Some(config)=head {self.head.projection=self.head.projection.with_adapter(config);}self
}
}