#[test]
fn test_scaled_rope_linear_scaling() {
let scaling = RopeScalingType::Linear { scale: 4.0 };
let scaled = ScaledRoPE::new(64, 10000.0, scaling).expect("test");
assert!((scaled.context_length_multiplier() - 4.0).abs() < 1e-6);
assert!((scaled.scaled_base() - 10000.0).abs() < 1e-6);
assert!((scaled.mscale() - 1.0).abs() < 1e-6);
}
#[test]
fn test_scaled_rope_ntk_scaling() {
let scaling = RopeScalingType::Ntk { scale: 4.0 };
let scaled = ScaledRoPE::new(64, 10000.0, scaling).expect("test");
assert!((scaled.context_length_multiplier() - 4.0).abs() < 1e-6);
assert!(scaled.scaled_base() > 10000.0);
assert!(scaled.scaled_base() > 40000.0);
assert!((scaled.mscale() - 1.0).abs() < 1e-6);
}
#[test]
fn test_scaled_rope_dynamic_ntk() {
let scaling = RopeScalingType::DynamicNtk {
original_max_len: 2048,
target_max_len: 8192,
};
let scaled = ScaledRoPE::new(64, 10000.0, scaling).expect("test");
assert!((scaled.context_length_multiplier() - 4.0).abs() < 1e-6);
assert!(scaled.scaled_base() > 40000.0);
}
#[test]
fn test_scaled_rope_yarn() {
let scaling = RopeScalingType::Yarn {
original_max_len: 2048,
target_max_len: 32768,
attn_factor: 0.0, beta_fast: 32.0,
beta_slow: 1.0,
};
let scaled = ScaledRoPE::new(64, 10000.0, scaling).expect("test");
assert!((scaled.context_length_multiplier() - 16.0).abs() < 1e-6);
assert!(scaled.mscale() > 1.0);
assert!(
(scaled.scaled_base() - 10000.0).abs() < 1e-3,
"YaRN must use the original base, got {}",
scaled.scaled_base()
);
}
#[test]
fn test_scaled_rope_yarn_extrapolation_uses_original_base() {
let dim = 64usize;
let base = 10000.0f32;
let scaling = RopeScalingType::Yarn {
original_max_len: 2048,
target_max_len: 8192,
attn_factor: 1.0,
beta_fast: 32.0,
beta_slow: 1.0,
};
let scaled = ScaledRoPE::new(dim, base, scaling).expect("test");
#[allow(clippy::cast_precision_loss)]
let original_base_inv_freq_1 = base.powf(-2.0 * 1.0 / (dim as f32));
#[allow(clippy::cast_precision_loss)]
let ntk_base = base * 4.0f32.powf((dim as f32) / ((dim as f32) - 2.0));
#[allow(clippy::cast_precision_loss)]
let ntk_inv_freq_1 = ntk_base.powf(-2.0 * 1.0 / (dim as f32));
let inv_freq = scaled.inv_freq();
assert!(inv_freq.len() > 1, "need at least 2 frequency pairs");
assert!(
(inv_freq[1] - original_base_inv_freq_1).abs() < 1e-5,
"PMAT-874: YaRN extrapolated dim 1 must use the ORIGINAL base \
(expected {original_base_inv_freq_1}, got {})",
inv_freq[1]
);
assert!(
(inv_freq[1] - ntk_inv_freq_1).abs() > 1e-4,
"PMAT-874: YaRN must NOT apply the NTK base modification \
(got {} which matches the buggy NTK value {ntk_inv_freq_1})",
inv_freq[1]
);
}
#[test]
fn test_scaled_rope_yarn_custom_attn_factor() {
let scaling = RopeScalingType::Yarn {
original_max_len: 2048,
target_max_len: 8192,
attn_factor: 1.5, beta_fast: 32.0,
beta_slow: 1.0,
};
let scaled = ScaledRoPE::new(64, 10000.0, scaling).expect("test");
assert!((scaled.mscale() - 1.5).abs() < 1e-6);
}
#[test]
fn test_scaled_rope_forward_no_scaling() {
let scaled = ScaledRoPE::new(4, 10000.0, RopeScalingType::None).expect("test");
let input = Tensor::from_vec(vec![4], vec![1.0, 0.0, 0.0, 1.0]).expect("test");
let output = scaled.forward(&input, 0).expect("test");
assert_eq!(output.shape(), &[4]);
}
#[test]
fn test_scaled_rope_forward_linear() {
let scaling = RopeScalingType::Linear { scale: 2.0 };
let scaled = ScaledRoPE::new(4, 10000.0, scaling).expect("test");
let input = Tensor::from_vec(vec![4], vec![1.0, 0.0, 0.0, 1.0]).expect("test");
let output = scaled.forward(&input, 10).expect("test");
assert_eq!(output.shape(), &[4]);
}
#[test]
fn test_scaled_rope_forward_ntk() {
let scaling = RopeScalingType::Ntk { scale: 4.0 };
let scaled = ScaledRoPE::new(4, 10000.0, scaling).expect("test");
let input = Tensor::from_vec(vec![4], vec![1.0, 0.0, 0.0, 1.0]).expect("test");
let output = scaled.forward(&input, 100).expect("test");
assert_eq!(output.shape(), &[4]);
let norm: f32 = output.data().iter().map(|x| x * x).sum::<f32>().sqrt();
assert!((norm - 2.0_f32.sqrt()).abs() < 0.1);
}
#[test]
fn test_scaled_rope_forward_yarn() {
let scaling = RopeScalingType::Yarn {
original_max_len: 2048,
target_max_len: 8192,
attn_factor: 1.0,
beta_fast: 32.0,
beta_slow: 1.0,
};
let scaled = ScaledRoPE::new(4, 10000.0, scaling).expect("test");
let input = Tensor::from_vec(vec![4], vec![1.0, 0.0, 0.0, 1.0]).expect("test");
let output = scaled.forward(&input, 5000).expect("test");
assert_eq!(output.shape(), &[4]);
}
#[test]
fn test_scaled_rope_zero_dim_error() {
let result = ScaledRoPE::new(0, 10000.0, RopeScalingType::None);
assert!(result.is_err());
}
#[test]
fn test_scaled_rope_odd_dim_error() {
let result = ScaledRoPE::new(63, 10000.0, RopeScalingType::None);
assert!(result.is_err());
}
#[test]
fn test_scaled_rope_dimension_mismatch() {
let scaled = ScaledRoPE::new(4, 10000.0, RopeScalingType::None).expect("test");
let input = Tensor::from_vec(vec![8], vec![0.0; 8]).expect("test");
let result = scaled.forward(&input, 0);
assert!(result.is_err());
}
#[test]
fn test_rope_scaling_type_default() {
let scaling = RopeScalingType::default();
assert_eq!(scaling, RopeScalingType::None);
}
#[test]
fn test_scaled_rope_with_default_base() {
let scaled = ScaledRoPE::with_default_base(64, RopeScalingType::None).expect("test");
assert!((scaled.original_base() - 10000.0).abs() < 1e-6);
}
#[test]
fn test_scaled_rope_inv_freq_length() {
let scaled = ScaledRoPE::new(128, 10000.0, RopeScalingType::None).expect("test");
assert_eq!(scaled.inv_freq().len(), 64); }
#[test]
fn test_alibi_creation() {
let alibi = ALiBi::new(8).expect("test");
assert_eq!(alibi.num_heads(), 8);
assert_eq!(alibi.slopes().len(), 8);
}
#[test]
fn test_alibi_zero_heads_error() {
let result = ALiBi::new(0);
assert!(result.is_err());
}
#[test]
fn test_alibi_slopes_power_of_2() {
let alibi = ALiBi::new(8).expect("test");
let slopes = alibi.slopes();
assert!((slopes[0] - 0.5).abs() < 1e-6); assert!((slopes[1] - 0.25).abs() < 1e-6); assert!((slopes[2] - 0.125).abs() < 1e-6); assert!((slopes[3] - 0.0625).abs() < 1e-6); }
#[test]
fn test_alibi_slopes_non_power_of_2() {
let alibi = ALiBi::new(6).expect("test");
let slopes = alibi.slopes();
assert_eq!(slopes.len(), 6);
assert!((slopes[0] - 0.25).abs() < 1e-6); assert!((slopes[1] - 0.0625).abs() < 1e-6); assert!((slopes[2] - 0.015_625).abs() < 1e-6); assert!((slopes[3] - 0.003_906_25).abs() < 1e-6);
assert!((slopes[4] - 0.5).abs() < 1e-6);
assert!((slopes[5] - 0.125).abs() < 1e-6);
}
#[test]
fn test_alibi_bias_shape() {
let alibi = ALiBi::new(4).expect("test");
let bias = alibi.get_bias(10).expect("test");
assert_eq!(bias.shape(), &[10, 10, 4]);
}
#[test]
fn test_alibi_bias_zero_seq_len_error() {
let alibi = ALiBi::new(4).expect("test");
let result = alibi.get_bias(0);
assert!(result.is_err());
}
#[test]
fn test_alibi_bias_diagonal_zero() {
let alibi = ALiBi::new(4).expect("test");
let bias = alibi.get_bias(5).expect("test");
for i in 0..5 {
for h in 0..4 {
let idx = i * 5 * 4 + i * 4 + h; let value = bias.data()[idx];
assert!(
value.abs() < 1e-6,
"Diagonal bias[{i}, {i}, {h}] should be 0, got {value}"
);
}
}
}
#[test]
fn test_alibi_bias_symmetry() {
let alibi = ALiBi::new(2).expect("test");
let bias = alibi.get_bias(4).expect("test");
for i in 0..4 {
for j in 0..4 {
for h in 0..2 {
let idx_ij = i * 4 * 2 + j * 2 + h;
let idx_ji = j * 4 * 2 + i * 2 + h;
let bias_ij = bias.data()[idx_ij];
let bias_ji = bias.data()[idx_ji];
assert!(
(bias_ij - bias_ji).abs() < 1e-6,
"Bias should be symmetric: [{i},{j},{h}]={bias_ij} vs [{j},{i},{h}]={bias_ji}"
);
}
}
}
}
#[test]
fn test_alibi_bias_computation() {
let alibi = ALiBi::new(2).expect("test");
let slopes = alibi.slopes();
let bias = alibi.get_bias(3).expect("test");
let idx = 2 * 2;
let expected = -slopes[0] * 2.0;
assert!(
(bias.data()[idx] - expected).abs() < 1e-6,
"Expected {expected}, got {}",
bias.data()[idx]
);
let idx = 3 * 2 + 2 * 2 + 1;
let expected = -slopes[1];
assert!(
(bias.data()[idx] - expected).abs() < 1e-6,
"Expected {expected}, got {}",
bias.data()[idx]
);
}
#[test]
fn test_alibi_bias_negative() {
let alibi = ALiBi::new(4).expect("test");
let bias = alibi.get_bias(10).expect("test");
for &value in bias.data() {
assert!(value <= 1e-6, "Bias should be non-positive, got {value}");
}
}
#[test]
fn test_alibi_bias_distance_proportional() {
let alibi = ALiBi::new(1).expect("test");
let slope = alibi.slopes()[0]; let bias = alibi.get_bias(5).expect("test");
let bias_01 = bias.data()[1];
let bias_02 = bias.data()[2];
let bias_03 = bias.data()[3];
assert!((bias_01 - (-slope)).abs() < 1e-6);
assert!((bias_02 - (-slope * 2.0)).abs() < 1e-6);
assert!((bias_03 - (-slope * 3.0)).abs() < 1e-6);
}
#[test]
fn test_alibi_single_head() {
let alibi = ALiBi::new(1).expect("test");
assert_eq!(alibi.num_heads(), 1);
assert_eq!(alibi.slopes().len(), 1);
assert!((alibi.slopes()[0] - 0.003_906_25).abs() < 1e-6);
}
#[test]
fn test_alibi_large_num_heads() {
let alibi = ALiBi::new(12).expect("test");
assert_eq!(alibi.num_heads(), 12);
assert_eq!(alibi.slopes().len(), 12);
for slope in alibi.slopes() {
assert!(*slope > 0.0, "Slope should be positive, got {slope}");
}
assert!((alibi.slopes()[0] - 0.5).abs() < 1e-6);
}
#[test]
fn test_alibi_bias_long_sequence() {
let alibi = ALiBi::new(8).expect("test");
let bias = alibi.get_bias(128).expect("test");
assert_eq!(bias.shape(), &[128, 128, 8]);
let near_bias = bias.data()[8]; let far_bias = bias.data()[100 * 8];
assert!(near_bias > far_bias); }