From 35496079b9be8565e1330ffb24900a8974539c8f Mon Sep 17 00:00:00 2001 From: James Dumay Date: Sun, 29 Mar 2026 16:51:43 +1100 Subject: [PATCH 01/15] Add trim support to ConcatKeyValueCache --- mlx-lm/src/cache.rs | 130 +++++++++++++++++++++++++++++++++++++++++++- 1 file changed, 129 insertions(+), 1 deletion(-) diff --git a/mlx-lm/src/cache.rs b/mlx-lm/src/cache.rs index da250618f..67e759485 100644 --- a/mlx-lm/src/cache.rs +++ b/mlx-lm/src/cache.rs @@ -1,4 +1,9 @@ -use mlx_rs::{error::Exception, ops::concatenate_axis, Array}; +use mlx_rs::{ + error::Exception, + ops::{concatenate_axis, indexing::IndexOp}, + transforms::eval, + Array, +}; // TODO: somehow move quantized methods to a separate trait? pub trait KeyValueCache { @@ -68,6 +73,50 @@ impl ConcatKeyValueCache { pub fn new() -> Self { Self::default() } + + pub fn trim_to(&mut self, token_count: i32) -> Result<(), Exception> { + let token_count = token_count.max(0); + match (&self.keys, &self.values) { + (Some(keys), Some(values)) => { + let current = keys.shape()[keys.shape().len() - 2]; + if token_count >= current { + self.offset = current; + return Ok(()); + } + + let trimmed_keys = slice_kv_prefix(keys, token_count)?; + let trimmed_values = slice_kv_prefix(values, token_count)?; + eval([&trimmed_keys, &trimmed_values])?; + let trimmed_keys = trimmed_keys.deep_clone(); + let trimmed_values = trimmed_values.deep_clone(); + + self.keys = Some(trimmed_keys); + self.values = Some(trimmed_values); + self.offset = token_count; + Ok(()) + } + _ => { + self.offset = 0; + Ok(()) + } + } + } + + pub fn trimmed_to(&self, token_count: i32) -> Result { + let mut clone = self.clone(); + clone.trim_to(token_count)?; + Ok(clone) + } +} + +fn slice_kv_prefix(array: &Array, token_count: i32) -> Result { + match array.shape().len() { + 4 => Ok(array.index((.., .., ..token_count, ..))), + 3 => Ok(array.index((.., ..token_count, ..))), + other => Err(Exception::custom(format!( + "unsupported KV cache rank {other} for prefix trim" + ))), + } } impl KeyValueCache for ConcatKeyValueCache { @@ -106,3 +155,82 @@ impl KeyValueCache for ConcatKeyValueCache { /// TODO: A generic KV Cache pub struct DefaultKeyValueCache {} + +#[cfg(test)] +mod tests { + use super::{ConcatKeyValueCache, KeyValueCache}; + use mlx_rs::Array; + use std::sync::{Mutex, MutexGuard, OnceLock}; + + fn test_guard() -> MutexGuard<'static, ()> { + static LOCK: OnceLock> = OnceLock::new(); + LOCK.get_or_init(|| Mutex::new(())) + .lock() + .expect("cache test lock poisoned") + } + + #[test] + fn trimmed_cache_keeps_prefix_and_offset() { + let _guard = test_guard(); + let mut cache = ConcatKeyValueCache::new(); + let keys = Array::from_slice(&[1f32, 2., 3., 4., 5., 6.], &[1, 1, 3, 2]); + let values = Array::from_slice(&[7f32, 8., 9., 10., 11., 12.], &[1, 1, 3, 2]); + let _ = cache.update_and_fetch(keys, values).expect("seed cache"); + + let trimmed = cache.trimmed_to(2).expect("trimmed clone"); + assert_eq!(trimmed.offset(), 2); + + let mut trimmed = trimmed; + let append_keys = Array::from_slice(&[13f32, 14.], &[1, 1, 1, 2]); + let append_values = Array::from_slice(&[15f32, 16.], &[1, 1, 1, 2]); + let (keys, values) = trimmed + .update_and_fetch(append_keys, append_values) + .expect("append after trim"); + + assert_eq!(keys.shape(), &[1, 1, 3, 2]); + assert_eq!(values.shape(), &[1, 1, 3, 2]); + assert_eq!(keys.as_slice::(), &[1., 2., 3., 4., 13., 14.]); + assert_eq!(values.as_slice::(), &[7., 8., 9., 10., 15., 16.]); + } + + #[test] + fn trim_to_larger_offset_is_noop() { + let _guard = test_guard(); + let mut cache = ConcatKeyValueCache::new(); + let keys = Array::from_slice(&[1f32, 2., 3., 4.], &[1, 1, 2, 2]); + let values = Array::from_slice(&[5f32, 6., 7., 8.], &[1, 1, 2, 2]); + let _ = cache.update_and_fetch(keys, values).expect("seed cache"); + + cache.trim_to(8).expect("trim beyond end"); + assert_eq!(cache.offset(), 2); + } + + #[test] + fn trimmed_to_does_not_mutate_original_cache() { + let _guard = test_guard(); + let mut cache = ConcatKeyValueCache::new(); + let keys = Array::from_slice(&[1f32, 2., 3., 4., 5., 6.], &[1, 1, 3, 2]); + let values = Array::from_slice(&[7f32, 8., 9., 10., 11., 12.], &[1, 1, 3, 2]); + let _ = cache.update_and_fetch(keys, values).expect("seed cache"); + + let trimmed = cache.trimmed_to(2).expect("trimmed clone"); + + assert_eq!(cache.offset(), 3); + assert_eq!(trimmed.offset(), 2); + + let mut original = cache; + let mut trimmed = trimmed; + let append_keys = Array::from_slice(&[13f32, 14.], &[1, 1, 1, 2]); + let append_values = Array::from_slice(&[15f32, 16.], &[1, 1, 1, 2]); + + let (original_keys, _) = original + .update_and_fetch(append_keys.deep_clone(), append_values.deep_clone()) + .expect("append original"); + let (trimmed_keys, _) = trimmed + .update_and_fetch(append_keys, append_values) + .expect("append trimmed"); + + assert_eq!(original_keys.as_slice::(), &[1., 2., 3., 4., 5., 6., 13., 14.]); + assert_eq!(trimmed_keys.as_slice::(), &[1., 2., 3., 4., 13., 14.]); + } +} From 6d95787e6f415638cb84380a87a70a6bd151fb42 Mon Sep 17 00:00:00 2001 From: James Dumay Date: Sun, 29 Mar 2026 18:18:46 +1100 Subject: [PATCH 02/15] Add qwen2 model support --- mlx-lm/src/lib.rs | 14 ++- mlx-lm/src/models/mod.rs | 1 + mlx-lm/src/models/qwen2.rs | 176 +++++++++++++++++++++++++++++++++++++ 3 files changed, 190 insertions(+), 1 deletion(-) create mode 100644 mlx-lm/src/models/qwen2.rs diff --git a/mlx-lm/src/lib.rs b/mlx-lm/src/lib.rs index a05ab58d3..dcc827ac6 100644 --- a/mlx-lm/src/lib.rs +++ b/mlx-lm/src/lib.rs @@ -7,7 +7,7 @@ pub mod utils; use mlx_rs::Array; -use crate::models::qwen3; +use crate::models::{qwen2, qwen3}; pub struct ModelInputBuilder<'a, C, T> { pub y: &'a Array, @@ -31,6 +31,18 @@ impl<'a, C> ModelInput<'a, C, Option> for qwen3::ModelInput<'a, C> { } } +impl<'a, C> ModelInput<'a, C, Option> for qwen2::ModelInput<'a, C> { + fn from_model_input_builder(builder: ModelInputBuilder<'a, C, Option>) -> Self { + let ModelInputBuilder { y, cache, state } = builder; + + Self { + inputs: y, + mask: state.as_ref(), + cache, + } + } +} + pub trait ModelOutput { fn logits(&self) -> &Array; } diff --git a/mlx-lm/src/models/mod.rs b/mlx-lm/src/models/mod.rs index 37d46974f..561eab29a 100644 --- a/mlx-lm/src/models/mod.rs +++ b/mlx-lm/src/models/mod.rs @@ -1,2 +1,3 @@ pub mod llama; +pub mod qwen2; pub mod qwen3; diff --git a/mlx-lm/src/models/qwen2.rs b/mlx-lm/src/models/qwen2.rs new file mode 100644 index 000000000..7bee4a782 --- /dev/null +++ b/mlx-lm/src/models/qwen2.rs @@ -0,0 +1,176 @@ +use std::path::Path; + +use mlx_rs::{error::Exception, module::ModuleParametersExt}; +use serde::Deserialize; +use tokenizers::Tokenizer; + +use crate::{error::Error, models::llama, utils::rope::FloatOrString}; + +pub type Attention = llama::Attention; +pub type Mlp = llama::Mlp; +pub type TransformerBlock = llama::TransformerBlock; +pub type Qwen2Model = llama::LlamaModel; +pub type Model = llama::Model; +pub type ModelInput<'a, C> = llama::ModelInput<'a, C>; +pub type Generate<'a, C> = llama::Generate<'a, C>; +pub type GenerateState<'a> = llama::GenerateState<'a>; + +#[derive(Debug, Clone, Deserialize)] +pub struct ModelArgs { + pub model_type: String, + pub hidden_size: i32, + pub num_hidden_layers: i32, + pub intermediate_size: i32, + pub num_attention_heads: i32, + pub rms_norm_eps: f32, + pub vocab_size: i32, + pub num_key_value_heads: i32, + #[serde(default = "default_max_position_embeddings")] + pub max_position_embeddings: i32, + #[serde(default = "default_rope_theta")] + pub rope_theta: f32, + #[serde(default)] + pub rope_traditional: bool, + #[serde(default)] + pub rope_scaling: Option>, + #[serde(default = "default_true")] + pub tie_word_embeddings: bool, +} + +fn default_true() -> bool { + true +} + +fn default_max_position_embeddings() -> i32 { + 32768 +} + +fn default_rope_theta() -> f32 { + 1_000_000.0 +} + +impl From for llama::ModelArgs { + fn from(value: ModelArgs) -> Self { + let head_dim = value.hidden_size / value.num_attention_heads; + Self { + model_type: value.model_type, + hidden_size: value.hidden_size, + num_hidden_layers: value.num_hidden_layers, + intermediate_size: value.intermediate_size, + num_attention_heads: value.num_attention_heads, + rms_norm_eps: value.rms_norm_eps, + vocab_size: value.vocab_size, + num_key_value_heads: value.num_key_value_heads, + max_position_embeddings: value.max_position_embeddings, + rope_theta: value.rope_theta, + head_dim, + tie_word_embeddings: value.tie_word_embeddings, + attention_bias: true, + mlp_bias: false, + rope_scaling: value.rope_scaling, + } + } +} + +pub fn load_qwen2_tokenizer(model_dir: impl AsRef) -> Result { + llama::load_llama_tokenizer(model_dir) +} + +pub fn get_qwen2_model_args(model_dir: impl AsRef) -> Result { + let model_args_filename = model_dir.as_ref().join("config.json"); + let file = std::fs::File::open(model_args_filename)?; + let model_args: ModelArgs = serde_json::from_reader(file)?; + + Ok(model_args) +} + +pub fn load_qwen2_model(model_dir: impl AsRef) -> Result { + let model_dir = model_dir.as_ref(); + let model_args = get_qwen2_model_args(model_dir)?; + let mut model = Model::new(model_args.into())?; + + let weights_index = model_dir.join("model.safetensors.index.json"); + if weights_index.exists() { + let json = std::fs::read_to_string(weights_index)?; + let weight_map: llama::WeightMap = serde_json::from_str(&json)?; + + let weight_files: std::collections::HashSet<&String> = + weight_map.weight_map.values().collect(); + for weight_file in weight_files { + let weights_filename = model_dir.join(weight_file); + model.load_safetensors(weights_filename)?; + } + } else { + let weights_filename = model_dir.join("model.safetensors"); + model.load_safetensors(weights_filename)?; + } + + Ok(model) +} + +pub fn sample(logits: &mlx_rs::Array, temp: f32) -> Result { + llama::sample(logits, temp) +} + +#[cfg(test)] +mod tests { + use mlx_rs::{ + ops::indexing::{IndexOp, NewAxis}, + transforms::eval, + Array, + }; + + use crate::{ + cache::ConcatKeyValueCache, + models::qwen2::{load_qwen2_model, load_qwen2_tokenizer}, + }; + + const CACHED_TEST_MODEL_DIR: &str = + "/Users/jdumay/.cache/huggingface/hub/models--mlx-community--Qwen2.5-0.5B-Instruct-bf16/snapshots/56d07e766edd7159fbe12ed12d9cf114bf38bf1e"; + + #[test] + #[ignore = "requires local model files"] + fn test_load_qwen2_model() { + let model = super::load_qwen2_model(CACHED_TEST_MODEL_DIR).unwrap(); + assert_eq!(model.model_type(), "qwen2"); + } + + #[test] + #[ignore = "requires local model files"] + fn test_load_tokenizer() { + let tokenizer = load_qwen2_tokenizer(CACHED_TEST_MODEL_DIR).unwrap(); + let _encoding = tokenizer.encode("Hello, world!", true).unwrap(); + } + + #[test] + #[ignore = "requires local model files"] + fn test_load_and_run_qwen2_with_concat_cache() { + let tokenizer = load_qwen2_tokenizer(CACHED_TEST_MODEL_DIR).unwrap(); + let mut model = load_qwen2_model(CACHED_TEST_MODEL_DIR).unwrap(); + + let encoding = tokenizer.encode("hello", true).unwrap(); + let prompt_tokens = Array::from(encoding.get_ids()).index(NewAxis); + let mut cache = Vec::new(); + + let mut tokens = Vec::new(); + let generate = super::Generate::::new( + &mut model, + &mut cache, + 0.0, + &prompt_tokens, + ); + for (token, ntoks) in generate.zip(0..10) { + let token = token.unwrap(); + tokens.push(token.clone()); + + if ntoks == 0 { + eval(&tokens).unwrap(); + } + } + + eval(&tokens).unwrap(); + let slice: Vec = tokens.drain(..).map(|t| t.item::()).collect(); + let s = tokenizer.decode(&slice, true).unwrap(); + assert!(!s.is_empty()); + } +} From 445e8ac3d667a55a109e355c6ebc116919c12c71 Mon Sep 17 00:00:00 2001 From: James Dumay Date: Sun, 29 Mar 2026 19:45:57 +1100 Subject: [PATCH 03/15] Implement qwen2 attention and cache behavior --- mlx-lm/src/models/qwen2.rs | 577 ++++++++++++++++++++++++++++++++++--- 1 file changed, 539 insertions(+), 38 deletions(-) diff --git a/mlx-lm/src/models/qwen2.rs b/mlx-lm/src/models/qwen2.rs index 7bee4a782..c32fee1bd 100644 --- a/mlx-lm/src/models/qwen2.rs +++ b/mlx-lm/src/models/qwen2.rs @@ -1,19 +1,33 @@ -use std::path::Path; +use std::{ + collections::{HashMap, HashSet}, + path::Path, +}; -use mlx_rs::{error::Exception, module::ModuleParametersExt}; +use mlx_rs::{ + argmax_axis, array, + builder::Builder, + categorical, + error::Exception, + macros::{ModuleParameters, Quantizable}, + module::{Module, ModuleParametersExt}, + nn, + ops::indexing::{IndexOp, NewAxis}, + quantization::MaybeQuantized, + Array, +}; use serde::Deserialize; +use serde_json::Value; use tokenizers::Tokenizer; -use crate::{error::Error, models::llama, utils::rope::FloatOrString}; - -pub type Attention = llama::Attention; -pub type Mlp = llama::Mlp; -pub type TransformerBlock = llama::TransformerBlock; -pub type Qwen2Model = llama::LlamaModel; -pub type Model = llama::Model; -pub type ModelInput<'a, C> = llama::ModelInput<'a, C>; -pub type Generate<'a, C> = llama::Generate<'a, C>; -pub type GenerateState<'a> = llama::GenerateState<'a>; +use crate::{ + cache::KeyValueCache, + error::Error, + utils::{ + create_attention_mask, + rope::{initialize_rope, FloatOrString, RopeVariant}, + AttentionMask, + }, +}; #[derive(Debug, Clone, Deserialize)] pub struct ModelArgs { @@ -32,7 +46,7 @@ pub struct ModelArgs { #[serde(default)] pub rope_traditional: bool, #[serde(default)] - pub rope_scaling: Option>, + pub rope_scaling: Option>, #[serde(default = "default_true")] pub tie_word_embeddings: bool, } @@ -49,31 +63,429 @@ fn default_rope_theta() -> f32 { 1_000_000.0 } -impl From for llama::ModelArgs { - fn from(value: ModelArgs) -> Self { - let head_dim = value.hidden_size / value.num_attention_heads; - Self { - model_type: value.model_type, - hidden_size: value.hidden_size, - num_hidden_layers: value.num_hidden_layers, - intermediate_size: value.intermediate_size, - num_attention_heads: value.num_attention_heads, - rms_norm_eps: value.rms_norm_eps, - vocab_size: value.vocab_size, - num_key_value_heads: value.num_key_value_heads, - max_position_embeddings: value.max_position_embeddings, - rope_theta: value.rope_theta, +#[derive(Debug, Clone, ModuleParameters, Quantizable)] +pub struct Attention { + pub n_heads: i32, + pub n_kv_heads: i32, + pub scale: f32, + + #[quantizable] + #[param] + pub q_proj: MaybeQuantized, + #[quantizable] + #[param] + pub k_proj: MaybeQuantized, + #[quantizable] + #[param] + pub v_proj: MaybeQuantized, + #[quantizable] + #[param] + pub o_proj: MaybeQuantized, + #[param] + pub rope: RopeVariant, +} + +impl Attention { + pub fn new(args: &ModelArgs) -> Result { + let dim = args.hidden_size; + let n_heads = args.num_attention_heads; + let n_kv_heads = args.num_key_value_heads; + + let head_dim = args.hidden_size / n_heads; + let scale = (head_dim as f32).sqrt().recip(); + + let q_proj = nn::LinearBuilder::new(dim, n_heads * head_dim) + .bias(true) + .build()?; + let k_proj = nn::LinearBuilder::new(dim, n_kv_heads * head_dim) + .bias(true) + .build()?; + let v_proj = nn::LinearBuilder::new(dim, n_kv_heads * head_dim) + .bias(true) + .build()?; + let o_proj = nn::LinearBuilder::new(n_heads * head_dim, dim) + .bias(false) + .build()?; + + let rope = initialize_rope( head_dim, - tie_word_embeddings: value.tie_word_embeddings, - attention_bias: true, - mlp_bias: false, - rope_scaling: value.rope_scaling, + args.rope_theta, + args.rope_traditional, + &args.rope_scaling, + args.max_position_embeddings, + )?; + + Ok(Self { + n_heads, + n_kv_heads, + scale, + q_proj: MaybeQuantized::Original(q_proj), + k_proj: MaybeQuantized::Original(k_proj), + v_proj: MaybeQuantized::Original(v_proj), + o_proj: MaybeQuantized::Original(o_proj), + rope, + }) + } +} + +pub struct AttentionInput<'a, C> { + pub x: &'a Array, + pub mask: Option<&'a Array>, + pub cache: Option<&'a mut C>, +} + +impl Module> for Attention +where + C: KeyValueCache, +{ + type Output = Array; + type Error = Exception; + + #[allow(non_snake_case)] + fn forward(&mut self, input: AttentionInput<'_, C>) -> Result { + let AttentionInput { x, mask, mut cache } = input; + + let shape = x.shape(); + let B = shape[0]; + let L = shape[1]; + + let queries = self.q_proj.forward(x)?; + let keys = self.k_proj.forward(x)?; + let values = self.v_proj.forward(x)?; + + let mut queries = queries + .reshape(&[B, L, self.n_heads, -1])? + .transpose_axes(&[0, 2, 1, 3])?; + let mut keys = keys + .reshape(&[B, L, self.n_kv_heads, -1])? + .transpose_axes(&[0, 2, 1, 3])?; + let mut values = values + .reshape(&[B, L, self.n_kv_heads, -1])? + .transpose_axes(&[0, 2, 1, 3])?; + + if let Some(cache) = cache.as_mut() { + let q_input = nn::RopeInputBuilder::new(&queries) + .offset(cache.offset()) + .build()?; + queries = self.rope.forward(q_input)?; + let k_input = nn::RopeInputBuilder::new(&keys) + .offset(cache.offset()) + .build()?; + keys = self.rope.forward(k_input)?; + + (keys, values) = cache.update_and_fetch(keys, values)?; + } else { + queries = self.rope.forward(nn::RopeInput::new(&queries))?; + keys = self.rope.forward(nn::RopeInput::new(&keys))?; + } + + let output = crate::utils::scaled_dot_product_attention( + queries, keys, values, cache, self.scale, mask, + )? + .transpose_axes(&[0, 2, 1, 3])? + .reshape(&[B, L, -1])?; + + self.o_proj.forward(&output) + } + + fn training_mode(&mut self, mode: bool) { + self.q_proj.training_mode(mode); + self.k_proj.training_mode(mode); + self.v_proj.training_mode(mode); + self.o_proj.training_mode(mode); + >::training_mode(&mut self.rope, mode); + } +} + +#[derive(Debug, Clone, ModuleParameters, Quantizable)] +pub struct Mlp { + #[quantizable] + #[param] + pub gate_proj: MaybeQuantized, + + #[quantizable] + #[param] + pub down_proj: MaybeQuantized, + + #[quantizable] + #[param] + pub up_proj: MaybeQuantized, +} + +impl Mlp { + pub fn new(dim: i32, hidden_dim: i32) -> Result { + let gate_proj = nn::LinearBuilder::new(dim, hidden_dim) + .bias(false) + .build()?; + let down_proj = nn::LinearBuilder::new(hidden_dim, dim) + .bias(false) + .build()?; + let up_proj = nn::LinearBuilder::new(dim, hidden_dim) + .bias(false) + .build()?; + + Ok(Self { + gate_proj: MaybeQuantized::Original(gate_proj), + down_proj: MaybeQuantized::Original(down_proj), + up_proj: MaybeQuantized::Original(up_proj), + }) + } +} + +impl Module<&Array> for Mlp { + type Output = Array; + type Error = Exception; + + fn forward(&mut self, input: &Array) -> Result { + let down_proj_input = + nn::silu(self.gate_proj.forward(input)?)?.multiply(self.up_proj.forward(input)?)?; + self.down_proj.forward(&down_proj_input) + } + + fn training_mode(&mut self, mode: bool) { + self.gate_proj.training_mode(mode); + self.down_proj.training_mode(mode); + self.up_proj.training_mode(mode); + } +} + +#[derive(Debug, Clone, ModuleParameters, Quantizable)] +pub struct TransformerBlock { + pub num_attention_heads: i32, + pub hidden_size: i32, + + #[quantizable] + #[param] + pub self_attn: Attention, + + #[quantizable] + #[param] + pub mlp: Mlp, + + #[param] + pub input_layernorm: nn::RmsNorm, + + #[param] + pub post_attention_layernorm: nn::RmsNorm, +} + +impl TransformerBlock { + pub fn new(args: &ModelArgs) -> Result { + let self_attn = Attention::new(args)?; + let mlp = Mlp::new(args.hidden_size, args.intermediate_size)?; + let input_layernorm = nn::RmsNormBuilder::new(args.hidden_size) + .eps(args.rms_norm_eps) + .build()?; + let post_attention_layernorm = nn::RmsNormBuilder::new(args.hidden_size) + .eps(args.rms_norm_eps) + .build()?; + + Ok(Self { + num_attention_heads: args.num_attention_heads, + hidden_size: args.hidden_size, + self_attn, + mlp, + input_layernorm, + post_attention_layernorm, + }) + } +} + +impl Module> for TransformerBlock +where + C: KeyValueCache, +{ + type Output = Array; + type Error = Exception; + + fn forward(&mut self, input: AttentionInput<'_, C>) -> Result { + let AttentionInput { x, mask, cache } = input; + + let self_attn_input = AttentionInput { + x: &self.input_layernorm.forward(x)?, + mask, + cache, + }; + let r = self.self_attn.forward(self_attn_input)?; + let h = x.add(r)?; + + let r = self + .mlp + .forward(&self.post_attention_layernorm.forward(&h)?)?; + h.add(r) + } + + fn training_mode(&mut self, mode: bool) { + >>::training_mode(&mut self.self_attn, mode); + self.mlp.training_mode(mode); + self.input_layernorm.training_mode(mode); + self.post_attention_layernorm.training_mode(mode); + } +} + +#[derive(Debug, Clone, ModuleParameters, Quantizable)] +pub struct Qwen2Model { + pub vocab_size: i32, + pub num_hidden_layers: i32, + + #[quantizable] + #[param] + pub embed_tokens: MaybeQuantized, + + #[quantizable] + #[param] + pub layers: Vec, + + #[param] + pub norm: nn::RmsNorm, +} + +impl Qwen2Model { + pub fn new(args: &ModelArgs) -> Result { + assert!(args.vocab_size.is_positive()); + + let embed_tokens = nn::Embedding::new(args.vocab_size, args.hidden_size)?; + let layers = (0..args.num_hidden_layers) + .map(|_| TransformerBlock::new(args)) + .collect::, _>>()?; + let norm = nn::RmsNormBuilder::new(args.hidden_size) + .eps(args.rms_norm_eps) + .build()?; + + Ok(Self { + vocab_size: args.vocab_size, + num_hidden_layers: args.num_hidden_layers, + embed_tokens: MaybeQuantized::Original(embed_tokens), + layers, + norm, + }) + } +} + +pub struct ModelInput<'a, C> { + pub inputs: &'a Array, + pub mask: Option<&'a Array>, + pub cache: &'a mut Vec>, +} + +impl Module> for Qwen2Model +where + C: KeyValueCache, +{ + type Output = Array; + type Error = Exception; + + fn forward(&mut self, input: ModelInput<'_, C>) -> Result { + let ModelInput { + inputs, + mask, + cache, + } = input; + + let mut h = self.embed_tokens.forward(inputs)?; + + let mask = match mask { + Some(mask) => Some(mask.clone()), + None => match create_attention_mask(&h, cache, None)? { + Some(AttentionMask::Array(a)) => Some(a), + Some(AttentionMask::Causal) => None, + None => None, + }, + }; + + if cache.is_empty() { + *cache = (0..self.layers.len()).map(|_| None).collect(); + } + + for (layer, c) in self.layers.iter_mut().zip(cache.iter_mut()) { + let layer_input = AttentionInput { + x: &h, + mask: mask.as_ref(), + cache: c.as_mut(), + }; + h = layer.forward(layer_input)?; + } + + self.norm.forward(&h) + } + + fn training_mode(&mut self, mode: bool) { + self.embed_tokens.training_mode(mode); + for layer in &mut self.layers { + >>::training_mode(layer, mode); + } + self.norm.training_mode(mode); + } +} + +#[derive(Debug, Clone, ModuleParameters, Quantizable)] +pub struct Model { + pub args: ModelArgs, + + #[quantizable] + #[param] + pub model: Qwen2Model, + + #[quantizable] + #[param] + pub lm_head: Option>, +} + +impl Model { + pub fn new(args: ModelArgs) -> Result { + let model = Qwen2Model::new(&args)?; + let lm_head = if !args.tie_word_embeddings { + Some(MaybeQuantized::Original( + nn::LinearBuilder::new(args.hidden_size, args.vocab_size) + .bias(false) + .build()?, + )) + } else { + None + }; + + Ok(Self { + args, + model, + lm_head, + }) + } + + pub fn model_type(&self) -> &str { + &self.args.model_type + } +} + +impl Module> for Model +where + C: KeyValueCache, +{ + type Output = Array; + type Error = Exception; + + fn forward(&mut self, input: ModelInput<'_, C>) -> Result { + let out = self.model.forward(input)?; + + match self.lm_head.as_mut() { + Some(lm_head) => lm_head.forward(&out), + None => match &mut self.model.embed_tokens { + MaybeQuantized::Original(embed_tokens) => embed_tokens.as_linear(&out), + MaybeQuantized::Quantized(q_embed_tokens) => q_embed_tokens.as_linear(&out), + }, + } + } + + fn training_mode(&mut self, mode: bool) { + >>::training_mode(&mut self.model, mode); + if let Some(lm_head) = &mut self.lm_head { + lm_head.training_mode(mode); } } } pub fn load_qwen2_tokenizer(model_dir: impl AsRef) -> Result { - llama::load_llama_tokenizer(model_dir) + let file = model_dir.as_ref().join("tokenizer.json"); + Tokenizer::from_file(file).map_err(Into::into) } pub fn get_qwen2_model_args(model_dir: impl AsRef) -> Result { @@ -84,18 +496,23 @@ pub fn get_qwen2_model_args(model_dir: impl AsRef) -> Result, + pub weight_map: HashMap, +} + pub fn load_qwen2_model(model_dir: impl AsRef) -> Result { let model_dir = model_dir.as_ref(); let model_args = get_qwen2_model_args(model_dir)?; - let mut model = Model::new(model_args.into())?; + let mut model = Model::new(model_args)?; let weights_index = model_dir.join("model.safetensors.index.json"); if weights_index.exists() { let json = std::fs::read_to_string(weights_index)?; - let weight_map: llama::WeightMap = serde_json::from_str(&json)?; + let weight_map: WeightMap = serde_json::from_str(&json)?; - let weight_files: std::collections::HashSet<&String> = - weight_map.weight_map.values().collect(); + let weight_files: HashSet<&String> = weight_map.weight_map.values().collect(); for weight_file in weight_files { let weights_filename = model_dir.join(weight_file); model.load_safetensors(weights_filename)?; @@ -108,8 +525,92 @@ pub fn load_qwen2_model(model_dir: impl AsRef) -> Result { Ok(model) } -pub fn sample(logits: &mlx_rs::Array, temp: f32) -> Result { - llama::sample(logits, temp) +pub fn sample(logits: &Array, temp: f32) -> Result { + match temp { + 0.0 => argmax_axis!(logits, -1), + _ => { + let logits = logits.multiply(array!(1.0 / temp))?; + categorical!(logits) + } + } +} + +pub struct Generate<'a, C> { + model: &'a mut Model, + cache: &'a mut Vec>, + temp: f32, + state: GenerateState<'a>, +} + +impl<'a, C> Generate<'a, C> +where + C: KeyValueCache, +{ + pub fn new( + model: &'a mut Model, + cache: &'a mut Vec>, + temp: f32, + prompt_token: &'a Array, + ) -> Self { + Self { + model, + cache, + temp, + state: GenerateState::Prefill { prompt_token }, + } + } +} + +pub enum GenerateState<'a> { + Prefill { prompt_token: &'a Array }, + Decode { y: Array }, +} + +macro_rules! tri { + ($expr:expr) => { + match $expr { + Ok(val) => val, + Err(e) => return Some(Err(e.into())), + } + }; +} + +impl<'a, C> Iterator for Generate<'a, C> +where + C: KeyValueCache, +{ + type Item = Result; + + fn next(&mut self) -> Option { + match &self.state { + GenerateState::Prefill { prompt_token } => { + let input = ModelInput { + inputs: prompt_token, + mask: None, + cache: self.cache, + }; + let logits = tri!(self.model.forward(input)); + let y = tri!(sample(&logits.index((.., -1, ..)), self.temp)); + self.state = GenerateState::Decode { y: y.clone() }; + + Some(Ok(y)) + } + GenerateState::Decode { y } => { + let inputs = y.index((.., NewAxis)); + let input = ModelInput { + inputs: &inputs, + mask: None, + cache: self.cache, + }; + let logits = tri!(self.model.forward(input)); + let y = tri!(sample(&logits.index((.., -1, ..)), self.temp)); + + self.state = GenerateState::Decode { y: y.clone() }; + + Some(Ok(y)) + } + } + } } #[cfg(test)] From 6ecc803fb493509cceb250c9328a0ca7ce8be422 Mon Sep 17 00:00:00 2001 From: James Dumay Date: Sun, 29 Mar 2026 20:24:36 +1100 Subject: [PATCH 04/15] Initialize qwen2 kv caches by default --- mlx-lm/src/models/qwen2.rs | 10 +++++----- 1 file changed, 5 insertions(+), 5 deletions(-) diff --git a/mlx-lm/src/models/qwen2.rs b/mlx-lm/src/models/qwen2.rs index c32fee1bd..43f41a797 100644 --- a/mlx-lm/src/models/qwen2.rs +++ b/mlx-lm/src/models/qwen2.rs @@ -370,7 +370,7 @@ pub struct ModelInput<'a, C> { impl Module> for Qwen2Model where - C: KeyValueCache, + C: KeyValueCache + Default, { type Output = Array; type Error = Exception; @@ -394,7 +394,7 @@ where }; if cache.is_empty() { - *cache = (0..self.layers.len()).map(|_| None).collect(); + *cache = (0..self.layers.len()).map(|_| Some(C::default())).collect(); } for (layer, c) in self.layers.iter_mut().zip(cache.iter_mut()) { @@ -458,7 +458,7 @@ impl Model { impl Module> for Model where - C: KeyValueCache, + C: KeyValueCache + Default, { type Output = Array; type Error = Exception; @@ -544,7 +544,7 @@ pub struct Generate<'a, C> { impl<'a, C> Generate<'a, C> where - C: KeyValueCache, + C: KeyValueCache + Default, { pub fn new( model: &'a mut Model, @@ -577,7 +577,7 @@ macro_rules! tri { impl<'a, C> Iterator for Generate<'a, C> where - C: KeyValueCache, + C: KeyValueCache + Default, { type Item = Result; From 4d8a7da2c74f9b7e8cb429a7d6d25953f5d4935b Mon Sep 17 00:00:00 2001 From: James Dumay Date: Sun, 29 Mar 2026 20:47:16 +1100 Subject: [PATCH 05/15] Use additive causal mask for qwen2 prefill --- mlx-lm/src/models/qwen2.rs | 2 +- 1 file changed, 1 insertion(+), 1 deletion(-) diff --git a/mlx-lm/src/models/qwen2.rs b/mlx-lm/src/models/qwen2.rs index 43f41a797..87842abde 100644 --- a/mlx-lm/src/models/qwen2.rs +++ b/mlx-lm/src/models/qwen2.rs @@ -386,7 +386,7 @@ where let mask = match mask { Some(mask) => Some(mask.clone()), - None => match create_attention_mask(&h, cache, None)? { + None => match create_attention_mask(&h, cache, Some(true))? { Some(AttentionMask::Array(a)) => Some(a), Some(AttentionMask::Causal) => None, None => None, From 73265dfb9533066a5f9a771b25ce86cb9350f39a Mon Sep 17 00:00:00 2001 From: James Dumay Date: Sun, 29 Mar 2026 21:19:51 +1100 Subject: [PATCH 06/15] Add KV cache for native qwen2 generation --- mlx-lm/src/cache.rs | 134 ++++++++++++++++++++++++++++++++++++- mlx-lm/src/models/qwen2.rs | 29 +++++++- 2 files changed, 159 insertions(+), 4 deletions(-) diff --git a/mlx-lm/src/cache.rs b/mlx-lm/src/cache.rs index 67e759485..0399b6c40 100644 --- a/mlx-lm/src/cache.rs +++ b/mlx-lm/src/cache.rs @@ -1,6 +1,10 @@ use mlx_rs::{ error::Exception, - ops::{concatenate_axis, indexing::IndexOp}, + ops::{ + concatenate_axis, + indexing::{IndexOp, TryIndexMutOp}, + zeros_dtype, + }, transforms::eval, Array, }; @@ -69,6 +73,21 @@ pub struct ConcatKeyValueCache { offset: i32, } +#[derive(Debug, Clone, Default)] +pub struct KVCache { + keys: Option, + values: Option, + offset: i32, +} + +impl KVCache { + const STEP: i32 = 256; + + pub fn new() -> Self { + Self::default() + } +} + impl ConcatKeyValueCache { pub fn new() -> Self { Self::default() @@ -153,12 +172,95 @@ impl KeyValueCache for ConcatKeyValueCache { } } +impl KeyValueCache for KVCache { + fn offset(&self) -> i32 { + self.offset + } + + fn max_size(&self) -> Option { + None + } + + fn update_and_fetch( + &mut self, + keys: Array, + values: Array, + ) -> Result<(Array, Array), Exception> { + let prev = self.offset; + let seq_len = keys.shape()[2]; + + let needs_resize = self + .keys + .as_ref() + .map(|cached| (prev + seq_len) > cached.shape()[2]) + .unwrap_or(true); + + if needs_resize { + let key_shape = keys.shape(); + let value_shape = values.shape(); + let expand_steps = ((Self::STEP + seq_len - 1) / Self::STEP) * Self::STEP; + let new_keys = zeros_dtype( + &[key_shape[0], key_shape[1], expand_steps, key_shape[3]], + keys.dtype(), + )?; + let new_values = zeros_dtype( + &[value_shape[0], value_shape[1], expand_steps, value_shape[3]], + values.dtype(), + )?; + + match (self.keys.take(), self.values.take()) { + (Some(existing_keys), Some(existing_values)) => { + let existing_keys = if prev < existing_keys.shape()[2] { + existing_keys.index((.., .., ..prev, ..)) + } else { + existing_keys + }; + let existing_values = if prev < existing_values.shape()[2] { + existing_values.index((.., .., ..prev, ..)) + } else { + existing_values + }; + self.keys = Some(concatenate_axis(&[existing_keys, new_keys], -2)?); + self.values = Some(concatenate_axis(&[existing_values, new_values], -2)?); + } + _ => { + self.keys = Some(new_keys); + self.values = Some(new_values); + } + } + } + + self.offset += seq_len; + let end = self.offset; + self.keys + .as_mut() + .expect("keys cache missing") + .try_index_mut((.., .., prev..end, ..), &keys)?; + self.values + .as_mut() + .expect("values cache missing") + .try_index_mut((.., .., prev..end, ..), &values)?; + + let keys = self + .keys + .as_ref() + .expect("keys cache missing") + .index((.., .., ..end, ..)); + let values = self + .values + .as_ref() + .expect("values cache missing") + .index((.., .., ..end, ..)); + Ok((keys, values)) + } +} + /// TODO: A generic KV Cache pub struct DefaultKeyValueCache {} #[cfg(test)] mod tests { - use super::{ConcatKeyValueCache, KeyValueCache}; + use super::{ConcatKeyValueCache, KVCache, KeyValueCache}; use mlx_rs::Array; use std::sync::{Mutex, MutexGuard, OnceLock}; @@ -230,7 +332,33 @@ mod tests { .update_and_fetch(append_keys, append_values) .expect("append trimmed"); - assert_eq!(original_keys.as_slice::(), &[1., 2., 3., 4., 5., 6., 13., 14.]); + assert_eq!( + original_keys.as_slice::(), + &[1., 2., 3., 4., 5., 6., 13., 14.] + ); assert_eq!(trimmed_keys.as_slice::(), &[1., 2., 3., 4., 13., 14.]); } + + #[test] + fn kv_cache_appends_with_preallocated_capacity() { + let _guard = test_guard(); + let mut cache = KVCache::new(); + let keys = Array::from_slice(&[1f32, 2., 3., 4.], &[1, 1, 2, 2]); + let values = Array::from_slice(&[5f32, 6., 7., 8.], &[1, 1, 2, 2]); + let (keys, values) = cache.update_and_fetch(keys, values).expect("seed cache"); + assert_eq!(cache.offset(), 2); + assert_eq!(keys.shape(), &[1, 1, 2, 2]); + assert_eq!(values.shape(), &[1, 1, 2, 2]); + + let append_keys = Array::from_slice(&[9f32, 10.], &[1, 1, 1, 2]); + let append_values = Array::from_slice(&[11f32, 12.], &[1, 1, 1, 2]); + let (keys, values) = cache + .update_and_fetch(append_keys, append_values) + .expect("append cache"); + assert_eq!(cache.offset(), 3); + assert_eq!(keys.shape(), &[1, 1, 3, 2]); + assert_eq!(values.shape(), &[1, 1, 3, 2]); + assert_eq!(keys.as_slice::(), &[1., 2., 3., 4., 9., 10.]); + assert_eq!(values.as_slice::(), &[5., 6., 7., 8., 11., 12.]); + } } diff --git a/mlx-lm/src/models/qwen2.rs b/mlx-lm/src/models/qwen2.rs index 87842abde..f909a49fd 100644 --- a/mlx-lm/src/models/qwen2.rs +++ b/mlx-lm/src/models/qwen2.rs @@ -622,7 +622,7 @@ mod tests { }; use crate::{ - cache::ConcatKeyValueCache, + cache::{ConcatKeyValueCache, KVCache}, models::qwen2::{load_qwen2_model, load_qwen2_tokenizer}, }; @@ -674,4 +674,31 @@ mod tests { let s = tokenizer.decode(&slice, true).unwrap(); assert!(!s.is_empty()); } + + #[test] + #[ignore = "requires local model files"] + fn test_load_and_run_qwen2_with_kv_cache() { + let tokenizer = load_qwen2_tokenizer(CACHED_TEST_MODEL_DIR).unwrap(); + let mut model = load_qwen2_model(CACHED_TEST_MODEL_DIR).unwrap(); + + let encoding = tokenizer.encode("hello", true).unwrap(); + let prompt_tokens = Array::from(encoding.get_ids()).index(NewAxis); + let mut cache = Vec::new(); + + let mut tokens = Vec::new(); + let generate = super::Generate::::new(&mut model, &mut cache, 0.0, &prompt_tokens); + for (token, ntoks) in generate.zip(0..10) { + let token = token.unwrap(); + tokens.push(token.clone()); + + if ntoks == 0 { + eval(&tokens).unwrap(); + } + } + + eval(&tokens).unwrap(); + let slice: Vec = tokens.drain(..).map(|t| t.item::()).collect(); + let s = tokenizer.decode(&slice, true).unwrap(); + assert!(!s.is_empty()); + } } From 4cdae74a169c049738cd84a96319a0096cf6e94e Mon Sep 17 00:00:00 2001 From: James Dumay Date: Sun, 29 Mar 2026 21:41:00 +1100 Subject: [PATCH 07/15] Use causal mask fast path for qwen2 prefill --- mlx-lm/src/models/qwen2.rs | 2 +- 1 file changed, 1 insertion(+), 1 deletion(-) diff --git a/mlx-lm/src/models/qwen2.rs b/mlx-lm/src/models/qwen2.rs index f909a49fd..362255ffc 100644 --- a/mlx-lm/src/models/qwen2.rs +++ b/mlx-lm/src/models/qwen2.rs @@ -386,7 +386,7 @@ where let mask = match mask { Some(mask) => Some(mask.clone()), - None => match create_attention_mask(&h, cache, Some(true))? { + None => match create_attention_mask(&h, cache, None)? { Some(AttentionMask::Array(a)) => Some(a), Some(AttentionMask::Causal) => None, None => None, From c3a1894c9e73f46dcd6aaba311ebc7cb98a31661 Mon Sep 17 00:00:00 2001 From: James Dumay Date: Sun, 29 Mar 2026 21:52:06 +1100 Subject: [PATCH 08/15] Preserve causal attention masks in qwen2 --- mlx-lm/src/models/qwen2.rs | 37 +++++++++++++++++++++++++++---------- 1 file changed, 27 insertions(+), 10 deletions(-) diff --git a/mlx-lm/src/models/qwen2.rs b/mlx-lm/src/models/qwen2.rs index 362255ffc..cdff20d4f 100644 --- a/mlx-lm/src/models/qwen2.rs +++ b/mlx-lm/src/models/qwen2.rs @@ -8,6 +8,7 @@ use mlx_rs::{ builder::Builder, categorical, error::Exception, + fast::ScaledDotProductAttentionMask, macros::{ModuleParameters, Quantizable}, module::{Module, ModuleParametersExt}, nn, @@ -130,7 +131,7 @@ impl Attention { pub struct AttentionInput<'a, C> { pub x: &'a Array, - pub mask: Option<&'a Array>, + pub mask: Option<&'a AttentionMask>, pub cache: Option<&'a mut C>, } @@ -179,9 +180,29 @@ where keys = self.rope.forward(nn::RopeInput::new(&keys))?; } - let output = crate::utils::scaled_dot_product_attention( - queries, keys, values, cache, self.scale, mask, - )? + let output = match mask { + Some(AttentionMask::Array(mask)) => { + crate::utils::scaled_dot_product_attention( + queries, + keys, + values, + cache, + self.scale, + Some(mask), + )? + } + Some(AttentionMask::Causal) => mlx_rs::fast::scaled_dot_product_attention( + queries, + keys, + values, + self.scale, + Some(ScaledDotProductAttentionMask::Causal), + None, + )?, + None => crate::utils::scaled_dot_product_attention( + queries, keys, values, cache, self.scale, None, + )?, + } .transpose_axes(&[0, 2, 1, 3])? .reshape(&[B, L, -1])?; @@ -364,7 +385,7 @@ impl Qwen2Model { pub struct ModelInput<'a, C> { pub inputs: &'a Array, - pub mask: Option<&'a Array>, + pub mask: Option<&'a AttentionMask>, pub cache: &'a mut Vec>, } @@ -386,11 +407,7 @@ where let mask = match mask { Some(mask) => Some(mask.clone()), - None => match create_attention_mask(&h, cache, None)? { - Some(AttentionMask::Array(a)) => Some(a), - Some(AttentionMask::Causal) => None, - None => None, - }, + None => create_attention_mask(&h, cache, None)?, }; if cache.is_empty() { From fa8ac7bf0de94db44cb44d94b021f903d3176f02 Mon Sep 17 00:00:00 2001 From: James Dumay Date: Sun, 29 Mar 2026 21:59:50 +1100 Subject: [PATCH 09/15] Fix qwen2 model input state type --- mlx-lm/src/lib.rs | 5 +++-- 1 file changed, 3 insertions(+), 2 deletions(-) diff --git a/mlx-lm/src/lib.rs b/mlx-lm/src/lib.rs index dcc827ac6..a82ffa282 100644 --- a/mlx-lm/src/lib.rs +++ b/mlx-lm/src/lib.rs @@ -8,6 +8,7 @@ pub mod utils; use mlx_rs::Array; use crate::models::{qwen2, qwen3}; +use crate::utils::AttentionMask; pub struct ModelInputBuilder<'a, C, T> { pub y: &'a Array, @@ -31,8 +32,8 @@ impl<'a, C> ModelInput<'a, C, Option> for qwen3::ModelInput<'a, C> { } } -impl<'a, C> ModelInput<'a, C, Option> for qwen2::ModelInput<'a, C> { - fn from_model_input_builder(builder: ModelInputBuilder<'a, C, Option>) -> Self { +impl<'a, C> ModelInput<'a, C, Option> for qwen2::ModelInput<'a, C> { + fn from_model_input_builder(builder: ModelInputBuilder<'a, C, Option>) -> Self { let ModelInputBuilder { y, cache, state } = builder; Self { From 153a8200f188ab6628bfaeb1e847aea79e944f6a Mon Sep 17 00:00:00 2001 From: James Dumay Date: Sun, 29 Mar 2026 22:11:44 +1100 Subject: [PATCH 10/15] Export attention mask type for qwen2 inputs --- mlx-lm/src/utils/mod.rs | 2 +- 1 file changed, 1 insertion(+), 1 deletion(-) diff --git a/mlx-lm/src/utils/mod.rs b/mlx-lm/src/utils/mod.rs index d6571b6e8..3b5dcfcf3 100644 --- a/mlx-lm/src/utils/mod.rs +++ b/mlx-lm/src/utils/mod.rs @@ -252,7 +252,7 @@ where } #[derive(Debug, Clone)] -pub(crate) enum AttentionMask { +pub enum AttentionMask { Array(Array), Causal, } From e957ecc35957512c2e41ea1ce05be4acab690ddf Mon Sep 17 00:00:00 2001 From: James Dumay Date: Sun, 29 Mar 2026 22:21:00 +1100 Subject: [PATCH 11/15] Force eval after qwen2 KV cache updates --- mlx-lm/src/cache.rs | 5 +++++ 1 file changed, 5 insertions(+) diff --git a/mlx-lm/src/cache.rs b/mlx-lm/src/cache.rs index 0399b6c40..74bbd9f8b 100644 --- a/mlx-lm/src/cache.rs +++ b/mlx-lm/src/cache.rs @@ -240,6 +240,10 @@ impl KeyValueCache for KVCache { .as_mut() .expect("values cache missing") .try_index_mut((.., .., prev..end, ..), &values)?; + eval([ + self.keys.as_ref().expect("keys cache missing"), + self.values.as_ref().expect("values cache missing"), + ])?; let keys = self .keys @@ -251,6 +255,7 @@ impl KeyValueCache for KVCache { .as_ref() .expect("values cache missing") .index((.., .., ..end, ..)); + eval([&keys, &values])?; Ok((keys, values)) } } From 59dcfb62d9205506c6026cedffa1a9be3a958af5 Mon Sep 17 00:00:00 2001 From: James Dumay Date: Sun, 29 Mar 2026 22:36:30 +1100 Subject: [PATCH 12/15] Clone qwen2 KV cache prefix views before reuse --- mlx-lm/src/cache.rs | 2 +- 1 file changed, 1 insertion(+), 1 deletion(-) diff --git a/mlx-lm/src/cache.rs b/mlx-lm/src/cache.rs index 74bbd9f8b..cbcc1b9c1 100644 --- a/mlx-lm/src/cache.rs +++ b/mlx-lm/src/cache.rs @@ -256,7 +256,7 @@ impl KeyValueCache for KVCache { .expect("values cache missing") .index((.., .., ..end, ..)); eval([&keys, &values])?; - Ok((keys, values)) + Ok((keys.deep_clone(), values.deep_clone())) } } From bb134fb9a08b2bb58520d9ced58cf8cf4f9cbd63 Mon Sep 17 00:00:00 2001 From: James Dumay Date: Mon, 30 Mar 2026 06:07:59 +1100 Subject: [PATCH 13/15] Fix native qwen2 cached decode --- mlx-lm/src/cache.rs | 5 +- mlx-lm/src/models/qwen2.rs | 864 ++++++++++++++++++++++++++++++++++++- 2 files changed, 858 insertions(+), 11 deletions(-) diff --git a/mlx-lm/src/cache.rs b/mlx-lm/src/cache.rs index cbcc1b9c1..5458e1655 100644 --- a/mlx-lm/src/cache.rs +++ b/mlx-lm/src/cache.rs @@ -93,6 +93,10 @@ impl ConcatKeyValueCache { Self::default() } + pub(crate) fn arrays(&self) -> (Option<&Array>, Option<&Array>) { + (self.keys.as_ref(), self.values.as_ref()) + } + pub fn trim_to(&mut self, token_count: i32) -> Result<(), Exception> { let token_count = token_count.max(0); match (&self.keys, &self.values) { @@ -164,7 +168,6 @@ impl KeyValueCache for ConcatKeyValueCache { } let shape = self.keys.as_ref().expect("Keys cannot be None").shape(); self.offset = shape[shape.len() - 2]; - Ok(( self.keys.clone().expect("Keys cannot be None"), self.values.clone().expect("Values cannot be None"), diff --git a/mlx-lm/src/models/qwen2.rs b/mlx-lm/src/models/qwen2.rs index cdff20d4f..a8ec5044f 100644 --- a/mlx-lm/src/models/qwen2.rs +++ b/mlx-lm/src/models/qwen2.rs @@ -8,7 +8,7 @@ use mlx_rs::{ builder::Builder, categorical, error::Exception, - fast::ScaledDotProductAttentionMask, + fast::{rope_dynamic, ScaledDotProductAttentionMask}, macros::{ModuleParameters, Quantizable}, module::{Module, ModuleParametersExt}, nn, @@ -129,6 +129,38 @@ impl Attention { } } +fn apply_cached_rope(rope: &mut RopeVariant, x: &Array, offset: i32) -> Result { + let seq_len = x.shape()[x.shape().len() - 2]; + + if seq_len != 1 { + let rope_input = nn::RopeInputBuilder::new(x).offset(offset).build()?; + return rope.forward(rope_input); + } + + let position_array = Array::from_int(offset); + + match rope { + RopeVariant::Default(rope) => rope_dynamic( + x, + rope.dimensions, + rope.traditional, + Some(rope.base), + rope.scale, + &position_array, + None::<&Array>, + ), + RopeVariant::Llama3(rope) => rope_dynamic( + x, + rope.dimensions, + rope.traditional, + None::, + rope.scale, + &position_array, + Some(&rope.freqs), + ), + } +} + pub struct AttentionInput<'a, C> { pub x: &'a Array, pub mask: Option<&'a AttentionMask>, @@ -165,14 +197,8 @@ where .transpose_axes(&[0, 2, 1, 3])?; if let Some(cache) = cache.as_mut() { - let q_input = nn::RopeInputBuilder::new(&queries) - .offset(cache.offset()) - .build()?; - queries = self.rope.forward(q_input)?; - let k_input = nn::RopeInputBuilder::new(&keys) - .offset(cache.offset()) - .build()?; - keys = self.rope.forward(k_input)?; + queries = apply_cached_rope(&mut self.rope, &queries, cache.offset())?; + keys = apply_cached_rope(&mut self.rope, &keys, cache.offset())?; (keys, values) = cache.update_and_fetch(keys, values)?; } else { @@ -633,14 +659,19 @@ where #[cfg(test)] mod tests { use mlx_rs::{ + builder::Builder, + fast::rope_dynamic, + nn::RopeInputBuilder, + module::Module, ops::indexing::{IndexOp, NewAxis}, transforms::eval, Array, }; use crate::{ - cache::{ConcatKeyValueCache, KVCache}, + cache::{ConcatKeyValueCache, KVCache, KeyValueCache}, models::qwen2::{load_qwen2_model, load_qwen2_tokenizer}, + utils::{create_causal_mask, AttentionMask}, }; const CACHED_TEST_MODEL_DIR: &str = @@ -718,4 +749,817 @@ mod tests { let s = tokenizer.decode(&slice, true).unwrap(); assert!(!s.is_empty()); } + + #[test] + #[ignore = "requires local model files"] + fn test_qwen2_concat_cache_benchmark_prompt_trace() { + let tokenizer = load_qwen2_tokenizer(CACHED_TEST_MODEL_DIR).unwrap(); + let mut model = load_qwen2_model(CACHED_TEST_MODEL_DIR).unwrap(); + + let prompt = "<|im_start|>system\nYou are a concise assistant.<|im_end|>\n<|im_start|>user\nContext: The native MLX path supports streaming, but some families behave inconsistently.\nTask: Give three short bullets on what is working, what is not, and the next fix.<|im_end|>\n<|im_start|>assistant\n"; + let encoding = tokenizer.encode(prompt, false).unwrap(); + let prompt_tokens = Array::from(encoding.get_ids()).index(NewAxis); + let mut cache = Vec::new(); + + let mut tokens = Vec::new(); + let generate = super::Generate::::new( + &mut model, + &mut cache, + 0.0, + &prompt_tokens, + ); + for token in generate.take(32) { + tokens.push(token.unwrap()); + } + + eval(&tokens).unwrap(); + let ids: Vec = tokens.iter().map(|t| t.item::()).collect(); + let text = tokenizer.decode(&ids, true).unwrap(); + eprintln!("qwen2-concat-trace ids={ids:?} text={text:?}"); + assert!(!text.is_empty()); + } + + #[test] + #[ignore = "requires local model files"] + fn test_qwen2_kv_cache_benchmark_prompt_trace() { + let tokenizer = load_qwen2_tokenizer(CACHED_TEST_MODEL_DIR).unwrap(); + let mut model = load_qwen2_model(CACHED_TEST_MODEL_DIR).unwrap(); + + let prompt = "<|im_start|>system\nYou are a concise assistant.<|im_end|>\n<|im_start|>user\nContext: The native MLX path supports streaming, but some families behave inconsistently.\nTask: Give three short bullets on what is working, what is not, and the next fix.<|im_end|>\n<|im_start|>assistant\n"; + let encoding = tokenizer.encode(prompt, false).unwrap(); + let prompt_tokens = Array::from(encoding.get_ids()).index(NewAxis); + let mut cache = Vec::new(); + + let mut tokens = Vec::new(); + let generate = + super::Generate::::new(&mut model, &mut cache, 0.0, &prompt_tokens); + for token in generate.take(32) { + tokens.push(token.unwrap()); + } + + eval(&tokens).unwrap(); + let ids: Vec = tokens.iter().map(|t| t.item::()).collect(); + let text = tokenizer.decode(&ids, true).unwrap(); + eprintln!("qwen2-kv-trace ids={ids:?} text={text:?}"); + assert!(!text.is_empty()); + } + + #[test] + #[ignore = "requires local model files"] + fn test_qwen2_concat_cache_step6_topk() { + let tokenizer = load_qwen2_tokenizer(CACHED_TEST_MODEL_DIR).unwrap(); + let mut model = load_qwen2_model(CACHED_TEST_MODEL_DIR).unwrap(); + + let prompt = "<|im_start|>system\nYou are a concise assistant.<|im_end|>\n<|im_start|>user\nContext: The native MLX path supports streaming, but some families behave inconsistently.\nTask: Give three short bullets on what is working, what is not, and the next fix.<|im_end|>\n<|im_start|>assistant\n"; + let encoding = tokenizer.encode(prompt, false).unwrap(); + let prompt_tokens = Array::from(encoding.get_ids()).index(NewAxis); + let mut cache = Vec::>::new(); + + let input = super::ModelInput { + inputs: &prompt_tokens, + mask: None, + cache: &mut cache, + }; + let logits = model.forward(input).unwrap(); + let mut step_logits = logits.index((.., -1, ..)); + eval([&step_logits]).unwrap(); + + let mut ids = Vec::new(); + for step in 0..8 { + let step_logits_f32 = step_logits.as_type::().unwrap(); + let flat = step_logits_f32.as_slice::(); + let mut ranked: Vec<(usize, f32)> = flat.iter().copied().enumerate().collect(); + ranked.sort_by(|a, b| b.1.total_cmp(&a.1)); + if step == 5 { + let top: Vec<(usize, f32)> = ranked.into_iter().take(10).collect(); + eprintln!("qwen2-step6-topk {top:?}"); + } + + let y = super::sample(&step_logits, 0.0).unwrap(); + let token_id = y.item::(); + ids.push(token_id); + let inputs = y.index((.., NewAxis)); + let input = super::ModelInput { + inputs: &inputs, + mask: None, + cache: &mut cache, + }; + let logits = model.forward(input).unwrap(); + step_logits = logits.index((.., -1, ..)); + eval([&step_logits]).unwrap(); + } + + let text = tokenizer.decode(&ids, true).unwrap(); + eprintln!("qwen2-step6-trace ids={ids:?} text={text:?}"); + assert!(!ids.is_empty()); + } + + #[test] + #[ignore = "requires local model files"] + fn test_qwen2_full_prefill_vs_incremental_trace() { + let tokenizer = load_qwen2_tokenizer(CACHED_TEST_MODEL_DIR).unwrap(); + let mut model = load_qwen2_model(CACHED_TEST_MODEL_DIR).unwrap(); + + let prompt = "<|im_start|>system\nYou are a concise assistant.<|im_end|>\n<|im_start|>user\nContext: The native MLX path supports streaming, but some families behave inconsistently.\nTask: Give three short bullets on what is working, what is not, and the next fix.<|im_end|>\n<|im_start|>assistant\n"; + let encoding = tokenizer.encode(prompt, false).unwrap(); + let prompt_ids = encoding.get_ids().to_vec(); + + let mut incremental_ids = Vec::new(); + let mut cache = Vec::>::new(); + let prompt_tokens = Array::from(prompt_ids.as_slice()).index(NewAxis); + let input = super::ModelInput { + inputs: &prompt_tokens, + mask: None, + cache: &mut cache, + }; + let logits = model.forward(input).unwrap(); + let mut step_logits = logits.index((.., -1, ..)); + eval([&step_logits]).unwrap(); + for _ in 0..12 { + let y = super::sample(&step_logits, 0.0).unwrap(); + let token_id = y.item::(); + incremental_ids.push(token_id); + let inputs = y.index((.., NewAxis)); + let input = super::ModelInput { + inputs: &inputs, + mask: None, + cache: &mut cache, + }; + let logits = model.forward(input).unwrap(); + step_logits = logits.index((.., -1, ..)); + eval([&step_logits]).unwrap(); + } + + let mut full_prefill_ids = Vec::new(); + for _ in 0..12 { + let all_ids: Vec = prompt_ids + .iter() + .copied() + .chain(full_prefill_ids.iter().copied()) + .collect(); + let tokens = Array::from(all_ids.as_slice()).index(NewAxis); + let mut fresh_cache = Vec::>::new(); + let input = super::ModelInput { + inputs: &tokens, + mask: None, + cache: &mut fresh_cache, + }; + let logits = model.forward(input).unwrap(); + let step_logits = logits.index((.., -1, ..)); + eval([&step_logits]).unwrap(); + let y = super::sample(&step_logits, 0.0).unwrap(); + full_prefill_ids.push(y.item::()); + } + + let incremental_text = tokenizer.decode(&incremental_ids, true).unwrap(); + let full_prefill_text = tokenizer.decode(&full_prefill_ids, true).unwrap(); + eprintln!("qwen2-incremental ids={incremental_ids:?} text={incremental_text:?}"); + eprintln!("qwen2-full-prefill ids={full_prefill_ids:?} text={full_prefill_text:?}"); + assert!(!incremental_ids.is_empty()); + assert!(!full_prefill_ids.is_empty()); + } + + #[test] + #[ignore = "requires local model files"] + fn test_qwen2_incremental_with_forced_causal_mask_trace() { + let tokenizer = load_qwen2_tokenizer(CACHED_TEST_MODEL_DIR).unwrap(); + let mut model = load_qwen2_model(CACHED_TEST_MODEL_DIR).unwrap(); + + let prompt = "<|im_start|>system\nYou are a concise assistant.<|im_end|>\n<|im_start|>user\nContext: The native MLX path supports streaming, but some families behave inconsistently.\nTask: Give three short bullets on what is working, what is not, and the next fix.<|im_end|>\n<|im_start|>assistant\n"; + let encoding = tokenizer.encode(prompt, false).unwrap(); + let prompt_tokens = Array::from(encoding.get_ids()).index(NewAxis); + let mut cache = Vec::>::new(); + + let input = super::ModelInput { + inputs: &prompt_tokens, + mask: None, + cache: &mut cache, + }; + let logits = model.forward(input).unwrap(); + let mut step_logits = logits.index((.., -1, ..)); + eval([&step_logits]).unwrap(); + + let mut ids = Vec::new(); + let causal = AttentionMask::Causal; + for _ in 0..12 { + let y = super::sample(&step_logits, 0.0).unwrap(); + let token_id = y.item::(); + ids.push(token_id); + let inputs = y.index((.., NewAxis)); + let input = super::ModelInput { + inputs: &inputs, + mask: Some(&causal), + cache: &mut cache, + }; + let logits = model.forward(input).unwrap(); + step_logits = logits.index((.., -1, ..)); + eval([&step_logits]).unwrap(); + } + + let text = tokenizer.decode(&ids, true).unwrap(); + eprintln!("qwen2-forced-causal ids={ids:?} text={text:?}"); + assert!(!ids.is_empty()); + } + + #[test] + #[ignore = "requires local model files"] + fn test_qwen2_incremental_vs_full_prefill_next_token_logits() { + let tokenizer = load_qwen2_tokenizer(CACHED_TEST_MODEL_DIR).unwrap(); + let mut model = load_qwen2_model(CACHED_TEST_MODEL_DIR).unwrap(); + + let prompt = "<|im_start|>system\nYou are a concise assistant.<|im_end|>\n<|im_start|>user\nContext: The native MLX path supports streaming, but some families behave inconsistently.\nTask: Give three short bullets on what is working, what is not, and the next fix.<|im_end|>\n<|im_start|>assistant\n"; + let encoding = tokenizer.encode(prompt, false).unwrap(); + let prompt_ids = encoding.get_ids().to_vec(); + let agreed_prefix = vec![12u32, 576, 9867, 19614, 55]; + + let prompt_tokens = Array::from(prompt_ids.as_slice()).index(NewAxis); + let mut cache = Vec::>::new(); + let input = super::ModelInput { + inputs: &prompt_tokens, + mask: None, + cache: &mut cache, + }; + let logits = model.forward(input).unwrap(); + let mut incremental_logits = logits.index((.., -1, ..)); + eval([&incremental_logits]).unwrap(); + for token_id in &agreed_prefix { + let y = Array::from_slice(&[*token_id], &[1]); + let inputs = y.index((.., NewAxis)); + let input = super::ModelInput { + inputs: &inputs, + mask: None, + cache: &mut cache, + }; + let logits = model.forward(input).unwrap(); + incremental_logits = logits.index((.., -1, ..)); + eval([&incremental_logits]).unwrap(); + } + + let all_ids: Vec = prompt_ids + .iter() + .copied() + .chain(agreed_prefix.iter().copied()) + .collect(); + let all_tokens = Array::from(all_ids.as_slice()).index(NewAxis); + let mut fresh_cache = Vec::>::new(); + let input = super::ModelInput { + inputs: &all_tokens, + mask: None, + cache: &mut fresh_cache, + }; + let logits = model.forward(input).unwrap(); + let full_prefill_logits = logits.index((.., -1, ..)); + eval([&full_prefill_logits]).unwrap(); + + let inc_f32 = incremental_logits.as_type::().unwrap(); + let full_f32 = full_prefill_logits.as_type::().unwrap(); + let inc = inc_f32.as_slice::(); + let full = full_f32.as_slice::(); + + let mut inc_ranked: Vec<(usize, f32)> = inc.iter().copied().enumerate().collect(); + let mut full_ranked: Vec<(usize, f32)> = full.iter().copied().enumerate().collect(); + inc_ranked.sort_by(|a, b| b.1.total_cmp(&a.1)); + full_ranked.sort_by(|a, b| b.1.total_cmp(&a.1)); + + let inc_top: Vec<(usize, f32)> = inc_ranked.into_iter().take(10).collect(); + let full_top: Vec<(usize, f32)> = full_ranked.into_iter().take(10).collect(); + eprintln!("qwen2-incremental-top {inc_top:?}"); + eprintln!("qwen2-full-prefill-top {full_top:?}"); + + let inc_best = Array::from_slice(&[inc_top[0].0 as u32], &[1]); + let full_best = Array::from_slice(&[full_top[0].0 as u32], &[1]); + let inc_best_text = tokenizer.decode(&[inc_top[0].0 as u32], true).unwrap(); + let full_best_text = tokenizer.decode(&[full_top[0].0 as u32], true).unwrap(); + eprintln!("qwen2-incremental-best {:?} {:?}", inc_best.item::(), inc_best_text); + eprintln!("qwen2-full-prefill-best {:?} {:?}", full_best.item::(), full_best_text); + + assert_eq!(inc_top[0].0, full_top[0].0); + } + + #[test] + #[ignore = "requires local model files"] + fn test_qwen2_chunk_prefill_then_single_step_vs_incremental_next_token_logits() { + let tokenizer = load_qwen2_tokenizer(CACHED_TEST_MODEL_DIR).unwrap(); + let mut model = load_qwen2_model(CACHED_TEST_MODEL_DIR).unwrap(); + + let prompt = "<|im_start|>system\nYou are a concise assistant.<|im_end|>\n<|im_start|>user\nContext: The native MLX path supports streaming, but some families behave inconsistently.\nTask: Give three short bullets on what is working, what is not, and the next fix.<|im_end|>\n<|im_start|>assistant\n"; + let encoding = tokenizer.encode(prompt, false).unwrap(); + let prompt_ids = encoding.get_ids().to_vec(); + let agreed_prefix = vec![12u32, 576, 9867, 19614, 55]; + + let prompt_tokens = Array::from(prompt_ids.as_slice()).index(NewAxis); + let mut incremental_cache = Vec::>::new(); + let input = super::ModelInput { + inputs: &prompt_tokens, + mask: None, + cache: &mut incremental_cache, + }; + let logits = model.forward(input).unwrap(); + let mut incremental_logits = logits.index((.., -1, ..)); + eval([&incremental_logits]).unwrap(); + for token_id in &agreed_prefix { + let y = Array::from_slice(&[*token_id], &[1]); + let inputs = y.index((.., NewAxis)); + let input = super::ModelInput { + inputs: &inputs, + mask: None, + cache: &mut incremental_cache, + }; + let logits = model.forward(input).unwrap(); + incremental_logits = logits.index((.., -1, ..)); + eval([&incremental_logits]).unwrap(); + } + + let prefixed_ids: Vec = prompt_ids + .iter() + .copied() + .chain(agreed_prefix[..agreed_prefix.len() - 1].iter().copied()) + .collect(); + let prefixed_tokens = Array::from(prefixed_ids.as_slice()).index(NewAxis); + let mut chunked_cache = Vec::>::new(); + let input = super::ModelInput { + inputs: &prefixed_tokens, + mask: None, + cache: &mut chunked_cache, + }; + let _ = model.forward(input).unwrap(); + + let last_prefix_token = + Array::from_slice(&[agreed_prefix[agreed_prefix.len() - 1]], &[1]).index(NewAxis); + let input = super::ModelInput { + inputs: &last_prefix_token, + mask: None, + cache: &mut chunked_cache, + }; + let logits = model.forward(input).unwrap(); + let chunked_logits = logits.index((.., -1, ..)); + eval([&chunked_logits]).unwrap(); + + let all_ids: Vec = prompt_ids + .iter() + .copied() + .chain(agreed_prefix.iter().copied()) + .collect(); + let all_tokens = Array::from(all_ids.as_slice()).index(NewAxis); + let mut full_prefill_cache = Vec::>::new(); + let input = super::ModelInput { + inputs: &all_tokens, + mask: None, + cache: &mut full_prefill_cache, + }; + let logits = model.forward(input).unwrap(); + let full_prefill_logits = logits.index((.., -1, ..)); + eval([&full_prefill_logits]).unwrap(); + + let collect_top = |logits: &Array| -> Vec<(usize, f32)> { + let logits_f32 = logits.as_type::().unwrap(); + let flat = logits_f32.as_slice::(); + let mut ranked: Vec<(usize, f32)> = flat.iter().copied().enumerate().collect(); + ranked.sort_by(|a, b| b.1.total_cmp(&a.1)); + ranked.into_iter().take(10).collect() + }; + + let incremental_top = collect_top(&incremental_logits); + let chunked_top = collect_top(&chunked_logits); + let full_prefill_top = collect_top(&full_prefill_logits); + + eprintln!("qwen2-incremental-after-prefix-top {incremental_top:?}"); + eprintln!("qwen2-chunked-after-prefix-top {chunked_top:?}"); + eprintln!("qwen2-full-prefill-after-prefix-top {full_prefill_top:?}"); + + let decode_best = |token_id: usize| -> String { + tokenizer.decode(&[token_id as u32], true).unwrap() + }; + + eprintln!( + "qwen2-best incremental={:?} chunked={:?} full_prefill={:?}", + (incremental_top[0].0, decode_best(incremental_top[0].0)), + (chunked_top[0].0, decode_best(chunked_top[0].0)), + (full_prefill_top[0].0, decode_best(full_prefill_top[0].0)), + ); + + assert_eq!(chunked_top[0].0, full_prefill_top[0].0); + assert_eq!(incremental_top[0].0, full_prefill_top[0].0); + } + + #[test] + #[ignore = "requires local model files"] + fn test_qwen2_chunk_prefill_single_step_with_explicit_array_mask() { + let tokenizer = load_qwen2_tokenizer(CACHED_TEST_MODEL_DIR).unwrap(); + let mut model = load_qwen2_model(CACHED_TEST_MODEL_DIR).unwrap(); + + let prompt = "<|im_start|>system\nYou are a concise assistant.<|im_end|>\n<|im_start|>user\nContext: The native MLX path supports streaming, but some families behave inconsistently.\nTask: Give three short bullets on what is working, what is not, and the next fix.<|im_end|>\n<|im_start|>assistant\n"; + let encoding = tokenizer.encode(prompt, false).unwrap(); + let prompt_ids = encoding.get_ids().to_vec(); + let agreed_prefix = vec![12u32, 576, 9867, 19614, 55]; + + let prefixed_ids: Vec = prompt_ids + .iter() + .copied() + .chain(agreed_prefix[..agreed_prefix.len() - 1].iter().copied()) + .collect(); + let prefixed_tokens = Array::from(prefixed_ids.as_slice()).index(NewAxis); + let mut chunked_cache = Vec::>::new(); + let input = super::ModelInput { + inputs: &prefixed_tokens, + mask: None, + cache: &mut chunked_cache, + }; + let _ = model.forward(input).unwrap(); + + let offset = chunked_cache[0].as_ref().unwrap().offset(); + let explicit_mask = create_causal_mask(1, Some(offset), None, None).unwrap(); + let explicit_mask = AttentionMask::Array(explicit_mask); + + let last_prefix_token = + Array::from_slice(&[agreed_prefix[agreed_prefix.len() - 1]], &[1]).index(NewAxis); + let input = super::ModelInput { + inputs: &last_prefix_token, + mask: Some(&explicit_mask), + cache: &mut chunked_cache, + }; + let logits = model.forward(input).unwrap(); + let masked_logits = logits.index((.., -1, ..)); + eval([&masked_logits]).unwrap(); + + let all_ids: Vec = prompt_ids + .iter() + .copied() + .chain(agreed_prefix.iter().copied()) + .collect(); + let all_tokens = Array::from(all_ids.as_slice()).index(NewAxis); + let mut full_prefill_cache = Vec::>::new(); + let input = super::ModelInput { + inputs: &all_tokens, + mask: None, + cache: &mut full_prefill_cache, + }; + let logits = model.forward(input).unwrap(); + let full_prefill_logits = logits.index((.., -1, ..)); + eval([&full_prefill_logits]).unwrap(); + + let collect_top = |logits: &Array| -> Vec<(usize, f32)> { + let logits_f32 = logits.as_type::().unwrap(); + let flat = logits_f32.as_slice::(); + let mut ranked: Vec<(usize, f32)> = flat.iter().copied().enumerate().collect(); + ranked.sort_by(|a, b| b.1.total_cmp(&a.1)); + ranked.into_iter().take(10).collect() + }; + + let masked_top = collect_top(&masked_logits); + let full_prefill_top = collect_top(&full_prefill_logits); + eprintln!("qwen2-explicit-array-mask-top {masked_top:?}"); + eprintln!("qwen2-full-prefill-after-prefix-top {full_prefill_top:?}"); + eprintln!( + "qwen2-best masked={:?} full_prefill={:?}", + (masked_top[0].0, tokenizer.decode(&[masked_top[0].0 as u32], true).unwrap()), + ( + full_prefill_top[0].0, + tokenizer + .decode(&[full_prefill_top[0].0 as u32], true) + .unwrap() + ), + ); + + assert_eq!(masked_top[0].0, full_prefill_top[0].0); + } + + #[test] + #[ignore = "requires local model files"] + fn test_qwen2_chunk_prefill_cache_state_vs_full_prefill_cache_state() { + let mut model = load_qwen2_model(CACHED_TEST_MODEL_DIR).unwrap(); + + let tokenizer = load_qwen2_tokenizer(CACHED_TEST_MODEL_DIR).unwrap(); + let prompt = "<|im_start|>system\nYou are a concise assistant.<|im_end|>\n<|im_start|>user\nContext: The native MLX path supports streaming, but some families behave inconsistently.\nTask: Give three short bullets on what is working, what is not, and the next fix.<|im_end|>\n<|im_start|>assistant\n"; + let encoding = tokenizer.encode(prompt, false).unwrap(); + let prompt_ids = encoding.get_ids().to_vec(); + let agreed_prefix = vec![12u32, 576, 9867, 19614, 55]; + + let prefixed_ids: Vec = prompt_ids + .iter() + .copied() + .chain(agreed_prefix[..agreed_prefix.len() - 1].iter().copied()) + .collect(); + let prefixed_tokens = Array::from(prefixed_ids.as_slice()).index(NewAxis); + let mut chunked_cache = Vec::>::new(); + let input = super::ModelInput { + inputs: &prefixed_tokens, + mask: None, + cache: &mut chunked_cache, + }; + let _ = model.forward(input).unwrap(); + let last_prefix_token = + Array::from_slice(&[agreed_prefix[agreed_prefix.len() - 1]], &[1]).index(NewAxis); + let input = super::ModelInput { + inputs: &last_prefix_token, + mask: None, + cache: &mut chunked_cache, + }; + let _ = model.forward(input).unwrap(); + + let all_ids: Vec = prompt_ids + .iter() + .copied() + .chain(agreed_prefix.iter().copied()) + .collect(); + let all_tokens = Array::from(all_ids.as_slice()).index(NewAxis); + let mut full_prefill_cache = Vec::>::new(); + let input = super::ModelInput { + inputs: &all_tokens, + mask: None, + cache: &mut full_prefill_cache, + }; + let _ = model.forward(input).unwrap(); + + let chunked_first = chunked_cache[0].as_ref().unwrap(); + let full_first = full_prefill_cache[0].as_ref().unwrap(); + let (chunked_keys, chunked_values) = chunked_first.arrays(); + let (full_keys, full_values) = full_first.arrays(); + let chunked_keys = chunked_keys.unwrap(); + let chunked_values = chunked_values.unwrap(); + let full_keys = full_keys.unwrap(); + let full_values = full_values.unwrap(); + + let chunked_last_key = chunked_keys.index((.., .., -1.., ..)); + let full_last_key = full_keys.index((.., .., -1.., ..)); + let chunked_last_value = chunked_values.index((.., .., -1.., ..)); + let full_last_value = full_values.index((.., .., -1.., ..)); + eval([ + &chunked_last_key, + &full_last_key, + &chunked_last_value, + &full_last_value, + ]) + .unwrap(); + + let max_abs_diff = |a: &Array, b: &Array| -> f32 { + let a = a.as_type::().unwrap(); + let b = b.as_type::().unwrap(); + a.subtract(&b) + .unwrap() + .abs() + .unwrap() + .max(false) + .unwrap() + .item::() + }; + + let key_diff = max_abs_diff(&chunked_last_key, &full_last_key); + let value_diff = max_abs_diff(&chunked_last_value, &full_last_value); + eprintln!("qwen2-first-layer-last-key-diff {key_diff}"); + eprintln!("qwen2-first-layer-last-value-diff {value_diff}"); + + assert!(key_diff < 0.1, "first-layer key mismatch: {key_diff}"); + assert!(value_diff < 1e-3, "first-layer value mismatch: {value_diff}"); + } + + #[test] + #[ignore = "requires local model files"] + fn test_qwen2_first_layer_key_matches_full_prefill_only_for_correct_rope_offset() { + let mut model = load_qwen2_model(CACHED_TEST_MODEL_DIR).unwrap(); + let tokenizer = load_qwen2_tokenizer(CACHED_TEST_MODEL_DIR).unwrap(); + + let prompt = "<|im_start|>system\nYou are a concise assistant.<|im_end|>\n<|im_start|>user\nContext: The native MLX path supports streaming, but some families behave inconsistently.\nTask: Give three short bullets on what is working, what is not, and the next fix.<|im_end|>\n<|im_start|>assistant\n"; + let encoding = tokenizer.encode(prompt, false).unwrap(); + let prompt_ids = encoding.get_ids().to_vec(); + let agreed_prefix = vec![12u32, 576, 9867, 19614, 55]; + + let all_ids: Vec = prompt_ids + .iter() + .copied() + .chain(agreed_prefix.iter().copied()) + .collect(); + let all_tokens = Array::from(all_ids.as_slice()).index(NewAxis); + let mut full_prefill_cache = Vec::>::new(); + let input = super::ModelInput { + inputs: &all_tokens, + mask: None, + cache: &mut full_prefill_cache, + }; + let _ = model.forward(input).unwrap(); + + let full_first = full_prefill_cache[0].as_ref().unwrap(); + let (full_keys, _) = full_first.arrays(); + let full_last_key = full_keys.unwrap().index((.., .., -1.., ..)); + eval([&full_last_key]).unwrap(); + + let offset = (prompt_ids.len() + agreed_prefix.len() - 1) as i32; + let last_prefix_token = + Array::from_slice(&[agreed_prefix[agreed_prefix.len() - 1]], &[1]).index(NewAxis); + let embedded = model.model.embed_tokens.forward(&last_prefix_token).unwrap(); + let normalized = model.model.layers[0] + .input_layernorm + .forward(&embedded) + .unwrap(); + let raw_keys = model.model.layers[0].self_attn.k_proj.forward(&normalized).unwrap(); + let shaped_keys = raw_keys + .reshape(&[1, 1, model.model.layers[0].self_attn.n_kv_heads, -1]) + .unwrap() + .transpose_axes(&[0, 2, 1, 3]) + .unwrap(); + + let max_abs_diff = |a: &Array, b: &Array| -> f32 { + let a = a.as_type::().unwrap(); + let b = b.as_type::().unwrap(); + a.subtract(&b) + .unwrap() + .abs() + .unwrap() + .max(false) + .unwrap() + .item::() + }; + + for candidate in [offset - 2, offset - 1, offset, offset + 1, offset + 2] { + let rope_input = RopeInputBuilder::new(&shaped_keys) + .offset(candidate) + .build() + .unwrap(); + let rotated = model.model.layers[0] + .self_attn + .rope + .forward(rope_input) + .unwrap(); + eval([&rotated]).unwrap(); + let diff = max_abs_diff(&rotated, &full_last_key); + eprintln!("qwen2-first-layer-key-offset candidate={candidate} diff={diff}"); + } + + assert!(offset >= 0); + } + + #[test] + #[ignore = "requires local model files"] + fn test_qwen2_first_layer_rope_sequence_vs_single_token_offset() { + let mut model = load_qwen2_model(CACHED_TEST_MODEL_DIR).unwrap(); + let tokenizer = load_qwen2_tokenizer(CACHED_TEST_MODEL_DIR).unwrap(); + + let prompt = "<|im_start|>system\nYou are a concise assistant.<|im_end|>\n<|im_start|>user\nContext: The native MLX path supports streaming, but some families behave inconsistently.\nTask: Give three short bullets on what is working, what is not, and the next fix.<|im_end|>\n<|im_start|>assistant\n"; + let encoding = tokenizer.encode(prompt, false).unwrap(); + let prompt_ids = encoding.get_ids().to_vec(); + let agreed_prefix = vec![12u32, 576, 9867, 19614, 55]; + let offset = prompt_ids.len() as i32; + + let prefix_tokens = Array::from(agreed_prefix.as_slice()).index(NewAxis); + let embedded_prefix = model.model.embed_tokens.forward(&prefix_tokens).unwrap(); + let normalized_prefix = model.model.layers[0] + .input_layernorm + .forward(&embedded_prefix) + .unwrap(); + let raw_prefix_keys = model.model.layers[0] + .self_attn + .k_proj + .forward(&normalized_prefix) + .unwrap() + .reshape(&[1, agreed_prefix.len() as i32, model.model.layers[0].self_attn.n_kv_heads, -1]) + .unwrap() + .transpose_axes(&[0, 2, 1, 3]) + .unwrap(); + + let seq_rope_input = RopeInputBuilder::new(&raw_prefix_keys) + .offset(offset) + .build() + .unwrap(); + let seq_rotated = model.model.layers[0] + .self_attn + .rope + .forward(seq_rope_input) + .unwrap(); + let seq_last = seq_rotated.index((.., .., -1.., ..)); + eval([&seq_last]).unwrap(); + + let last_token = Array::from_slice(&[agreed_prefix[agreed_prefix.len() - 1]], &[1]).index(NewAxis); + let embedded_last = model.model.embed_tokens.forward(&last_token).unwrap(); + let normalized_last = model.model.layers[0] + .input_layernorm + .forward(&embedded_last) + .unwrap(); + let raw_last_key = model.model.layers[0] + .self_attn + .k_proj + .forward(&normalized_last) + .unwrap() + .reshape(&[1, 1, model.model.layers[0].self_attn.n_kv_heads, -1]) + .unwrap() + .transpose_axes(&[0, 2, 1, 3]) + .unwrap(); + + let single_rope_input = RopeInputBuilder::new(&raw_last_key) + .offset(offset + agreed_prefix.len() as i32 - 1) + .build() + .unwrap(); + let single_rotated = model.model.layers[0] + .self_attn + .rope + .forward(single_rope_input) + .unwrap(); + eval([&single_rotated]).unwrap(); + + let max_abs_diff = |a: &Array, b: &Array| -> f32 { + let a = a.as_type::().unwrap(); + let b = b.as_type::().unwrap(); + a.subtract(&b) + .unwrap() + .abs() + .unwrap() + .max(false) + .unwrap() + .item::() + }; + + let diff = max_abs_diff(&seq_last, &single_rotated); + eprintln!("qwen2-first-layer-rope-seq-vs-single diff={diff}"); + assert!(diff > 1.0, "expected scalar-offset rope path to diverge, got {diff}"); + } + + #[test] + #[ignore = "requires local model files"] + fn test_qwen2_first_layer_rope_dynamic_single_token_offset() { + let mut model = load_qwen2_model(CACHED_TEST_MODEL_DIR).unwrap(); + let tokenizer = load_qwen2_tokenizer(CACHED_TEST_MODEL_DIR).unwrap(); + + let prompt = "<|im_start|>system\nYou are a concise assistant.<|im_end|>\n<|im_start|>user\nContext: The native MLX path supports streaming, but some families behave inconsistently.\nTask: Give three short bullets on what is working, what is not, and the next fix.<|im_end|>\n<|im_start|>assistant\n"; + let encoding = tokenizer.encode(prompt, false).unwrap(); + let prompt_ids = encoding.get_ids().to_vec(); + let agreed_prefix = vec![12u32, 576, 9867, 19614, 55]; + let offset = prompt_ids.len() as i32; + + let prefix_tokens = Array::from(agreed_prefix.as_slice()).index(NewAxis); + let embedded_prefix = model.model.embed_tokens.forward(&prefix_tokens).unwrap(); + let normalized_prefix = model.model.layers[0] + .input_layernorm + .forward(&embedded_prefix) + .unwrap(); + let raw_prefix_keys = model.model.layers[0] + .self_attn + .k_proj + .forward(&normalized_prefix) + .unwrap() + .reshape(&[1, agreed_prefix.len() as i32, model.model.layers[0].self_attn.n_kv_heads, -1]) + .unwrap() + .transpose_axes(&[0, 2, 1, 3]) + .unwrap(); + + let seq_rope_input = RopeInputBuilder::new(&raw_prefix_keys) + .offset(offset) + .build() + .unwrap(); + let seq_rotated = model.model.layers[0] + .self_attn + .rope + .forward(seq_rope_input) + .unwrap(); + let seq_last = seq_rotated.index((.., .., -1.., ..)); + eval([&seq_last]).unwrap(); + + let last_token = Array::from_slice(&[agreed_prefix[agreed_prefix.len() - 1]], &[1]).index(NewAxis); + let embedded_last = model.model.embed_tokens.forward(&last_token).unwrap(); + let normalized_last = model.model.layers[0] + .input_layernorm + .forward(&embedded_last) + .unwrap(); + let raw_last_key = model.model.layers[0] + .self_attn + .k_proj + .forward(&normalized_last) + .unwrap() + .reshape(&[1, 1, model.model.layers[0].self_attn.n_kv_heads, -1]) + .unwrap() + .transpose_axes(&[0, 2, 1, 3]) + .unwrap(); + + let dynamic_rotated = match &model.model.layers[0].self_attn.rope { + crate::utils::rope::RopeVariant::Default(rope) => { + let dynamic_offset = Array::from_int(offset + agreed_prefix.len() as i32 - 1); + rope_dynamic( + &raw_last_key, + rope.dimensions, + rope.traditional, + Some(rope.base), + rope.scale, + &dynamic_offset, + None::<&Array>, + ) + .unwrap() + } + _ => panic!("expected default rope variant for qwen2"), + }; + eval([&dynamic_rotated]).unwrap(); + + let max_abs_diff = |a: &Array, b: &Array| -> f32 { + let a = a.as_type::().unwrap(); + let b = b.as_type::().unwrap(); + a.subtract(&b) + .unwrap() + .abs() + .unwrap() + .max(false) + .unwrap() + .item::() + }; + + let diff = max_abs_diff(&seq_last, &dynamic_rotated); + eprintln!("qwen2-first-layer-rope-seq-vs-dynamic-single diff={diff}"); + assert!(diff < 0.1, "rope seq vs dynamic single mismatch: {diff}"); + } } From 9647f60c9e0cb6362613f80d3359803be573019c Mon Sep 17 00:00:00 2001 From: James Dumay Date: Mon, 30 Mar 2026 06:38:30 +1100 Subject: [PATCH 14/15] Reuse qwen2 rope offset arrays per decode step --- mlx-lm/src/models/qwen2.rs | 47 +++++++++++++++++++++++++++++++++----- 1 file changed, 41 insertions(+), 6 deletions(-) diff --git a/mlx-lm/src/models/qwen2.rs b/mlx-lm/src/models/qwen2.rs index a8ec5044f..dee32a1ec 100644 --- a/mlx-lm/src/models/qwen2.rs +++ b/mlx-lm/src/models/qwen2.rs @@ -129,7 +129,12 @@ impl Attention { } } -fn apply_cached_rope(rope: &mut RopeVariant, x: &Array, offset: i32) -> Result { +fn apply_cached_rope( + rope: &mut RopeVariant, + x: &Array, + offset: i32, + offset_array: Option<&Array>, +) -> Result { let seq_len = x.shape()[x.shape().len() - 2]; if seq_len != 1 { @@ -137,7 +142,14 @@ fn apply_cached_rope(rope: &mut RopeVariant, x: &Array, offset: i32) -> Result offset_array, + None => { + owned_offset = Array::from_int(offset); + &owned_offset + } + }; match rope { RopeVariant::Default(rope) => rope_dynamic( @@ -165,6 +177,7 @@ pub struct AttentionInput<'a, C> { pub x: &'a Array, pub mask: Option<&'a AttentionMask>, pub cache: Option<&'a mut C>, + pub rope_offset_array: Option<&'a Array>, } impl Module> for Attention @@ -176,7 +189,12 @@ where #[allow(non_snake_case)] fn forward(&mut self, input: AttentionInput<'_, C>) -> Result { - let AttentionInput { x, mask, mut cache } = input; + let AttentionInput { + x, + mask, + mut cache, + rope_offset_array, + } = input; let shape = x.shape(); let B = shape[0]; @@ -197,8 +215,8 @@ where .transpose_axes(&[0, 2, 1, 3])?; if let Some(cache) = cache.as_mut() { - queries = apply_cached_rope(&mut self.rope, &queries, cache.offset())?; - keys = apply_cached_rope(&mut self.rope, &keys, cache.offset())?; + queries = apply_cached_rope(&mut self.rope, &queries, cache.offset(), rope_offset_array)?; + keys = apply_cached_rope(&mut self.rope, &keys, cache.offset(), rope_offset_array)?; (keys, values) = cache.update_and_fetch(keys, values)?; } else { @@ -346,12 +364,18 @@ where type Error = Exception; fn forward(&mut self, input: AttentionInput<'_, C>) -> Result { - let AttentionInput { x, mask, cache } = input; + let AttentionInput { + x, + mask, + cache, + rope_offset_array, + } = input; let self_attn_input = AttentionInput { x: &self.input_layernorm.forward(x)?, mask, cache, + rope_offset_array, }; let r = self.self_attn.forward(self_attn_input)?; let h = x.add(r)?; @@ -430,6 +454,16 @@ where } = input; let mut h = self.embed_tokens.forward(inputs)?; + let rope_offset_array = if inputs.shape()[1] == 1 { + let offset = cache + .first() + .and_then(|c| c.as_ref()) + .map(|c| c.offset()) + .unwrap_or(0); + Some(Array::from_int(offset)) + } else { + None + }; let mask = match mask { Some(mask) => Some(mask.clone()), @@ -445,6 +479,7 @@ where x: &h, mask: mask.as_ref(), cache: c.as_mut(), + rope_offset_array: rope_offset_array.as_ref(), }; h = layer.forward(layer_input)?; } From 58a056b936c9a0f1455b7b723745a30c6ac40130 Mon Sep 17 00:00:00 2001 From: James Dumay Date: Mon, 30 Mar 2026 07:49:25 +1100 Subject: [PATCH 15/15] Add KV padded-buffer corruption repro --- mlx-lm/src/cache.rs | 269 +++++++++++++++++++++++++++++++++++++++++--- 1 file changed, 254 insertions(+), 15 deletions(-) diff --git a/mlx-lm/src/cache.rs b/mlx-lm/src/cache.rs index 5458e1655..3a09ca947 100644 --- a/mlx-lm/src/cache.rs +++ b/mlx-lm/src/cache.rs @@ -86,6 +86,10 @@ impl KVCache { pub fn new() -> Self { Self::default() } + + pub(crate) fn arrays(&self) -> (Option<&Array>, Option<&Array>) { + (self.keys.as_ref(), self.values.as_ref()) + } } impl ConcatKeyValueCache { @@ -192,22 +196,31 @@ impl KeyValueCache for KVCache { let prev = self.offset; let seq_len = keys.shape()[2]; + if self.keys.is_none() && self.values.is_none() { + self.offset = seq_len; + self.keys = Some(keys.deep_clone()); + self.values = Some(values.deep_clone()); + return Ok((keys, values)); + } + let needs_resize = self .keys .as_ref() .map(|cached| (prev + seq_len) > cached.shape()[2]) .unwrap_or(true); + let mut wrote_current = false; if needs_resize { let key_shape = keys.shape(); let value_shape = values.shape(); let expand_steps = ((Self::STEP + seq_len - 1) / Self::STEP) * Self::STEP; - let new_keys = zeros_dtype( - &[key_shape[0], key_shape[1], expand_steps, key_shape[3]], + let total_capacity = prev + expand_steps; + let mut zero_pad_keys = zeros_dtype( + &[key_shape[0], key_shape[1], expand_steps - seq_len, key_shape[3]], keys.dtype(), )?; - let new_values = zeros_dtype( - &[value_shape[0], value_shape[1], expand_steps, value_shape[3]], + let mut zero_pad_values = zeros_dtype( + &[value_shape[0], value_shape[1], expand_steps - seq_len, value_shape[3]], values.dtype(), )?; @@ -223,26 +236,47 @@ impl KeyValueCache for KVCache { } else { existing_values }; - self.keys = Some(concatenate_axis(&[existing_keys, new_keys], -2)?); - self.values = Some(concatenate_axis(&[existing_values, new_values], -2)?); + self.keys = Some(concatenate_axis( + &[existing_keys, keys.deep_clone(), zero_pad_keys], + -2, + )?); + self.values = Some(concatenate_axis( + &[existing_values, values.deep_clone(), zero_pad_values], + -2, + )?); + wrote_current = true; } _ => { + let mut new_keys = zeros_dtype( + &[key_shape[0], key_shape[1], total_capacity, key_shape[3]], + keys.dtype(), + )?; + let mut new_values = zeros_dtype( + &[value_shape[0], value_shape[1], total_capacity, value_shape[3]], + values.dtype(), + )?; + new_keys.try_index_mut((.., .., ..seq_len, ..), &keys)?; + new_values.try_index_mut((.., .., ..seq_len, ..), &values)?; + eval([&new_keys, &new_values])?; self.keys = Some(new_keys); self.values = Some(new_values); + wrote_current = true; } } } self.offset += seq_len; let end = self.offset; - self.keys - .as_mut() - .expect("keys cache missing") - .try_index_mut((.., .., prev..end, ..), &keys)?; - self.values - .as_mut() - .expect("values cache missing") - .try_index_mut((.., .., prev..end, ..), &values)?; + if !wrote_current { + self.keys + .as_mut() + .expect("keys cache missing") + .try_index_mut((.., .., prev..end, ..), &keys)?; + self.values + .as_mut() + .expect("values cache missing") + .try_index_mut((.., .., prev..end, ..), &values)?; + } eval([ self.keys.as_ref().expect("keys cache missing"), self.values.as_ref().expect("values cache missing"), @@ -269,7 +303,10 @@ pub struct DefaultKeyValueCache {} #[cfg(test)] mod tests { use super::{ConcatKeyValueCache, KVCache, KeyValueCache}; - use mlx_rs::Array; + use mlx_rs::{ + ops::{concatenate_axis, indexing::{IndexOp, TryIndexMutOp}, zeros_dtype}, + Array, + }; use std::sync::{Mutex, MutexGuard, OnceLock}; fn test_guard() -> MutexGuard<'static, ()> { @@ -369,4 +406,206 @@ mod tests { assert_eq!(keys.as_slice::(), &[1., 2., 3., 4., 9., 10.]); assert_eq!(values.as_slice::(), &[5., 6., 7., 8., 11., 12.]); } + + #[test] + fn kv_cache_matches_concat_for_repeated_single_token_appends() { + let _guard = test_guard(); + let mut concat = ConcatKeyValueCache::new(); + let mut kv = KVCache::new(); + + let prefix_len = 220usize; + let heads = 2usize; + let dim = 64usize; + let prefix_elems = heads * prefix_len * dim; + let prefix_keys: Vec = (0..prefix_elems).map(|i| i as f32 / 1000.0).collect(); + let prefix_values: Vec = (0..prefix_elems).map(|i| i as f32 / 2000.0).collect(); + let prefix_keys = Array::from_slice(&prefix_keys, &[1, heads as i32, prefix_len as i32, dim as i32]); + let prefix_values = + Array::from_slice(&prefix_values, &[1, heads as i32, prefix_len as i32, dim as i32]); + + let (concat_keys, concat_values) = concat + .update_and_fetch(prefix_keys.deep_clone(), prefix_values.deep_clone()) + .expect("seed concat"); + let (kv_keys, kv_values) = kv + .update_and_fetch(prefix_keys, prefix_values) + .expect("seed kv"); + let seed_key_match = concat_keys + .all_close(&kv_keys, 1e-6, 1e-6, None) + .unwrap() + .item::(); + let seed_value_match = concat_values + .all_close(&kv_values, 1e-6, 1e-6, None) + .unwrap() + .item::(); + if !seed_key_match || !seed_value_match { + let key_diff = (concat_keys.subtract(&kv_keys).unwrap()) + .abs() + .unwrap() + .max(false) + .unwrap() + .item::(); + let value_diff = (concat_values.subtract(&kv_values).unwrap()) + .abs() + .unwrap() + .max(false) + .unwrap() + .item::(); + panic!( + "seed mismatch: concat_shape={:?} kv_shape={:?} key_diff={key_diff} value_diff={value_diff}", + concat_keys.shape(), + kv_keys.shape() + ); + } + + for step in 0..8usize { + let token_keys: Vec = (0..(heads * dim)) + .map(|i| (10_000 + step * heads * dim + i) as f32 / 1000.0) + .collect(); + let token_values: Vec = (0..(heads * dim)) + .map(|i| (20_000 + step * heads * dim + i) as f32 / 1000.0) + .collect(); + let token_keys = + Array::from_slice(&token_keys, &[1, heads as i32, 1, dim as i32]); + let token_values = + Array::from_slice(&token_values, &[1, heads as i32, 1, dim as i32]); + + let (concat_keys, concat_values) = concat + .update_and_fetch(token_keys.deep_clone(), token_values.deep_clone()) + .expect("append concat"); + let (kv_keys, kv_values) = kv + .update_and_fetch(token_keys, token_values) + .expect("append kv"); + + let key_match = concat_keys + .all_close(&kv_keys, 1e-6, 1e-6, None) + .expect("compare key arrays") + .item::(); + let value_match = concat_values + .all_close(&kv_values, 1e-6, 1e-6, None) + .expect("compare value arrays") + .item::(); + + if !key_match || !value_match { + let concat_key_slice = concat_keys.as_slice::(); + let kv_key_slice = kv_keys.as_slice::(); + let compare_from = concat_key_slice.len().saturating_sub(16); + let concat_tail = &concat_key_slice[compare_from..]; + let kv_tail = &kv_key_slice[compare_from..]; + let key_diff = (concat_keys.subtract(&kv_keys).unwrap()) + .abs() + .unwrap() + .max(false) + .unwrap() + .item::(); + let value_diff = (concat_values.subtract(&kv_values).unwrap()) + .abs() + .unwrap() + .max(false) + .unwrap() + .item::(); + panic!( + "kv cache diverged from concat at step {step}: key_diff={key_diff} value_diff={value_diff} concat_tail={concat_tail:?} kv_tail={kv_tail:?}" + ); + } + } + } + + #[test] + fn direct_index_mut_large_prefix_then_single_append_preserves_tail() { + let _guard = test_guard(); + let heads = 2usize; + let dim = 64usize; + let prefix_len = 220usize; + let total_capacity = 476usize; + + let prefix_keys: Vec = (0..(heads * prefix_len * dim)) + .map(|i| i as f32 / 1000.0) + .collect(); + let token_keys: Vec = (0..(heads * dim)) + .map(|i| (10_000 + i) as f32 / 1000.0) + .collect(); + + let prefix_keys = + Array::from_slice(&prefix_keys, &[1, heads as i32, prefix_len as i32, dim as i32]); + let token_keys = Array::from_slice(&token_keys, &[1, heads as i32, 1, dim as i32]); + let mut buffer = zeros_dtype( + &[1, heads as i32, total_capacity as i32, dim as i32], + prefix_keys.dtype(), + ) + .unwrap(); + + buffer + .try_index_mut((.., .., ..prefix_len as i32, ..), &prefix_keys) + .unwrap(); + buffer + .try_index_mut( + (.., .., prefix_len as i32..(prefix_len as i32 + 1), ..), + &token_keys, + ) + .unwrap(); + mlx_rs::transforms::eval([&buffer]).unwrap(); + + let live = buffer.index((.., .., ..(prefix_len as i32 + 1), ..)); + mlx_rs::transforms::eval([&live]).unwrap(); + let live_tail = &live.as_slice::()[live.as_slice::().len() - 16..]; + let token_tail = &token_keys.as_slice::()[token_keys.as_slice::().len() - 16..]; + assert_eq!(live_tail, token_tail); + } + + #[test] + fn direct_concat_prefix_token_zeropad_preserves_tail() { + let _guard = test_guard(); + let heads = 2usize; + let dim = 64usize; + let prefix_len = 220usize; + let pad_len = 255usize; + + let prefix_keys: Vec = (0..(heads * prefix_len * dim)) + .map(|i| i as f32 / 1000.0) + .collect(); + let token_keys: Vec = (0..(heads * dim)) + .map(|i| (10_000 + i) as f32 / 1000.0) + .collect(); + + let prefix_keys = + Array::from_slice(&prefix_keys, &[1, heads as i32, prefix_len as i32, dim as i32]); + let token_keys = Array::from_slice(&token_keys, &[1, heads as i32, 1, dim as i32]); + let zero_pad = zeros_dtype( + &[1, heads as i32, pad_len as i32, dim as i32], + prefix_keys.dtype(), + ) + .unwrap(); + + let buffer = concatenate_axis(&[prefix_keys, token_keys.deep_clone(), zero_pad], -2).unwrap(); + mlx_rs::transforms::eval([&buffer]).unwrap(); + let live = buffer.index((.., .., ..(prefix_len as i32 + 1), ..)); + mlx_rs::transforms::eval([&live]).unwrap(); + let live_tail = &live.as_slice::()[live.as_slice::().len() - 16..]; + let token_tail = &token_keys.as_slice::()[token_keys.as_slice::().len() - 16..]; + assert_eq!(live_tail, token_tail); + } + + #[test] + fn direct_concat_prefix_and_token_preserves_tail() { + let _guard = test_guard(); + let heads = 2usize; + let dim = 64usize; + let prefix_len = 220usize; + + let prefix_keys: Vec = (0..(heads * prefix_len * dim)) + .map(|i| i as f32 / 1000.0) + .collect(); + let token_keys: Vec = (0..(heads * dim)) + .map(|i| (10_000 + i) as f32 / 1000.0) + .collect(); + + let prefix_keys = + Array::from_slice(&prefix_keys, &[1, heads as i32, prefix_len as i32, dim as i32]); + let token_keys = Array::from_slice(&token_keys, &[1, heads as i32, 1, dim as i32]); + let buffer = concatenate_axis(&[prefix_keys, token_keys.deep_clone()], -2).unwrap(); + mlx_rs::transforms::eval([&buffer]).unwrap(); + let live_tail = &buffer.as_slice::()[buffer.as_slice::().len() - 16..]; + let token_tail = &token_keys.as_slice::()[token_keys.as_slice::().len() - 16..]; + assert_eq!(live_tail, token_tail); + } }