ruda-nn 0.21.46

Ruda neural network layers, activation modules and losses.
use ruda_model::tensor::{Bool,Tensor,backend::Backend};
use crate::{Dropout,attention::{DenseAttentionMask,DenseAttentionOptions,PackedSequenceLayout,PackedAttentionOptions,PackedDocumentAttentionMask},cache::ProjectedKvCache};
use super::{ProjectedGroupedQueryAttention,TransformerProjection,DenseTransformerNorm};
use super::dense::try_residual_branch;

pub(super) fn attention_branch<B:Backend,P:TransformerProjection<B>,F>(attention:&ProjectedGroupedQueryAttention<B,P>,norm:&DenseTransformerNorm<B>,dropout:&Dropout,
    norm_first:bool,input:Tensor<B,3>,masks:DenseAttentionMask<B>,options:DenseAttentionOptions,positions:F) -> Result<Tensor<B,3>,P::Error>
    where F:FnOnce(Tensor<B,4>,Tensor<B,4>)->(Tensor<B,4>,Tensor<B,4>) {
    try_residual_branch(input,norm,dropout,norm_first,|source| {
        let (query,key,value)=attention.project(source.clone(),source.clone(),source)?;
        let shape=(query.dims(),key.dims());let (query,key)=positions(query,key);
        assert_eq!((query.dims(),key.dims()),shape,"positions changed head geometry");attention.forward_projected(query,key,value,masks,options)
    })
}
pub(super) fn packed_attention_branch<B:Backend,P:TransformerProjection<B>,F>(attention:&ProjectedGroupedQueryAttention<B,P>,norm:&DenseTransformerNorm<B>,dropout:&Dropout,
    norm_first:bool,input:Tensor<B,2>,layout:&PackedSequenceLayout,masks:Option<&[PackedDocumentAttentionMask<B>]>,options:PackedAttentionOptions,positions:F)
    -> Result<Tensor<B,2>,P::Error> where F:FnOnce(Tensor<B,3>,Tensor<B,3>)->(Tensor<B,3>,Tensor<B,3>) {
    assert_eq!(input.dims()[0],layout.tokens(),"packed document boundaries differ from actual rows");
    try_residual_branch(input,norm,dropout,norm_first,|source| {
        let (query,key,value)=attention.project_packed(source.clone(),source.clone(),source)?;
        let shape=(query.dims(),key.dims());let (query,key)=positions(query,key);
        assert_eq!((query.dims(),key.dims()),shape,"packed positions changed geometry");
        if let Some(masks)=masks {attention.forward_packed_masked_projected(query,key,value,layout,layout,masks,options)}
        else {attention.forward_packed_projected(query,key,value,layout,layout,options)}
    })
}
pub(super) fn cached_attention_branch<B:Backend,P:TransformerProjection<B>,F>(attention:&ProjectedGroupedQueryAttention<B,P>,norm:&DenseTransformerNorm<B>,dropout:&Dropout,
    norm_first:bool,input:Tensor<B,3>,visible:Option<Tensor<B,2,Bool>>,cache:&mut ProjectedKvCache<B>,masks:DenseAttentionMask<B>,options:DenseAttentionOptions,positions:F)
    -> Result<Tensor<B,3>,P::Error> where F:FnOnce(Tensor<B,4>,Tensor<B,4>,usize)->(Tensor<B,4>,Tensor<B,4>) {
    try_residual_branch(input,norm,dropout,norm_first,|source|attention.forward_cached_with_positions(source,visible,cache,masks,options,positions))
}