use alloc::vec;
use ruda_model::config::Config;
use ruda_model::module::{Content, DisplaySettings, Module, ModuleDisplay};
use ruda_model::tensor::{FloatDType, Int};
use ruda_model::tensor::Tensor;
use ruda_model::tensor::backend::Backend;
use core::ops::Range;
#[cfg(not(feature = "std"))]
#[allow(unused_imports)]
use num_traits::Float as _;
#[derive(Config, Debug)]
pub struct RotaryEncodingConfig {
pub max_sequence_length: usize,
pub d_model: usize,
#[config(default = "10000.0")]
pub theta: f32,
}
impl RotaryEncodingConfig {
pub fn init<B: Backend>(&self, device: &B::Device) -> RotaryEncoding<B> {
self.initialize(|x| x, device)
}
pub fn init_with_frequency_scaling<B: Backend>(
&self,
scaling: impl Fn(Tensor<B, 1>) -> Tensor<B, 1>,
device: &B::Device,
) -> RotaryEncoding<B> {
self.initialize(scaling, device)
}
fn initialize<B: Backend>(
&self,
scaling: impl Fn(Tensor<B, 1>) -> Tensor<B, 1>,
device: &B::Device,
) -> RotaryEncoding<B> {
assert_eq!(
self.d_model % 2,
0,
"The input embedding dimension must be even"
);
assert!(
self.theta > 0.0,
"Theta parameter must be positive (default: 10000)."
);
let exponent = Tensor::<B, 1, Int>::arange_step(0..self.d_model as i64, 2, device)
.float()
.div_scalar(self.d_model as f32);
let theta = exponent.mul_scalar(self.theta.ln()).exp().recip();
let theta = scaling(theta);
let freq_complex =
RotaryEncoding::compute_rotary_frequencies(0..self.max_sequence_length, theta.clone());
RotaryEncoding {
freq_complex,
theta,
start_offset: 0,
}
}
}
#[derive(Module, Debug)]
#[module(custom_display)]
pub struct RotaryEncoding<B: Backend> {
pub freq_complex: Tensor<B, 3>,
pub theta: Tensor<B, 1>,
start_offset: usize,
}
impl<B: Backend> ModuleDisplay for RotaryEncoding<B> {
fn custom_settings(&self) -> Option<DisplaySettings> {
DisplaySettings::new()
.with_new_line_after_attribute(false)
.optional()
}
fn custom_content(&self, content: Content) -> Option<Content> {
let [max_sequence_length, d_model, _] = self.freq_complex.shape().dims();
content
.add("d_model", &d_model)
.add("max_sequence_length", &max_sequence_length)
.optional()
}
}
#[allow(clippy::single_range_in_vec_init)]
impl<B: Backend> RotaryEncoding<B> {
pub fn forward<const D: usize>(&self, x: Tensor<B, D>) -> Tensor<B, D> {
self.apply(x, 0)
}
pub fn forward_with_compute_dtype<const D: usize>(
&self,
x: Tensor<B, D>,
dtype: FloatDType,
) -> Tensor<B, D> {
self.apply_with_compute_dtype(x, 0, dtype)
}
pub fn apply_with_compute_dtype<const D: usize>(
&self,
x: Tensor<B, D>,
start: usize,
dtype: FloatDType,
) -> Tensor<B, D> {
let storage_dtype = x.dtype();
let dtype: ruda_model::tensor::DType = dtype.into();
self.apply(x.cast(dtype), start).cast(storage_dtype)
}
pub fn apply<const D: usize>(&self, x: Tensor<B, D>, start: usize) -> Tensor<B, D> {
assert!(
D >= 2,
"Input tensor must have at least 2 dimensions for sequence length and hidden dimension"
);
let device = x.device();
let input_shape = x.shape();
let dtype = x.dtype();
let (seq_len, d_model) = (x.dims()[D - 2], x.dims()[D - 1]);
let dummy_dim_size = input_shape.num_elements() / (seq_len * d_model);
let sign_tensor =
Tensor::<B, 2>::from_floats([[1.0, 0.0, 0.0, 1.0], [0.0, -1.0, 1.0, 0.0]], &device)
.cast(dtype);
let out: Tensor<B, 4> = x
.reshape([dummy_dim_size, seq_len, d_model / 2, 2])
.matmul(sign_tensor.unsqueeze())
.reshape([dummy_dim_size, seq_len, d_model, 2])
* self
.freq_complex
.clone()
.slice([start..start + seq_len])
.cast(dtype)
.unsqueeze();
out.sum_dim(-1).reshape(input_shape)
}
pub fn shift(&mut self, start: usize) {
let max_seq_len = self.freq_complex.dims()[0];
assert!(
start > self.start_offset,
"Shift start position must be monotonically increasing"
);
let current_end = self.start_offset + max_seq_len;
if start >= current_end {
let new_freqs =
Self::compute_rotary_frequencies(start..start + max_seq_len, self.theta.clone());
self.freq_complex
.inplace(|freqs| freqs.slice_assign([0..max_seq_len], new_freqs));
} else {
let num_keep = current_end - start;
let start_rel = start - self.start_offset;
let tail_freqs = self.freq_complex.clone().slice([start_rel..max_seq_len]);
self.freq_complex
.inplace(|freqs| freqs.slice_assign([0..num_keep], tail_freqs));
let new_freqs = Self::compute_rotary_frequencies(
current_end..start + max_seq_len,
self.theta.clone(),
);
self.freq_complex
.inplace(|freqs| freqs.slice_assign([num_keep..max_seq_len], new_freqs));
}
self.start_offset = start;
}
fn compute_rotary_frequencies(range: Range<usize>, theta: Tensor<B, 1>) -> Tensor<B, 3> {
let d_model = theta.dims()[0] * 2;
let num_positions = range.end - range.start;
let frequencies: Tensor<B, 2> =
Tensor::<B, 1, Int>::arange(range.start as i64..range.end as i64, &theta.device())
.float()
.unsqueeze()
.transpose()
.repeat_dim(1, d_model / 2)
* theta.unsqueeze();
let p_cos = frequencies.clone().cos();
let p_sin = frequencies.sin();
Tensor::cat(vec![p_cos, p_sin], 1)
.reshape([num_positions, 2, d_model / 2])
.transpose()
.unsqueeze_dim::<4>(2)
.repeat_dim(2, 2)
.reshape([num_positions, d_model, 2])
}
}
#[cfg(test)]
mod tests {
use super::*;
use crate::TestBackend;
use ruda_model::tensor::{Tolerance, ops::FloatElem};
type FT = FloatElem<TestBackend>;
#[cfg(feature = "std")]
#[test]
fn rotary_encoding_half_storage_and_fp32_compute_preserve_shifted_gradients() {
use crate::TestAutodiffBackend as B;
use ruda_model::tensor::DType;
let device = Default::default();
let mut module = RotaryEncodingConfig::new(10, 4).init::<B>(&device);
module.shift(2);
for dtype in [DType::F16, DType::BF16] {
let input = Tensor::<B, 3>::from_floats(
[[[1., 2., 3., 4.], [5., 6., 7., 8.]]], &device,
).cast(dtype).require_grad();
let native = module.apply(input.clone(), 1);
assert_eq!(native.dtype(), dtype);
assert_eq!(native.dims(), input.dims());
let reference = module.apply(input.clone().cast(DType::F32), 1).cast(dtype);
native.cast(DType::F32).to_data().assert_approx_eq::<f32>(
&reference.clone().cast(DType::F32).to_data(), Tolerance::absolute(0.125));
let output = module.apply_with_compute_dtype(input.clone(), 1, FloatDType::F32);
assert_eq!(output.dtype(), dtype);
output.clone().cast(DType::F32).to_data().assert_approx_eq::<f32>(
&reference.clone().cast(DType::F32).to_data(), Tolerance::absolute(0.));
let expected_grads = reference.square().sum().backward();
let grads = output.square().sum().backward();
input.grad(&grads).unwrap().cast(DType::F32).to_data().assert_approx_eq::<f32>(
&input.grad(&expected_grads).unwrap().cast(DType::F32).to_data(), Tolerance::absolute(0.));
assert_eq!(module.freq_complex.dtype(), DType::F32);
assert_eq!(module.theta.dtype(), DType::F32);
}
}
#[test]
fn test_rotary_encoding_forward() {
let device = Default::default();
let rotary_encoding = RotaryEncodingConfig::new(10, 4).init::<TestBackend>(&device);
let input = Tensor::<TestBackend, 3>::from_floats(
[
[[1.0, 2.0, 3.0, 4.0], [5.0, 6.0, 7.0, 8.0]],
[[9.0, 10.0, 11.0, 12.0], [13.0, 14.0, 15.0, 16.0]],
],
&device,
);
let input = input.unsqueeze::<4>();
let output = rotary_encoding.forward(input);
let expected_output = Tensor::<TestBackend, 3>::from_floats(
[
[
[1.0000, 2.0000, 3.0000, 4.0000],
[-2.3473, 7.4492, 6.9197, 8.0696],
],
[
[9.0000, 10.0000, 11.0000, 12.0000],
[-4.7567, 18.5034, 14.8393, 16.1492],
],
],
&device,
);
output
.squeeze_dim::<3>(0)
.to_data()
.assert_approx_eq::<FT>(&expected_output.to_data(), Tolerance::default());
}
#[test]
fn test_rotary_encoding_3d() {
let device = Default::default();
let rotary_encoding = RotaryEncodingConfig::new(10, 4).init::<TestBackend>(&device);
let input = Tensor::<TestBackend, 3>::from_floats(
[
[[1.0, 2.0, 3.0, 4.0], [5.0, 6.0, 7.0, 8.0]],
[[9.0, 10.0, 11.0, 12.0], [13.0, 14.0, 15.0, 16.0]],
],
&device,
);
let output = rotary_encoding.forward(input);
let expected_output = Tensor::<TestBackend, 3>::from_floats(
[
[
[1.0000, 2.0000, 3.0000, 4.0000],
[-2.3473, 7.4492, 6.9197, 8.0696],
],
[
[9.0000, 10.0000, 11.0000, 12.0000],
[-4.7567, 18.5034, 14.8393, 16.1492],
],
],
&device,
);
output
.to_data()
.assert_approx_eq::<FT>(&expected_output.to_data(), Tolerance::default());
}
#[test]
fn test_zero_input_rotary_encoding_forward() {
let device = Default::default();
let rotary_encoding = RotaryEncodingConfig::new(10, 4).init::<TestBackend>(&device);
let input = Tensor::<TestBackend, 4>::zeros([1, 2, 2, 4], &device);
let output = rotary_encoding.forward(input);
let expected_output = Tensor::<TestBackend, 3>::from_floats(
[
[
[0.0000, 0.0000, 0.0000, 0.0000],
[0.0000, 0.0000, 0.0000, 0.0000],
],
[
[0.0000, 0.0000, 0.0000, 0.0000],
[0.0000, 0.0000, 0.0000, 0.0000],
],
],
&device,
);
output
.squeeze_dim::<3>(0)
.to_data()
.assert_approx_eq::<FT>(&expected_output.to_data(), Tolerance::default());
}
#[test]
#[should_panic]
fn test_valid_input_hidden_dim() {
let d_model = 15;
let device = Default::default();
let pe = RotaryEncodingConfig::new(10, d_model).init::<TestBackend>(&device);
let input = Tensor::<TestBackend, 3>::zeros([1, 5, d_model], &device);
let _output = pe.forward(input);
}
#[test]
fn test_rotary_encoding_frequencies() {
let device = Default::default();
let rotary_encoding = RotaryEncodingConfig::new(2, 8).init::<TestBackend>(&device);
let expected_freqs = Tensor::<TestBackend, 3>::from_floats(
[
[
[1.0000, 0.0000],
[1.0000, 0.0000],
[1.0000, 0.0000],
[1.0000, 0.0000],
],
[
[5.4030e-01, 8.4147e-01],
[9.9500e-01, 9.9833e-02],
[9.9995e-01, 9.9998e-03],
[9.9999e-01, 9.9999e-04],
],
],
&device,
)
.unsqueeze_dim::<4>(2)
.repeat_dim(2, 2)
.reshape([2, 8, 2]);
rotary_encoding
.freq_complex
.to_data()
.assert_approx_eq::<FT>(&expected_freqs.to_data(), Tolerance::default());
}
fn apply_freq_scaling_by_parts<B: Backend>(freqs: Tensor<B, 1>) -> Tensor<B, 1> {
let scale_factor = 8.;
let low_freq_factor = 1.;
let high_freq_factor = 4.;
let old_context_len = 8192.;
let low_freq_wavelen = old_context_len / low_freq_factor;
let high_freq_wavelen = old_context_len / high_freq_factor;
let wavelen = freqs.clone().recip().mul_scalar(2. * core::f32::consts::PI);
let cond = wavelen.clone().greater_equal_elem(high_freq_wavelen);
let smooth = wavelen
.clone()
.recip()
.mul_scalar(old_context_len)
.sub_scalar(low_freq_factor)
.div_scalar(high_freq_factor - low_freq_factor);
let new_freqs = smooth
.clone()
.neg()
.add_scalar(1.)
.mul(freqs.clone().div_scalar(scale_factor))
.add(smooth.clone().mul(freqs.clone()));
let new_freqs = freqs.clone().mask_where(cond, new_freqs);
let cond = wavelen.clone().greater_elem(low_freq_wavelen);
let new_freqs = new_freqs.mask_where(cond, freqs.clone().div_scalar(scale_factor));
let cond = wavelen.lower_elem(high_freq_wavelen);
new_freqs.mask_where(cond, freqs)
}
#[test]
fn test_rotary_encoding_with_frequency_scaling() {
let device = Default::default();
let rotary_encoding = RotaryEncodingConfig::new(2, 8)
.init_with_frequency_scaling::<TestBackend>(apply_freq_scaling_by_parts, &device);
let expected_freqs = Tensor::<TestBackend, 3>::from_floats(
[
[
[1.0000, 0.0000],
[1.0000, 0.0000],
[1.0000, 0.0000],
[1.0000, 0.0000],
],
[
[5.4030e-01, 8.4148e-01],
[9.9500e-01, 9.9833e-02],
[9.9995e-01, 9.9998e-03],
[1.0000, 2.1361e-04],
],
],
&device,
)
.unsqueeze_dim::<4>(2)
.repeat_dim(2, 2)
.reshape([2, 8, 2]);
rotary_encoding
.freq_complex
.to_data()
.assert_approx_eq::<FT>(&expected_freqs.to_data(), Tolerance::default());
}
#[test]
fn test_rotary_encoding_shift_full() {
let device = Default::default();
let rotary_encoding = RotaryEncodingConfig::new(10, 4).init::<TestBackend>(&device);
let input = Tensor::<TestBackend, 3>::from_floats(
[
[[1.0, 2.0, 3.0, 4.0], [5.0, 6.0, 7.0, 8.0]],
[[9.0, 10.0, 11.0, 12.0], [13.0, 14.0, 15.0, 16.0]],
],
&device,
)
.unsqueeze::<4>();
let expected_output = rotary_encoding.apply(input.clone(), 6);
let mut rotary_encoding = RotaryEncodingConfig::new(4, 4).init::<TestBackend>(&device);
rotary_encoding.shift(6);
let output = rotary_encoding.apply(input, 0);
output
.into_data()
.assert_approx_eq::<FT>(&expected_output.into_data(), Tolerance::default());
}
#[test]
fn test_rotary_encoding_shift() {
let device = Default::default();
let rotary_encoding = RotaryEncodingConfig::new(10, 4).init::<TestBackend>(&device);
let input = Tensor::<TestBackend, 3>::from_floats(
[
[[1.0, 2.0, 3.0, 4.0], [5.0, 6.0, 7.0, 8.0]],
[[9.0, 10.0, 11.0, 12.0], [13.0, 14.0, 15.0, 16.0]],
],
&device,
)
.unsqueeze::<4>();
let expected_output = rotary_encoding.apply(input.clone(), 2);
let mut rotary_encoding = RotaryEncodingConfig::new(4, 4).init::<TestBackend>(&device);
rotary_encoding.shift(2);
let output = rotary_encoding.apply(input, 0);
output
.into_data()
.assert_approx_eq::<FT>(&expected_output.into_data(), Tolerance::default());
}
#[test]
fn test_rotary_encoding_shift_multiple() {
let device = Default::default();
let mut rotary_encoding = RotaryEncodingConfig::new(4, 4).init::<TestBackend>(&device);
rotary_encoding.shift(2);
rotary_encoding.shift(5);
}
#[test]
#[should_panic = "Shift start position must be monotonically increasing"]
fn test_rotary_encoding_shift_should_increase() {
let device = Default::default();
let mut rotary_encoding = RotaryEncodingConfig::new(4, 4).init::<TestBackend>(&device);
rotary_encoding.shift(6);
rotary_encoding.shift(4); }
#[test]
fn display() {
let config = RotaryEncodingConfig::new(10, 4);
let pe = config.init::<TestBackend>(&Default::default());
assert_eq!(
alloc::format!("{pe}"),
"RotaryEncoding {d_model: 4, max_sequence_length: 10}"
);
}
}