use crate::engine::{
Cancellation,
array::{Array, Dtype},
error::{Error, Result},
ops::{self, AttentionMask},
};
pub fn attention_mask_for(seq_len: i32) -> AttentionMask {
if seq_len == 1 {
AttentionMask::None
} else {
AttentionMask::Causal
}
}
pub(super) fn forward_embeds_chunked(
input_ids: &Array,
embeds: &Array,
chunk_tokens: usize,
cancellation: Cancellation<'_>,
mut forward: impl FnMut(&Array, Array) -> Result<Array>,
) -> Result<Array> {
let input_shape = input_ids.shape();
let embed_shape = embeds.shape();
if input_shape.len() != 2
|| embed_shape.len() != 3
|| input_shape[0] != embed_shape[0]
|| input_shape[1] != embed_shape[1]
|| input_shape[0] <= 0
|| input_shape[1] <= 0
|| embed_shape[2] <= 0
{
return Err(Error::Model(format!(
"media prefill expected input [B, L] and embeddings [B, L, H], got \
{input_shape:?} and {embed_shape:?}"
)));
}
if chunk_tokens == 0 {
return Err(Error::Model(
"media prefill chunk size must be positive".to_string(),
));
}
let sequence_tokens = usize::try_from(input_shape[1])
.map_err(|_| Error::Model("media prefill sequence length exceeds usize".to_string()))?;
let chunk_tokens = if cancellation.is_cooperative() {
chunk_tokens
} else {
sequence_tokens
};
let mut output = None;
for start in (0..sequence_tokens).step_by(chunk_tokens) {
cancellation.checkpoint()?;
let end = start.saturating_add(chunk_tokens).min(sequence_tokens);
let start = i32::try_from(start)
.map_err(|_| Error::Model("media prefill offset exceeds i32".to_string()))?;
let end = i32::try_from(end)
.map_err(|_| Error::Model("media prefill offset exceeds i32".to_string()))?;
let ids = ops::slice(input_ids, &[0, start], &[input_shape[0], end])?;
let chunk_embeds = ops::slice(
embeds,
&[0, start, 0],
&[embed_shape[0], end, embed_shape[2]],
)?;
let logits = forward(&ids, chunk_embeds)?;
if usize::try_from(end)
.map_err(|_| Error::Model("media prefill offset exceeds usize".to_string()))?
< sequence_tokens
{
eval_last_logits(&logits)?;
}
output = Some(logits);
cancellation.checkpoint()?;
}
output.ok_or_else(|| Error::Model("cannot prefill an empty media prompt".to_string()))
}
fn eval_last_logits(logits: &Array) -> Result<()> {
let shape = logits.shape();
if shape.len() != 3 || shape[0] <= 0 || shape[1] <= 0 || shape[2] <= 0 {
return Err(Error::Model(format!(
"media prefill produced invalid logits shape {shape:?}"
)));
}
let last = ops::slice(
logits,
&[0, shape[1] - 1, 0],
&[shape[0], shape[1], shape[2]],
)?;
last.eval()
}
#[derive(Debug, Clone, Copy)]
pub struct RopeConfig {
pub dims: i32,
pub base: f32,
pub traditional: bool,
pub scale: f32,
}
impl RopeConfig {
pub fn new(dims: i32, base: f32) -> Self {
RopeConfig {
dims,
base,
traditional: false,
scale: 1.0,
}
}
pub fn apply(&self, x: &Array, offset: i32) -> Result<Array> {
ops::rope(
x,
self.dims,
self.traditional,
Some(self.base),
self.scale,
offset,
None,
)
}
}
pub fn yarn_freqs(
dims: i32,
base: f32,
factor: f32,
original_max_position_embeddings: i32,
beta_fast: f32,
beta_slow: f32,
) -> Vec<f32> {
let dims_f = f64::from(dims);
let base_f = f64::from(base);
let orig = f64::from(original_max_position_embeddings);
let correction_dim = |num_rotations: f64| -> f64 {
dims_f * (orig / (num_rotations * 2.0 * std::f64::consts::PI)).ln() / (2.0 * base_f.ln())
};
let low = correction_dim(f64::from(beta_fast)).floor().max(0.0);
let mut high = correction_dim(f64::from(beta_slow))
.ceil()
.min(dims_f - 1.0);
if low == high {
high += 0.001; }
(0..dims / 2)
.map(|i| {
let freq_extra = base_f.powf(f64::from(2 * i) / dims_f);
let freq_inter = f64::from(factor) * freq_extra;
let ramp = ((f64::from(i) - low) / (high - low)).clamp(0.0, 1.0);
let mask = 1.0 - ramp;
let freq = (freq_inter * freq_extra) / (freq_inter * mask + freq_extra * (1.0 - mask));
freq as f32
})
.collect()
}
#[derive(Clone)]
pub struct YarnRope {
dims: i32,
freqs: Array,
mscale_vec: Option<Array>,
}
impl YarnRope {
#[allow(
clippy::too_many_arguments,
reason = "constructor exposes the complete published YaRN parameterization"
)]
pub fn new(
dims: i32,
head_dim: i32,
base: f32,
factor: f32,
original_max_position_embeddings: i32,
beta_fast: f32,
beta_slow: f32,
attention_factor: f32,
) -> Result<Self> {
let freqs = yarn_freqs(
dims,
base,
factor,
original_max_position_embeddings,
beta_fast,
beta_slow,
);
let freqs = Array::from_slice(&freqs, &[dims / 2])?;
let mscale_vec = if (attention_factor - 1.0).abs() > f32::EPSILON {
let mut v = vec![1.0f32; head_dim as usize];
v[..dims as usize].fill(attention_factor);
Some(Array::from_slice(&v, &[head_dim])?)
} else {
None
};
Ok(YarnRope {
dims,
freqs,
mscale_vec,
})
}
pub fn apply(&self, x: &Array, offset: i32) -> Result<Array> {
let x = match &self.mscale_vec {
Some(v) => {
let v = ops::astype(v, x.dtype())?;
ops::multiply(x, &v)?
}
None => x.clone(),
};
ops::rope(&x, self.dims, false, None, 1.0, offset, Some(&self.freqs))
}
}
pub fn split_heads(x: &Array, batch: i32, seq: i32, heads: i32) -> Result<Array> {
let reshaped = ops::reshape(x, &[batch, seq, heads, -1])?;
ops::transpose_axes(&reshaped, &[0, 2, 1, 3])
}
pub fn merge_heads(x: &Array, batch: i32, seq: i32) -> Result<Array> {
let t = ops::transpose_axes(x, &[0, 2, 1, 3])?;
ops::reshape(&t, &[batch, seq, -1])
}
pub fn repeat_kv_heads(x: &Array, n_repeats: i32) -> Result<Array> {
if n_repeats == 1 {
return Ok(x.clone());
}
let shape = x.shape();
let (b, h, l, d) = (shape[0], shape[1], shape[2], shape[3]);
let expanded = ops::expand_dims(x, 2)?;
let broadcasted = ops::broadcast_to(&expanded, &[b, h, n_repeats, l, d])?;
ops::reshape(&broadcasted, &[b, h * n_repeats, l, d])
}
pub fn splice_media_features(
h: &Array,
input_ids: &Array,
mut features: Vec<Array>,
placeholder_token_id: i32,
modality: &str,
) -> Result<Array> {
let features = if features.len() == 1 {
features.remove(0)
} else {
let refs: Vec<&Array> = features.iter().collect();
ops::concatenate(&refs, 1)?
};
let features = ops::astype(&features, h.dtype())?;
let placeholder = ops::astype(&Array::scalar_i32(placeholder_token_id)?, input_ids.dtype())?;
let mask = ops::equal(input_ids, &placeholder)?;
let mask_count_arr = ops::sum_axes(
&ops::reshape(&ops::astype(&mask, Dtype::Int32)?, &[-1])?,
&[0],
false,
)?;
let mask_count = mask_count_arr.item_f32()? as i32;
let feature_count = features.dim(1);
if mask_count != feature_count {
return Err(Error::Model(format!(
"{modality} token count ({mask_count}) does not match {modality} \
feature count ({feature_count}); check that {modality} placeholder \
expansion produced the right number of tokens"
)));
}
let mask_expanded = ops::broadcast_to(&ops::expand_dims(&mask, -1)?, &h.shape())?;
masked_scatter(h, &mask_expanded, &features)
}
pub fn masked_scatter(input: &Array, mask: &Array, source: &Array) -> Result<Array> {
let input_shape = input.shape();
let mask_flat = ops::reshape(&ops::astype(mask, Dtype::Int32)?, &[-1])?;
let input_flat = ops::reshape(input, &[-1])?;
let source_flat = ops::reshape(source, &[-1])?;
let source_size = source_flat.dim(0);
let idx = ops::subtract(&ops::cumsum(&mask_flat, 0)?, &Array::scalar_i32(1)?)?;
let idx = ops::maximum(&idx, &Array::scalar_i32(0)?)?;
let idx = ops::minimum(&idx, &Array::scalar_i32((source_size - 1).max(0))?)?;
let idx = ops::astype(&idx, Dtype::UInt32)?;
let aligned = ops::take(&source_flat, &idx)?;
let mask_bool = ops::astype(&mask_flat, Dtype::Bool)?;
let result = ops::where_cond(&mask_bool, &aligned, &input_flat)?;
ops::reshape(&result, &input_shape)
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn yarn_freqs_laguna_m1_values() {
let dims = 128;
let base = 500000.0f32;
let factor = 32.0f32;
let freqs = yarn_freqs(dims, base, factor, 4096, 64.0, 1.0);
assert_eq!(freqs.len(), 64);
assert!((freqs[0] - 1.0).abs() < 1e-6);
let extra = |i: i32| (base as f64).powf(f64::from(2 * i) / 128.0);
let interp_63 = (32.0 * extra(63)) as f32;
assert!((freqs[63] - interp_63).abs() / interp_63 < 1e-4);
let mask = 1.0 - (20.0 - 11.0) / (32.0 - 11.0);
let e = extra(20);
let expected_20 = ((32.0 * e * e) / (32.0 * e * mask + e * (1.0 - mask))) as f32;
assert!((freqs[20] - expected_20).abs() / expected_20 < 1e-4);
for w in freqs.windows(2) {
assert!(w[1] > w[0]);
}
}
#[test]
fn yarn_freqs_laguna_s21_values() {
let base = 500000.0f32;
let freqs = yarn_freqs(64, base, 128.0, 8192, 32.0, 1.0);
assert_eq!(freqs.len(), 32);
assert!((freqs[0] - 1.0).abs() < 1e-6);
let interp_tail = (128.0 * (base as f64).powf(62.0 / 64.0)) as f32;
assert!((freqs[31] - interp_tail).abs() / interp_tail < 1e-4);
for w in freqs.windows(2) {
assert!(w[1] > w[0]);
}
}
}