use alloc::vec::Vec;
use core::ops::Range;
use ruda_autodiff::{Autodiff,checkpoint::strategy::CheckpointStrategy,tensor_parallel as region};
use ruda_model::tensor::{Bool,DType,Int,IntDType,Tensor,TensorPrimitive,ElementConversion,backend::Backend};
use region::BroadcastTensorCollective;
use crate::loss::LossTerms;
#[derive(Clone,Debug,PartialEq,Eq)]
pub struct VocabParallelLossLayout {
offsets: Vec<usize>,
vocabulary_size: usize,
}
impl VocabParallelLossLayout {
pub fn new(widths: Vec<usize>,vocabulary_size: usize) -> Self {
assert!(!widths.is_empty() && vocabulary_size > 0,"parallel loss needs ranks and real classes");
let mut offsets = Vec::with_capacity(widths.len()+1);
offsets.push(0usize);
for width in widths {
assert!(width > 0,"parallel loss storage intervals must be nonempty");
let next = offsets.last().copied().unwrap().checked_add(width).expect("parallel vocabulary storage overflow");
assert!(next <= i64::MAX as usize,"parallel vocabulary storage exceeds signed index range");
offsets.push(next);
}
assert!(vocabulary_size <= *offsets.last().unwrap(),"logical vocabulary exceeds actual shard storage");
Self {offsets,vocabulary_size}
}
pub fn world_size(&self) -> usize {self.offsets.len()-1}
pub fn vocabulary_size(&self) -> usize {self.vocabulary_size}
pub fn storage_size(&self) -> usize {*self.offsets.last().unwrap()}
pub fn interval(&self,rank: usize) -> Range<usize> {
assert!(rank < self.world_size(),"parallel vocabulary rank is out of range");
self.offsets[rank]..self.offsets[rank+1]
}
}
#[derive(Clone,Debug)]
pub struct VocabParallelCrossEntropy {
layout: VocabParallelLossLayout,
label_smoothing: f64,
ignore_index: Option<i64>,
}
struct Statistics<B: Backend,S: CheckpointStrategy> {
shifted: Tensor<Autodiff<B,S>,2>,
log_sum: Tensor<Autodiff<B,S>,1>,
valid: Tensor<Autodiff<B,S>,1,Bool>,
weights: Option<Tensor<Autodiff<B,S>,2>>,
}
fn floating(dtype: DType) {
assert!(matches!(dtype,DType::F16|DType::BF16|DType::F32|DType::Flex32|DType::F64),
"vocabulary loss requires native floating operands");
}
fn product<B: Backend,S: CheckpointStrategy>(value: Tensor<Autodiff<B,S>,2>,coefficient: Tensor<Autodiff<B,S>,2>)
-> Tensor<Autodiff<B,S>,2> {
let zero_infinity = coefficient.clone().equal_elem(0).bool_and(value.clone().is_inf());
value.mask_fill(zero_infinity,0)*coefficient
}
impl VocabParallelCrossEntropy {
pub fn new(layout: VocabParallelLossLayout,label_smoothing: f64,ignore_index: Option<i64>) -> Self {
assert!(label_smoothing.is_finite() && (0.0..=1.0).contains(&label_smoothing),"invalid full-vocabulary smoothing");
Self {layout,label_smoothing,ignore_index}
}
pub fn layout(&self) -> &VocabParallelLossLayout {&self.layout}
fn visibility<B: Backend,S: CheckpointStrategy>(&self,logits: &Tensor<Autodiff<B,S>,2>,
visible: Option<Tensor<Autodiff<B,S>,1,Bool>>) -> Tensor<Autodiff<B,S>,1,Bool> {
if let Some(visible) = visible {
assert_eq!(visible.dims(),[logits.dims()[0]],"parallel loss visibility and actual rows differ");
assert_eq!(visible.device(),logits.device(),"parallel loss visibility must share the device");
visible
} else {Tensor::<Autodiff<B,S>,1,Bool>::zeros([logits.dims()[0]],&logits.device()).bool_not()}
}
fn statistics<B,S,C>(&self,logits: Tensor<Autodiff<B,S>,2>,valid: Tensor<Autodiff<B,S>,1,Bool>,
class_weights: Option<Tensor<Autodiff<B,S>,1>>,target_dtype: Option<DType>,communicator: &C)
-> Result<Statistics<B,S>,C::Error>
where B: Backend,S: CheckpointStrategy,C: BroadcastTensorCollective<B> {
floating(logits.dtype());
let [rows,width] = logits.dims();
assert_eq!(communicator.world_size() as usize,self.layout.world_size(),"vocabulary layout and communicator size differ");
let interval = self.layout.interval(communicator.rank() as usize);
assert_eq!(interval.len(),width,"local logits do not match this rank's actual vocabulary storage");
let mut double = logits.dtype() == DType::F64 || target_dtype == Some(DType::F64);
if let Some(weights) = &class_weights {
floating(weights.dtype());
assert_eq!(weights.dims(),[width],"local class weights and actual vocabulary columns differ");
assert_eq!(weights.device(),logits.device(),"parallel class weights must share the logits device");
double |= weights.dtype() == DType::F64;
}
let flag = Tensor::<B,1>::full([1],if double {1.} else {0.},(&logits.device(),DType::F32));
let flag = communicator.all_reduce_sum(flag.into_primitive().tensor())?;
let double = Tensor::<B,1>::from_primitive(TensorPrimitive::Float(flag)).greater_elem(0).any().into_scalar().elem::<bool>();
let dtype = if double {DType::F64} else {DType::F32};
let logits = logits.cast(dtype);
let real = Tensor::<Autodiff<B,S>,1,Int>::arange(interval.start as i64..interval.end as i64,(&logits.device(),DType::I64))
.lower_elem(self.layout.vocabulary_size as i64).reshape([1,width]).expand([rows,width]);
let selected = valid.clone().reshape([rows,1]).expand([rows,width]).bool_and(real.clone());
let weights = class_weights.map(|weights|weights.cast(dtype).reshape([1,width]).expand([rows,width]).mask_fill(selected.clone().bool_not(),0));
if rows == 0 {
let log_sum = logits.clone().sum_dim(1).reshape([0]);
return Ok(Statistics {shifted:logits,log_sum,valid,weights});
}
let logits = logits.mask_fill(valid.clone().bool_not().reshape([rows,1]).expand([rows,width]),0)
.mask_fill(real.clone().bool_not(),f64::NEG_INFINITY);
let local_max = logits.clone().detach().max_dim(1).inner();
let maxima = communicator.all_gather_float(local_max.into_primitive().tensor())?;
let maximum = Tensor::<B,2>::from_primitive(TensorPrimitive::Float(maxima))
.reshape([self.layout.world_size(),rows]).max_dim(0).reshape([rows,1]);
let shifted = logits-Tensor::<Autodiff<B,S>,2>::from_inner(maximum);
let sums = region::reduce_from_region(shifted.clone().exp().sum_dim(1),communicator.clone())?;
let log_sum = sums.log().reshape([rows]);
Ok(Statistics {shifted:shifted.mask_fill(selected.bool_not(),0),log_sum,valid,weights})
}
pub fn forward_terms<B,S,C>(&self,logits: Tensor<Autodiff<B,S>,2>,labels: Tensor<Autodiff<B,S>,1,Int>,communicator: C,
visible: Option<Tensor<Autodiff<B,S>,1,Bool>>,class_weights: Option<Tensor<Autodiff<B,S>,1>>)
-> Result<LossTerms<Autodiff<B,S>>,C::Error>
where B: Backend,S: CheckpointStrategy,C: BroadcastTensorCollective<B> {
assert_eq!(labels.dims(),[logits.dims()[0]],"parallel loss labels and actual rows differ");
assert_eq!(labels.device(),logits.device(),"parallel loss labels must share the device");
let [rows,width] = logits.dims();
let labels = labels.cast(IntDType::I64);
let mut valid = self.visibility(&logits,visible);
if let Some(ignore) = self.ignore_index {valid = valid.bool_and(labels.clone().equal_elem(ignore).bool_not());}
let statistics = self.statistics(logits,valid,class_weights,None,&communicator)?;
let Statistics {shifted,log_sum,valid,weights} = statistics;
let dtype = log_sum.dtype();
if rows == 0 {return Ok(LossTerms {values:log_sum,normalizers:Tensor::zeros([0],(&labels.device(),dtype)),valid});}
let invalid = labels.clone().lower_elem(0).bool_or(labels.clone().greater_equal_elem(self.layout.vocabulary_size as i64)).bool_and(valid.clone());
let invalid = invalid.any().float().cast(DType::F32).inner();
let invalid = communicator.all_reduce_sum(invalid.into_primitive().tensor())?;
assert!(!Tensor::<B,1>::from_primitive(TensorPrimitive::Float(invalid)).greater_elem(0).any().into_scalar().elem::<bool>(),
"selected target lies outside the logical vocabulary");
let interval = self.layout.interval(communicator.rank() as usize);
let owned = labels.clone().greater_equal_elem(interval.start as i64)
.bool_and(labels.clone().lower_elem(interval.end.min(self.layout.vocabulary_size) as i64)).bool_and(valid.clone());
let indices = labels.sub_scalar(interval.start as i64).mask_fill(owned.clone().bool_not(),0).reshape([rows,1]);
let normalizers = if let Some(weights) = &weights {
let selected = weights.clone().gather(1,indices.clone()).reshape([rows]).mask_fill(owned.clone().bool_not(),0);
region::reduce_from_region(selected,communicator.clone())?
} else {valid.clone().float().cast(dtype)};
let selected = shifted.clone().gather(1,indices).mask_fill(owned.bool_not().reshape([rows,1]),0);
let selected = region::reduce_from_region(selected,communicator.clone())?.reshape([rows]);
let mut values = if self.label_smoothing == 1. {log_sum.zeros_like()} else {
let loss = (log_sum.clone()-selected).reshape([rows,1]);
product(loss,normalizers.clone().reshape([rows,1])).reshape([rows]).mul_scalar(1.-self.label_smoothing)
};
if self.label_smoothing > 0. {
let (mass,weighted) = if let Some(weights) = weights {
let mass = region::reduce_from_region(weights.clone().sum_dim(1),communicator.clone())?.reshape([rows]);
let weighted = region::reduce_from_region(product(shifted,weights).sum_dim(1),communicator)?.reshape([rows]);
(mass,weighted)
} else {
let mass = log_sum.full_like(self.layout.vocabulary_size as f64);
let weighted = region::reduce_from_region(shifted.sum_dim(1),communicator)?.reshape([rows]);
(mass,weighted)
};
let uniform = log_sum*mass-weighted;
values = values+uniform.mul_scalar(self.label_smoothing/self.layout.vocabulary_size as f64);
}
Ok(LossTerms {values:values.mask_fill(valid.clone().bool_not(),0),normalizers,valid})
}
pub fn forward_soft_terms<B,S,C>(&self,logits: Tensor<Autodiff<B,S>,2>,targets: Tensor<Autodiff<B,S>,2>,communicator: C,
visible: Option<Tensor<Autodiff<B,S>,1,Bool>>,class_weights: Option<Tensor<Autodiff<B,S>,1>>)
-> Result<LossTerms<Autodiff<B,S>>,C::Error>
where B: Backend,S: CheckpointStrategy,C: BroadcastTensorCollective<B> {
assert_eq!(logits.dims(),targets.dims(),"parallel soft labels and actual local logits differ");
assert_eq!(logits.device(),targets.device(),"parallel soft labels must share the device");
floating(targets.dtype());
let [rows,width] = logits.dims();
let valid = self.visibility(&logits,visible);
let Statistics {shifted,log_sum,valid,weights} = self.statistics(logits,valid,class_weights,Some(targets.dtype()),&communicator)?;
let dtype = log_sum.dtype();
if rows == 0 {return Ok(LossTerms {values:log_sum+targets.cast(dtype).sum_dim(1).reshape([0]),
normalizers:Tensor::zeros([0],(&valid.device(),dtype)),valid});}
let interval = self.layout.interval(communicator.rank() as usize);
let real = Tensor::<Autodiff<B,S>,1,Int>::arange(interval.start as i64..interval.end as i64,(&targets.device(),DType::I64))
.lower_elem(self.layout.vocabulary_size as i64).reshape([1,width]).expand([rows,width]);
let selected = valid.clone().reshape([rows,1]).expand([rows,width]).bool_and(real);
let targets = targets.cast(dtype).mask_fill(selected.clone().bool_not(),0);
let targets = targets.mul_scalar(1.-self.label_smoothing).add_scalar(self.label_smoothing/self.layout.vocabulary_size as f64)
.mask_fill(selected.bool_not(),0);
let coefficients = if let Some(weights) = weights {targets*weights} else {targets};
let mass = region::reduce_from_region(coefficients.clone().sum_dim(1),communicator.clone())?.reshape([rows]);
let weighted = region::reduce_from_region(product(shifted,coefficients).sum_dim(1),communicator)?.reshape([rows]);
let values = (product(log_sum.reshape([rows,1]),mass.reshape([rows,1])).reshape([rows])-weighted).mask_fill(valid.clone().bool_not(),0);
Ok(LossTerms {values,normalizers:valid.clone().float().cast(dtype),valid})
}
pub fn forward_token_terms<B,S,C>(&self,logits: Tensor<Autodiff<B,S>,3>,labels: Tensor<Autodiff<B,S>,2,Int>,communicator: C,
visible: Option<Tensor<Autodiff<B,S>,2,Bool>>,class_weights: Option<Tensor<Autodiff<B,S>,1>>)
-> Result<LossTerms<Autodiff<B,S>,2>,C::Error>
where B: Backend,S: CheckpointStrategy,C: BroadcastTensorCollective<B> {
let [batch,tokens,width] = logits.dims();
assert_eq!(labels.dims(),[batch,tokens],"parallel token targets differ from actual axes");
let rows = batch.checked_mul(tokens).expect("parallel token row count overflow");
let visible = visible.map(|mask| {assert_eq!(mask.dims(),[batch,tokens],"parallel token visibility differs");mask.reshape([rows])});
self.forward_terms(logits.reshape([rows,width]),labels.reshape([rows]),communicator,visible,class_weights)
.map(|terms|terms.reshape([batch,tokens]))
}
pub fn forward_soft_token_terms<B,S,C>(&self,logits: Tensor<Autodiff<B,S>,3>,targets: Tensor<Autodiff<B,S>,3>,communicator: C,
visible: Option<Tensor<Autodiff<B,S>,2,Bool>>,class_weights: Option<Tensor<Autodiff<B,S>,1>>)
-> Result<LossTerms<Autodiff<B,S>,2>,C::Error>
where B: Backend,S: CheckpointStrategy,C: BroadcastTensorCollective<B> {
let [batch,tokens,width] = logits.dims();
assert_eq!(targets.dims(),[batch,tokens,width],"parallel soft token targets differ from actual axes");
let rows = batch.checked_mul(tokens).expect("parallel soft token row count overflow");
let visible = visible.map(|mask| {assert_eq!(mask.dims(),[batch,tokens],"parallel soft token visibility differs");mask.reshape([rows])});
self.forward_soft_terms(logits.reshape([rows,width]),targets.reshape([rows,width]),communicator,visible,class_weights)
.map(|terms|terms.reshape([batch,tokens]))
}
}