use super::checkpoint::{read_buffer, read_u32, write_buffer, write_u32}; use super::*; use crate::engine::qwen::{QwenModel, QwenTensor}; use std::ffi::CStr; const HIDDEN: u32 = 2_560; const HC: u32 = 4; const HC_WIDTH: u32 = HIDDEN * HC; const HC_RANK: u32 = 320; const LAYERS: usize = 48; const GDN_HEADS_K: u32 = 16; const GDN_HEADS_V: u32 = 48; const HEAD_DIM: u32 = 128; const GDN_QKV: u32 = (GDN_HEADS_K * 2 + GDN_HEADS_V) * HEAD_DIM; const GDN_VALUE: u32 = GDN_HEADS_V * HEAD_DIM; const GDN_CONTROLS: u32 = GDN_VALUE + GDN_HEADS_V * 2; const ATTN_HEADS: u32 = 24; const ATTN_KV_HEADS: u32 = 2; const ATTN_DIM: u32 = 256; const ATTN_WIDTH: u32 = ATTN_HEADS * ATTN_DIM; const ATTN_KV_WIDTH: u32 = ATTN_KV_HEADS * ATTN_DIM; const EXPERTS: u32 = 512; const EXPERTS_USED: usize = 10; const EXPERT_WIDTH: u32 = 640; const VOCAB: u32 = 248_320; const DENSE_BUDGET: u32 = 2_048; const PLE_HEADS: usize = 16; const PLE_HEAD_DIM: u32 = 160; const PLE_HISTORY: usize = 2; const PLE_CONV_STATE: u32 = 9; const EOS_TOKEN: i32 = 248_044; const PLE_ROW_BYTES: usize = 100; const CHECKPOINT_MAGIC: &[u8; 8] = b"DS4QWN01"; const CHECKPOINT_VERSION: u32 = 2; const CHECKPOINT_CHUNK: usize = 8 * 1024 * 1024; #[derive(Clone, Copy)] struct Weight<'a> { tensor: &'a QwenTensor, expert: Option, } struct Affine<'a> { packed: Weight<'a>, scales: Weight<'a>, biases: Weight<'a>, bits: u32, group: u32, } enum LayerState { Gdn { conv: Buffer, recurrent: Buffer }, Attention { kv: Buffer }, } struct Scratch { hidden: Buffer, hc: Buffer, hc_norm: Buffer, hc_mix: Buffer, rank: Buffer, block: Buffer, injection: Buffer, qkv: Buffer, controls: Buffer, gdn_raw: Buffer, gdn_out: Buffer, router: Buffer, route_ids: Buffer, route_weights: Buffer, gate: Buffer, up: Buffer, mid: Buffer, expert: Buffer, moe: Buffer, shared: Buffer, q_packed: Buffer, q: Buffer, q_gate: Buffer, k: Buffer, k_rope: Buffer, v: Buffer, attention: Buffer, ple_packed: Buffer, ple_scales: Buffer, ple_biases: Buffer, ple_embedding: Buffer, ple_key: Buffer, ple_value: Buffer, ple_gated: Buffer, ple_norm: Buffer, ple_output: Buffer, logits: Buffer, } impl Scratch { fn new() -> Result { Ok(Self { hidden: Buffer::floats(HIDDEN.into())?, hc: Buffer::floats(HC_WIDTH.into())?, hc_norm: Buffer::floats(HC_WIDTH.into())?, hc_mix: Buffer::floats(HC_WIDTH.into())?, rank: Buffer::floats(HC_RANK.into())?, block: Buffer::floats(HIDDEN.into())?, injection: Buffer::floats(HC.into())?, qkv: Buffer::floats(GDN_QKV.into())?, controls: Buffer::floats(GDN_CONTROLS.into())?, gdn_raw: Buffer::floats(GDN_VALUE.into())?, gdn_out: Buffer::floats(GDN_VALUE.into())?, router: Buffer::floats(EXPERTS.into())?, route_ids: Buffer::bytes((EXPERTS_USED * 4) as u64)?, route_weights: Buffer::floats(EXPERTS_USED as u64)?, gate: Buffer::floats(EXPERT_WIDTH.into())?, up: Buffer::floats(EXPERT_WIDTH.into())?, mid: Buffer::floats(EXPERT_WIDTH.into())?, expert: Buffer::floats(HIDDEN.into())?, moe: Buffer::floats(HIDDEN.into())?, shared: Buffer::floats(HIDDEN.into())?, q_packed: Buffer::floats((ATTN_WIDTH * 2).into())?, q: Buffer::floats(ATTN_WIDTH.into())?, q_gate: Buffer::floats(ATTN_WIDTH.into())?, k: Buffer::floats(ATTN_KV_WIDTH.into())?, k_rope: Buffer::floats(ATTN_KV_WIDTH.into())?, v: Buffer::floats(ATTN_KV_WIDTH.into())?, attention: Buffer::floats(ATTN_WIDTH.into())?, ple_packed: Buffer::bytes((PLE_HEADS as u64) * 80)?, ple_scales: Buffer::bytes((PLE_HEADS as u64) * 10)?, ple_biases: Buffer::bytes((PLE_HEADS as u64) * 10)?, ple_embedding: Buffer::floats(HIDDEN.into())?, ple_key: Buffer::floats(HC_WIDTH.into())?, ple_value: Buffer::floats(HIDDEN.into())?, ple_gated: Buffer::floats(HC_WIDTH.into())?, ple_norm: Buffer::floats(HC_WIDTH.into())?, ple_output: Buffer::floats(HC_WIDTH.into())?, logits: Buffer::floats(VOCAB.into())?, }) } } struct PleContract { multipliers: [i64; 3], sizes: [i64; PLE_HEADS], offsets: [i64; PLE_HEADS], } struct PleState { history: [i32; PLE_HISTORY], conv: Buffer, } pub(in crate::engine) struct QwenExecutor { model: QwenModel, states: Vec, ple_contract: PleContract, ple_state: PleState, scratch: Scratch, logits: Vec, tokens: Vec, position: u32, context: u32, checkpoint_tag: [u8; 32], _context: Context, } pub(in crate::engine) struct QwenResidentState { states: Vec, ple_state: PleState, logits: Vec, tokens: Vec, position: u32, checkpoint_tag: [u8; 32], } impl QwenExecutor { pub(super) fn open(model: QwenModel, context: u32) -> Result { let ple_contract = ple_contract(&model)?; let native = Context::open_qwen(model.memory().admission)?; let states = allocate_states(context)?; Ok(Self { model, states, ple_contract, ple_state: allocate_ple_state()?, scratch: Scratch::new()?, logits: vec![0.0; VOCAB as usize], tokens: Vec::new(), position: 0, context, checkpoint_tag: [0; 32], _context: native, }) } pub(super) fn eval(&mut self, token: i32) -> Result<(), String> { if token < 0 || token as u32 >= VOCAB { return Err(format!("token {token} is outside the Qwen vocabulary")); } if self.position >= self.context { return Err(format!( "the Qwen executor supports {} tokens per session", self.context )); } if self.position + 1 > DENSE_BUDGET { return Err( "Qwen sparse QSA selection is required beyond 2048 tokens; native QSA belongs to issue #97" .into(), ); } self.begin_token(token)?; for layer in 0..LAYERS { if layer == 1 { self.ple(token)?; } self.encode_layer(layer)?; } self.final_output()?; self.tokens.push(token); self.position += 1; Ok(()) } fn begin_token(&self, token: i32) -> Result<(), String> { let embedding = self.affine("language_model.model.embed_tokens", HIDDEN, VOCAB, None)?; let mut args = args(); args.u[0] = HIDDEN; args.u[2] = embedding.bits; args.u[3] = embedding.group; args.u[4] = token as u32; let commands = Commands::begin()?; self.dispatch( c"kernel_qwen_affine_embedding", &self.scratch.hidden, None, None, None, &[ self.view(embedding.packed), self.view(embedding.scales), self.view(embedding.biases), ], &args, HIDDEN, 1, )?; args.u[0] = HIDDEN; self.dispatch( c"kernel_qwen_repeat4", &self.scratch.hc, Some(&self.scratch.hidden), None, None, &[], &args, HC_WIDTH, 1, )?; commands.finish() } fn ple(&mut self, token: i32) -> Result<(), String> { let rows = ple_rows(&self.ple_contract, self.ple_state.history, token)?; let packed = self.model.tensor("ngram.weight")?; let scales = self.model.tensor("ngram.scales")?; let biases = self.model.tensor("ngram.biases")?; let mut packed_stage = [0_u8; PLE_HEADS * 80]; let mut scales_stage = [0_u8; PLE_HEADS * 10]; let mut biases_stage = [0_u8; PLE_HEADS * 10]; let packed_bytes = self.model.tensor_bytes(packed)?; let scales_bytes = self.model.tensor_bytes(scales)?; let biases_bytes = self.model.tensor_bytes(biases)?; for (head, &row) in rows.iter().enumerate() { let bytes = gather_ple_row(row, packed_bytes, scales_bytes, biases_bytes)?; packed_stage[head * 80..][..80].copy_from_slice(&bytes[..80]); scales_stage[head * 10..][..10].copy_from_slice(&bytes[80..90]); biases_stage[head * 10..][..10].copy_from_slice(&bytes[90..]); } self.scratch.ple_packed.write(0, &packed_stage)?; self.scratch.ple_scales.write(0, &scales_stage)?; self.scratch.ple_biases.write(0, &biases_stage)?; let commands = Commands::begin()?; let mut dequant = args(); dequant.u[0] = PLE_HEAD_DIM; dequant.u[1] = PLE_HEADS as u32; dequant.u[2] = 4; dequant.u[3] = 32; self.dispatch( c"kernel_qwen_ple_dequant", &self.scratch.ple_embedding, Some(&self.scratch.ple_packed), Some(&self.scratch.ple_scales), Some(&self.scratch.ple_biases), &[], &dequant, HIDDEN, 1, )?; let prefix = "language_model.model.layers.1.ple"; self.bf16_mv( self.weight(&format!("{prefix}.key_proj.weight"))?, &self.scratch.ple_embedding, &self.scratch.ple_key, HIDDEN, HC_WIDTH, )?; self.bf16_mv( self.weight(&format!("{prefix}.value_proj.weight"))?, &self.scratch.ple_embedding, &self.scratch.ple_value, HIDDEN, HIDDEN, )?; for (input, output, name) in [ (&self.scratch.ple_key, &self.scratch.ple_key, "norm_key"), (&self.scratch.hc, &self.scratch.hc_norm, "norm_query"), ] { let mut norm = args(); norm.u[0] = HC_WIDTH; norm.u[1] = HIDDEN; norm.f[0] = 1.0e-6; self.dispatch( c"kernel_qwen_zero_rms", output, Some(input), None, None, &[self.view(self.weight(&format!("{prefix}.{name}.weight"))?)], &norm, HC, 1, )?; } let mut gate = args(); gate.u[0] = HIDDEN; self.dispatch( c"kernel_qwen_ple_gate", &self.scratch.ple_gated, Some(&self.scratch.ple_key), Some(&self.scratch.hc_norm), Some(&self.scratch.ple_value), &[], &gate, HC, 1, )?; let mut norm = args(); norm.u[0] = HC_WIDTH; norm.u[1] = HIDDEN; norm.f[0] = 1.0e-6; self.dispatch( c"kernel_qwen_zero_rms", &self.scratch.ple_norm, Some(&self.scratch.ple_gated), None, None, &[self.view(self.weight(&format!("{prefix}.norm_conv.weight"))?)], &norm, HC, 1, )?; let mut conv = args(); conv.u[0] = HC_WIDTH; self.dispatch( c"kernel_qwen_ple_conv", &self.scratch.ple_output, Some(&self.scratch.ple_gated), Some(&self.scratch.ple_norm), Some(&self.ple_state.conv), &[self.view(self.weight(&format!("{prefix}.conv_weight"))?)], &conv, HC_WIDTH, 1, )?; let mut add = args(); add.u[0] = HC_WIDTH; self.dispatch( c"kernel_qwen_add", &self.scratch.hc_norm, Some(&self.scratch.hc), Some(&self.scratch.ple_output), None, &[], &add, HC_WIDTH, 1, )?; self.scratch.hc.copy_from( 0, &self.scratch.hc_norm, 0, u64::from(HC_WIDTH) * 4, "committing Qwen PLE injection", )?; commands.finish()?; self.ple_state.history = advance_ple_history(self.ple_state.history, token); Ok(()) } fn encode_layer(&mut self, layer: usize) -> Result<(), String> { let prefix = format!("language_model.model.layers.{layer}"); let commands = Commands::begin()?; self.hyper_read(&format!("{prefix}.attn_hyper_connection"))?; match &self.states[layer] { LayerState::Gdn { .. } => self.gdn(&prefix, layer)?, LayerState::Attention { .. } => self.attention(&prefix, layer)?, } self.hyper_write()?; self.hyper_read(&format!("{prefix}.mlp_hyper_connection"))?; self.affine_mv_into( &self.affine(&format!("{prefix}.mlp.gate"), HIDDEN, EXPERTS, None)?, &self.scratch.block, &self.scratch.router, HIDDEN, EXPERTS, )?; let mut route_args = args(); route_args.u[0] = EXPERTS; self.dispatch( c"kernel_qwen_route_top10", &self.scratch.route_ids, Some(&self.scratch.route_weights), Some(&self.scratch.router), None, &[], &route_args, 1, 1, )?; commands.finish()?; let mut ids = [0_i32; EXPERTS_USED]; let mut weights = [0.0_f32; EXPERTS_USED]; self.scratch.route_ids.read_i32(&mut ids)?; self.scratch.route_weights.read_f32(&mut weights)?; self.scratch.moe.fill(0.0, HIDDEN.into())?; let commands = Commands::begin()?; for (&expert, &weight) in ids.iter().zip(&weights) { if !(0..EXPERTS as i32).contains(&expert) || !weight.is_finite() || weight < 0.0 { let mut router = vec![0.0; EXPERTS as usize]; self.scratch.router.read_f32(&mut router)?; let non_finite = router.iter().filter(|value| !value.is_finite()).count(); let mut block = vec![0.0; HIDDEN as usize]; self.scratch.block.read_f32(&mut block)?; let block_non_finite = block.iter().filter(|value| !value.is_finite()).count(); return Err(format!( "Qwen layer {layer} router produced invalid expert {expert} with weight {weight} ({non_finite} non-finite logits, {block_non_finite} non-finite inputs)" )); } self.expert(&prefix, expert as u32, weight)?; } self.shared_expert(&prefix)?; let mut inject_args = args(); inject_args.u[0] = HIDDEN; self.dispatch( c"kernel_qwen_hyper_inject", &self.scratch.hc_norm, Some(&self.scratch.hc), Some(&self.scratch.moe), Some(&self.scratch.injection), &[], &inject_args, HC_WIDTH, 1, )?; self.scratch.hc.copy_from( 0, &self.scratch.hc_norm, 0, u64::from(HC_WIDTH) * 4, "committing Qwen MoE hyper streams", )?; commands.finish() } fn hyper_read(&self, prefix: &str) -> Result<(), String> { let norm = self.weight(&format!("{prefix}.hc_norm.weight"))?; let down = self.weight(&format!("{prefix}.input_mix_weight_down.weight"))?; let up = self.weight(&format!("{prefix}.input_mix_weight_up.weight"))?; let inject = self.weight(&format!("{prefix}.block_inject_weight.weight"))?; let mut rms = args(); rms.u[0] = HC_WIDTH; rms.u[1] = HIDDEN; rms.f[0] = 1.0e-6; self.dispatch( c"kernel_qwen_zero_rms", &self.scratch.hc_norm, Some(&self.scratch.hc), None, None, &[self.view(norm)], &rms, HC, 1, )?; self.bf16_mv( down, &self.scratch.hc_norm, &self.scratch.rank, HC_WIDTH, HC_RANK, )?; let mut unary = args(); unary.u[0] = HC_RANK; self.dispatch( c"kernel_qwen_silu_div4", &self.scratch.rank, Some(&self.scratch.rank), None, None, &[], &unary, HC_RANK, 1, )?; self.bf16_mv( up, &self.scratch.rank, &self.scratch.hc_mix, HC_RANK, HC_WIDTH, )?; unary.u[0] = HC_WIDTH; self.dispatch( c"kernel_qwen_sigmoid", &self.scratch.hc_mix, Some(&self.scratch.hc_mix), None, None, &[], &unary, HC_WIDTH, 1, )?; let mut mix = args(); mix.u[0] = HIDDEN; self.dispatch( c"kernel_qwen_hyper_mix", &self.scratch.block, Some(&self.scratch.hc_norm), Some(&self.scratch.hc_mix), None, &[], &mix, HIDDEN, 1, )?; self.bf16_mv( inject, &self.scratch.hc_norm, &self.scratch.injection, HC_WIDTH, HC, )?; unary.u[0] = HC; self.dispatch( c"kernel_qwen_sigmoid2_div4", &self.scratch.injection, Some(&self.scratch.injection), None, None, &[], &unary, HC, 1, ) } fn hyper_write(&self) -> Result<(), String> { let mut values = args(); values.u[0] = HIDDEN; self.dispatch( c"kernel_qwen_hyper_inject", &self.scratch.hc_norm, Some(&self.scratch.hc), Some(&self.scratch.hidden), Some(&self.scratch.injection), &[], &values, HC_WIDTH, 1, )?; self.scratch.hc.copy_from( 0, &self.scratch.hc_norm, 0, u64::from(HC_WIDTH) * 4, "committing Qwen attention hyper streams", ) } fn gdn(&self, prefix: &str, layer: usize) -> Result<(), String> { let LayerState::Gdn { conv, recurrent } = &self.states[layer] else { return Err("Qwen GDN graph received attention state".into()); }; self.affine_mv_into( &self.affine( &format!("{prefix}.linear_attn.in_proj_qkv"), HIDDEN, GDN_QKV, None, )?, &self.scratch.block, &self.scratch.qkv, HIDDEN, GDN_QKV, )?; for (name, offset, width) in [ ("in_proj_z", 0_u64, GDN_VALUE), ("in_proj_b", u64::from(GDN_VALUE), GDN_HEADS_V), ("in_proj_a", u64::from(GDN_VALUE + GDN_HEADS_V), GDN_HEADS_V), ] { let target = self .scratch .controls .view(offset * 4, u64::from(width) * 4)?; self.affine_mv_into( &self.affine(&format!("{prefix}.linear_attn.{name}"), HIDDEN, width, None)?, &self.scratch.block, &target, HIDDEN, width, )?; } let mut conv_args = args(); conv_args.u[0] = GDN_QKV; self.dispatch( c"kernel_qwen_conv_silu", &self.scratch.qkv, Some(&self.scratch.qkv), Some(conv), None, &[self.view(self.weight(&format!("{prefix}.linear_attn.conv1d.weight"))?)], &conv_args, GDN_QKV, 1, )?; let mut step = args(); step.u[0] = HEAD_DIM; step.u[1] = GDN_HEADS_K; step.u[2] = GDN_HEADS_V; step.u[3] = GDN_VALUE; step.u[4] = GDN_VALUE + GDN_HEADS_V; step.f[0] = 1.0e-6; self.dispatch( c"kernel_qwen_gdn_step", &self.scratch.gdn_raw, Some(&self.scratch.gdn_out), Some(&self.scratch.controls), Some(recurrent), &[ self.view(self.weight(&format!("{prefix}.linear_attn.A_log"))?), self.view(self.weight(&format!("{prefix}.linear_attn.dt_bias"))?), ], &step, HEAD_DIM, GDN_HEADS_V, )?; let mut gate = args(); gate.u[0] = HEAD_DIM; gate.u[1] = GDN_HEADS_V; gate.f[0] = 1.0e-6; self.dispatch( c"kernel_qwen_gdn_norm_gate", &self.scratch.gdn_out, Some(&self.scratch.gdn_raw), Some(&self.scratch.controls), None, &[self.view(self.weight(&format!("{prefix}.linear_attn.norm.weight"))?)], &gate, GDN_HEADS_V, 1, )?; self.affine_mv_into( &self.affine( &format!("{prefix}.linear_attn.out_proj"), GDN_VALUE, HIDDEN, None, )?, &self.scratch.gdn_out, &self.scratch.hidden, GDN_VALUE, HIDDEN, ) } fn attention(&self, prefix: &str, layer: usize) -> Result<(), String> { if self.position + 1 > DENSE_BUDGET { return Err( "Qwen sparse QSA selection is required beyond 2048 tokens; native QSA belongs to issue #97" .into(), ); } let LayerState::Attention { kv } = &self.states[layer] else { return Err("Qwen attention graph received GDN state".into()); }; self.affine_mv_into( &self.affine( &format!("{prefix}.self_attn.q_proj"), HIDDEN, ATTN_WIDTH * 2, None, )?, &self.scratch.block, &self.scratch.q_packed, HIDDEN, ATTN_WIDTH * 2, )?; self.affine_mv_into( &self.affine( &format!("{prefix}.self_attn.k_proj"), HIDDEN, ATTN_KV_WIDTH, None, )?, &self.scratch.block, &self.scratch.k, HIDDEN, ATTN_KV_WIDTH, )?; self.affine_mv_into( &self.affine( &format!("{prefix}.self_attn.v_proj"), HIDDEN, ATTN_KV_WIDTH, None, )?, &self.scratch.block, &self.scratch.v, HIDDEN, ATTN_KV_WIDTH, )?; let mut split = args(); split.u[0] = ATTN_HEADS; split.u[1] = ATTN_DIM; self.dispatch( c"kernel_qwen_split_q_gate", &self.scratch.q, Some(&self.scratch.q_packed), Some(&self.scratch.q_gate), None, &[], &split, ATTN_WIDTH, 1, )?; self.head_norm_rope( &self.scratch.q, &self.scratch.attention, self.weight(&format!("{prefix}.self_attn.q_norm.weight"))?, ATTN_HEADS, )?; self.head_norm_rope( &self.scratch.k, &self.scratch.k_rope, self.weight(&format!("{prefix}.self_attn.k_norm.weight"))?, ATTN_KV_HEADS, )?; let mut store = args(); store.u[0] = ATTN_KV_WIDTH; store.u[1] = self.position; self.dispatch( c"kernel_qwen_store_kv_bf16", kv, Some(&self.scratch.k_rope), Some(&self.scratch.v), None, &[], &store, ATTN_KV_WIDTH, 1, )?; let mut dense = args(); dense.u[0] = ATTN_HEADS; dense.u[1] = ATTN_KV_HEADS; dense.u[2] = ATTN_DIM; dense.u[3] = self.position + 1; self.dispatch( c"kernel_qwen_dense_attention", &self.scratch.q, Some(&self.scratch.attention), Some(kv), None, &[], &dense, ATTN_HEADS, 1, )?; let mut gate = args(); gate.u[0] = ATTN_WIDTH; self.dispatch( c"kernel_qwen_gate_attention", &self.scratch.attention, Some(&self.scratch.q), Some(&self.scratch.q_gate), None, &[], &gate, ATTN_WIDTH, 1, )?; self.affine_mv_into( &self.affine( &format!("{prefix}.self_attn.o_proj"), ATTN_WIDTH, HIDDEN, None, )?, &self.scratch.attention, &self.scratch.hidden, ATTN_WIDTH, HIDDEN, ) } fn head_norm_rope( &self, input: &Buffer, output: &Buffer, weight: Weight<'_>, heads: u32, ) -> Result<(), String> { let mut values = args(); values.u[0] = ATTN_DIM; values.u[1] = 64; values.u[2] = heads; values.u[3] = self.position; values.f[0] = 1.0e-6; values.f[1] = 10_000_000.0; self.dispatch( c"kernel_qwen_head_norm_rope", output, Some(input), None, None, &[self.view(weight)], &values, ATTN_DIM, heads, ) } fn expert(&self, prefix: &str, expert: u32, weight: f32) -> Result<(), String> { for (name, target) in [ ("gate_proj", &self.scratch.gate), ("up_proj", &self.scratch.up), ] { self.affine_mv_into( &self.affine( &format!("{prefix}.mlp.switch_mlp.{name}"), HIDDEN, EXPERT_WIDTH, Some(expert), )?, &self.scratch.block, target, HIDDEN, EXPERT_WIDTH, )?; } let mut swiglu = args(); swiglu.u[0] = EXPERT_WIDTH; self.dispatch( c"kernel_qwen_swiglu", &self.scratch.mid, Some(&self.scratch.gate), Some(&self.scratch.up), None, &[], &swiglu, EXPERT_WIDTH, 1, )?; self.affine_mv_into( &self.affine( &format!("{prefix}.mlp.switch_mlp.down_proj"), EXPERT_WIDTH, HIDDEN, Some(expert), )?, &self.scratch.mid, &self.scratch.expert, EXPERT_WIDTH, HIDDEN, )?; let mut accumulate = args(); accumulate.u[0] = HIDDEN; accumulate.f[0] = weight; self.dispatch( c"kernel_qwen_accumulate", &self.scratch.moe, Some(&self.scratch.expert), None, None, &[], &accumulate, HIDDEN, 1, ) } fn shared_expert(&self, prefix: &str) -> Result<(), String> { for (name, target) in [ ("gate_proj", &self.scratch.gate), ("up_proj", &self.scratch.up), ] { self.affine_mv_into( &self.affine( &format!("{prefix}.mlp.shared_expert.{name}"), HIDDEN, EXPERT_WIDTH, None, )?, &self.scratch.block, target, HIDDEN, EXPERT_WIDTH, )?; } let mut swiglu = args(); swiglu.u[0] = EXPERT_WIDTH; self.dispatch( c"kernel_qwen_swiglu", &self.scratch.mid, Some(&self.scratch.gate), Some(&self.scratch.up), None, &[], &swiglu, EXPERT_WIDTH, 1, )?; self.affine_mv_into( &self.affine( &format!("{prefix}.mlp.shared_expert.down_proj"), EXPERT_WIDTH, HIDDEN, None, )?, &self.scratch.mid, &self.scratch.shared, EXPERT_WIDTH, HIDDEN, )?; self.affine_mv_into( &self.affine(&format!("{prefix}.mlp.shared_expert_gate"), HIDDEN, 1, None)?, &self.scratch.block, &self.scratch.gate, HIDDEN, 1, )?; let mut accumulate = args(); accumulate.u[0] = HIDDEN; self.dispatch( c"kernel_qwen_accumulate_sigmoid_scalar", &self.scratch.moe, Some(&self.scratch.shared), Some(&self.scratch.gate), None, &[], &accumulate, HIDDEN, 1, ) } fn final_output(&mut self) -> Result<(), String> { let commands = Commands::begin()?; self.final_mix()?; self.affine_mv_into( &self.affine("language_model.lm_head", HIDDEN, VOCAB, None)?, &self.scratch.block, &self.scratch.logits, HIDDEN, VOCAB, )?; commands.finish()?; self.scratch.logits.read_f32(&mut self.logits) } fn final_mix(&self) -> Result<(), String> { let prefix = "language_model.model.hyper_connection_mixer"; let norm = self.weight(&format!("{prefix}.hc_norm.weight"))?; let down = self.weight(&format!("{prefix}.input_mix_weight_down.weight"))?; let up = self.weight(&format!("{prefix}.input_mix_weight_up.weight"))?; let mut rms = args(); rms.u[0] = HC_WIDTH; rms.u[1] = HIDDEN; rms.f[0] = 1.0e-6; self.dispatch( c"kernel_qwen_zero_rms", &self.scratch.hc_norm, Some(&self.scratch.hc), None, None, &[self.view(norm)], &rms, HC, 1, )?; self.bf16_mv( down, &self.scratch.hc_norm, &self.scratch.rank, HC_WIDTH, HC_RANK, )?; let mut unary = args(); unary.u[0] = HC_RANK; self.dispatch( c"kernel_qwen_silu_div4", &self.scratch.rank, Some(&self.scratch.rank), None, None, &[], &unary, HC_RANK, 1, )?; self.bf16_mv( up, &self.scratch.rank, &self.scratch.hc_mix, HC_RANK, HC_WIDTH, )?; unary.u[0] = HC_WIDTH; self.dispatch( c"kernel_qwen_sigmoid", &self.scratch.hc_mix, Some(&self.scratch.hc_mix), None, None, &[], &unary, HC_WIDTH, 1, )?; let mut mix = args(); mix.u[0] = HIDDEN; self.dispatch( c"kernel_qwen_hyper_mix", &self.scratch.block, Some(&self.scratch.hc_norm), Some(&self.scratch.hc_mix), None, &[], &mix, HIDDEN, 1, ) } fn affine( &self, prefix: &str, in_dim: u32, out_dim: u32, expert: Option, ) -> Result, String> { let packed = self.weight(&format!("{prefix}.weight"))?; let scales = self.weight(&format!("{prefix}.scales"))?; let biases = self.weight(&format!("{prefix}.biases"))?; let bits = packed .tensor .quant_bits .ok_or_else(|| format!("{prefix} is not affine quantized"))?; let group = packed .tensor .group_size .ok_or_else(|| format!("{prefix} has no affine group size"))? as u32; if !matches!(bits, 2 | 4 | 8) || !in_dim.is_multiple_of(group) { return Err(format!("{prefix} has an incompatible affine layout")); } let expected = if expert.is_some() { vec![ EXPERTS as u64, out_dim as u64, (in_dim / (32 / bits)) as u64, ] } else { vec![out_dim as u64, (in_dim / (32 / bits)) as u64] }; if packed.tensor.shape != expected { return Err(format!("{prefix} has an incompatible executor shape")); } let expected_parameters = if expert.is_some() { vec![EXPERTS as u64, out_dim as u64, (in_dim / group) as u64] } else { vec![out_dim as u64, (in_dim / group) as u64] }; if packed.tensor.dtype != "U32" || scales.tensor.dtype != "BF16" || biases.tensor.dtype != "BF16" || scales.tensor.shape != expected_parameters || biases.tensor.shape != expected_parameters { return Err(format!("{prefix} has incompatible affine parameters")); } Ok(Affine { packed: Weight { tensor: packed.tensor, expert, }, scales: Weight { tensor: scales.tensor, expert, }, biases: Weight { tensor: biases.tensor, expert, }, bits, group, }) } fn weight(&self, name: &str) -> Result, String> { self.model.tensor(name).map(|tensor| Weight { tensor, expert: None, }) } fn view(&self, weight: Weight<'_>) -> QwenWeightView { let (map, _) = self.model.map(weight.tensor.map); let experts = weight.expert.map_or(1, |_| EXPERTS as u64); let bytes = (weight.tensor.range.end - weight.tensor.range.start) / experts; let offset = weight.tensor.range.start + u64::from(weight.expert.unwrap_or(0)) * bytes; QwenWeightView { map: map.as_ptr().cast(), size: map.len() as u64, offset, bytes, } } fn affine_mv_into( &self, weight: &Affine<'_>, input: &Buffer, output: &Buffer, in_dim: u32, out_dim: u32, ) -> Result<(), String> { let mut values = args(); values.u[0] = in_dim; values.u[1] = out_dim; values.u[2] = weight.bits; values.u[3] = weight.group; self.dispatch( c"kernel_qwen_affine_mv", output, Some(input), None, None, &[ self.view(weight.packed), self.view(weight.scales), self.view(weight.biases), ], &values, out_dim, 1, ) } fn bf16_mv( &self, weight: Weight<'_>, input: &Buffer, output: &Buffer, in_dim: u32, out_dim: u32, ) -> Result<(), String> { if weight.tensor.dtype != "BF16" || weight.tensor.shape != [out_dim as u64, in_dim as u64] { return Err(format!( "{} has an incompatible BF16 matrix shape", weight.tensor.name )); } let mut values = args(); values.u[0] = in_dim; values.u[1] = out_dim; self.dispatch( c"kernel_qwen_bf16_mv", output, Some(input), None, None, &[self.view(weight)], &values, out_dim, 1, ) } #[allow(clippy::too_many_arguments)] fn dispatch( &self, kernel: &CStr, out: &Buffer, a: Option<&Buffer>, b: Option<&Buffer>, c: Option<&Buffer>, weights: &[QwenWeightView], values: &QwenKernelArgs, grid_x: u32, grid_y: u32, ) -> Result<(), String> { dispatch_qwen(kernel, out, a, b, c, weights, values, grid_x, grid_y) } pub(super) fn prefill( &mut self, tokens: &[i32], mut progress: impl FnMut(u32) -> bool, ) -> Result { let mut completed = 0; for &token in tokens { if !progress(completed) { break; } self.eval(token)?; completed += 1; } progress(completed); Ok(completed as usize) } pub(super) fn logits(&self) -> &[f32] { &self.logits } pub(super) fn execution_stats(&self) -> ExecutionStats { ExecutionStats::default() } pub(super) fn model(&self) -> &QwenModel { &self.model } pub(super) fn context(&self) -> u32 { self.context } pub(super) fn position(&self) -> u32 { self.position } pub(super) fn tokens(&self) -> &[i32] { &self.tokens } pub(super) fn checkpoint_tag(&self) -> [u8; 32] { self.checkpoint_tag } pub(super) fn note_checkpoint_tag(&mut self, tag: [u8; 32]) { self.checkpoint_tag = tag; } pub(super) fn reset(&mut self) -> Result<(), String> { self.states = allocate_states(self.context)?; self.ple_state = allocate_ple_state()?; self.logits.fill(0.0); self.tokens.clear(); self.position = 0; self.checkpoint_tag = [0; 32]; Ok(()) } pub(super) fn align_prompt(&mut self, tokens: &[i32]) -> Result { if !tokens.starts_with(&self.tokens) { self.reset()?; } Ok(self.tokens.len()) } fn blank_resident(&self) -> Result { Ok(QwenResidentState { states: allocate_states(self.context)?, ple_state: allocate_ple_state()?, logits: vec![0.0; VOCAB as usize], tokens: Vec::new(), position: 0, checkpoint_tag: [0; 32], }) } pub(super) fn swap_resident_state( &mut self, state: &mut Option, ) -> Result<(), String> { let mut incoming = state.take().map_or_else(|| self.blank_resident(), Ok)?; std::mem::swap(&mut self.states, &mut incoming.states); std::mem::swap(&mut self.ple_state, &mut incoming.ple_state); std::mem::swap(&mut self.logits, &mut incoming.logits); std::mem::swap(&mut self.tokens, &mut incoming.tokens); std::mem::swap(&mut self.position, &mut incoming.position); std::mem::swap(&mut self.checkpoint_tag, &mut incoming.checkpoint_tag); *state = Some(incoming); Ok(()) } pub(in crate::engine) fn save_checkpoint( &mut self, path: &Path, tag: [u8; 32], progress: &mut impl FnMut(u64), ) -> Result<(), String> { if let Some(parent) = path.parent() { fs::create_dir_all(parent).map_err(|error| error.to_string())?; } let temporary = path.with_extension("tmp"); let mut file = File::create(&temporary).map_err(|error| error.to_string())?; file.write_all(CHECKPOINT_MAGIC) .map_err(|error| error.to_string())?; for value in [ CHECKPOINT_VERSION, self.context, self.position, VOCAB, LAYERS as u32, ] { write_u32(&mut file, value)?; } file.write_all(&self.model.checkpoint_identity()) .map_err(|error| error.to_string())?; file.write_all(&tag).map_err(|error| error.to_string())?; for token in self.ple_state.history { write_u32(&mut file, token as u32)?; } for &token in &self.tokens { write_u32(&mut file, token as u32)?; } for &logit in &self.logits { write_u32(&mut file, logit.to_bits())?; } let mut chunk = vec![0; CHECKPOINT_CHUNK]; write_buffer( &mut file, &self.ple_state.conv, 0, u64::from(HC_WIDTH) * PLE_CONV_STATE as u64 * 2, &mut chunk, progress, )?; for state in &self.states { match state { LayerState::Gdn { conv, recurrent } => { write_buffer( &mut file, conv, 0, u64::from(GDN_QKV) * 3 * 2, &mut chunk, progress, )?; write_buffer( &mut file, recurrent, 0, u64::from(GDN_HEADS_V) * HEAD_DIM as u64 * HEAD_DIM as u64 * 4, &mut chunk, progress, )?; } LayerState::Attention { kv } => { write_buffer( &mut file, kv, 0, u64::from(self.position) * ATTN_KV_WIDTH as u64 * 4, &mut chunk, progress, )?; } } } file.sync_all().map_err(|error| error.to_string())?; fs::rename(temporary, path).map_err(|error| error.to_string())?; self.checkpoint_tag = tag; Ok(()) } pub(in crate::engine) fn load_checkpoint( &mut self, path: &Path, progress: &mut impl FnMut(u64), ) -> Result { let mut file = match File::open(path) { Ok(file) => file, Err(error) if error.kind() == std::io::ErrorKind::NotFound => return Ok(false), Err(error) => return Err(error.to_string()), }; let mut magic = [0; 8]; file.read_exact(&mut magic) .map_err(|error| error.to_string())?; if &magic != CHECKPOINT_MAGIC { return Err("Qwen checkpoint has an invalid signature".into()); } for expected in [CHECKPOINT_VERSION, self.context] { if read_u32(&mut file)? != expected { return Err("Qwen checkpoint does not match the current executor".into()); } } let position = read_u32(&mut file)?; if position > self.context || read_u32(&mut file)? != VOCAB || read_u32(&mut file)? != LAYERS as u32 { return Err("Qwen checkpoint shape is invalid".into()); } let mut identity = [0; 32]; file.read_exact(&mut identity) .map_err(|error| error.to_string())?; if identity != self.model.checkpoint_identity() { return Err("Qwen checkpoint model identity changed".into()); } let mut tag = [0; 32]; file.read_exact(&mut tag) .map_err(|error| error.to_string())?; let mut ple_history = [0; PLE_HISTORY]; for token in &mut ple_history { let value = read_u32(&mut file)?; if value >= VOCAB { return Err("Qwen checkpoint PLE history is invalid".into()); } *token = value as i32; } let mut tokens = Vec::with_capacity(position as usize); for _ in 0..position { let token = read_u32(&mut file)?; if token >= VOCAB { return Err("Qwen checkpoint token is invalid".into()); } tokens.push(token as i32); } let mut logits = Vec::with_capacity(VOCAB as usize); for _ in 0..VOCAB { logits.push(f32::from_bits(read_u32(&mut file)?)); } self.reset()?; let mut chunk = vec![0; CHECKPOINT_CHUNK]; read_buffer( &mut file, &self.ple_state.conv, 0, u64::from(HC_WIDTH) * PLE_CONV_STATE as u64 * 2, &mut chunk, progress, )?; for state in &self.states { match state { LayerState::Gdn { conv, recurrent } => { read_buffer( &mut file, conv, 0, u64::from(GDN_QKV) * 3 * 2, &mut chunk, progress, )?; read_buffer( &mut file, recurrent, 0, u64::from(GDN_HEADS_V) * HEAD_DIM as u64 * HEAD_DIM as u64 * 4, &mut chunk, progress, )?; } LayerState::Attention { kv } => read_buffer( &mut file, kv, 0, u64::from(position) * ATTN_KV_WIDTH as u64 * 4, &mut chunk, progress, )?, } } let mut trailing = [0]; if file .read(&mut trailing) .map_err(|error| error.to_string())? != 0 { self.reset()?; return Err("Qwen checkpoint has trailing data".into()); } self.position = position; self.tokens = tokens; self.logits = logits; self.ple_state.history = ple_history; self.checkpoint_tag = tag; Ok(true) } } fn allocate_states(context: u32) -> Result, String> { (0..LAYERS) .map(|layer| { if layer % 4 == 3 { Ok(LayerState::Attention { kv: Buffer::bytes(u64::from(context) * ATTN_KV_WIDTH as u64 * 4)?, }) } else { let conv = Buffer::bytes(u64::from(GDN_QKV) * 3 * 2)?; let recurrent = Buffer::floats(u64::from(GDN_HEADS_V) * HEAD_DIM as u64 * HEAD_DIM as u64)?; conv.fill(0.0, u64::from(GDN_QKV) * 3 / 2)?; recurrent.fill( 0.0, u64::from(GDN_HEADS_V) * HEAD_DIM as u64 * HEAD_DIM as u64, )?; Ok(LayerState::Gdn { conv, recurrent }) } }) .collect() } fn allocate_ple_state() -> Result { let conv = Buffer::bytes(u64::from(HC_WIDTH) * PLE_CONV_STATE as u64 * 2)?; conv.fill(0.0, u64::from(HC_WIDTH) * PLE_CONV_STATE as u64 / 2)?; Ok(PleState { history: [EOS_TOKEN; PLE_HISTORY], conv, }) } fn ple_contract(model: &QwenModel) -> Result { let multipliers = read_i64_array::<3>( model, "language_model.model.layers.1.ple.ple_embedding.layer_multipliers", )?; let sizes = read_i64_array::( model, "language_model.model.layers.1.ple.ple_embedding.ngram_heads_vocab_sizes", )?; let offsets = read_i64_array::( model, "language_model.model.layers.1.ple.ple_embedding.ngram_heads_offsets", )?; let expected_multipliers = official_ple_multipliers(); let mut expected_sizes = [0; PLE_HEADS]; let mut expected_offsets = [0; PLE_HEADS]; let mut total = 0_i64; let mut prime = 19_999_999_i64; for head in 0..PLE_HEADS { prime = next_prime(prime); expected_sizes[head] = prime; expected_offsets[head] = total; total += prime; } if multipliers != expected_multipliers || sizes != expected_sizes || offsets != expected_offsets { return Err("Qwen PLE hash parameters do not match the official contract".into()); } for (name, shape, bits, group) in [ ("ngram.weight", [320_001_536, 20], Some(4), Some(32)), ("ngram.scales", [320_001_536, 5], Some(4), Some(32)), ("ngram.biases", [320_001_536, 5], Some(4), Some(32)), ] { let tensor = model.tensor(name)?; if tensor.shape != shape || tensor.quant_bits != bits || tensor.group_size != group { return Err(format!("{name} does not match the Qwen PLE row layout")); } } Ok(PleContract { multipliers, sizes, offsets, }) } fn read_i64_array(model: &QwenModel, name: &str) -> Result<[i64; N], String> { let tensor = model.tensor(name)?; if tensor.dtype != "I64" || tensor.shape != [N as u64] { return Err(format!("{name} does not match the Qwen PLE integer layout")); } let bytes = model.tensor_bytes(tensor)?; if bytes.len() != N * 8 { return Err(format!("{name} has an invalid byte length")); } Ok(std::array::from_fn(|index| { i64::from_le_bytes(bytes[index * 8..index * 8 + 8].try_into().unwrap()) })) } fn official_ple_multipliers() -> [i64; 3] { const GAMMA: u64 = 0x9e37_79b9_7f4a_7c15; const M1: u64 = 0xbf58_476d_1ce4_e5b9; const M2: u64 = 0x94d0_49bb_1331_11eb; let bound = (i64::MAX / VOCAB as i64 / 2) as u64; std::array::from_fn(|index| { let mut value = 1234_u64.wrapping_add(GAMMA.wrapping_mul(index as u64 + 1)); value = value.wrapping_add(GAMMA); value = (value ^ (value >> 30)).wrapping_mul(M1); value = (value ^ (value >> 27)).wrapping_mul(M2); value ^= value >> 31; (2 * (value % bound) + 1) as i64 }) } fn next_prime(mut value: i64) -> i64 { loop { value += 1; if value % 2 != 0 && (3..=((value as f64).sqrt() as i64)) .step_by(2) .all(|divisor| value % divisor != 0) { return value; } } } fn ple_rows( contract: &PleContract, history: [i32; PLE_HISTORY], token: i32, ) -> Result<[u64; PLE_HEADS], String> { if token < 0 || token as u32 >= VOCAB { return Err(format!("token {token} is outside the Qwen vocabulary")); } let shifted = [token as i64, history[1] as i64, history[0] as i64]; let two = shifted[0].wrapping_mul(contract.multipliers[0]) ^ shifted[1].wrapping_mul(contract.multipliers[1]); let three = two ^ shifted[2].wrapping_mul(contract.multipliers[2]); Ok(std::array::from_fn(|head| { let mixed = if head < 8 { two } else { three }; (contract.offsets[head] + mixed.rem_euclid(contract.sizes[head])) as u64 })) } fn advance_ple_history(history: [i32; PLE_HISTORY], token: i32) -> [i32; PLE_HISTORY] { if token == EOS_TOKEN { [EOS_TOKEN; PLE_HISTORY] } else { [history[1], token] } } fn gather_ple_row( row: u64, packed: &[u8], scales: &[u8], biases: &[u8], ) -> Result<[u8; PLE_ROW_BYTES], String> { let row = usize::try_from(row).map_err(|_| "Qwen PLE row exceeds this platform")?; let mut value = [0; PLE_ROW_BYTES]; value[..80].copy_from_slice(ple_row_slice(packed, row, 80)?); value[80..90].copy_from_slice(ple_row_slice(scales, row, 10)?); value[90..].copy_from_slice(ple_row_slice(biases, row, 10)?); Ok(value) } fn ple_row_slice(data: &[u8], row: usize, width: usize) -> Result<&[u8], String> { let start = row .checked_mul(width) .ok_or_else(|| "Qwen PLE row offset overflows".to_owned())?; data.get(start..start + width) .ok_or_else(|| format!("Qwen PLE row {row} is truncated")) } fn args() -> QwenKernelArgs { QwenKernelArgs::default() } #[allow(clippy::too_many_arguments)] fn dispatch_qwen( kernel: &CStr, out: &Buffer, a: Option<&Buffer>, b: Option<&Buffer>, c: Option<&Buffer>, weights: &[QwenWeightView], values: &QwenKernelArgs, grid_x: u32, grid_y: u32, ) -> Result<(), String> { call( unsafe { ds4_gpu_qwen_dispatch( kernel.as_ptr(), out.raw(), a.map_or(std::ptr::null(), |buffer| buffer.raw().cast_const()), b.map_or(std::ptr::null(), |buffer| buffer.raw().cast_const()), c.map_or(std::ptr::null(), |buffer| buffer.raw().cast_const()), weights.as_ptr(), weights.len() as u32, values, grid_x, grid_y, ) }, kernel.to_str().unwrap_or("running a Qwen Metal kernel"), ) } #[cfg(test)] mod tests { use super::*; use memmap2::MmapOptions; use sha2::{Digest, Sha256}; use std::fs; use std::path::PathBuf; fn bf16(value: f32) -> u16 { let bits = value.to_bits(); ((bits + 0x7fff + ((bits >> 16) & 1)) >> 16) as u16 } fn close(actual: f32, expected: f32) { assert!( (actual - expected).abs() <= 2.0e-4 * expected.abs().max(1.0), "{actual} != {expected}" ); } #[test] fn qwen_ple_hash_contract_matches_golden_boundaries() { assert_eq!( official_ple_multipliers(), [23_703_573_157_769, 20_109_073_645_365, 8_052_911_324_071] ); let mut contract = PleContract { multipliers: official_ple_multipliers(), sizes: [0; PLE_HEADS], offsets: [0; PLE_HEADS], }; let mut prime = 19_999_999; let mut offset = 0; for head in 0..PLE_HEADS { prime = next_prime(prime); contract.sizes[head] = prime; contract.offsets[head] = offset; offset += prime; } let initial = ple_rows(&contract, [EOS_TOKEN; 2], 1).unwrap(); let repeated = ple_rows(&contract, [1, 1], 1).unwrap(); let boundary = ple_rows(&contract, [EOS_TOKEN, 42], 43).unwrap(); assert_eq!( initial, [ 16_121_432, 28_938_500, 59_087_997, 73_487_090, 81_148_277, 104_500_129, 120_276_032, 149_373_875, 176_283_436, 184_305_849, 216_528_839, 231_080_079, 257_961_536, 266_068_568, 289_043_455, 305_959_965, ] ); assert_eq!( repeated, [ 6_868_091, 38_325_817, 54_054_700, 68_075_137, 82_949_816, 101_241_419, 138_678_867, 155_262_032, 176_541_251, 196_154_476, 215_703_237, 234_413_824, 254_220_543, 274_027_268, 293_962_951, 313_640_732, ] ); assert_eq!( boundary, [ 18_529_343, 23_547_650, 56_056_978, 73_570_159, 88_581_601, 113_585_506, 121_091_299, 151_099_148, 175_585_266, 184_439_538, 216_587_431, 222_082_137, 250_284_866, 278_847_169, 281_781_121, 317_050_322, ] ); assert_eq!(contract.offsets[0], 0); assert_eq!(contract.offsets[8], 160_000_374); assert_eq!(contract.offsets[15] + contract.sizes[15], 320_001_446); for (head, &row) in initial.iter().enumerate() { assert!( (contract.offsets[head]..contract.offsets[head] + contract.sizes[head]) .contains(&(row as i64)) ); } assert_eq!(320_001_536 % 128, 0); let history = advance_ple_history([EOS_TOKEN; 2], 1); assert_eq!(advance_ple_history(history, 2), [1, 2]); assert_eq!(advance_ple_history([1, 2], EOS_TOKEN), [EOS_TOKEN; 2]); let hash_chunks = |chunk: usize| { let mut history = [EOS_TOKEN; PLE_HISTORY]; [1, 2, EOS_TOKEN, 3, 4] .chunks(chunk) .flat_map(|tokens| { tokens .iter() .map(|&token| { let rows = ple_rows(&contract, history, token).unwrap(); history = advance_ple_history(history, token); rows }) .collect::>() }) .collect::>() }; assert_eq!(hash_chunks(1), hash_chunks(2)); assert_eq!(hash_chunks(1), hash_chunks(5)); } #[test] fn qwen_ple_gather_copies_only_the_requested_row() { let packed = (0..240).map(|value| value as u8).collect::>(); let scales = (0..30).map(|value| (value + 17) as u8).collect::>(); let biases = (0..30).map(|value| (value + 47) as u8).collect::>(); let direct = gather_ple_row(1, &packed, &scales, &biases).unwrap(); assert_eq!(&direct[..80], &packed[80..160]); assert_eq!(&direct[80..90], &scales[10..20]); assert_eq!(&direct[90..], &biases[10..20]); assert_eq!( gather_ple_row(1, &packed, &scales, &biases).unwrap(), direct ); assert!( gather_ple_row(3, &packed, &scales, &biases) .unwrap_err() .contains("truncated") ); } #[test] #[ignore = "requires Apple Metal"] fn qwen_metal_primitives_match_reference_vectors() { configure_sources().unwrap(); let _context = Context::open_qwen(0).unwrap(); let input = Buffer::floats(64).unwrap(); input.write_f32(&vec![1.0; 64]).unwrap(); let output = Buffer::floats(64).unwrap(); let path = std::env::temp_dir().join(format!( "ds4-qwen96-weights-{}-{}", std::process::id(), std::thread::current().name().unwrap_or("test") )); // Safetensors headers are not guaranteed to align the data region. let mut bytes = vec![0]; for _ in 0..8 { bytes.extend_from_slice(&0x3333_3333_u32.to_le_bytes()); } for value in [ 0.5, -1.0, 0.0, 0.0, 1.0, 2.0, 3.0, 4.0, 1.0, 1.0, 0.0, 0.0, 0.0, 0.0, ] { bytes.extend_from_slice(&bf16(value).to_le_bytes()); } fs::write(&path, bytes).unwrap(); let file = File::open(&path).unwrap(); // SAFETY: this test owns the read-only file for the lifetime of the mapping. let map = unsafe { MmapOptions::new().map(&file).unwrap() }; let view = |offset, size| QwenWeightView { map: map.as_ptr().cast(), size: map.len() as u64, offset, bytes: size, }; let mut affine = args(); affine.u[0] = 64; affine.u[1] = 1; affine.u[2] = 4; affine.u[3] = 64; dispatch_qwen( c"kernel_qwen_affine_mv", &output, Some(&input), None, None, &[view(1, 32), view(33, 2), view(35, 2)], &affine, 1, 1, ) .unwrap(); let mut scalar = [0.0]; output.read_f32(&mut scalar).unwrap(); close(scalar[0], 32.0); let conv_input = Buffer::floats(1).unwrap(); conv_input.write_f32(&[2.0]).unwrap(); let conv_state = Buffer::bytes(6).unwrap(); conv_state.write(0, &[0; 6]).unwrap(); let conv_output = Buffer::floats(1).unwrap(); let mut conv = args(); conv.u[0] = 1; dispatch_qwen( c"kernel_qwen_conv_silu", &conv_output, Some(&conv_input), Some(&conv_state), None, &[view(41, 8)], &conv, 1, 1, ) .unwrap(); conv_output.read_f32(&mut scalar).unwrap(); close(scalar[0], 8.0 / (1.0 + (-8.0_f32).exp())); assert_eq!( { let mut state = [0; 6]; conv_state.read(0, &mut state).unwrap(); state }, [0, 0, 0, 0, 0, 64] ); let qkv = Buffer::floats(6).unwrap(); qkv.write_f32(&[3.0, 4.0, 0.0, 2.0, 5.0, 7.0]).unwrap(); let controls = Buffer::floats(4).unwrap(); controls.write_f32(&[0.0; 4]).unwrap(); let recurrent = Buffer::floats(4).unwrap(); recurrent.write_f32(&[1.0, 2.0, 3.0, 4.0]).unwrap(); let raw = Buffer::floats(2).unwrap(); let gated = Buffer::floats(2).unwrap(); let mut step = args(); step.u[0] = 2; step.u[1] = 1; step.u[2] = 1; step.u[3] = 2; step.u[4] = 3; step.f[0] = 1.0e-6; dispatch_qwen( c"kernel_qwen_gdn_step", &raw, Some(&qkv), Some(&controls), Some(&recurrent), &[view(37, 2), view(39, 2)], &step, 2, 1, ) .unwrap(); let mut raw_values = [0.0; 2]; raw.read_f32(&mut raw_values).unwrap(); close(raw_values[0], 2.0506096); close(raw_values[1], 2.9698484); let mut norm = args(); norm.u[0] = 2; norm.u[1] = 1; norm.f[0] = 1.0e-6; dispatch_qwen( c"kernel_qwen_gdn_norm_gate", &gated, Some(&raw), Some(&controls), None, &[view(49, 4)], &norm, 1, 1, ) .unwrap(); let mut gated_values = [0.0; 2]; gated.read_f32(&mut gated_values).unwrap(); let rms = ((raw_values[0].powi(2) + raw_values[1].powi(2)) * 0.5 + 1.0e-6).sqrt(); close(gated_values[0], raw_values[0] / rms * 0.5); close(gated_values[1], raw_values[1] / rms * 0.5); let router = Buffer::floats(EXPERTS.into()).unwrap(); let mut logits = vec![-100.0; EXPERTS as usize]; for (index, value) in logits.iter_mut().take(EXPERTS_USED).enumerate() { *value = index as f32; } router.write_f32(&logits).unwrap(); let ids = Buffer::bytes((EXPERTS_USED * 4) as u64).unwrap(); let weights = Buffer::floats(EXPERTS_USED as u64).unwrap(); let mut route = args(); route.u[0] = EXPERTS; dispatch_qwen( c"kernel_qwen_route_top10", &ids, Some(&weights), Some(&router), None, &[], &route, 1, 1, ) .unwrap(); let mut selected = [0; EXPERTS_USED]; ids.read_i32(&mut selected).unwrap(); assert_eq!(selected, [9, 8, 7, 6, 5, 4, 3, 2, 1, 0]); let mut probabilities = [0.0; EXPERTS_USED]; weights.read_f32(&mut probabilities).unwrap(); close(probabilities.iter().sum(), 1.0); let norm_input = Buffer::floats(4).unwrap(); norm_input.write_f32(&[1.0, 2.0, 3.0, 4.0]).unwrap(); let norm_output = Buffer::floats(4).unwrap(); let mut zero_norm = args(); zero_norm.u[0] = 4; zero_norm.u[1] = 2; zero_norm.f[0] = 1.0e-6; dispatch_qwen( c"kernel_qwen_zero_rms", &norm_output, Some(&norm_input), None, None, &[view(53, 8)], &zero_norm, 2, 1, ) .unwrap(); let mut normalized = [0.0; 4]; norm_output.read_f32(&mut normalized).unwrap(); for (actual, expected) in normalized.into_iter().zip([ 1.0 / (2.5_f32 + 1.0e-6).sqrt(), 2.0 / (2.5_f32 + 1.0e-6).sqrt(), 3.0 / (12.5_f32 + 1.0e-6).sqrt(), 4.0 / (12.5_f32 + 1.0e-6).sqrt(), ]) { close(actual, expected); } let hyper_input = Buffer::floats(8).unwrap(); hyper_input .write_f32(&[1.0, 2.0, 3.0, 4.0, 5.0, 6.0, 7.0, 8.0]) .unwrap(); let hyper_weights = Buffer::floats(8).unwrap(); hyper_weights.write_f32(&[1.0; 8]).unwrap(); let mixed = Buffer::floats(2).unwrap(); let mut hyper = args(); hyper.u[0] = 2; dispatch_qwen( c"kernel_qwen_hyper_mix", &mixed, Some(&hyper_input), Some(&hyper_weights), None, &[], &hyper, 2, 1, ) .unwrap(); let mut mixed_values = [0.0; 2]; mixed.read_f32(&mut mixed_values).unwrap(); assert_eq!(mixed_values, [4.0, 5.0]); let injection = Buffer::floats(4).unwrap(); injection.write_f32(&[1.0, 2.0, 3.0, 4.0]).unwrap(); let combined = Buffer::floats(8).unwrap(); dispatch_qwen( c"kernel_qwen_hyper_inject", &combined, Some(&hyper_input), Some(&mixed), Some(&injection), &[], &hyper, 8, 1, ) .unwrap(); let mut combined_values = [0.0; 8]; combined.read_f32(&mut combined_values).unwrap(); assert_eq!( combined_values, [5.0, 7.0, 11.0, 14.0, 17.0, 21.0, 23.0, 28.0] ); let expert_gate = Buffer::floats(2).unwrap(); expert_gate.write_f32(&[0.0, 2.0]).unwrap(); let expert_up = Buffer::floats(2).unwrap(); expert_up.write_f32(&[4.0, 3.0]).unwrap(); let expert = Buffer::floats(2).unwrap(); let mut moe = args(); moe.u[0] = 2; dispatch_qwen( c"kernel_qwen_swiglu", &expert, Some(&expert_gate), Some(&expert_up), None, &[], &moe, 2, 1, ) .unwrap(); let accumulated = Buffer::floats(2).unwrap(); accumulated.fill(0.0, 2).unwrap(); moe.f[0] = 0.25; dispatch_qwen( c"kernel_qwen_accumulate", &accumulated, Some(&expert), None, None, &[], &moe, 2, 1, ) .unwrap(); let shared = Buffer::floats(2).unwrap(); shared.write_f32(&[2.0, 4.0]).unwrap(); let shared_gate = Buffer::floats(1).unwrap(); shared_gate.write_f32(&[0.0]).unwrap(); dispatch_qwen( c"kernel_qwen_accumulate_sigmoid_scalar", &accumulated, Some(&shared), Some(&shared_gate), None, &[], &moe, 2, 1, ) .unwrap(); let mut moe_values = [0.0; 2]; accumulated.read_f32(&mut moe_values).unwrap(); close(moe_values[0], 1.0); close( moe_values[1], 0.25 * (2.0 / (1.0 + (-2.0_f32).exp()) * 3.0) + 2.0, ); let rope_input = Buffer::floats(4).unwrap(); rope_input.write_f32(&[1.0, 2.0, 3.0, 4.0]).unwrap(); let rope_output = Buffer::floats(4).unwrap(); let mut rope = args(); rope.u[0] = 4; rope.u[1] = 4; rope.u[2] = 1; rope.u[3] = 1; rope.f[0] = 1.0e-6; rope.f[1] = 10_000.0; dispatch_qwen( c"kernel_qwen_head_norm_rope", &rope_output, Some(&rope_input), None, None, &[view(53, 8)], &rope, 4, 1, ) .unwrap(); let rms = (7.5_f32 + 1.0e-6).sqrt(); let theta = [1.0_f32, 0.01]; let normalized = [1.0 / rms, 2.0 / rms, 3.0 / rms, 4.0 / rms]; let expected_rope = [ normalized[0] * theta[0].cos() - normalized[2] * theta[0].sin(), normalized[1] * theta[1].cos() - normalized[3] * theta[1].sin(), normalized[2] * theta[0].cos() + normalized[0] * theta[0].sin(), normalized[3] * theta[1].cos() + normalized[1] * theta[1].sin(), ]; let mut actual_rope = [0.0; 4]; rope_output.read_f32(&mut actual_rope).unwrap(); for (actual, expected) in actual_rope.into_iter().zip(expected_rope) { close(actual, expected); } let cache = Buffer::bytes(16).unwrap(); let key = Buffer::floats(2).unwrap(); let value = Buffer::floats(2).unwrap(); let mut store = args(); store.u[0] = 2; for (position, (keys, values)) in [([1.0, 0.0], [2.0, 3.0]), ([0.0, 1.0], [5.0, 7.0])] .into_iter() .enumerate() { key.write_f32(&keys).unwrap(); value.write_f32(&values).unwrap(); store.u[1] = position as u32; dispatch_qwen( c"kernel_qwen_store_kv_bf16", &cache, Some(&key), Some(&value), None, &[], &store, 2, 1, ) .unwrap(); } let query = Buffer::floats(2).unwrap(); query.write_f32(&[1.0, 0.0]).unwrap(); let attention = Buffer::floats(2).unwrap(); let mut dense = args(); dense.u[0] = 1; dense.u[1] = 1; dense.u[2] = 2; dense.u[3] = 2; dispatch_qwen( c"kernel_qwen_dense_attention", &attention, Some(&query), Some(&cache), None, &[], &dense, 1, 1, ) .unwrap(); let first = (1.0_f32 / 2.0_f32.sqrt()).exp(); let probability = first / (first + 1.0); let mut actual_attention = [0.0; 2]; attention.read_f32(&mut actual_attention).unwrap(); close( actual_attention[0], probability * 2.0 + (1.0 - probability) * 5.0, ); close( actual_attention[1], probability * 3.0 + (1.0 - probability) * 7.0, ); let ple_packed = Buffer::bytes(80).unwrap(); ple_packed.write(0, &[0x33; 80]).unwrap(); let ple_scales = Buffer::bytes(10).unwrap(); ple_scales .write(0, &bf16(0.5).to_le_bytes().repeat(5)) .unwrap(); let ple_biases = Buffer::bytes(10).unwrap(); ple_biases .write(0, &bf16(-1.0).to_le_bytes().repeat(5)) .unwrap(); let ple_embedding = Buffer::floats(160).unwrap(); let mut dequant = args(); dequant.u[0] = 160; dequant.u[1] = 1; dequant.u[2] = 4; dequant.u[3] = 32; dispatch_qwen( c"kernel_qwen_ple_dequant", &ple_embedding, Some(&ple_packed), Some(&ple_scales), Some(&ple_biases), &[], &dequant, 160, 1, ) .unwrap(); let mut embedding = [0.0; 160]; ple_embedding.read_f32(&mut embedding).unwrap(); assert!(embedding.into_iter().all(|value| value == 0.5)); let ple_key = Buffer::floats(8).unwrap(); ple_key.write_f32(&[1.0; 8]).unwrap(); let ple_query = Buffer::floats(8).unwrap(); ple_query.write_f32(&[1.0; 8]).unwrap(); let ple_value = Buffer::floats(2).unwrap(); ple_value.write_f32(&[2.0, 3.0]).unwrap(); let ple_gated = Buffer::floats(8).unwrap(); let mut gate = args(); gate.u[0] = 2; dispatch_qwen( c"kernel_qwen_ple_gate", &ple_gated, Some(&ple_key), Some(&ple_query), Some(&ple_value), &[], &gate, 4, 1, ) .unwrap(); let expected_gate = 1.0 / (1.0 + (-2.0_f32.sqrt().sqrt()).exp()); let mut gated = [0.0; 8]; ple_gated.read_f32(&mut gated).unwrap(); for stream in 0..4 { close(gated[stream * 2], 2.0 * expected_gate); close(gated[stream * 2 + 1], 3.0 * expected_gate); } let ple_normalized = Buffer::floats(1).unwrap(); ple_normalized.write_f32(&[2.0]).unwrap(); let ple_gate_value = Buffer::floats(1).unwrap(); ple_gate_value.write_f32(&[0.5]).unwrap(); let ple_state = Buffer::bytes(18).unwrap(); ple_state.write(0, &[0; 18]).unwrap(); let ple_output = Buffer::floats(1).unwrap(); let mut conv = args(); conv.u[0] = 1; dispatch_qwen( c"kernel_qwen_ple_conv", &ple_output, Some(&ple_gate_value), Some(&ple_normalized), Some(&ple_state), &[view(41, 8)], &conv, 1, 1, ) .unwrap(); ple_output.read_f32(&mut scalar).unwrap(); close(scalar[0], 0.5 + 8.0 / (1.0 + (-8.0_f32).exp())); let mut ple_history = [0; 18]; ple_state.read(0, &mut ple_history).unwrap(); assert_eq!(&ple_history[16..], &bf16(2.0).to_le_bytes()); drop(map); drop(file); fs::remove_file(path).unwrap(); } #[test] #[ignore = "requires the pinned 105 GB Qwen artifact set and Apple Metal"] fn qwen_core_boundary_and_checkpoint_are_stable() { configure_sources().unwrap(); let root = std::env::var_os("DS4SERVER_QWEN38_SOURCE") .map(PathBuf::from) .expect("set DS4SERVER_QWEN38_SOURCE to the pinned artifact directory"); let model = QwenModel::open(&root, 4).unwrap(); let residency_before = model.mapped_residency().unwrap(); let mut executor = QwenExecutor::open(model, 4).unwrap(); executor.eval(1).unwrap(); let residency_after = executor.model().mapped_residency().unwrap(); assert!(residency_after.0 <= executor.model().memory().resident_core); assert!(residency_after.1 <= executor.model().memory().mapped_ple); eprintln!( "Qwen mapped residency core/PLE before {:?}, after {:?}", residency_before, residency_after ); assert!(executor.logits.iter().all(|value| value.is_finite())); let mut digest = Sha256::new(); for value in &executor.logits { digest.update(value.to_bits().to_le_bytes()); } let digest: [u8; 32] = digest.finalize().into(); assert_eq!( digest, [ 137, 244, 133, 253, 201, 214, 196, 144, 249, 130, 28, 63, 124, 75, 32, 40, 16, 148, 148, 123, 5, 50, 90, 165, 101, 44, 223, 164, 62, 54, 164, 86, ] ); let reference = [ executor.logits[0], executor.logits[1], executor.logits[1000], executor.logits[VOCAB as usize - 1], ]; let next = executor .logits .iter() .enumerate() .max_by(|a, b| a.1.total_cmp(b.1)) .unwrap() .0 as i32; for (actual, expected) in reference .into_iter() .zip([6.406_557, 2.082_818_3, -3.058_045_6, -0.121_010_3]) { close(actual, expected); } assert_eq!(next, 89_648); let checkpoint = std::env::temp_dir().join(format!( "ds4-qwen96-checkpoint-{}-{}", std::process::id(), std::thread::current().name().unwrap_or("test") )); executor .save_checkpoint(&checkpoint, [7; 32], &mut |_| {}) .unwrap(); executor.eval(next).unwrap(); let continued = [ executor.logits[0], executor.logits[1], executor.logits[1000], executor.logits[VOCAB as usize - 1], ]; let LayerState::Gdn { conv, recurrent } = &executor.states[0] else { unreachable!() }; let mut continued_conv = vec![0; GDN_QKV as usize * 3 * 2]; conv.read(0, &mut continued_conv).unwrap(); let mut continued_recurrent = vec![0; 4 * 1024]; recurrent.read(0, &mut continued_recurrent).unwrap(); let mut continued_ple = vec![0; HC_WIDTH as usize * PLE_CONV_STATE as usize * 2]; executor.ple_state.conv.read(0, &mut continued_ple).unwrap(); assert!(executor.load_checkpoint(&checkpoint, &mut |_| {}).unwrap()); assert_eq!(executor.position, 1); assert_eq!(executor.tokens, [1]); assert_eq!(executor.ple_state.history, [EOS_TOKEN, 1]); assert_eq!(executor.checkpoint_tag, [7; 32]); assert_eq!( [ executor.logits[0], executor.logits[1], executor.logits[1000], executor.logits[VOCAB as usize - 1], ], reference ); executor.eval(next).unwrap(); for (actual, expected) in [ executor.logits[0], executor.logits[1], executor.logits[1000], executor.logits[VOCAB as usize - 1], ] .into_iter() .zip(continued) { close(actual, expected); } let LayerState::Gdn { conv, recurrent } = &executor.states[0] else { unreachable!() }; let mut resumed_conv = vec![0; continued_conv.len()]; conv.read(0, &mut resumed_conv).unwrap(); assert_eq!(resumed_conv, continued_conv); let mut resumed_recurrent = vec![0; continued_recurrent.len()]; recurrent.read(0, &mut resumed_recurrent).unwrap(); assert_eq!(resumed_recurrent, continued_recurrent); let mut resumed_ple = vec![0; continued_ple.len()]; executor.ple_state.conv.read(0, &mut resumed_ple).unwrap(); assert_eq!(resumed_ple, continued_ple); executor.reset().unwrap(); assert_eq!(executor.position, 0); assert_eq!(executor.ple_state.history, [EOS_TOKEN; PLE_HISTORY]); let mut reset_ple = vec![1; continued_ple.len()]; executor.ple_state.conv.read(0, &mut reset_ple).unwrap(); assert!(reset_ple.into_iter().all(|byte| byte == 0)); executor.eval(1).unwrap(); for (actual, expected) in [ executor.logits[0], executor.logits[1], executor.logits[1000], executor.logits[VOCAB as usize - 1], ] .into_iter() .zip(reference) { close(actual, expected); } executor.reset().unwrap(); executor.eval(EOS_TOKEN).unwrap(); assert_eq!(executor.ple_state.history, [EOS_TOKEN; PLE_HISTORY]); executor.position = DENSE_BUDGET; executor.context = DENSE_BUDGET + 1; let error = executor.eval(1).unwrap_err(); assert!(error.contains("issue #97")); fs::remove_file(checkpoint).unwrap(); } }