use super::*;
use crate::{attention::{DenseAttentionMask,DenseAttentionOptions,PackedSequenceLayout},cache::TransformerKvCache,
loss::{CausalCrossEntropyConfig,CausalLoss}};
use ruda_model::tensor::Bool;
pub type FullyShardedGreedySelection<B> = crate::tensor_parallel::VocabParallelGreedySelection<B>;
pub type FullyShardedTopKSelection<B> = crate::tensor_parallel::VocabParallelTopKSelection<B>;
fn last_rows<B:Backend>(hidden:Tensor<B,3>) -> Tensor<B,2> {
let [batch,tokens,width]=hidden.dims();assert!(tokens>0,"native sharded last-row projection needs an actual token");
hidden.slice([0..batch,tokens-1..tokens,0..width]).reshape([batch,width])
}
impl<B:Backend> FullyShardedTransformerHead<B> {
pub fn forward_last_inference<C:BroadcastTensorCollective<B>>(&self,hidden:Tensor<B,3>,communicator:C)
-> Result<Tensor<B,2>,C::Error> {self.forward_inference(last_rows(hidden),communicator)}
pub fn forward_topk_inference<C:BroadcastTensorCollective<B>>(&self,hidden:Tensor<B,2>,communicator:C,k:usize,
visible:Option<Tensor<B,1,Bool>>) -> Result<FullyShardedTopKSelection<B>,C::Error> {
assert!(k<=self.classes(),"sharded head top-k exceeds actual output classes");
let logits=self.forward_inference(hidden,communicator)?;Ok(crate::tensor_parallel::full_logits_topk(logits,k,visible))
}
pub fn forward_topk_last_inference<C:BroadcastTensorCollective<B>>(&self,hidden:Tensor<B,3>,communicator:C,k:usize,
visible:Option<Tensor<B,1,Bool>>) -> Result<FullyShardedTopKSelection<B>,C::Error> {
self.forward_topk_inference(last_rows(hidden),communicator,k,visible)
}
pub fn forward_greedy_inference<C:BroadcastTensorCollective<B>>(&self,hidden:Tensor<B,2>,communicator:C,
visible:Option<Tensor<B,1,Bool>>) -> Result<FullyShardedGreedySelection<B>,C::Error> {
let rows=hidden.dims()[0];let selected=self.forward_topk_inference(hidden,communicator,1,visible)?;
Ok(FullyShardedGreedySelection {indices:selected.indices.reshape([rows]),valid:selected.valid.reshape([rows])})
}
pub fn forward_greedy_last_inference<C:BroadcastTensorCollective<B>>(&self,hidden:Tensor<B,3>,communicator:C,
visible:Option<Tensor<B,1,Bool>>) -> Result<FullyShardedGreedySelection<B>,C::Error> {
self.forward_greedy_inference(last_rows(hidden),communicator,visible)
}
pub fn forward_causal_loss_inference<C:BroadcastTensorCollective<B>>(&self,hidden:Tensor<B,3>,labels:Tensor<B,2,Int>,
criterion:&CausalCrossEntropyConfig,label_smoothing:f64,communicator:C) -> Result<CausalLoss<B>,C::Error> {
let head=self.gather_inference(communicator)?;
Ok(criterion.forward_hidden_with_smoothing(hidden,labels,|rows|head.forward(rows),label_smoothing))
}
pub fn forward_packed_causal_loss_inference<C:BroadcastTensorCollective<B>>(&self,hidden:Tensor<B,2>,labels:Tensor<B,1,Int>,
layout:&PackedSequenceLayout,criterion:&CausalCrossEntropyConfig,label_smoothing:f64,communicator:C) -> Result<CausalLoss<B>,C::Error> {
let head=self.gather_inference(communicator)?;
Ok(criterion.forward_packed_hidden_with_smoothing(hidden,labels,layout,|rows|head.forward(rows),label_smoothing))
}
}
impl<B:Backend> FullyShardedTransformerModel<B> {
pub fn forward_topk_last_inference_with<C,F>(&self,input:FullyShardedTransformerInput<B>,communicator:C,layer:F,
k:usize,visible:Option<Tensor<B,1,Bool>>) -> Result<FullyShardedTopKSelection<B>,C::Error>
where C:BroadcastTensorCollective<B>,F:FnMut(usize,&FullyShardedTransformerBlock<B>,Tensor<B,3>)->Result<Tensor<B,3>,C::Error> {
assert!(input.tokens.dims()[1]>0,"native sharded prompt candidates need actual input tokens");
let hidden=self.forward_hidden_inference_with(input,communicator.clone(),layer)?;
self.head.forward_topk_last_inference(hidden,communicator,k,visible)
}
pub fn forward_cached_topk_inference<C,P>(&self,input:FullyShardedTransformerInput<B>,token_visibility:Option<Tensor<B,2,Bool>>,
cache:&mut TransformerKvCache<B>,masks:DenseAttentionMask<B>,options:DenseAttentionOptions,communicator:C,positions:P,
k:usize,row_visibility:Option<Tensor<B,1,Bool>>) -> Result<FullyShardedTopKSelection<B>,C::Error>
where C:BroadcastTensorCollective<B>,P:FnMut(usize,Tensor<B,4>,Tensor<B,4>,usize)->(Tensor<B,4>,Tensor<B,4>) {
assert!(input.tokens.dims()[1]>0,"native sharded cached candidates need actual new tokens");
let hidden=self.forward_cached_hidden_inference(input,token_visibility,cache,masks,options,communicator.clone(),positions)?;
self.head.forward_topk_last_inference(hidden,communicator,k,row_visibility)
}
pub fn forward_causal_loss_inference_with<C,F>(&self,input:FullyShardedTransformerInput<B>,labels:Tensor<B,2,Int>,
criterion:&CausalCrossEntropyConfig,label_smoothing:f64,communicator:C,layer:F) -> Result<CausalLoss<B>,C::Error>
where C:BroadcastTensorCollective<B>,F:FnMut(usize,&FullyShardedTransformerBlock<B>,Tensor<B,3>)->Result<Tensor<B,3>,C::Error> {
assert_eq!(input.tokens.dims(),labels.dims(),"sharded evaluation token/label geometry differs");
let hidden=self.forward_hidden_inference_with(input,communicator.clone(),layer)?;
self.head.forward_causal_loss_inference(hidden,labels,criterion,label_smoothing,communicator)
}
pub fn forward_packed_causal_loss_inference_with<C,F>(&self,input:FullyShardedTransformerInput<B,1>,labels:Tensor<B,1,Int>,
layout:&PackedSequenceLayout,criterion:&CausalCrossEntropyConfig,label_smoothing:f64,communicator:C,layer:F) -> Result<CausalLoss<B>,C::Error>
where C:BroadcastTensorCollective<B>,F:FnMut(usize,&FullyShardedTransformerBlock<B>,Tensor<B,2>)->Result<Tensor<B,2>,C::Error> {
assert_eq!(input.tokens.dims(),labels.dims(),"sharded packed evaluation token/label geometry differs");
let hidden=self.forward_packed_hidden_inference_with(input,layout,communicator.clone(),layer)?;
self.head.forward_packed_causal_loss_inference(hidden,labels,layout,criterion,label_smoothing,communicator)
}
}
impl<B:Backend> FullyShardedEncoderDecoderModel<B> {
pub fn forward_topk_last_inference_with<C,E,F>(&self,source:FullyShardedTransformerInput<B>,target:FullyShardedTransformerInput<B>,
communicator:C,encoder:E,decoder:F,k:usize,visible:Option<Tensor<B,1,Bool>>) -> Result<FullyShardedTopKSelection<B>,C::Error>
where C:BroadcastTensorCollective<B>,E:FnMut(usize,&FullyShardedTransformerBlock<B>,Tensor<B,3>)->Result<Tensor<B,3>,C::Error>,
F:FnMut(usize,&FullyShardedEncoderDecoderLayer<B>,Tensor<B,3>,Tensor<B,3>)->Result<Tensor<B,3>,C::Error> {
assert!(target.tokens.dims()[1]>0,"native sharded paired candidates need actual target tokens");
let hidden=self.forward_hidden_inference_with(source,target,communicator.clone(),encoder,decoder)?;
self.head.forward_topk_last_inference(hidden,communicator,k,visible)
}
pub fn forward_causal_loss_inference_with<C,E,F>(&self,source:FullyShardedTransformerInput<B>,target:FullyShardedTransformerInput<B>,
labels:Tensor<B,2,Int>,criterion:&CausalCrossEntropyConfig,label_smoothing:f64,communicator:C,encoder:E,decoder:F)
-> Result<CausalLoss<B>,C::Error>
where C:BroadcastTensorCollective<B>,E:FnMut(usize,&FullyShardedTransformerBlock<B>,Tensor<B,3>)->Result<Tensor<B,3>,C::Error>,
F:FnMut(usize,&FullyShardedEncoderDecoderLayer<B>,Tensor<B,3>,Tensor<B,3>)->Result<Tensor<B,3>,C::Error> {
assert_eq!(target.tokens.dims(),labels.dims(),"sharded paired evaluation target/label geometry differs");
let hidden=self.forward_hidden_inference_with(source,target,communicator.clone(),encoder,decoder)?;
self.head.forward_causal_loss_inference(hidden,labels,criterion,label_smoothing,communicator)
}
pub fn forward_packed_causal_loss_inference_with<C,E,F>(&self,source:FullyShardedTransformerInput<B,1>,target:FullyShardedTransformerInput<B,1>,
source_layout:&PackedSequenceLayout,target_layout:&PackedSequenceLayout,labels:Tensor<B,1,Int>,criterion:&CausalCrossEntropyConfig,
label_smoothing:f64,communicator:C,encoder:E,decoder:F) -> Result<CausalLoss<B>,C::Error>
where C:BroadcastTensorCollective<B>,E:FnMut(usize,&FullyShardedTransformerBlock<B>,Tensor<B,2>)->Result<Tensor<B,2>,C::Error>,
F:FnMut(usize,&FullyShardedEncoderDecoderLayer<B>,Tensor<B,2>,Tensor<B,2>)->Result<Tensor<B,2>,C::Error> {
assert_eq!(target.tokens.dims(),labels.dims(),"sharded paired packed evaluation target/label geometry differs");
let hidden=self.forward_packed_hidden_inference_with(source,target,source_layout,target_layout,communicator.clone(),encoder,decoder)?;
self.head.forward_packed_causal_loss_inference(hidden,labels,target_layout,criterion,label_smoothing,communicator)
}
}