use std::collections::HashMap;
use std::sync::{Mutex, OnceLock};
use metal::MTLResourceOptions;
use crate::buffer::MlxBuffer;
use crate::dtypes::DType;
use crate::error::{MlxError, Result};
#[derive(Debug, Clone, Copy, PartialEq, Eq, Hash)]
struct RopeFreqCacheKey {
device_ptr: usize,
freq_base_bits: u32,
denominator: u32,
pair_count: u32,
}
static ROPE_FREQ_CACHE: OnceLock<Mutex<HashMap<RopeFreqCacheKey, MlxBuffer>>> = OnceLock::new();
fn build_inv_freqs(freq_base: f32, denominator: u32, pair_count: u32) -> Vec<f32> {
(0..pair_count)
.map(|pair| {
let dim_ratio = (2 * pair) as f32 / denominator as f32;
1.0_f32 / freq_base.powf(dim_ratio)
})
.collect()
}
pub(crate) fn with_inv_freqs<R>(
device: &metal::DeviceRef,
freq_base: f32,
denominator: u32,
pair_count: u32,
f: impl FnOnce(&MlxBuffer) -> R,
) -> Result<R> {
if !freq_base.is_finite() || freq_base <= 0.0 {
return Err(MlxError::InvalidArgument(format!(
"RoPE freq_base must be finite and positive, got {freq_base}"
)));
}
if denominator == 0 || pair_count == 0 {
return Err(MlxError::InvalidArgument(format!(
"RoPE frequency denominator and pair_count must be > 0, got {denominator} and {pair_count}"
)));
}
let key = RopeFreqCacheKey {
device_ptr: device as *const metal::DeviceRef as usize,
freq_base_bits: freq_base.to_bits(),
denominator,
pair_count,
};
let inv_freqs = {
let mut cache = ROPE_FREQ_CACHE
.get_or_init(|| Mutex::new(HashMap::new()))
.lock()
.map_err(|_| {
MlxError::CommandBufferError("RoPE frequency cache lock poisoned".into())
})?;
cache
.entry(key)
.or_insert_with(|| {
let values = build_inv_freqs(freq_base, denominator, pair_count);
let byte_len = std::mem::size_of_val(values.as_slice());
let metal_buf = device.new_buffer_with_data(
values.as_ptr().cast(),
byte_len as u64,
MTLResourceOptions::StorageModeShared,
);
MlxBuffer::from_raw(metal_buf, DType::F32, vec![pair_count as usize])
})
.clone()
};
Ok(f(&inv_freqs))
}
#[cfg(test)]
mod tests {
use super::{build_inv_freqs, with_inv_freqs, ROPE_FREQ_CACHE};
fn cache_len() -> usize {
match ROPE_FREQ_CACHE
.get_or_init(|| std::sync::Mutex::new(std::collections::HashMap::new()))
.lock()
{
Ok(cache) => cache.len(),
Err(poisoned) => poisoned.into_inner().len(),
}
}
#[test]
fn host_schedule_matches_f32_rope_definition() {
let got = build_inv_freqs(1_000_000.0, 16, 8);
for (pair, &freq) in got.iter().enumerate() {
let ratio = (2 * pair) as f32 / 16.0_f32;
assert_eq!(freq, 1.0_f32 / 1_000_000.0_f32.powf(ratio));
}
}
#[test]
fn process_cache_survives_encoding_thread_exit() {
let before = cache_len();
let worker = std::thread::spawn(|| {
let Ok(device) = crate::MlxDevice::new() else {
return false;
};
with_inv_freqs(device.metal_device(), 123_456.75_f32, 14, 7, |_| ()).is_ok()
});
assert!(matches!(worker.join(), Ok(true)));
assert_eq!(cache_len(), before + 1);
}
}