ruda-nn 0.21.46

Ruda neural network layers, activation modules and losses.
use super::*;
use alloc::{collections::BTreeMap,format,vec::Vec};
use crate::{Linear,transformer::AdaptedProjection};
use super::super::{TensorParallelAdaptedEncoderDecoderStack,TensorParallelAdaptedDecoderCrossAttention};
use ruda_model::{module::ParamId,record::RecorderError};

pub(super) trait AdapterTree<B:Backend> {
    fn adapters(&self) -> Vec<&Linear<B>>;
    fn map_adapters<F:FnMut(&mut Linear<B>)>(&mut self,apply:&mut F);
}

fn projection<'a,B:Backend>(projection:&'a AdaptedProjection<B>,values:&mut Vec<&'a Linear<B>>) {
    if let AdaptedProjection::LoRA(layer) = projection {values.push(&layer.adapter_a);values.push(&layer.adapter_b);}
}
fn map_projection<B:Backend,F:FnMut(&mut Linear<B>)>(projection:&mut AdaptedProjection<B>,apply:&mut F) {
    if let AdaptedProjection::LoRA(layer) = projection {apply(&mut layer.adapter_a);apply(&mut layer.adapter_b);}
}

impl<B:Backend> AdapterTree<B> for TensorParallelAdaptedStackLayer<B> {
    fn adapters(&self) -> Vec<&Linear<B>> {
        let mut values = Vec::new();
        if let Self::Adapted(block) = self {
            let attention = &block.attention.local;let feed = &block.feed_forward.local;
            for item in [&attention.query,&attention.key,&attention.value,&attention.output,&feed.up,&feed.down] {projection(item,&mut values);}
            if let Some(gate) = &feed.gate {projection(gate,&mut values);}
        }
        values
    }
    fn map_adapters<F:FnMut(&mut Linear<B>)>(&mut self,apply:&mut F) {
        if let Self::Adapted(block) = self {
            let attention = &mut block.attention.local;let feed = &mut block.feed_forward.local;
            for item in [&mut attention.query,&mut attention.key,&mut attention.value,&mut attention.output,&mut feed.up,&mut feed.down] {map_projection(item,apply);}
            if let Some(gate) = &mut feed.gate {map_projection(gate,apply);}
        }
    }
}

impl<B:Backend> AdapterTree<B> for TensorParallelAdaptedTransformerStack<B> {
    fn adapters(&self) -> Vec<&Linear<B>> {self.layers.iter().flat_map(|layer|layer.adapters()).collect()}
    fn map_adapters<F:FnMut(&mut Linear<B>)>(&mut self,apply:&mut F) {for layer in &mut self.layers {layer.map_adapters(apply);}}
}

impl<B:Backend> AdapterTree<B> for TensorParallelAdaptedDecoderCrossAttention<B> {
    fn adapters(&self) -> Vec<&Linear<B>> {
        let mut values = Vec::new();
        if let Self::Adapted(block) = self {
            let attention = &block.attention.local;
            for item in [&attention.query,&attention.key,&attention.value,&attention.output] {projection(item,&mut values);}
        }
        values
    }
    fn map_adapters<F:FnMut(&mut Linear<B>)>(&mut self,apply:&mut F) {
        if let Self::Adapted(block) = self {
            let attention = &mut block.attention.local;
            for item in [&mut attention.query,&mut attention.key,&mut attention.value,&mut attention.output] {map_projection(item,apply);}
        }
    }
}

impl<B:Backend> AdapterTree<B> for TensorParallelAdaptedEncoderDecoderStack<B> {
    fn adapters(&self) -> Vec<&Linear<B>> {
        let mut values = Vec::new();for layer in &self.layers {values.extend(layer.backbone.adapters());values.extend(layer.cross_attention.adapters());}values
    }
    fn map_adapters<F:FnMut(&mut Linear<B>)>(&mut self,apply:&mut F) {
        for layer in &mut self.layers {layer.backbone.map_adapters(apply);layer.cross_attention.map_adapters(apply);}
    }
}

impl<B:Backend> AdapterTree<B> for TensorParallelOutputHead<B> {
    fn adapters(&self) -> Vec<&Linear<B>> {
        match self {
            Self::AdaptedLinear(head)=>alloc::vec![&head.local.projection.adapter_a,&head.local.projection.adapter_b],
            Self::AdaptedVocabulary(head)=>alloc::vec![&head.projection.adapter_a,&head.projection.adapter_b],
            _=>Vec::new(),
        }
    }
    fn map_adapters<F:FnMut(&mut Linear<B>)>(&mut self,apply:&mut F) {
        match self {
            Self::AdaptedLinear(head)=>{apply(&mut head.local.projection.adapter_a);apply(&mut head.local.projection.adapter_b);},
            Self::AdaptedVocabulary(head)=>{apply(&mut head.projection.adapter_a);apply(&mut head.projection.adapter_b);},
            _=>{},
        }
    }
}

impl<B:Backend> AdapterTree<B> for TensorParallelTransformerModel<B> {
    fn adapters(&self) -> Vec<&Linear<B>> {let mut values = self.backbone.adapters();values.extend(self.head.adapters());values}
    fn map_adapters<F:FnMut(&mut Linear<B>)>(&mut self,apply:&mut F) {self.backbone.map_adapters(apply);self.head.map_adapters(apply);}
}
impl<B:Backend> AdapterTree<B> for TensorParallelEncoderDecoderModel<B> {
    fn adapters(&self) -> Vec<&Linear<B>> {let mut values = self.encoder.adapters();values.extend(self.decoder.adapters());values.extend(self.head.adapters());values}
    fn map_adapters<F:FnMut(&mut Linear<B>)>(&mut self,apply:&mut F) {self.encoder.map_adapters(apply);self.decoder.map_adapters(apply);self.head.map_adapters(apply);}
}

fn invalid(reason:&str) -> RecorderError {RecorderError::Unknown(format!("Invalid parallel model adapter record: {reason}"))}

pub(super) fn capture<B:Backend,M:AdapterTree<B>>(model:&M) -> Result<Vec<usize>,RecorderError> {
    let mut seen:BTreeMap<(ParamId,bool),(usize,Tensor<B,2>)> = BTreeMap::new();let mut aliases = Vec::new();
    for adapter in model.adapters() {
        let value = adapter.weight.val();let key = (adapter.weight.id,value.is_require_grad());let index = aliases.len();
        if let Some((canonical,previous)) = seen.get(&key) {
            if previous.dims() != value.dims() || previous.dtype() != value.dtype() || previous.device() != value.device() {
                return Err(invalid("shared adapter IDs have incompatible geometry/storage/device"));
            }
            aliases.push(*canonical);
        } else {seen.insert(key,(index,value));aliases.push(index);}
    }
    Ok(aliases)
}

pub(super) fn rejoin<B:Backend,M:AdapterTree<B>>(mut model:M,expected:&[usize]) -> Result<M,RecorderError> {
    if capture(&model)?.as_slice() != expected {return Err(invalid("restored adapter identities no longer match original parameter sharing"));}
    let values:Vec<_> = model.adapters().iter().map(|adapter|adapter.weight.val()).collect();let mut index = 0usize;
    model.map_adapters(&mut |adapter| {let value = values[expected[index]].clone();adapter.weight = adapter.weight.clone().map(|_|value);index += 1;});
    Ok(model)
}