use super::*;
use crate::loss::{CausalCrossEntropyConfig,LossTerms};
use crate::attention::{PackedSequenceLayout,DenseAttentionMask,DenseAttentionOptions,PackedAttentionOptions};
use ruda_autodiff::collective::{CollectiveScope,ScopedTensorCollective,ScopedCollectiveError};
use ruda_model::tensor::{IntDType,TensorData,ElementConversion};
mod paired;
mod weighted;
pub use weighted::*;
pub struct FullyShardedLoss<B:Backend,S:CheckpointStrategy> {
pub loss_sum:Tensor<Autodiff<B,S>,1>,
pub local_count:Tensor<B,1,Int>,
pub global_count:u64,
pub global_loss_sum:Tensor<B,1>,
}
impl<B:Backend,S:CheckpointStrategy> FullyShardedLoss<B,S> {
pub fn mean(&self) -> Tensor<Autodiff<B,S>,1> {self.normalized(self.global_count)}
pub fn normalized(&self,global_window_count:u64) -> Tensor<Autodiff<B,S>,1> {
self.loss_sum.clone()/global_window_count.max(1) as f64
}
pub fn global_mean(&self) -> Tensor<B,1> {self.global_loss_sum.clone()/self.global_count.max(1) as f64}
}
pub fn complete_fully_sharded_loss<B,S,C>(scope:&CollectiveScope<B,S>,loss_sum:Tensor<Autodiff<B,S>,1>,
local_count:Tensor<Autodiff<B,S>,1,Int>,communicator:C) -> Result<FullyShardedLoss<B,S>,ScopedCollectiveError<C::Error>>
where B:Backend,S:CheckpointStrategy,C:BroadcastTensorCollective<B> {
let loss_sum=scope.complete(loss_sum,communicator.clone())?;
if local_count.dims()!=[1] || local_count.device()!=loss_sum.device() {return Err(ScopedCollectiveError::Protocol("actual loss count shape/device differs"));}
let local_count=local_count.inner().cast(IntDType::I64);let device=loss_sum.device();
let mask=Tensor::<B,1,Int>::from_data(TensorData::new(alloc::vec![65535_i64],[1]),(&device,DType::I64));
let mut parts=Vec::with_capacity(4);
for word in 0..4 {
let part=local_count.clone().bitwise_right_shift_scalar(((word*16) as i32).elem()).bitwise_and(mask.clone());
parts.push(part.cast(FloatDType::F32));
}
let words=Tensor::<B,1>::cat(parts,0);
let words=if communicator.world_size()==1 {words} else {Tensor::from_primitive(TensorPrimitive::Float(
communicator.all_gather_float(words.into_primitive().tensor()).map_err(ScopedCollectiveError::Collective)?))};
let expected=(communicator.world_size() as usize).checked_mul(4).ok_or(ScopedCollectiveError::Protocol("count gather size overflows"))?;
if words.dims()!=[expected] || words.dtype()!=DType::F32 || words.device()!=device {return Err(ScopedCollectiveError::Protocol("count transport changed original shape/storage/device"));}
let words=words.into_data().to_vec::<f32>().map_err(|_|ScopedCollectiveError::Protocol("count metadata cannot be decoded"))?;
let mut global_count=0u64;
for rank in 0..communicator.world_size() as usize {
let mut count=0u64;
for word in 0..4 {
let part=words[rank*4+word];
if !part.is_finite() || part<0.0 || part>65535.0 || part as u32 as f32!=part {return Err(ScopedCollectiveError::Protocol("invalid exact count word"));}
count|=(part as u64)<<(word*16);
}
if count>i64::MAX as u64 {return Err(ScopedCollectiveError::Protocol("a native effective count is negative or overflowed"));}
global_count=global_count.checked_add(count).ok_or(ScopedCollectiveError::Protocol("global effective count overflows u64"))?;
}
let work=if loss_sum.dtype()==DType::F64 {DType::F64} else {DType::F32};
let metric=loss_sum.clone().inner().cast(work);
let global_loss_sum=if communicator.world_size()==1 {metric} else {Tensor::<B,1>::from_primitive(TensorPrimitive::Float(
communicator.all_reduce_sum(metric.into_primitive().tensor()).map_err(ScopedCollectiveError::Collective)?))};
if global_loss_sum.dims()!=[1] || global_loss_sum.dtype()!=work || global_loss_sum.device()!=device {return Err(ScopedCollectiveError::Protocol("loss metric transport changed original shape/storage/device"));}
Ok(FullyShardedLoss {loss_sum,local_count,global_count,global_loss_sum})
}
impl<B:Backend,S:CheckpointStrategy> FullyShardedTransformerModel<Autodiff<B,S>> {
pub fn forward_causal_with<C,F>(&self,input:FullyShardedTransformerInput<Autodiff<B,S>>,labels:Tensor<Autodiff<B,S>,2,Int>,
criterion:&CausalCrossEntropyConfig,label_smoothing:f64,communicator:C,mut layer:F)
-> Result<FullyShardedLoss<B,S>,ScopedCollectiveError<C::Error>>
where C:BroadcastTensorCollective<B>,F:FnMut(usize,&FullyShardedTransformerBlock<Autodiff<B,S>>,Tensor<Autodiff<B,S>,3>,ScopedTensorCollective<C,B,S>)->Result<Tensor<Autodiff<B,S>,3>,C::Error> {
assert_eq!(input.tokens.dims(),labels.dims(),"actual sharded causal input/label geometry differs");
let scope=CollectiveScope::new();let transport=scope.bind(communicator.clone());
let hidden=self.forward_hidden_with(input,transport.clone(),|index,block,hidden|layer(index,block,hidden,transport.clone())).map_err(ScopedCollectiveError::Collective)?;
let head=self.head.gather(transport).map_err(ScopedCollectiveError::Collective)?;
let loss=criterion.forward_hidden_with_smoothing(hidden,labels,|rows|head.forward(rows),label_smoothing);
complete_fully_sharded_loss(&scope,loss.loss_sum,loss.valid_tokens,communicator)
}
pub fn forward_causal_with_positions<C,P>(&self,input:FullyShardedTransformerInput<Autodiff<B,S>>,labels:Tensor<Autodiff<B,S>,2,Int>,
masks:DenseAttentionMask<Autodiff<B,S>>,options:DenseAttentionOptions,criterion:&CausalCrossEntropyConfig,label_smoothing:f64,communicator:C,mut positions:P)
-> Result<FullyShardedLoss<B,S>,ScopedCollectiveError<C::Error>>
where C:BroadcastTensorCollective<B>,P:FnMut(usize,Tensor<Autodiff<B,S>,4>,Tensor<Autodiff<B,S>,4>)->(Tensor<Autodiff<B,S>,4>,Tensor<Autodiff<B,S>,4>) {
self.forward_causal_with(input,labels,criterion,label_smoothing,communicator,|index,block,hidden,transport|
block.forward(hidden,masks.clone(),options,transport,|query,key|positions(index,query,key)))
}
pub fn forward_packed_causal_with<C,F>(&self,input:FullyShardedTransformerInput<Autodiff<B,S>,1>,labels:Tensor<Autodiff<B,S>,1,Int>,
layout:&PackedSequenceLayout,criterion:&CausalCrossEntropyConfig,label_smoothing:f64,communicator:C,mut layer:F)
-> Result<FullyShardedLoss<B,S>,ScopedCollectiveError<C::Error>>
where C:BroadcastTensorCollective<B>,F:FnMut(usize,&FullyShardedTransformerBlock<Autodiff<B,S>>,Tensor<Autodiff<B,S>,2>,ScopedTensorCollective<C,B,S>)->Result<Tensor<Autodiff<B,S>,2>,C::Error> {
assert_eq!(input.tokens.dims(),labels.dims(),"actual sharded packed causal input/label geometry differs");
let scope=CollectiveScope::new();let transport=scope.bind(communicator.clone());
let hidden=self.forward_packed_hidden_with(input,layout,transport.clone(),|index,block,hidden|layer(index,block,hidden,transport.clone())).map_err(ScopedCollectiveError::Collective)?;
let head=self.head.gather(transport).map_err(ScopedCollectiveError::Collective)?;
let loss=criterion.forward_packed_hidden_with_smoothing(hidden,labels,layout,|rows|head.forward(rows),label_smoothing);
complete_fully_sharded_loss(&scope,loss.loss_sum,loss.valid_tokens,communicator)
}
pub fn forward_packed_causal_with_positions<C,P>(&self,input:FullyShardedTransformerInput<Autodiff<B,S>,1>,labels:Tensor<Autodiff<B,S>,1,Int>,
layout:&PackedSequenceLayout,options:PackedAttentionOptions,criterion:&CausalCrossEntropyConfig,label_smoothing:f64,communicator:C,mut positions:P)
-> Result<FullyShardedLoss<B,S>,ScopedCollectiveError<C::Error>>
where C:BroadcastTensorCollective<B>,P:FnMut(usize,Tensor<Autodiff<B,S>,3>,Tensor<Autodiff<B,S>,3>)->(Tensor<Autodiff<B,S>,3>,Tensor<Autodiff<B,S>,3>) {
self.forward_packed_causal_with(input,labels,layout,criterion,label_smoothing,communicator,|index,block,hidden,transport|
block.forward_packed(hidden,layout,options,transport,|query,key|positions(index,query,key)))
}
}