Skip to main content

luma_tensor/ops/
nn.rs

1use crate::{Device, Dim, Float, FloatMeta, Storage, Tensor};
2
3impl<D: Device> Tensor<D, Float> {
4    /// Softmax over `dim`. Records `Op::Softmax`.
5    pub fn softmax<Dm: Dim>(&self, dim: Dm) -> crate::Result<Self> {
6        let dim = dim.to_index(self.shape(), "softmax")?;
7        let storage = D::f_softmax(&*self.storage_read()?, self.layout(), dim)?;
8        let meta = FloatMeta::on_softmax(self, dim);
9        assert_eq!(self.dtype(), storage.dtype());
10        Ok(Self::from_storage(storage, self.shape().clone(), meta))
11    }
12
13    /// RMSNorm over the last dim: `x / sqrt(mean(x^2)+eps) * weight`.
14    /// Records `Op::RmsNorm`.
15    pub fn rms_norm(&self, weight: &Self, eps: f64) -> crate::Result<Self> {
16        let storage = D::f_rms_norm(&*self.storage_read()?, self.layout(), &*weight.storage_read()?, weight.layout(), eps)?;
17        let meta = FloatMeta::on_rms_norm(self, weight, eps);
18        assert_eq!(self.dtype(), storage.dtype());
19        Ok(Self::from_storage(storage, self.shape().clone(), meta))
20    }
21}