mod checkpoint; mod glm; mod gpu; mod hotlist; mod profile; use glm::GlmExecutor; use gpu::*; use profile::ExpertProfile; use super::gguf::{F16, F32, Gguf, IQ2_XXS, MXFP4, Q2_K, Q4_K, Q8_0, Tensor as GgufTensor}; use super::validation::{DsparkConfig, SupportKind, dspark_config}; use super::{Model, ModelFamily, Rng, exact_delta_sample}; use crate::model::ModelChoice; use crate::settings::{ EngineSpeculativeSettings, EngineSsdSettings, EngineSteeringSettings, ReasoningMode, }; use sha2::{Digest, Sha256}; use std::env; use std::ffi::{CStr, c_char, c_void}; use std::fs::{self, File}; use std::io::{Read, Write}; use std::os::unix::fs::FileExt; use std::path::Path; use std::ptr::NonNull; use std::sync::Arc; use std::thread::JoinHandle; use std::time::{Duration, Instant, UNIX_EPOCH}; const CHECKPOINT_MAGIC: &[u8; 8] = b"DS4RKV01"; const CHECKPOINT_VERSION: u32 = 2; const CHECKPOINT_IO_CHUNK: usize = 8 * 1024 * 1024; const DEFAULT_PREFILL_CHUNK: u32 = 4096; unsafe extern "C" { fn getenv(name: *const c_char) -> *mut c_char; } fn environment_present(name: &CStr) -> bool { // SAFETY: every caller passes a static, NUL-terminated C string and only // checks whether the process environment contains it. !unsafe { getenv(name.as_ptr()) }.is_null() } const SOURCES: [(&str, &str); 19] = [ ("DS4_METAL_FLASH_ATTN_SOURCE", "flash_attn.metal"), ("DS4_METAL_DENSE_SOURCE", "dense.metal"), ("DS4_METAL_MOE_SOURCE", "moe.metal"), ("DS4_METAL_DSV4_HC_SOURCE", "dsv4_hc.metal"), ("DS4_METAL_UNARY_SOURCE", "unary.metal"), ("DS4_METAL_DSV4_KV_SOURCE", "dsv4_kv.metal"), ("DS4_METAL_DSV4_ROPE_SOURCE", "dsv4_rope.metal"), ("DS4_METAL_DSV4_MISC_SOURCE", "dsv4_misc.metal"), ("DS4_METAL_ARGSORT_SOURCE", "argsort.metal"), ("DS4_METAL_CPY_SOURCE", "cpy.metal"), ("DS4_METAL_CONCAT_SOURCE", "concat.metal"), ("DS4_METAL_GET_ROWS_SOURCE", "get_rows.metal"), ("DS4_METAL_SUM_ROWS_SOURCE", "sum_rows.metal"), ("DS4_METAL_SOFTMAX_SOURCE", "softmax.metal"), ("DS4_METAL_REPEAT_SOURCE", "repeat.metal"), ("DS4_METAL_GLU_SOURCE", "glu.metal"), ("DS4_METAL_NORM_SOURCE", "norm.metal"), ("DS4_METAL_BIN_SOURCE", "bin.metal"), ("DS4_METAL_SET_ROWS_SOURCE", "set_rows.metal"), ]; // The Metal boundary uses this only to decide whether diagnostic logs get ANSI // color. App output is never a terminal, so Rust owns the one non-Metal symbol // that ds4_metal.m expects without pulling in ds4.c. #[unsafe(no_mangle)] extern "C" fn ds4_log_is_tty(_stream: *mut c_void) -> bool { false } pub(crate) fn configure_sources() -> Result<(), String> { let bundled = env::current_exe().ok().and_then(|path| { path.parent()? .parent() .map(|path| path.join("Resources/metal")) }); let checkout = Path::new(env!("CARGO_MANIFEST_DIR")).join("metal"); let directory = bundled .filter(|path| path.join("dense.metal").is_file()) .unwrap_or(checkout); for (name, file) in SOURCES { let path = directory.join(file); if !path.is_file() { return Err(format!( "required Metal kernel is missing: {}", path.display() )); } // SAFETY: main calls this before Iced starts any application threads. unsafe { env::set_var(name, path) }; } Ok(()) } #[derive(Clone, Copy)] struct Weight { offset: u64, kind: u32, bytes: u64, dims: [u64; 3], } impl Weight { fn bind(model: &Gguf, name: &str) -> Result { let tensor = model.tensor(name)?; Ok(Self::from_tensor(tensor)) } fn optional(model: &Gguf, name: &str) -> Option { model.tensors.get(name).map(Self::from_tensor) } fn from_tensor(tensor: &GgufTensor) -> Self { let mut dims = [1; 3]; for (to, from) in dims.iter_mut().zip(&tensor.dims) { *to = *from; } Self { offset: tensor.offset, kind: tensor.kind, bytes: tensor.bytes, dims, } } } struct Layer { hc_attn_fn: Weight, hc_attn_scale: Weight, hc_attn_base: Weight, attn_norm: Weight, attn_q_a: Weight, attn_q_a_norm: Weight, attn_q_b: Weight, attn_kv: Weight, attn_kv_norm: Weight, attn_sinks: Weight, attn_output_a: Weight, attn_output_b: Weight, attn_compressor: Option, indexer: Option, hc_ffn_fn: Weight, hc_ffn_scale: Weight, hc_ffn_base: Weight, ffn_norm: Weight, router: Weight, router_bias: Option, router_hash: Option, expert_gate: Weight, expert_up: Weight, expert_down: Weight, shared_gate: Weight, shared_up: Weight, shared_down: Weight, } #[derive(Clone, Copy)] struct CompressorWeights { ape: Weight, kv: Weight, gate: Weight, norm: Weight, } #[derive(Clone, Copy)] struct IndexerWeights { q: Weight, proj: Weight, compressor: CompressorWeights, } impl Layer { fn bind(model: &Gguf, shape: super::Shape, index: u32) -> Result { let layer = Self::bind_prefix(model, shape, index, &format!("blk.{index}"))?; for (name, weight) in [ ("attention q", layer.attn_q_b), ("attention output A", layer.attn_output_a), ("attention output B", layer.attn_output_b), ("shared gate", layer.shared_gate), ("shared up", layer.shared_up), ("shared down", layer.shared_down), ] { if weight.kind != Q8_0 { return Err(format!("DeepSeek Metal path requires Q8_0 {name} weights")); } } if let Some(indexer) = layer.indexer && (!matches!(indexer.q.kind, F16 | Q8_0) || indexer.proj.kind != F16) { return Err("DeepSeek Metal path requires F16/Q8 indexer weights".into()); } Ok(layer) } fn bind_prefix( model: &Gguf, shape: super::Shape, index: u32, prefix: &str, ) -> Result { let required = |suffix: &str| Weight::bind(model, &format!("{prefix}.{suffix}")); let optional = |suffix: &str| Weight::optional(model, &format!("{prefix}.{suffix}")); let attn_compressor = (compression_ratio(shape, index) != 0) .then(|| { Ok::<_, String>(CompressorWeights { ape: required("attn_compressor_ape.weight")?, kv: required("attn_compressor_kv.weight")?, gate: required("attn_compressor_gate.weight")?, norm: required("attn_compressor_norm.weight")?, }) }) .transpose()?; let indexer = (compression_ratio(shape, index) == 4) .then(|| { Ok::<_, String>(IndexerWeights { q: required("indexer.attn_q_b.weight")?, proj: required("indexer.proj.weight")?, compressor: CompressorWeights { ape: required("indexer_compressor_ape.weight")?, kv: required("indexer_compressor_kv.weight")?, gate: required("indexer_compressor_gate.weight")?, norm: required("indexer_compressor_norm.weight")?, }, }) }) .transpose()?; Ok(Self { hc_attn_fn: required("hc_attn_fn.weight")?, hc_attn_scale: required("hc_attn_scale.weight")?, hc_attn_base: required("hc_attn_base.weight")?, attn_norm: required("attn_norm.weight")?, attn_q_a: required("attn_q_a.weight")?, attn_q_a_norm: required("attn_q_a_norm.weight")?, attn_q_b: required("attn_q_b.weight")?, attn_kv: required("attn_kv.weight")?, attn_kv_norm: required("attn_kv_a_norm.weight")?, attn_sinks: required("attn_sinks.weight")?, attn_output_a: required("attn_output_a.weight")?, attn_output_b: required("attn_output_b.weight")?, attn_compressor, indexer, hc_ffn_fn: required("hc_ffn_fn.weight")?, hc_ffn_scale: required("hc_ffn_scale.weight")?, hc_ffn_base: required("hc_ffn_base.weight")?, ffn_norm: required("ffn_norm.weight")?, router: required("ffn_gate_inp.weight")?, router_bias: optional("exp_probs_b.bias"), router_hash: optional("ffn_gate_tid2eid.weight"), expert_gate: required("ffn_gate_exps.weight")?, expert_up: required("ffn_up_exps.weight")?, expert_down: required("ffn_down_exps.weight")?, shared_gate: required("ffn_gate_shexp.weight")?, shared_up: required("ffn_up_shexp.weight")?, shared_down: required("ffn_down_shexp.weight")?, }) } } struct Weights { token_embedding: Weight, output_hc_base: Weight, output_hc_fn: Weight, output_hc_scale: Weight, output_norm: Weight, output: Weight, layers: Vec, } struct DsparkStageWeights { block: Layer, main_proj: Option, main_norm: Option, norm: Option, hc_head_base: Option, hc_head_fn: Option, hc_head_scale: Option, markov_w1: Option, markov_w2: Option, confidence: Option, } impl DsparkStageWeights { fn bind(model: &Gguf, shape: super::Shape, stage: u32) -> Result { let prefix = format!("mtp.{stage}"); Ok(Self { block: Layer::bind_prefix(model, shape, 1, &prefix)?, main_proj: Weight::optional(model, &format!("{prefix}.main_proj.weight")), main_norm: Weight::optional(model, &format!("{prefix}.main_norm.weight")), norm: Weight::optional(model, &format!("{prefix}.norm.weight")), hc_head_base: Weight::optional(model, &format!("{prefix}.hc_head_base.weight")), hc_head_fn: Weight::optional(model, &format!("{prefix}.hc_head_fn.weight")), hc_head_scale: Weight::optional(model, &format!("{prefix}.hc_head_scale.weight")), markov_w1: Weight::optional(model, &format!("{prefix}.markov_head.markov_w1.weight")), markov_w2: Weight::optional(model, &format!("{prefix}.markov_head.markov_w2.weight")), confidence: Weight::optional(model, &format!("{prefix}.confidence_head.proj.weight")), }) } } struct Dspark { config: DsparkConfig, weights: Vec, mean_weights: Buffer, mean_rows: Buffer, target_hidden: Buffer, target_hidden_batch: Buffer, packed_target_hidden: Buffer, stage0_proj: Buffer, main_x: Buffer, stage_input_hc: Buffer, stage_output_hc: Buffer, draft_tokens: Buffer, raw_caches: Vec, scratch: BatchScratch, logits: Buffer, capture_mask: u32, cache_start: u32, cache_len: u32, confidence_threshold: f32, strict: bool, drafted: u64, accepted: u64, scheduler_cycles: u32, scheduler_accepted: u32, scheduler_no_draft: u32, scheduler_skip: u32, scheduler_lifetime_accepted: u32, scheduler_long_accept_seen: bool, last_confidence: Option, } fn dspark_scheduler_pause(cycles: u32, accepted: u32, no_draft: u32) -> u32 { if cycles == 0 { return 0; } let low_acceptance = u64::from(accepted) * 1_000 < u64::from(cycles) * 1_500; let many_no_draft = no_draft * 2 >= cycles; if many_no_draft { 4 } else if low_acceptance { 2 } else { 0 } } impl Dspark { fn new( model: &Model, support: &Gguf, session: &Session, settings: EngineSpeculativeSettings, quality: bool, ) -> Result { let config = dspark_config(support)?; let shape = model.shape; let rows = config.block_size + 1; if rows > session.prefill_cap { return Err(format!( "DSpark block needs {rows} rows, but prefill workspace has {}", session.prefill_cap )); } let weights = (0..config.stages) .map(|stage| DsparkStageWeights::bind(support, shape, stage)) .collect::, _>>()?; let hc_dim = shape.hc * shape.embd; let mean_weights = Buffer::floats(shape.hc)?; let mean_rows = Buffer::floats(u64::from(session.prefill_cap) * shape.hc)?; let mean = vec![1.0 / shape.hc as f32; shape.hc as usize]; let repeated = vec![1.0 / shape.hc as f32; (session.prefill_cap as u64 * shape.hc) as usize]; mean_weights.write_f32(&mean)?; mean_rows.write_f32(&repeated)?; Ok(Self { target_hidden: Buffer::floats(config.target_layers.len() as u64 * shape.embd)?, target_hidden_batch: Buffer::floats( config.target_layers.len() as u64 * u64::from(session.prefill_cap) * shape.embd, )?, packed_target_hidden: Buffer::floats( u64::from(session.prefill_cap) * config.target_layers.len() as u64 * shape.embd, )?, stage0_proj: Buffer::floats(shape.embd)?, main_x: Buffer::floats(shape.embd)?, stage_input_hc: Buffer::floats(u64::from(rows) * hc_dim)?, stage_output_hc: Buffer::floats(u64::from(config.block_size) * hc_dim)?, draft_tokens: Buffer::bytes(u64::from(config.block_size) * 4)?, raw_caches: (0..config.stages) .map(|_| Buffer::floats(u64::from(session.raw_cap) * shape.head_dim)) .collect::>()?, scratch: BatchScratch::allocate(model, session.context, rows, false)?, logits: Buffer::floats(u64::from(config.block_size) * shape.vocab)?, config, weights, mean_weights, mean_rows, capture_mask: 0, cache_start: 0, cache_len: 0, confidence_threshold: if settings.dspark_exact_sampling && !settings.dspark_confidence_threshold_set { settings.dspark_confidence_threshold.max(0.8) } else { settings.dspark_confidence_threshold }, strict: settings.dspark_strict || quality, drafted: 0, accepted: 0, scheduler_cycles: 0, scheduler_accepted: 0, scheduler_no_draft: 0, scheduler_skip: 0, scheduler_lifetime_accepted: 0, scheduler_long_accept_seen: false, last_confidence: None, }) } fn scheduler_should_skip(&mut self) -> bool { if self.scheduler_skip == 0 { return false; } self.scheduler_skip -= 1; true } fn scheduler_note(&mut self, accepted: u32, no_draft: bool) { self.scheduler_cycles += 1; self.scheduler_accepted = self.scheduler_accepted.saturating_add(accepted); self.scheduler_lifetime_accepted = self.scheduler_lifetime_accepted.saturating_add(accepted); self.scheduler_long_accept_seen |= accepted > 2; self.scheduler_no_draft += u32::from(no_draft); if no_draft { let skip = if self.scheduler_lifetime_accepted == 0 && self .last_confidence .is_some_and(|confidence| confidence <= 0.5) { 7 } else if self.scheduler_lifetime_accepted != 0 && !self.scheduler_long_accept_seen { 4 } else { 3 }; self.scheduler_skip = self.scheduler_skip.max(skip); } if self.scheduler_cycles >= 4 { self.scheduler_skip = self.scheduler_skip.max(dspark_scheduler_pause( self.scheduler_cycles, self.scheduler_accepted, self.scheduler_no_draft, )); self.scheduler_cycles = 0; self.scheduler_accepted = 0; self.scheduler_no_draft = 0; } } fn target_slot(&self, layer: u32) -> Option { self.config .target_layers .iter() .position(|target| *target == layer) .map(|slot| slot as u32) } fn begin_capture(&mut self) { self.capture_mask = 0; } fn commit_proposed_prefix(&mut self, rows: u32, raw_cap: u32) { let added = rows.min(raw_cap); let total = self.cache_len.saturating_add(added); if total > raw_cap { self.cache_start = (self.cache_start + total - raw_cap) % raw_cap; } self.cache_len = total.min(raw_cap); } fn capture_decode( &mut self, layer: u32, hc: &Buffer, shape: super::Shape, ) -> Result<(), String> { let Some(slot) = self.target_slot(layer) else { return Ok(()); }; let target = self .target_hidden .view(u64::from(slot) * shape.embd * 4, shape.embd * 4)?; call( unsafe { ds4_gpu_hc_weighted_sum_tensor( target.raw(), hc.raw(), self.mean_weights.raw(), shape.embd as u32, shape.hc as u32, ) }, "capturing DSpark target hidden state", )?; self.capture_mask |= 1 << slot; Ok(()) } fn capture_batch( &mut self, layer: u32, hc: &Buffer, rows: u32, prefill_cap: u32, shape: super::Shape, ) -> Result<(), String> { let Some(slot) = self.target_slot(layer) else { return Ok(()); }; let batch = self.target_hidden_batch.view( u64::from(slot) * u64::from(prefill_cap) * shape.embd * 4, u64::from(rows) * shape.embd * 4, )?; call( unsafe { ds4_gpu_hc_weighted_sum_tensor( batch.raw(), hc.raw(), self.mean_rows.raw(), shape.embd as u32, shape.hc as u32, ) }, "capturing batched DSpark target hidden states", )?; self.target_hidden.copy_from( u64::from(slot) * shape.embd * 4, &batch, u64::from(rows - 1) * shape.embd * 4, shape.embd * 4, "capturing final DSpark target hidden state", )?; self.capture_mask |= 1 << slot; Ok(()) } fn capture_complete(&self) -> bool { self.capture_mask == (1_u32 << self.config.target_layers.len()) - 1 } fn seed_batch_cache( &mut self, support: &Gguf, pos: u32, rows: u32, prefill_cap: u32, raw_cap: u32, shape: super::Shape, ) -> Result<(), String> { if !self.capture_complete() { return Err("DSpark target-layer capture is incomplete".into()); } let stage0 = self.weights.first().ok_or("DSpark has no stages")?; let main_proj = stage0 .main_proj .ok_or("DSpark stage 0 projection is missing")?; let main_norm = stage0.main_norm.ok_or("DSpark stage 0 norm is missing")?; let map = support.map_ptr().cast(); let size = support.len(); let input = self.config.target_layers.len() as u64 * shape.embd; let projected = self .target_hidden_batch .view(0, u64::from(rows) * shape.embd * 4)?; let norm = self .packed_target_hidden .view(0, u64::from(rows) * shape.embd * 4)?; let kv_bytes = u64::from(rows) * shape.head_dim * 4; let kv_raw = self.target_hidden_batch.view(0, kv_bytes)?; let kv = self.target_hidden_batch.view(kv_bytes, kv_bytes)?; let commands = Commands::begin()?; call( unsafe { ds4_gpu_pack_slot_rows_f32_tensor( self.packed_target_hidden.raw(), self.target_hidden_batch.raw(), rows, shape.embd as u32, self.config.target_layers.len() as u32, prefill_cap, ) }, "packing DSpark target hidden states", )?; matmul_rows( &projected, main_proj, input, shape.embd, &self.packed_target_hidden, rows, map, size, )?; call( unsafe { ds4_gpu_rms_norm_weight_rows_tensor( norm.raw(), projected.raw(), map, size, main_norm.offset, shape.embd as u32, rows, shape.rms_epsilon, ) }, "normalizing DSpark target hidden states", )?; for (stage, cache) in self.weights.iter().zip(&self.raw_caches) { matmul_rows( &kv_raw, stage.block.attn_kv, shape.embd, shape.head_dim, &norm, rows, map, size, )?; call( unsafe { ds4_gpu_rms_norm_weight_rows_tensor( kv.raw(), kv_raw.raw(), map, size, stage.block.attn_kv_norm.offset, shape.head_dim as u32, rows, shape.rms_epsilon, ) }, "normalizing DSpark target KV", )?; call( unsafe { ds4_gpu_rope_tail_tensor( kv.raw(), rows, 1, shape.head_dim as u32, shape.rot as u32, pos, 0, false, shape.rope_base, 1.0, 0.0, 1.0, shape.rope_beta_fast, shape.rope_beta_slow, ) }, "applying DSpark target KV RoPE", )?; call( unsafe { ds4_gpu_dsv4_fp8_kv_quantize_tensor( kv.raw(), rows, shape.head_dim as u32, shape.rot as u32, ) }, "quantizing DSpark target KV", )?; call( unsafe { ds4_gpu_store_raw_kv_batch_tensor( cache.raw(), kv.raw(), raw_cap, pos, rows, shape.head_dim as u32, ) }, "seeding DSpark target KV cache", )?; } commands.finish()?; self.cache_start = pos % raw_cap; self.cache_len = rows; Ok(()) } fn seed_current_cache( &mut self, support: &Gguf, pos: u32, raw_cap: u32, shape: super::Shape, ) -> Result<(), String> { if !self.capture_complete() { return Err("DSpark target-layer capture is incomplete".into()); } let stage0 = self.weights.first().ok_or("DSpark has no stages")?; let main_proj = stage0 .main_proj .ok_or("DSpark stage 0 projection is missing")?; let main_norm = stage0.main_norm.ok_or("DSpark stage 0 norm is missing")?; let map = support.map_ptr().cast(); let size = support.len(); let input = self.config.target_layers.len() as u64 * shape.embd; let commands = Commands::begin()?; matmul( &self.stage0_proj, main_proj, input, shape.embd, &self.target_hidden, map, size, )?; call( unsafe { ds4_gpu_rms_norm_weight_tensor( self.main_x.raw(), self.stage0_proj.raw(), map, size, main_norm.offset, shape.embd as u32, shape.rms_epsilon, ) }, "normalizing current DSpark target state", )?; for (stage, cache) in self.weights.iter().zip(&self.raw_caches) { matmul( &self.scratch.kv_raw, stage.block.attn_kv, shape.embd, shape.head_dim, &self.main_x, map, size, )?; call( unsafe { ds4_gpu_rms_norm_weight_rows_tensor( self.scratch.kv.raw(), self.scratch.kv_raw.raw(), map, size, stage.block.attn_kv_norm.offset, shape.head_dim as u32, 1, shape.rms_epsilon, ) }, "normalizing current DSpark target KV", )?; call( unsafe { ds4_gpu_rope_tail_tensor( self.scratch.kv.raw(), 1, 1, shape.head_dim as u32, shape.rot as u32, pos, 0, false, shape.rope_base, 1.0, 0.0, 1.0, shape.rope_beta_fast, shape.rope_beta_slow, ) }, "applying current DSpark target KV RoPE", )?; call( unsafe { ds4_gpu_dsv4_fp8_kv_quantize_tensor( self.scratch.kv.raw(), 1, shape.head_dim as u32, shape.rot as u32, ) }, "quantizing current DSpark target KV", )?; call( unsafe { ds4_gpu_store_raw_kv_batch_tensor( cache.raw(), self.scratch.kv.raw(), raw_cap, pos, 1, shape.head_dim as u32, ) }, "storing current DSpark target KV", )?; } commands.finish()?; let append = (self.cache_start + self.cache_len) % raw_cap; if self.cache_len == 0 || append != pos % raw_cap { self.cache_start = pos % raw_cap; self.cache_len = 1; } else if self.cache_len < raw_cap { self.cache_len += 1; } else { self.cache_start = (self.cache_start + 1) % raw_cap; } Ok(()) } fn propose( &mut self, base: &Model, base_weights: &Weights, token: i32, pos: u32, raw_cap: u32, ) -> Result, String> { if self.strict || !self.capture_complete() { return Ok(Vec::new()); } let support = base .support .as_ref() .ok_or("DSpark support model is missing")?; let shape = base.shape; let rows = self.config.block_size + 1; if rows > raw_cap { return Err("DSpark block exceeds its raw-cache capacity".into()); } let max_support = raw_cap - rows; if self.cache_len > max_support { let discard = self.cache_len - max_support; self.cache_start = (self.cache_start + discard) % raw_cap; self.cache_len = max_support; } let map = support.map_ptr().cast(); let size = support.len(); let stage0 = self.weights.first().ok_or("DSpark has no stages")?; let main_proj = stage0 .main_proj .ok_or("DSpark stage 0 projection is missing")?; let main_norm = stage0.main_norm.ok_or("DSpark stage 0 norm is missing")?; let input = self.config.target_layers.len() as u64 * shape.embd; let mut ids = vec![self.config.noise_token as i32; self.config.block_size as usize]; ids[0] = token; self.draft_tokens.write_i32(&ids)?; self.scratch.tokens.write_i32(&ids)?; let commands = Commands::begin()?; matmul( &self.stage0_proj, main_proj, input, shape.embd, &self.target_hidden, map, size, )?; call( unsafe { ds4_gpu_rms_norm_weight_tensor( self.main_x.raw(), self.stage0_proj.raw(), map, size, main_norm.offset, shape.embd as u32, shape.rms_epsilon, ) }, "normalizing DSpark stage-0 projection", )?; call( unsafe { ds4_gpu_repeat_hc_tensor( self.stage_input_hc.raw(), self.main_x.raw(), shape.embd as u32, shape.hc as u32, ) }, "expanding DSpark target hidden state", )?; let draft_hc = self.stage_input_hc.view( shape.hc * shape.embd * 4, u64::from(self.config.block_size) * shape.hc * shape.embd * 4, )?; call( unsafe { ds4_gpu_embed_tokens_hc_tensor( draft_hc.raw(), self.draft_tokens.raw(), base.main.map_ptr().cast(), base.main.len(), base_weights.token_embedding.offset, shape.vocab as u32, self.config.block_size, shape.embd as u32, shape.hc as u32, ) }, "embedding DSpark draft block", )?; for stage in 0..self.weights.len() { self.eval_stage(support, stage, pos, raw_cap, shape)?; if stage + 1 < self.weights.len() { self.stage_input_hc.copy_from( shape.hc * shape.embd * 4, &self.scratch.next_hc, 0, u64::from(self.config.block_size) * shape.hc * shape.embd * 4, "feeding the next DSpark stage", )?; } } self.stage_output_hc.copy_from( 0, &self.scratch.next_hc, 0, u64::from(self.config.block_size) * shape.hc * shape.embd * 4, "capturing DSpark stage output", )?; self.eval_output_heads(commands, base, support, base_weights, token) } fn eval_stage( &mut self, support: &Gguf, stage_index: usize, pos: u32, raw_cap: u32, shape: super::Shape, ) -> Result<(), String> { let stage = &self.weights[stage_index]; let w = &stage.block; let map = support.map_ptr().cast(); let size = support.len(); let draft = self.config.block_size; let rows = draft + 1; let hc_dim = shape.hc * shape.embd; let mix_hc = 2 * shape.hc + shape.hc * shape.hc; let q_dim = shape.heads * shape.head_dim; let group_dim = shape.head_dim * (shape.heads / shape.out_groups); let draft_norm = self .scratch .norm .view(shape.embd * 4, u64::from(draft) * shape.embd * 4)?; let draft_input = self .stage_input_hc .view(hc_dim * 4, u64::from(draft) * hc_dim * 4)?; let draft_split = self .scratch .hc_split .view(mix_hc * 4, u64::from(draft) * mix_hc * 4)?; let after_attention_hc = self .scratch .after_attention_hc .view(0, u64::from(draft) * hc_dim * 4)?; let target_kv = self.scratch.kv.view(0, shape.head_dim * 4)?; let draft_kv = self .scratch .kv .view(shape.head_dim * 4, u64::from(draft) * shape.head_dim * 4)?; call( unsafe { ds4_gpu_rms_norm_plain_rows_tensor( self.scratch.flat_hc.raw(), self.stage_input_hc.raw(), hc_dim as u32, rows, shape.rms_epsilon, ) }, "normalizing DSpark attention HC rows", )?; matmul_rows( &self.scratch.hc_mix, w.hc_attn_fn, hc_dim, mix_hc, &self.scratch.flat_hc, rows, map, size, )?; call( unsafe { ds4_gpu_hc_split_weighted_sum_norm_tensor( self.scratch.current.raw(), self.scratch.norm.raw(), self.scratch.hc_split.raw(), self.scratch.hc_mix.raw(), self.stage_input_hc.raw(), map, size, w.hc_attn_scale.offset, w.hc_attn_base.offset, w.attn_norm.offset, shape.embd as u32, shape.hc as u32, shape.hc_sinkhorn as u32, shape.hc_epsilon, shape.rms_epsilon, ) }, "mixing DSpark attention HC rows", )?; matmul_rows( &self.scratch.q_rank, w.attn_q_a, shape.embd, shape.lora_q, &draft_norm, draft, map, size, )?; call( unsafe { ds4_gpu_rms_norm_weight_rows_tensor( self.scratch.q_rank_norm.raw(), self.scratch.q_rank.raw(), map, size, w.attn_q_a_norm.offset, shape.lora_q as u32, draft, shape.rms_epsilon, ) }, "normalizing DSpark Q rank", )?; matmul_rows( &self.scratch.q, w.attn_q_b, shape.lora_q, q_dim, &self.scratch.q_rank_norm, draft, map, size, )?; call( unsafe { ds4_gpu_head_rms_norm_tensor( self.scratch.q.raw(), draft, shape.heads as u32, shape.head_dim as u32, shape.rms_epsilon, ) }, "normalizing DSpark Q heads", )?; call( unsafe { ds4_gpu_rope_tail_tensor( self.scratch.q.raw(), draft, shape.heads as u32, shape.head_dim as u32, shape.rot as u32, pos, 0, false, shape.rope_base, 1.0, 0.0, 1.0, shape.rope_beta_fast, shape.rope_beta_slow, ) }, "applying DSpark Q RoPE", )?; matmul_rows( &self.scratch.kv_raw, w.attn_kv, shape.embd, shape.head_dim, &self.scratch.norm, rows, map, size, )?; call( unsafe { ds4_gpu_rms_norm_weight_rows_tensor( self.scratch.kv.raw(), self.scratch.kv_raw.raw(), map, size, w.attn_kv_norm.offset, shape.head_dim as u32, rows, shape.rms_epsilon, ) }, "normalizing DSpark KV rows", )?; for (view, count) in [(&target_kv, 1), (&draft_kv, draft)] { call( unsafe { ds4_gpu_rope_tail_tensor( view.raw(), count, 1, shape.head_dim as u32, shape.rot as u32, pos, 0, false, shape.rope_base, 1.0, 0.0, 1.0, shape.rope_beta_fast, shape.rope_beta_slow, ) }, "applying DSpark KV RoPE", )?; } call( unsafe { ds4_gpu_dsv4_fp8_kv_quantize_tensor( self.scratch.kv.raw(), rows, shape.head_dim as u32, shape.rot as u32, ) }, "quantizing DSpark KV rows", )?; let append = (self.cache_start + self.cache_len) % raw_cap; call( unsafe { ds4_gpu_store_raw_kv_batch_tensor( self.raw_caches[stage_index].raw(), self.scratch.kv.raw(), raw_cap, append, rows, shape.head_dim as u32, ) }, "storing DSpark stage KV rows", )?; call( unsafe { ds4_gpu_attention_noncausal_raw_batch_heads_tensor( self.scratch.heads.raw(), map, size, w.attn_sinks.offset, self.scratch.q.raw(), self.raw_caches[stage_index].raw(), draft, self.cache_len + rows, raw_cap, self.cache_start, shape.heads as u32, shape.head_dim as u32, ) }, "running DSpark noncausal attention", )?; call( unsafe { ds4_gpu_rope_tail_tensor( self.scratch.heads.raw(), draft, shape.heads as u32, shape.head_dim as u32, shape.rot as u32, pos, 0, true, shape.rope_base, 1.0, 0.0, 1.0, shape.rope_beta_fast, shape.rope_beta_slow, ) }, "applying inverse DSpark attention RoPE", )?; call( unsafe { ds4_gpu_attention_output_q8_batch_tensor( self.scratch.attention_out.raw(), self.scratch.attention_low.raw(), self.scratch.attention_group_tmp.raw(), self.scratch.attention_low_tmp.raw(), map, size, w.attn_output_a.offset, w.attn_output_b.offset, group_dim, shape.lora_o, shape.out_groups as u32, shape.embd, self.scratch.heads.raw(), draft, ) }, "projecting DSpark attention output", )?; call( unsafe { ds4_gpu_hc_expand_split_tensor( after_attention_hc.raw(), self.scratch.attention_out.raw(), draft_input.raw(), draft_split.raw(), shape.embd as u32, shape.hc as u32, ) }, "expanding DSpark attention HC", )?; self.eval_stage_ffn(stage_index, shape, support)?; Ok(()) } fn eval_stage_ffn( &mut self, stage_index: usize, shape: super::Shape, support: &Gguf, ) -> Result<(), String> { let w = &self.weights[stage_index].block; let map = support.map_ptr().cast(); let size = support.len(); let rows = self.config.block_size; let hc_dim = shape.hc * shape.embd; let mix_hc = 2 * shape.hc + shape.hc * shape.hc; let hc_bytes = u64::from(rows) * hc_dim * 4; let mix_bytes = u64::from(rows) * mix_hc * 4; let embd_bytes = u64::from(rows) * shape.embd * 4; let after_attention_hc = self.scratch.after_attention_hc.view(0, hc_bytes)?; let flat_hc = self.scratch.flat_hc.view(0, hc_bytes)?; let hc_mix = self.scratch.hc_mix.view(0, mix_bytes)?; let hc_split = self.scratch.hc_split.view(0, mix_bytes)?; let current = self.scratch.current.view(0, embd_bytes)?; let norm = self.scratch.norm.view(0, embd_bytes)?; let next_hc = self.scratch.next_hc.view(0, hc_bytes)?; call( unsafe { ds4_gpu_rms_norm_plain_rows_tensor( flat_hc.raw(), after_attention_hc.raw(), hc_dim as u32, rows, shape.rms_epsilon, ) }, "normalizing DSpark FFN HC rows", )?; matmul_rows( &hc_mix, w.hc_ffn_fn, hc_dim, mix_hc, &flat_hc, rows, map, size, )?; call( unsafe { ds4_gpu_hc_split_weighted_sum_norm_tensor( current.raw(), norm.raw(), hc_split.raw(), hc_mix.raw(), after_attention_hc.raw(), map, size, w.hc_ffn_scale.offset, w.hc_ffn_base.offset, w.ffn_norm.offset, shape.embd as u32, shape.hc as u32, shape.hc_sinkhorn as u32, shape.hc_epsilon, shape.rms_epsilon, ) }, "mixing DSpark FFN HC rows", )?; matmul_rows( &self.scratch.router_logits, w.router, shape.embd, shape.experts, &norm, rows, map, size, )?; call( unsafe { ds4_gpu_router_select_batch_tensor( self.scratch.router_selected.raw(), self.scratch.router_weights.raw(), self.scratch.router_probs.raw(), map, size, w.router_bias.map_or(0, |weight| weight.offset), 0, 0, 0, 0, w.router_bias.is_some(), false, self.scratch.router_logits.raw(), self.scratch.tokens.raw(), shape.experts as u32, shape.experts_used as u32, shape.expert_weight_scale, rows, ) }, "routing DSpark experts", )?; for (out, weight) in [ (&self.scratch.shared_gate, w.shared_gate), (&self.scratch.shared_up, w.shared_up), ] { matmul_rows( out, weight, shape.embd, shape.ff_expert, &norm, rows, map, size, )?; } call( unsafe { ds4_gpu_swiglu_tensor( self.scratch.shared_mid.raw(), self.scratch.shared_gate.raw(), self.scratch.shared_up.raw(), rows * shape.ff_expert as u32, shape.swiglu_clamp, 1.0, ) }, "activating DSpark shared expert", )?; matmul_rows( &self.scratch.shared_out, w.shared_down, shape.ff_expert, shape.embd, &self.scratch.shared_mid, rows, map, size, )?; let gate_row = w.expert_gate.bytes / (w.expert_gate.dims[1] * w.expert_gate.dims[2]); let down_row = w.expert_down.bytes / (w.expert_down.dims[1] * w.expert_down.dims[2]); let mut mid_f16 = false; call( unsafe { ds4_gpu_routed_moe_batch_tensor( self.scratch.routed_out.raw(), self.scratch.routed_gate.raw(), self.scratch.routed_up.raw(), self.scratch.routed_mid.raw(), self.scratch.routed_experts.raw(), map, size, w.expert_gate.offset, w.expert_up.offset, w.expert_down.offset, w.expert_gate.kind, w.expert_down.kind, w.expert_gate.dims[1] * gate_row, gate_row, w.expert_down.dims[1] * down_row, down_row, w.expert_gate.dims[0] as u32, w.expert_down.dims[0] as u32, w.expert_down.dims[1] as u32, self.scratch.router_selected.raw(), self.scratch.router_weights.raw(), shape.experts as u32, shape.experts_used as u32, shape.swiglu_clamp, norm.raw(), stage_index as u32, rows, &mut mid_f16, false, ) }, "running DSpark routed experts", )?; call( unsafe { ds4_gpu_hc_expand_add_split_tensor( next_hc.raw(), self.scratch.routed_out.raw(), self.scratch.shared_out.raw(), after_attention_hc.raw(), hc_split.raw(), shape.embd as u32, shape.hc as u32, ) }, "expanding DSpark FFN HC", ) } fn eval_output_heads( &mut self, commands: Commands, base: &Model, support: &Gguf, base_weights: &Weights, first_token: i32, ) -> Result, String> { let shape = base.shape; let draft = self.config.block_size; let hc_dim = shape.hc * shape.embd; let final_stage = self.weights.last().ok_or("DSpark has no final stage")?; let norm = final_stage.norm.ok_or("DSpark final norm is missing")?; let head_base = final_stage .hc_head_base .ok_or("DSpark final HC base is missing")?; let head_fn = final_stage .hc_head_fn .ok_or("DSpark final HC projection is missing")?; let head_scale = final_stage .hc_head_scale .ok_or("DSpark final HC scale is missing")?; let markov_w1 = final_stage.markov_w1.ok_or("DSpark Markov W1 is missing")?; let markov_w2 = final_stage.markov_w2.ok_or("DSpark Markov W2 is missing")?; let map = support.map_ptr().cast(); let size = support.len(); let output_pre = self .scratch .hc_mix .view(0, u64::from(draft) * shape.hc * 4)?; let output_weights = self .scratch .hc_split .view(0, u64::from(draft) * shape.hc * 4)?; let output_embedding = self .scratch .current .view(0, u64::from(draft) * shape.embd * 4)?; let output_norm = self .scratch .norm .view(0, u64::from(draft) * shape.embd * 4)?; call( unsafe { ds4_gpu_rms_norm_plain_rows_tensor( self.scratch.flat_hc.raw(), self.stage_output_hc.raw(), hc_dim as u32, draft, shape.rms_epsilon, ) }, "normalizing DSpark output HC rows", )?; matmul_rows( &output_pre, head_fn, hc_dim, shape.hc, &self.scratch.flat_hc, draft, map, size, )?; call( unsafe { ds4_gpu_output_hc_weights_tensor( output_weights.raw(), output_pre.raw(), map, size, head_scale.offset, head_base.offset, shape.hc as u32, shape.hc_epsilon, ) }, "computing DSpark output HC weights", )?; call( unsafe { ds4_gpu_hc_weighted_sum_tensor( output_embedding.raw(), self.stage_output_hc.raw(), output_weights.raw(), shape.embd as u32, shape.hc as u32, ) }, "collapsing DSpark output HC rows", )?; call( unsafe { ds4_gpu_rms_norm_weight_rows_tensor( output_norm.raw(), output_embedding.raw(), map, size, norm.offset, shape.embd as u32, draft, shape.rms_epsilon, ) }, "normalizing DSpark output rows", )?; matmul_rows( &self.logits, base_weights.output, shape.embd, shape.vocab, &output_norm, draft, base.main.map_ptr().cast(), base.main.len(), )?; commands.finish()?; let confidence = final_stage .confidence .ok_or("DSpark confidence head is missing")?; let mut logits = vec![0.0; (u64::from(draft) * shape.vocab) as usize]; let mut hidden = vec![0.0; (u64::from(draft) * shape.embd) as usize]; self.logits.read_f32(&mut logits)?; output_norm.read_f32(&mut hidden)?; let mut proposals = Vec::with_capacity(draft as usize); let mut previous = first_token as u32; self.last_confidence = None; for row in 0..draft as usize { let state = dense_row(support, markov_w1, previous)?; let mut features = Vec::with_capacity(shape.embd as usize + state.len()); features.extend_from_slice( &hidden[row * shape.embd as usize..(row + 1) * shape.embd as usize], ); features.extend_from_slice(&state); let confidence_logit = dense_dot(support, confidence, 0, &features)?; let confidence_value = if confidence_logit >= 0.0 { 1.0 / (1.0 + (-confidence_logit).exp()) } else { let value = confidence_logit.exp(); value / (1.0 + value) }; if row == 0 { self.last_confidence = Some(confidence_logit); } if self.confidence_threshold > 0.0 && confidence_value < self.confidence_threshold { break; } let row_logits = &logits[row * shape.vocab as usize..(row + 1) * shape.vocab as usize]; let best = dense_argmax(support, markov_w2, &state, row_logits)?; proposals.push(best); previous = best as u32; } self.drafted += proposals.len() as u64; Ok(proposals) } } struct Steering { directions: Buffer, attention_scale: f32, ffn_scale: f32, width: u32, identity: [u8; 32], } impl Steering { fn load(model: &Model, settings: EngineSteeringSettings) -> Result, String> { if settings.attention_scale == 0.0 && settings.ffn_scale == 0.0 { return Ok(None); } let path = settings .file .as_deref() .ok_or("directional steering needs a direction-vector file")?; let bytes = fs::read(path) .map_err(|error| format!("Cannot read directional steering file {path}: {error}"))?; let expected = model .shape .layers .checked_mul(model.shape.embd as u32) .and_then(|values| values.checked_mul(4)) .ok_or("directional steering size overflow")? as usize; if bytes.len() != expected { return Err(format!( "directional steering file has {} bytes, expected {expected}", bytes.len() )); } let directions = Buffer::bytes(expected as u64)?; directions.write(0, &bytes)?; let mut hash = Sha256::new(); hash.update(b"DS4Server directional steering v1"); hash.update(settings.attention_scale.to_bits().to_le_bytes()); hash.update(settings.ffn_scale.to_bits().to_le_bytes()); hash.update(&bytes); Ok(Some(Self { directions, attention_scale: settings.attention_scale, ffn_scale: settings.ffn_scale, width: model.shape.embd as u32, identity: hash.finalize().into(), })) } fn apply(&self, output: &Buffer, layer: u32, rows: u32, attention: bool) -> Result<(), String> { let scale = if attention { self.attention_scale } else { self.ffn_scale }; if scale == 0.0 { return Ok(()); } call( unsafe { ds4_gpu_directional_steering_project_tensor( output.raw(), self.directions.raw(), layer, self.width, rows, scale, ) }, "applying directional steering", ) } } struct SsdPlan { model_spans: Vec<(u64, u64)>, resident_bytes: u64, cache_experts: u32, per_expert_bytes: u64, admission_bytes: u64, preload_experts: u32, preload_by_layer: Vec>, preload_seeded: Vec, static_decode_map_current: std::sync::atomic::AtomicBool, cold: bool, loader: Option, selected_requests: std::cell::Cell, selected_experts: std::cell::Cell, selected_wait_ns: std::cell::Cell, } struct DeepSeekModelSpans { ranges: Vec<(u64, u64)>, max_tensor_bytes: u64, } struct PrefillPread { layer: u32, workers: Vec>>, } impl PrefillPread { fn start(model: &Model, layer: u32, spans: &DeepSeekModelSpans) -> Result { const THREADS: usize = 8; const CHUNK: usize = 1024 * 1024; let file = Arc::new(File::open(model.main.path()).map_err(|error| { format!( "Cannot open {} for SSD prefill: {error}", model.main.path().display() ) })?); let mut ranges = vec![Vec::new(); THREADS]; for &(offset, bytes) in &spans.ranges { let part = bytes.div_ceil(THREADS as u64); let mut consumed = 0; for worker_ranges in &mut ranges { if consumed == bytes { break; } let size = (bytes - consumed).min(part); worker_ranges.push((offset + consumed, size)); consumed += size; } } let workers = ranges .into_iter() .map(|ranges| { let file = Arc::clone(&file); std::thread::spawn(move || { let mut buffer = vec![0_u8; CHUNK]; for (offset, bytes) in ranges { let mut read = 0; while read < bytes { let wanted = (bytes - read).min(CHUNK as u64) as usize; let count = loop { match file.read_at(&mut buffer[..wanted], offset + read) { Err(error) if error.kind() == std::io::ErrorKind::Interrupted => {} result => break result.map_err(|error| error.to_string())?, } }; if count == 0 { return Err("unexpected EOF during SSD prefill".into()); } read += count as u64; } } Ok(()) }) }) .collect(); Ok(Self { layer, workers }) } fn finish(mut self) -> Result<(), String> { for worker in self.workers.drain(..) { worker .join() .map_err(|_| format!("SSD prefill worker panicked for layer {}", self.layer))??; } Ok(()) } } impl Drop for PrefillPread { fn drop(&mut self) { for worker in self.workers.drain(..) { let _ = worker.join(); } } } fn prefill_pread_enabled() -> bool { !environment_present(c"DS4_METAL_DISABLE_STREAMING_PREFILL_LAYER_PREAD") && !environment_present(c"DS4_METAL_DISABLE_STREAMING_PREFILL_LAYER_PREPARE") } fn prefill_selected_addr(rows: u32, shape: super::Shape, weights: &Layer) -> bool { unsafe { ds4_gpu_stream_prefill_batch_selected_addr_enabled( rows, shape.experts as u32, shape.experts_used as u32, weights.expert_gate.kind, weights.expert_down.kind, ) != 0 } } fn start_prefill_pread( model: &Model, ssd: &SsdPlan, weights: &Layer, layer: u32, rows: u32, ) -> Result { let spans = deepseek_layer_model_spans( model, weights, layer, prefill_selected_addr(rows, model.shape, weights), ssd.per_expert_bytes, )?; PrefillPread::start(model, layer, &spans) } #[derive(Clone, Copy)] struct SelectedLoadJob { selected: usize, count: u32, event: u64, model_map: usize, model_size: u64, layer: u32, total_experts: u32, gate_offset: u64, up_offset: u64, down_offset: u64, gate_expert_bytes: u64, down_expert_bytes: u64, } struct SelectedLoadResult { ids: [i32; 8], loaded: bool, job: SelectedLoadJob, } struct SelectedLoadWorker { jobs: std::sync::mpsc::SyncSender>, results: std::sync::mpsc::Receiver>, thread: Option>, } impl SelectedLoadWorker { fn new() -> Self { let (jobs, incoming) = std::sync::mpsc::sync_channel::>(1); let (outgoing, results) = std::sync::mpsc::sync_channel(1); let thread = std::thread::spawn(move || { unsafe { ds4_gpu_stream_expert_cache_note_service_thread() }; while let Ok(Some(job)) = incoming.recv() { let result = load_selected_experts(job); if outgoing.send(result).is_err() { break; } } }); Self { jobs, results, thread: Some(thread), } } fn begin(&self, selected: &Buffer, table: StreamExpertTable, count: u32) -> Result<(), String> { if count == 0 || count > 8 { return Err("SSD expert loading supports one to eight selected experts".into()); } let mut event = 0; call( unsafe { ds4_gpu_signal_selected_readback_ready(&mut event) }, "signalling SSD expert selection", )?; self.jobs .send(Some(SelectedLoadJob { selected: selected.raw() as usize, count, event, model_map: table.model_map as usize, model_size: table.model_size, layer: table.layer, total_experts: table.total_experts, gate_offset: table.gate_offset, up_offset: table.up_offset, down_offset: table.down_offset, gate_expert_bytes: table.gate_expert_bytes, down_expert_bytes: table.down_expert_bytes, })) .map_err(|_| "SSD expert loader stopped unexpectedly".to_string()) } fn finish(&self, set_override: bool) -> Result<(), String> { call( unsafe { ds4_gpu_flush_commands() }, "overlapping SSD expert loading", )?; let result = self .results .recv() .map_err(|_| "SSD expert loader stopped unexpectedly".to_string())??; if !result.loaded { let table = result.job.table(); let count = result.job.count as usize; call( unsafe { ds4_gpu_stream_expert_cache_begin_selected_load( &table, result.ids.as_ptr(), count as u32, ) }, "retrying selected SSD expert loading", )?; } if set_override { let count = result.job.count as usize; call( unsafe { ds4_gpu_routed_moe_set_selected_override(result.ids.as_ptr(), count as u32) }, "selecting streamed experts", )?; } Ok(()) } } impl Drop for SelectedLoadWorker { fn drop(&mut self) { let _ = self.jobs.send(None); if let Some(thread) = self.thread.take() { let _ = thread.join(); } } } fn load_selected_experts(job: SelectedLoadJob) -> Result { call( unsafe { ds4_gpu_wait_selected_readback_ready( job.event, c"selected-id async expert load".as_ptr(), ) }, "waiting for SSD expert selection", )?; let mut ids = [0_i32; 8]; let count = job.count as usize; call( unsafe { ds4_gpu_tensor_read( job.selected as *const GpuTensor, 0, ids.as_mut_ptr().cast(), u64::from(job.count) * 4, ) }, "reading selected SSD experts", )?; if ids[..count] .iter() .any(|expert| *expert < 0 || *expert as u32 >= job.total_experts) { return Err(format!( "SSD router selected an expert outside 0..{} at layer {}", job.total_experts, job.layer )); } let table = job.table(); let loaded = unsafe { ds4_gpu_stream_expert_cache_begin_selected_load(&table, ids.as_ptr(), job.count) } != 0; Ok(SelectedLoadResult { ids, loaded, job }) } impl SelectedLoadJob { fn table(self) -> StreamExpertTable { StreamExpertTable { model_map: self.model_map as *const c_void, model_size: self.model_size, layer: self.layer, total_experts: self.total_experts, gate_offset: self.gate_offset, up_offset: self.up_offset, down_offset: self.down_offset, gate_expert_bytes: self.gate_expert_bytes, down_expert_bytes: self.down_expert_bytes, } } } impl Weights { fn bind(model: &Model) -> Result { if model.shape.family != ModelFamily::DeepSeek { return Err("the DeepSeek Metal executor received a different model family".into()); } let main = &model.main; let output = Weight::bind(main, "output.weight")?; if output.kind != Q8_0 { return Err("DeepSeek Metal path requires a Q8_0 output weight".into()); } let layers = (0..model.shape.layers) .map(|index| Layer::bind(main, model.shape, index)) .collect::>()?; Ok(Self { token_embedding: Weight::bind(main, "token_embd.weight")?, output_hc_base: Weight::bind(main, "output_hc_base.weight")?, output_hc_fn: Weight::bind(main, "output_hc_fn.weight")?, output_hc_scale: Weight::bind(main, "output_hc_scale.weight")?, output_norm: Weight::bind(main, "output_norm.weight")?, output, layers, }) } } fn mxfp4_decode_fast_lookup_allowed( model: &Model, weights: &Weights, quality: bool, ssd_streaming: bool, ) -> bool { !quality && !ssd_streaming && model.support_kind.is_none() && weights.layers.get(4).is_some_and(|layer| { layer.expert_gate.kind == MXFP4 && layer.expert_up.kind == MXFP4 && layer.expert_down.kind == MXFP4 }) && unsafe { ds4_gpu_device_is_pre_m5_apple_silicon() } != 0 && !environment_present(c"DS4_METAL_DISABLE_PRE_M5_DECODE_PIPELINE_FAST_LOOKUP") } impl SsdPlan { fn new( model: &Model, weights: &Weights, settings: EngineSsdSettings, context: u32, prefill_chunk: u32, ) -> Result { let first = weights .layers .first() .ok_or("model has no routed expert layers")?; let gate = first.expert_gate.bytes / model.shape.experts; let down = first.expert_down.bytes / model.shape.experts; let per_expert_bytes = gate .checked_mul(2) .and_then(|bytes| bytes.checked_add(down)) .ok_or("routed expert size overflow")?; if gate == 0 || down == 0 { return Err("routed expert tensors have an invalid layout".into()); } let resident_spans = streaming_model_spans(model, weights, per_expert_bytes)?; let resident_bytes = resident_spans.iter().try_fold(0_u64, |total, (_, bytes)| { total .checked_add(*bytes) .ok_or("resident weight size overflow") })?; let model_spans = resident_spans; let runtime_bytes = estimated_deepseek_runtime_bytes(model.shape, context, prefill_chunk); let recommended = unsafe { ds4_gpu_recommended_working_set_size() }; if recommended == 0 && settings.cache_experts == 0 && settings.cache_bytes == 0 { return Err( "Metal did not report a working-set size; set an explicit SSD cache budget".into(), ); } let prefill_headroom = weights .layers .iter() .map(|layer| { layer .expert_gate .bytes .saturating_add(layer.expert_up.bytes) .saturating_add(layer.expert_down.bytes) }) .max() .unwrap_or(0) .saturating_mul(2); let max_experts = model .shape .layers .saturating_mul(model.shape.experts as u32); let (cache_experts, reserved_headroom) = if settings.cache_experts != 0 { (settings.cache_experts.min(max_experts), 0) } else { let total = if settings.cache_bytes != 0 { const GIB: u64 = 1024 * 1024 * 1024; if recommended == 0 { settings.cache_bytes } else { let safe = recommended .saturating_mul(7) .checked_div(8) .unwrap_or(recommended) .saturating_sub(runtime_bytes) / GIB * GIB; settings.cache_bytes.min(safe.max(GIB)) } } else { let percent = env::var("DS4_SSD_AUTO_CACHE_PCT") .ok() .and_then(|value| value.parse::().ok()) .filter(|value| (50..=95).contains(value)) .unwrap_or(80); recommended .saturating_mul(percent) .checked_div(100) .unwrap_or(recommended) .saturating_sub(resident_bytes) }; if total <= prefill_headroom { return Err(format!( "SSD cache budget is too small: two routed prefill layers need {:.2} GiB", prefill_headroom as f64 / 1_073_741_824.0 )); } ( u32::try_from((total - prefill_headroom) / per_expert_bytes) .unwrap_or(u32::MAX) .min(max_experts), prefill_headroom, ) }; if cache_experts == 0 { return Err("SSD streaming has no memory for an expert cache".into()); } let cache_bytes = per_expert_bytes.saturating_mul(u64::from(cache_experts)); let admission_bytes = resident_bytes .saturating_add(runtime_bytes) .saturating_add(cache_bytes) .saturating_add(reserved_headroom); Ok(Self { model_spans, resident_bytes, cache_experts, per_expert_bytes, admission_bytes, preload_experts: if settings.cold { 0 } else if settings.preload_experts != 0 { settings.preload_experts.min(cache_experts) } else { let cap = env::var("DS4_METAL_STREAMING_EXPERT_AUTO_PRELOAD_CAP") .ok() .and_then(|value| value.parse::().ok()) .unwrap_or(4096); if cap == 0 { cache_experts } else { cache_experts.min(cap) } }, preload_by_layer: vec![Vec::new(); model.shape.layers as usize], preload_seeded: (0..model.shape.layers) .map(|_| std::sync::atomic::AtomicBool::new(false)) .collect(), static_decode_map_current: std::sync::atomic::AtomicBool::new(false), cold: settings.cold, loader: None, selected_requests: std::cell::Cell::new(0), selected_experts: std::cell::Cell::new(0), selected_wait_ns: std::cell::Cell::new(0), }) } fn configure(&mut self, model: &Model, _weights: &Weights) -> Result<(), String> { unsafe { ds4_gpu_set_streaming_expert_cache_expert_bytes(self.per_expert_bytes); ds4_gpu_set_streaming_expert_cache_budget(self.cache_experts); } if self.cold || self.preload_experts == 0 || env::var_os("DS4_METAL_DISABLE_STREAMING_EXPERT_HOTLIST").is_some() { self.loader = Some(SelectedLoadWorker::new()); return Ok(()); } let mut loaded = 0_u32; let hotlist = match model.shape.model { ModelChoice::DeepSeekV4Flash0731 => hotlist::FLASH, ModelChoice::DeepSeekV4Pro => hotlist::PRO, ModelChoice::Glm52 => unreachable!("GLM uses its dedicated executor"), }; for &(layer, expert) in hotlist { if loaded == self.preload_experts { break; } if u32::from(layer) >= model.shape.layers || u64::from(expert) >= model.shape.experts { continue; } self.preload_by_layer[layer as usize] .push((i32::from(expert), self.preload_experts - loaded)); loaded += 1; } self.loader = Some(SelectedLoadWorker::new()); Ok(()) } fn seed_mapped_layer( &self, model: &Model, layer: &Layer, index: usize, gpu_copy: bool, ) -> Result { if self.preload_seeded[index].load(std::sync::atomic::Ordering::Acquire) { return Ok(true); } let entries = &self.preload_by_layer[index]; if entries.is_empty() { self.preload_seeded[index].store(true, std::sync::atomic::Ordering::Release); return Ok(true); } let ids = entries.iter().map(|entry| entry.0).collect::>(); let priorities = entries.iter().map(|entry| entry.1).collect::>(); let table = stream_expert_table(model, layer, index as u32, model.shape.experts); let seeded = unsafe { if gpu_copy { ds4_gpu_stream_expert_cache_seed_experts_gpu_copy( &table, ids.as_ptr(), priorities.as_ptr(), ids.len() as u32, ) } else { ds4_gpu_stream_expert_cache_seed_experts( &table, ids.as_ptr(), priorities.as_ptr(), ids.len() as u32, ) } } != 0; if seeded { self.preload_seeded[index].store(true, std::sync::atomic::Ordering::Release); } Ok(seeded) } fn begin_selected( &self, selected: &Buffer, table: StreamExpertTable, count: u32, ) -> Result<(), String> { self.selected_requests .set(self.selected_requests.get().saturating_add(1)); self.selected_experts .set(self.selected_experts.get().saturating_add(u64::from(count))); self.loader .as_ref() .ok_or("SSD expert loader was not initialized")? .begin(selected, table, count) } fn finish_selected(&self, set_override: bool) -> Result<(), String> { let started = Instant::now(); let result = self .loader .as_ref() .ok_or("SSD expert loader was not initialized")? .finish(set_override); self.selected_wait_ns.set( self.selected_wait_ns .get() .saturating_add(u64::try_from(started.elapsed().as_nanos()).unwrap_or(u64::MAX)), ); result } } fn stream_expert_table( model: &Model, layer: &Layer, index: u32, experts: u64, ) -> StreamExpertTable { StreamExpertTable { model_map: model.main.map_ptr().cast(), model_size: model.main.len(), layer: index, total_experts: experts as u32, gate_offset: layer.expert_gate.offset, up_offset: layer.expert_up.offset, down_offset: layer.expert_down.offset, gate_expert_bytes: layer.expert_gate.bytes / experts, down_expert_bytes: layer.expert_down.bytes / experts, } } fn streaming_model_spans( model: &Model, weights: &Weights, slab_bytes: u64, ) -> Result, String> { let mut streamed = Vec::new(); for layer in &weights.layers { let bytes = (layer.expert_gate.bytes + layer.expert_up.bytes + layer.expert_down.bytes) / model.shape.experts; if bytes == slab_bytes { streamed.extend([ (layer.expert_gate.offset, layer.expert_gate.bytes), (layer.expert_up.offset, layer.expert_up.bytes), (layer.expert_down.offset, layer.expert_down.bytes), ]); } } let mut spans = model .main .tensors .values() .filter(|tensor| { !streamed .iter() .any(|&(offset, bytes)| tensor.offset == offset && tensor.bytes == bytes) }) .map(|tensor| (tensor.offset, tensor.bytes)) .collect::>(); spans.sort_unstable_by_key(|span| span.0); let mut merged: Vec<(u64, u64)> = Vec::new(); for (offset, bytes) in spans { let end = offset.checked_add(bytes).ok_or("model span overflow")?; if let Some((previous_offset, previous_bytes)) = merged.last_mut() { let previous_end = previous_offset.saturating_add(*previous_bytes); if offset <= previous_end { *previous_bytes = previous_end.max(end) - *previous_offset; continue; } } merged.push((offset, bytes)); } if merged.is_empty() { return Err("SSD streaming found no resident model tensors".into()); } Ok(merged) } fn deepseek_model_spans( model: &Model, purpose: &str, mut include: impl FnMut(&str) -> bool, ) -> Result { const ISOLATED_Q4_BYTES: u64 = 2 * 1024 * 1024 * 1024; let q4_groups = std::env::var("DS4_METAL_Q4_PRO_MAP_GROUPS") .ok() .and_then(|value| value.parse::().ok()) .filter(|groups| *groups > 0 && *groups <= 384 && 384 % *groups == 0) .unwrap_or(1); let mut spans = Vec::new(); let mut max_tensor_bytes = 0; for (_, tensor) in model.main.tensors.iter().filter(|(name, _)| include(name)) { if tensor.bytes == 0 { continue; } let isolate = matches!(tensor.kind, Q4_K | MXFP4) && tensor.bytes >= ISOLATED_Q4_BYTES; let groups = if isolate && tensor.dims.len() == 3 && tensor.dims[2] == 384 && tensor.bytes.is_multiple_of(u64::from(q4_groups)) { q4_groups } else { 1 }; let bytes = tensor.bytes / u64::from(groups); max_tensor_bytes = max_tensor_bytes.max(bytes); for group in 0..groups { spans.push((tensor.offset + u64::from(group) * bytes, bytes, isolate)); } } let ranges = finish_deepseek_model_spans(spans, purpose)?; Ok(DeepSeekModelSpans { ranges, max_tensor_bytes, }) } fn finish_deepseek_model_spans( mut spans: Vec<(u64, u64, bool)>, purpose: &str, ) -> Result, String> { spans.sort_unstable_by_key(|span| span.0); let mut merged: Vec<(u64, u64, bool)> = Vec::new(); for (offset, bytes, isolate) in spans { let end = offset .checked_add(bytes) .ok_or_else(|| format!("{purpose} model span overflow"))?; if let Some((previous_offset, previous_bytes, previous_isolate)) = merged.last_mut() { let previous_end = previous_offset.saturating_add(*previous_bytes); if !isolate && !*previous_isolate && offset <= previous_end { *previous_bytes = previous_end.max(end) - *previous_offset; continue; } } merged.push((offset, bytes, isolate)); } if merged.is_empty() { return Err(format!("{purpose} has no model tensors")); } Ok(merged .into_iter() .map(|(offset, bytes, _)| (offset, bytes)) .collect()) } fn deepseek_token_model_spans(model: &Model) -> Result { deepseek_model_spans(model, "DeepSeek token embedding", |name| { name == "token_embd.weight" }) } fn deepseek_layer_model_spans( model: &Model, weights: &Layer, layer: u32, decode_only: bool, slab_bytes: u64, ) -> Result { let prefix = format!("blk.{layer}."); let stream_experts = decode_only && (weights.expert_gate.bytes + weights.expert_up.bytes + weights.expert_down.bytes) / model.shape.experts == slab_bytes; deepseek_model_spans(model, &format!("DeepSeek layer {layer}"), |name| { (name.starts_with(&prefix) && (!stream_experts || (!name.ends_with("ffn_gate_exps.weight") && !name.ends_with("ffn_up_exps.weight") && !name.ends_with("ffn_down_exps.weight")))) || (!decode_only && layer == 0 && name == "token_embd.weight") }) } fn deepseek_output_model_spans(model: &Model) -> Result { deepseek_model_spans(model, "DeepSeek output head", |name| { name.starts_with("output") }) } fn deepseek_mtp_base_model_spans(model: &Model) -> Result { deepseek_model_spans(model, "DeepSeek MTP base weights", |name| { name == "token_embd.weight" || name.starts_with("output") }) } fn install_deepseek_model_spans( model: &Model, spans: &DeepSeekModelSpans, purpose: &str, ) -> Result<(), String> { let (offsets, sizes): (Vec<_>, Vec<_>) = spans.ranges.iter().copied().unzip(); call( unsafe { ds4_gpu_set_model_map_spans( model.main.map_ptr().cast(), model.main.len(), offsets.as_ptr(), sizes.as_ptr(), spans.ranges.len() as u32, spans.max_tensor_bytes, ) }, purpose, ) } fn install_support_model_map(model: &Model, purpose: &str) -> Result<(), String> { let support = model.support.as_ref().ok_or("support model is missing")?; call( unsafe { ds4_gpu_set_model_map_range( support.map_ptr().cast(), support.len(), support.data_offset(), support.len() - support.data_offset(), support.max_tensor_bytes(), ) }, purpose, ) } fn install_speculative_model_maps(model: &Model, purpose: &str) -> Result<(), String> { install_deepseek_model_spans(model, &deepseek_mtp_base_model_spans(model)?, purpose)?; install_support_model_map(model, purpose) } fn estimated_deepseek_runtime_bytes(shape: super::Shape, context: u32, prefill: u32) -> u64 { let prefill_cap = effective_prefill_cap(context, prefill); let raw_cap = effective_raw_cap(shape, context, prefill_cap); let raw = u64::from(shape.layers) .saturating_mul(raw_cap.into()) .saturating_mul(shape.head_dim) .saturating_mul(4); let compressed = (0..shape.layers).fold(0_u64, |total, layer| { let ratio = compression_ratio(shape, layer); if ratio == 0 { return total; } let rows = u64::from(context / ratio + 2); total .saturating_add(rows.saturating_mul(shape.head_dim).saturating_mul(2)) .saturating_add(if ratio == 4 { rows.saturating_mul(shape.indexer_head_dim) .saturating_mul(4) } else { 0 }) }); let min_ratio = (0..shape.layers) .map(|layer| compression_ratio(shape, layer)) .filter(|ratio| *ratio != 0) .min() .unwrap_or(context.max(1)); let comp_cap = context / min_ratio + 2; let attention_stage = u64::from(prefill_cap / min_ratio + 2) .max(2) .saturating_mul(shape.head_dim) .saturating_mul(4); let scratch = 2_u64 .saturating_mul(comp_cap.into()) .saturating_mul(prefill_cap.into()) .saturating_mul(4) .saturating_add(attention_stage); raw.saturating_add(compressed).saturating_add(scratch) } fn resident_deepseek_admission_bytes( model: &Model, context: u32, prefill: u32, ) -> Result { resident_deepseek_admission_for_weights( model.main.len() - model.main.data_offset(), model.shape, context, prefill, ) } fn resident_deepseek_admission_for_weights( weight_bytes: u64, shape: super::Shape, context: u32, prefill: u32, ) -> Result { weight_bytes .checked_add(estimated_deepseek_runtime_bytes(shape, context, prefill)) .and_then(|bytes| bytes.checked_add(512 * 1024 * 1024)) .ok_or_else(|| "DeepSeek runtime memory size overflow".into()) } fn effective_prefill_cap(context: u32, requested: u32) -> u32 { if requested == 0 { context.clamp(1, DEFAULT_PREFILL_CHUNK) } else { context.min(requested).max(1) } } fn effective_raw_cap(shape: super::Shape, context: u32, prefill_cap: u32) -> u32 { let window = (shape.sliding_window as u32).min(context).max(1); let wanted = u64::from(window) .saturating_add(prefill_cap.into()) .min(context.into()) .max(1) .div_ceil(256) .saturating_mul(256) .min(8192); u32::try_from(wanted).unwrap_or(8192).max(window) } struct Scratch { current_hc: Buffer, next_hc: Buffer, flat_hc: Buffer, hc_mix: Buffer, hc_split: Buffer, current: Buffer, norm: Buffer, q_rank: Buffer, q_rank_norm: Buffer, q: Buffer, kv_raw: Buffer, kv: Buffer, compressed_kv: Buffer, compressed_score: Buffer, index_compressed_kv: Buffer, index_compressed_score: Buffer, compressed_stage: Buffer, indexer_q: Buffer, indexer_weights: Buffer, indexer_scores: Buffer, indexer_selected: Buffer, heads: Buffer, attention_low: Buffer, attention_out: Buffer, router_logits: Buffer, router_probs: Buffer, router_selected: Buffer, router_weights: Buffer, routed_gate: Buffer, routed_up: Buffer, routed_mid: Buffer, routed_experts: Buffer, routed_out: Buffer, shared_gate: Buffer, shared_up: Buffer, shared_mid: Buffer, shared_out: Buffer, output_pre: Buffer, output_weights: Buffer, output_embedding: Buffer, output_norm: Buffer, logits: Buffer, } struct BatchScratch { tokens: Buffer, current_hc: Buffer, next_hc: Buffer, after_attention_hc: Buffer, flat_hc: Buffer, hc_mix: Buffer, hc_split: Buffer, current: Buffer, norm: Buffer, q_rank: Buffer, q_rank_norm: Buffer, q: Buffer, q_half: Buffer, kv_raw: Buffer, kv: Buffer, compressed_kv: Buffer, compressed_score: Buffer, compressed_stage: Buffer, indexer_q: Buffer, indexer_weights: Buffer, indexer_scores: Buffer, indexer_selected: Buffer, heads: Buffer, attention_low: Buffer, attention_group_tmp: Buffer, attention_low_tmp: Buffer, attention_out: Buffer, router_logits: Buffer, router_probs: Buffer, router_selected: Buffer, router_weights: Buffer, routed_gate: Buffer, routed_up: Buffer, routed_mid: Buffer, routed_experts: Buffer, routed_out: Buffer, shared_gate: Buffer, shared_up: Buffer, shared_mid: Buffer, shared_out: Buffer, output_logits: Option, } impl BatchScratch { fn allocate(model: &Model, pos: u32, rows: u32, output_logits: bool) -> Result { let shape = model.shape; let output_rows = if output_logits && rows > 1 && rows < 8 { 8 } else { rows }; let rows = u64::from(output_rows); let hc_dim = shape.hc * shape.embd; let mix_hc = 2 * shape.hc + shape.hc * shape.hc; let q_dim = shape.heads * shape.head_dim; let group_dim = shape.head_dim * (shape.heads / shape.out_groups); let low_dim = shape.out_groups * shape.lora_o; let routed = shape.experts_used * shape.ff_expert; Ok(Self { tokens: Buffer::bytes(rows * 4)?, current_hc: Buffer::floats(rows * hc_dim)?, next_hc: Buffer::floats(rows * hc_dim)?, after_attention_hc: Buffer::floats(rows * hc_dim)?, flat_hc: Buffer::floats(rows * hc_dim)?, hc_mix: Buffer::floats(rows * mix_hc)?, hc_split: Buffer::floats(rows * mix_hc)?, current: Buffer::floats(rows * shape.embd)?, norm: Buffer::floats(rows * shape.embd)?, q_rank: Buffer::floats(rows * shape.lora_q)?, q_rank_norm: Buffer::floats(rows * shape.lora_q)?, q: Buffer::floats(rows * q_dim)?, q_half: Buffer::bytes(rows * q_dim.max(shape.embd) * 2)?, kv_raw: Buffer::floats(rows * shape.head_dim)?, kv: Buffer::floats(rows * shape.head_dim)?, compressed_kv: Buffer::floats(rows * 2 * shape.head_dim)?, compressed_score: Buffer::floats(rows * 2 * shape.head_dim)?, compressed_stage: Buffer::floats(rows * shape.head_dim)?, indexer_q: Buffer::floats(rows * shape.indexer_heads * shape.indexer_head_dim)?, indexer_weights: Buffer::floats(rows * shape.indexer_heads)?, indexer_scores: Buffer::floats( rows * (u64::from(pos) + rows).div_ceil(4).saturating_add(2), )?, indexer_selected: Buffer::bytes(rows * shape.indexer_top_k * 4)?, heads: Buffer::floats(rows * q_dim)?, attention_low: Buffer::floats(rows * low_dim)?, attention_group_tmp: Buffer::floats(rows * group_dim)?, attention_low_tmp: Buffer::floats(rows * shape.lora_o)?, attention_out: Buffer::floats(rows * shape.embd)?, router_logits: Buffer::floats(rows * shape.experts)?, router_probs: Buffer::floats(rows * shape.experts)?, router_selected: Buffer::bytes(rows * shape.experts_used * 4)?, router_weights: Buffer::floats(rows * shape.experts_used)?, routed_gate: Buffer::floats(rows * routed)?, routed_up: Buffer::floats(rows * routed)?, routed_mid: Buffer::floats(rows * routed)?, routed_experts: Buffer::floats(rows * shape.experts_used * shape.embd)?, routed_out: Buffer::floats(rows * shape.embd)?, shared_gate: Buffer::floats(rows * shape.ff_expert)?, shared_up: Buffer::floats(rows * shape.ff_expert)?, shared_mid: Buffer::floats(rows * shape.ff_expert)?, shared_out: Buffer::floats(rows * shape.embd)?, output_logits: output_logits .then(|| Buffer::floats(rows * shape.vocab)) .transpose()?, }) } } impl Scratch { fn allocate(model: &Model, context: u32) -> Result { let shape = model.shape; let hc_dim = shape.hc * shape.embd; let mix_hc = 2 * shape.hc + shape.hc * shape.hc; let q_dim = shape.heads * shape.head_dim; let low_dim = shape.out_groups * shape.lora_o; let routed = shape.experts_used * shape.ff_expert; Ok(Self { current_hc: Buffer::floats(hc_dim)?, next_hc: Buffer::floats(hc_dim)?, flat_hc: Buffer::floats(hc_dim)?, hc_mix: Buffer::floats(mix_hc)?, hc_split: Buffer::floats(mix_hc)?, current: Buffer::floats(shape.embd)?, norm: Buffer::floats(shape.embd)?, q_rank: Buffer::floats(shape.lora_q)?, q_rank_norm: Buffer::floats(shape.lora_q)?, q: Buffer::floats(q_dim)?, kv_raw: Buffer::floats(shape.head_dim)?, kv: Buffer::floats(shape.head_dim)?, compressed_kv: Buffer::floats(2 * shape.head_dim)?, compressed_score: Buffer::floats(2 * shape.head_dim)?, index_compressed_kv: Buffer::floats(2 * shape.indexer_head_dim)?, index_compressed_score: Buffer::floats(2 * shape.indexer_head_dim)?, compressed_stage: Buffer::floats(shape.head_dim)?, indexer_q: Buffer::floats(shape.indexer_heads * shape.indexer_head_dim)?, indexer_weights: Buffer::floats(shape.indexer_heads)?, indexer_scores: Buffer::floats(u64::from(context / 4 + 2))?, indexer_selected: Buffer::bytes(shape.indexer_top_k * 4)?, heads: Buffer::floats(q_dim)?, attention_low: Buffer::floats(low_dim)?, attention_out: Buffer::floats(shape.embd)?, router_logits: Buffer::floats(shape.experts)?, router_probs: Buffer::floats(shape.experts)?, router_selected: Buffer::bytes(shape.experts_used * 4)?, router_weights: Buffer::floats(shape.experts_used)?, routed_gate: Buffer::floats(routed)?, routed_up: Buffer::floats(routed)?, routed_mid: Buffer::floats(routed)?, routed_experts: Buffer::floats(shape.experts_used * shape.embd)?, routed_out: Buffer::floats(shape.embd)?, shared_gate: Buffer::floats(shape.ff_expert)?, shared_up: Buffer::floats(shape.ff_expert)?, shared_mid: Buffer::floats(shape.ff_expert)?, shared_out: Buffer::floats(shape.embd)?, output_pre: Buffer::floats(shape.hc)?, output_weights: Buffer::floats(shape.hc)?, output_embedding: Buffer::floats(shape.embd)?, output_norm: Buffer::floats(shape.embd)?, logits: Buffer::floats(shape.vocab)?, }) } } /// Rust-owned DeepSeek Flash session. The ObjC boundary only schedules the /// unchanged Metal kernels; token/session/graph state lives here. pub(super) struct Session { scratch: Scratch, layers: Vec, context: u32, prefill_cap: u32, raw_cap: u32, position: u32, } struct LayerState { raw_cache: Buffer, compression: Option, indexer: Option, } struct CompressionState { ratio: u32, cache: Buffer, state_kv: Buffer, state_score: Buffer, rows: u32, } struct CompressionFrontier { state_kv: Buffer, state_score: Buffer, bytes: u64, rows: u32, } struct LayerFrontier { compression: Option, indexer: Option, } struct SpecFrontier { layers: Vec, position: u32, token_len: usize, logits: Vec, dspark_target_hidden: Option, dspark_capture_mask: u32, dspark_cache_start: u32, dspark_cache_len: u32, } struct SpecPrefixFrontier { layers: Vec, } struct BatchVerification { tops: Vec, logits: Vec>, prefixes: Vec, } fn capture_compression_frontier( state: &CompressionState, bytes: u64, purpose: &str, ) -> Result { let state_kv = Buffer::bytes(bytes)?; let state_score = Buffer::bytes(bytes)?; state_kv.copy_from(0, &state.state_kv, 0, bytes, purpose)?; state_score.copy_from(0, &state.state_score, 0, bytes, purpose)?; Ok(CompressionFrontier { state_kv, state_score, bytes, rows: state.rows, }) } fn restore_compression_frontier( state: &mut CompressionState, saved: &CompressionFrontier, purpose: &str, ) -> Result<(), String> { state .state_kv .copy_from(0, &saved.state_kv, 0, saved.bytes, purpose)?; state .state_score .copy_from(0, &saved.state_score, 0, saved.bytes, purpose)?; state.rows = saved.rows; Ok(()) } impl LayerState { fn allocate(model: &Model, index: u32, context: u32, raw_cap: u32) -> Result { let shape = model.shape; let ratio = compression_ratio(shape, index); let compression = if ratio == 0 { None } else { let coefficient = if ratio == 4 { 2 } else { 1 }; let width = coefficient * shape.head_dim; let state_rows = coefficient * ratio as u64; let state_kv = Buffer::floats(width * state_rows)?; let state_score = Buffer::floats(width * state_rows)?; state_kv.fill(0.0, width * state_rows)?; state_score.fill(f32::NEG_INFINITY, width * state_rows)?; Some(CompressionState { ratio, cache: Buffer::bytes((context / ratio + 2) as u64 * shape.head_dim * 2)?, state_kv, state_score, rows: 0, }) }; let indexer = if ratio == 4 { let width = 2 * shape.indexer_head_dim; let state_rows = 2 * u64::from(ratio); let state_kv = Buffer::floats(width * state_rows)?; let state_score = Buffer::floats(width * state_rows)?; state_kv.fill(0.0, width * state_rows)?; state_score.fill(f32::NEG_INFINITY, width * state_rows)?; Some(CompressionState { ratio, cache: Buffer::floats(u64::from(context / ratio + 2) * shape.indexer_head_dim)?, state_kv, state_score, rows: 0, }) } else { None }; Ok(Self { raw_cache: Buffer::floats(raw_cap as u64 * shape.head_dim)?, compression, indexer, }) } } impl Session { fn new(model: &Model, context: u32, prefill_chunk: u32) -> Result { if context == 0 { return Err("context must contain at least one token".into()); } let prefill_cap = effective_prefill_cap(context, prefill_chunk); let raw_cap = effective_raw_cap(model.shape, context, prefill_cap); let layers = (0..model.shape.layers) .map(|index| LayerState::allocate(model, index, context, raw_cap)) .collect::>()?; Ok(Self { scratch: Scratch::allocate(model, context)?, layers, context, prefill_cap, raw_cap, position: 0, }) } } // SAFETY: declaration order is required because Rust drops fields in order. // `session` must release every Buffer before `_context` calls ds4_gpu_cleanup(), // and `_context` must drop before `model` unmaps memory wrapped without copying // by `ds4_gpu_cleanup` in native/metal/ds4_metal.m. This intentionally differs from // DS4's `ds4.c` consumes this exact field order; do not reorder it. #[derive(Clone, Copy, Default)] pub(super) struct ExecutionStats { pub(super) speculative_mode: u8, pub(super) speculative_cycles: u64, pub(super) drafted_tokens: u64, pub(super) accepted_draft_tokens: u64, pub(super) verifier_passes: u64, pub(super) verifier_ms: u64, pub(super) ssd_enabled: bool, pub(super) ssd_resident_bytes: u64, pub(super) ssd_cache_bytes: u64, pub(super) ssd_cache_experts: u64, pub(super) ssd_cache_entries: u64, pub(super) ssd_preloaded_experts: u64, pub(super) ssd_cache_hits: u64, pub(super) ssd_cache_misses: u64, pub(super) ssd_cache_evictions: u64, pub(super) ssd_cache_wraps: u64, pub(super) ssd_buffer_allocs: u64, pub(super) ssd_buffer_reuses: u64, pub(super) ssd_pread_bytes: u64, pub(super) ssd_pread_ms: u64, pub(super) ssd_evict_advise_bytes: u64, pub(super) ssd_willneed_advise_bytes: u64, pub(super) ssd_selected_requests: u64, pub(super) ssd_requested_bytes: u64, pub(super) ssd_wait_ms: u64, } pub(super) struct DeepSeekExecutor { weights: Weights, session: Session, dspark: Option, steering: Option, ssd: Option, profile: Option, logits: Vec, tokens: Vec, quality: bool, power_percent: u8, prefill_layer_average: Vec, decode_average: f64, speculative_cycles: u64, verifier_passes: u64, verifier_ns: u64, checkpoint_tag: [u8; 32], model_modified: (u64, u32), model_identity: [u8; 32], mxfp4_decode_fast_lookup: bool, _context: Context, model: Model, speculative: EngineSpeculativeSettings, } pub(super) struct DeepSeekResidentState { session: Session, dspark: Option, logits: Vec, tokens: Vec, checkpoint_tag: [u8; 32], } impl DeepSeekExecutor { #[allow(dead_code)] #[allow(clippy::too_many_arguments)] pub(super) fn open( model: Model, context: u32, quality: bool, prefill_chunk: u32, power_percent: u8, speculative: EngineSpeculativeSettings, ssd: EngineSsdSettings, steering: EngineSteeringSettings, ) -> Result { Self::open_profile( model, context, quality, prefill_chunk, power_percent, speculative, ssd, steering, None, ) } #[allow(clippy::too_many_arguments)] pub(super) fn open_profile( model: Model, context: u32, quality: bool, prefill_chunk: u32, power_percent: u8, speculative: EngineSpeculativeSettings, ssd: EngineSsdSettings, steering: EngineSteeringSettings, expert_profile_path: Option<&str>, ) -> Result { let weights = Weights::bind(&model)?; let mut ssd_plan = ssd .enabled .then(|| SsdPlan::new(&model, &weights, ssd, context, prefill_chunk)) .transpose()?; let initial_ssd_spans = ssd_plan .as_ref() .map(|_| deepseek_token_model_spans(&model)) .transpose()?; let spans = initial_ssd_spans .as_ref() .map(|spans| (spans.ranges.as_slice(), spans.max_tensor_bytes)); let admission = if let Some(plan) = &ssd_plan { plan.admission_bytes } else { resident_deepseek_admission_bytes(&model, context, prefill_chunk)? }; let context_handle = Context::open(&model, quality, ssd.enabled, admission, spans)?; let mxfp4_decode_fast_lookup = mxfp4_decode_fast_lookup_allowed(&model, &weights, quality, ssd.enabled); let steering = Steering::load(&model, steering)?; let session = Session::new( &model, context, if prefill_chunk == 0 { DEFAULT_PREFILL_CHUNK } else { prefill_chunk }, )?; let dspark = match ( model.support_kind, model.support.as_ref(), speculative.dspark, ) { (Some(SupportKind::DSpark), Some(support), true) => Some(Dspark::new( &model, support, &session, speculative, quality, )?), _ => None, }; if let Some(plan) = &mut ssd_plan { plan.configure(&model, &weights)?; } let model_modified = fs::metadata(model.main.path()) .and_then(|metadata| metadata.modified()) .ok() .and_then(|modified| modified.duration_since(UNIX_EPOCH).ok()) .map(|duration| (duration.as_secs(), duration.subsec_nanos())) .unwrap_or_default(); let model_identity = if let Some(steering) = &steering { let mut hash = Sha256::new(); hash.update(model.checkpoint_identity()); hash.update(steering.identity); hash.finalize().into() } else { model.checkpoint_identity() }; let profile = ExpertProfile::new( expert_profile_path, model.shape.model, model.shape.layers, model.shape.experts, model.shape.experts_used, )?; Ok(Self { weights, session, dspark, steering, ssd: ssd_plan, profile, logits: vec![0.0; model.shape.vocab as usize], tokens: Vec::new(), quality, power_percent: if power_percent == 0 { 100 } else { power_percent }, prefill_layer_average: vec![0.0; model.shape.layers as usize], decode_average: 0.0, speculative_cycles: 0, verifier_passes: 0, verifier_ns: 0, checkpoint_tag: [0; 32], model_modified, model_identity, mxfp4_decode_fast_lookup, _context: context_handle, model, speculative, }) } pub(super) fn eval(&mut self, token: i32) -> Result<(), String> { self.eval_target(token)?; self.seed_dspark_current_cache() } fn seed_dspark_current_cache(&mut self) -> Result<(), String> { if let (Some(_), Some(ssd)) = (&self.dspark, &self.ssd) { install_speculative_model_maps(&self.model, "DSpark support mapping")?; ssd.static_decode_map_current .store(false, std::sync::atomic::Ordering::Release); } if let (Some(dspark), Some(support)) = (&mut self.dspark, self.model.support.as_ref()) { dspark.seed_current_cache( support, self.session.position - 1, self.session.raw_cap, self.model.shape, )?; } Ok(()) } fn eval_target(&mut self, token: i32) -> Result<(), String> { if token < 0 || token as u64 >= self.model.shape.vocab { return Err(format!("token {token} is outside the vocabulary")); } if self.session.position >= self.session.context { return Err(format!( "the Rust Metal executor currently supports {} tokens per session", self.session.context )); } let started = Instant::now(); if let Some(dspark) = &mut self.dspark { dspark.begin_capture(); } let fast_lookup = self.mxfp4_decode_fast_lookup && (self.session.position >= 2048 || !environment_present( c"DS4_METAL_DISABLE_PRE_M5_DECODE_EARLY_PIPELINE_FAST_LOOKUP", )); let previous_fast_lookup = unsafe { ds4_gpu_set_decode_pipeline_fast_lookup(i32::from(fast_lookup)) }; let encoded = if self.ssd.is_some() { self.encode_streaming_token(token as u32) } else { (|| { let commands = Commands::begin()?; self.encode_token(token as u32)?; commands.finish() })() }; unsafe { ds4_gpu_set_decode_pipeline_fast_lookup(previous_fast_lookup) }; encoded?; self.session.scratch.logits.read_f32(&mut self.logits)?; self.session.position += 1; self.tokens.push(token); if let Some(profile) = &self.profile { profile.write()?; } throttle( &mut self.decode_average, started.elapsed(), self.power_percent, ); Ok(()) } fn snapshot_spec_frontier(&self) -> Result { let shape = self.model.shape; let commands = Commands::begin()?; let layers = self .session .layers .iter() .map(|layer| { let compression = layer .compression .as_ref() .map(|state| { let coefficient = if state.ratio == 4 { 2 } else { 1 }; capture_compression_frontier( state, coefficient * coefficient * state.ratio as u64 * shape.head_dim * 4, "saving speculative compressor state", ) }) .transpose()?; let indexer = layer .indexer .as_ref() .map(|state| { capture_compression_frontier( state, 4 * state.ratio as u64 * shape.indexer_head_dim * 4, "saving speculative indexer state", ) }) .transpose()?; Ok::<_, String>(LayerFrontier { compression, indexer, }) }) .collect::, _>>()?; let dspark_target_hidden = self .dspark .as_ref() .map(|dspark| { let bytes = dspark.config.target_layers.len() as u64 * shape.embd * 4; let saved = Buffer::bytes(bytes)?; saved.copy_from( 0, &dspark.target_hidden, 0, bytes, "saving speculative DSpark target state", )?; Ok::<_, String>(saved) }) .transpose()?; commands.finish()?; Ok(SpecFrontier { layers, position: self.session.position, token_len: self.tokens.len(), logits: self.logits.clone(), dspark_target_hidden, dspark_capture_mask: self.dspark.as_ref().map_or(0, |value| value.capture_mask), dspark_cache_start: self.dspark.as_ref().map_or(0, |value| value.cache_start), dspark_cache_len: self.dspark.as_ref().map_or(0, |value| value.cache_len), }) } fn restore_spec_frontier(&mut self, frontier: &SpecFrontier) -> Result<(), String> { if frontier.layers.len() != self.session.layers.len() { return Err("speculative frontier layer count changed".into()); } let commands = Commands::begin()?; for (layer, saved) in self.session.layers.iter_mut().zip(&frontier.layers) { match (&mut layer.compression, &saved.compression) { (Some(state), Some(saved)) => restore_compression_frontier( state, saved, "restoring speculative compressor state", )?, (None, None) => {} _ => return Err("speculative compressor layout changed".into()), } match (&mut layer.indexer, &saved.indexer) { (Some(state), Some(saved)) => restore_compression_frontier( state, saved, "restoring speculative indexer state", )?, (None, None) => {} _ => return Err("speculative indexer layout changed".into()), } } if let (Some(dspark), Some(saved)) = (&mut self.dspark, &frontier.dspark_target_hidden) { let bytes = dspark.config.target_layers.len() as u64 * self.model.shape.embd * 4; dspark.target_hidden.copy_from( 0, saved, 0, bytes, "restoring speculative DSpark target state", )?; dspark.capture_mask = frontier.dspark_capture_mask; dspark.cache_start = frontier.dspark_cache_start; dspark.cache_len = frontier.dspark_cache_len; } commands.finish()?; self.session.position = frontier.position; self.tokens.truncate(frontier.token_len); self.logits.clone_from(&frontier.logits); Ok(()) } fn commit_spec_prefix( &mut self, baseline: &SpecFrontier, prefix: &SpecPrefixFrontier, proposals: &[i32], logits: &[f32], ) -> Result<(), String> { let count = u32::try_from(proposals.len()).map_err(|_| "speculative prefix is too large")?; if count == 0 || prefix.layers.len() != self.session.layers.len() { return Err("invalid speculative prefix frontier".into()); } let commands = Commands::begin()?; for (layer, saved) in self.session.layers.iter_mut().zip(&prefix.layers) { match (&mut layer.compression, &saved.compression) { (Some(state), Some(saved)) => restore_compression_frontier( state, saved, "committing speculative compressor prefix", )?, (None, None) => {} _ => return Err("speculative compressor prefix layout changed".into()), } match (&mut layer.indexer, &saved.indexer) { (Some(state), Some(saved)) => restore_compression_frontier( state, saved, "committing speculative indexer prefix", )?, (None, None) => {} _ => return Err("speculative indexer prefix layout changed".into()), } } if let Some(dspark) = &mut self.dspark { let row = u64::from(count - 1); for slot in 0..dspark.config.target_layers.len() as u64 { dspark.target_hidden.copy_from( slot * self.model.shape.embd * 4, &dspark.target_hidden_batch, (slot * u64::from(self.session.prefill_cap) + row) * self.model.shape.embd * 4, self.model.shape.embd * 4, "committing speculative DSpark target prefix", )?; } dspark.capture_mask = (1_u32 << dspark.config.target_layers.len()) - 1; dspark.cache_start = baseline.dspark_cache_start; dspark.cache_len = baseline.dspark_cache_len; dspark.commit_proposed_prefix(count, self.session.raw_cap); } commands.finish()?; self.session.position = baseline.position + count; self.tokens.truncate(baseline.token_len); self.tokens.extend_from_slice(proposals); self.logits.clone_from_slice(logits); Ok(()) } fn verify_target_suffix( &mut self, proposals: &[i32], cancelled: &std::sync::atomic::AtomicBool, ) -> Result, String> { if proposals.is_empty() || argmax(&self.logits) != proposals[0] || cancelled.load(std::sync::atomic::Ordering::Relaxed) { return Ok(Vec::new()); } let started = Instant::now(); if self.quality || proposals.len() == 1 || self .ssd .as_ref() .is_some_and(|ssd| u64::from(ssd.cache_experts) < self.model.shape.experts) { let mut accepted = Vec::new(); for &proposal in proposals { if argmax(&self.logits) != proposal || cancelled.load(std::sync::atomic::Ordering::Relaxed) { break; } self.eval_target(proposal)?; self.verifier_passes += 1; accepted.push(proposal); } self.verifier_ns = self .verifier_ns .saturating_add(u64::try_from(started.elapsed().as_nanos()).unwrap_or(u64::MAX)); return Ok(accepted); } let frontier = self.snapshot_spec_frontier()?; let verification = match self.eval_batch_tops(proposals) { Ok(verification) => { self.verifier_passes += 1; verification } Err(error) => { self.restore_spec_frontier(&frontier)?; self.verifier_ns = self.verifier_ns.saturating_add( u64::try_from(started.elapsed().as_nanos()).unwrap_or(u64::MAX), ); return Err(error); } }; let mut commit = 1_usize; while commit < proposals.len() && verification.tops[commit - 1] == proposals[commit] { commit += 1; } if commit == proposals.len() { self.verifier_ns = self .verifier_ns .saturating_add(u64::try_from(started.elapsed().as_nanos()).unwrap_or(u64::MAX)); return Ok(proposals.to_vec()); } if let (Some(prefix), Some(logits)) = ( verification.prefixes.get(commit - 1), verification.logits.get(commit - 1), ) { self.commit_spec_prefix(&frontier, prefix, &proposals[..commit], logits)?; self.verifier_ns = self .verifier_ns .saturating_add(u64::try_from(started.elapsed().as_nanos()).unwrap_or(u64::MAX)); return Ok(proposals[..commit].to_vec()); } self.restore_spec_frontier(&frontier)?; if let Some(dspark) = &mut self.dspark { dspark.commit_proposed_prefix(1, self.session.raw_cap); } for (index, &proposal) in proposals[..commit].iter().enumerate() { if index == 0 && self.dspark.is_some() { self.eval_target(proposal)?; } else { self.eval(proposal)?; } self.verifier_passes += 1; } self.verifier_ns = self .verifier_ns .saturating_add(u64::try_from(started.elapsed().as_nanos()).unwrap_or(u64::MAX)); Ok(proposals[..commit].to_vec()) } #[allow(clippy::too_many_arguments)] fn verify_target_suffix_stochastic( &mut self, proposals: &[i32], temperature: f32, top_p: f32, min_p: f32, top_k: i32, rng: &mut Rng, cancelled: &std::sync::atomic::AtomicBool, ) -> Result<(Vec, usize), String> { if proposals.is_empty() || cancelled.load(std::sync::atomic::Ordering::Relaxed) { return Ok((Vec::new(), 0)); } let started = Instant::now(); if self.quality || proposals.len() == 1 || self .ssd .as_ref() .is_some_and(|ssd| u64::from(ssd.cache_experts) < self.model.shape.experts) { let mut emitted = Vec::new(); let mut accepted = 0; for &proposal in proposals { let (token, was_draft) = exact_delta_sample( &self.logits, proposal, temperature, top_p, min_p, top_k, rng, ); self.eval_target(token)?; self.verifier_passes += 1; emitted.push(token); if !was_draft { break; } accepted += 1; if cancelled.load(std::sync::atomic::Ordering::Relaxed) { break; } } self.verifier_ns = self .verifier_ns .saturating_add(u64::try_from(started.elapsed().as_nanos()).unwrap_or(u64::MAX)); return Ok((emitted, accepted)); } let (first, accepted_first) = exact_delta_sample( &self.logits, proposals[0], temperature, top_p, min_p, top_k, rng, ); if !accepted_first { if let Some(dspark) = &mut self.dspark { dspark.commit_proposed_prefix(1, self.session.raw_cap); } self.eval_target(first)?; self.verifier_passes += 1; self.verifier_ns = self .verifier_ns .saturating_add(u64::try_from(started.elapsed().as_nanos()).unwrap_or(u64::MAX)); return Ok((vec![first], 0)); } let frontier = self.snapshot_spec_frontier()?; let verification = match self.eval_batch_tops(proposals) { Ok(verification) => { self.verifier_passes += 1; verification } Err(error) => { self.restore_spec_frontier(&frontier)?; return Err(error); } }; let mut accepted = 1; let mut replacement = None; for (index, &proposal) in proposals.iter().enumerate().skip(1) { let (token, was_draft) = exact_delta_sample( &verification.logits[index - 1], proposal, temperature, top_p, min_p, top_k, rng, ); if !was_draft { replacement = Some(token); break; } accepted += 1; } if replacement.is_none() { self.verifier_ns = self .verifier_ns .saturating_add(u64::try_from(started.elapsed().as_nanos()).unwrap_or(u64::MAX)); return Ok((proposals.to_vec(), accepted)); } let prefix = verification .prefixes .get(accepted - 1) .ok_or("missing stochastic verifier prefix")?; let logits = verification .logits .get(accepted - 1) .ok_or("missing stochastic verifier logits")?; self.commit_spec_prefix(&frontier, prefix, &proposals[..accepted], logits)?; let replacement = replacement.expect("replacement disappeared"); self.eval_target(replacement)?; self.verifier_passes += 1; let mut emitted = proposals[..accepted].to_vec(); emitted.push(replacement); self.verifier_ns = self .verifier_ns .saturating_add(u64::try_from(started.elapsed().as_nanos()).unwrap_or(u64::MAX)); Ok((emitted, accepted)) } pub(super) fn eval_speculative_greedy( &mut self, first_token: i32, max_tokens: u32, reasoning: ReasoningMode, cancelled: &std::sync::atomic::AtomicBool, ) -> Result, String> { if self.dspark.is_some() { self.speculative_cycles += 1; } self.eval_target(first_token)?; let mut accepted = vec![first_token]; if self.dspark.as_ref().is_some_and(|dspark| dspark.strict) { return Ok(accepted); } if self.dspark.is_some() { let scheduler_skip = self .dspark .as_mut() .is_some_and(Dspark::scheduler_should_skip); if max_tokens < 10 || scheduler_skip { self.seed_dspark_current_cache()?; return Ok(accepted); } if self.ssd.is_some() { install_speculative_model_maps(&self.model, "DSpark support mapping")?; } let mut dspark = self.dspark.take().expect("DSpark disappeared"); let proposals = dspark.propose( &self.model, &self.weights, first_token, self.session.position.saturating_sub(1), self.session.raw_cap, ); self.dspark = Some(dspark); let mut proposals = proposals?; proposals.truncate( max_tokens .saturating_sub(1) .min(self.session.context.saturating_sub(self.session.position)) as usize, ); if let Some(stop) = proposals .iter() .position(|token| self.model.is_stop_token_for_reasoning(*token, reasoning)) { proposals.truncate(stop + 1); } let no_draft = proposals.is_empty(); let verified = self.verify_target_suffix(&proposals, cancelled)?; if verified.is_empty() { self.dspark .as_mut() .expect("DSpark disappeared") .commit_proposed_prefix(1, self.session.raw_cap); } accepted.extend_from_slice(&verified); let dspark = self.dspark.as_mut().expect("DSpark disappeared"); dspark.accepted += verified.len() as u64; dspark.scheduler_note(verified.len() as u32, no_draft); return Ok(accepted); } Ok(accepted) } #[allow(clippy::too_many_arguments)] fn eval_speculative_sampled( &mut self, first_token: i32, max_tokens: u32, reasoning: ReasoningMode, temperature: f32, top_p: f32, min_p: f32, top_k: i32, rng: &mut Rng, cancelled: &std::sync::atomic::AtomicBool, ) -> Result, String> { if self.dspark.is_none() { self.eval_target(first_token)?; return Ok(vec![first_token]); } if !self.speculative.dspark_exact_sampling { return self.eval_speculative_greedy(first_token, max_tokens, reasoning, cancelled); } self.speculative_cycles += 1; self.eval_target(first_token)?; let mut emitted = vec![first_token]; if self.dspark.as_ref().is_some_and(|dspark| dspark.strict) || max_tokens <= 1 || cancelled.load(std::sync::atomic::Ordering::Relaxed) { return Ok(emitted); } let scheduler_skip = self .dspark .as_mut() .is_some_and(Dspark::scheduler_should_skip); if max_tokens < 10 || scheduler_skip { self.seed_dspark_current_cache()?; return Ok(emitted); } if self.ssd.is_some() { install_speculative_model_maps(&self.model, "DSpark support mapping")?; } let mut dspark = self.dspark.take().expect("DSpark disappeared"); let proposals = dspark.propose( &self.model, &self.weights, first_token, self.session.position.saturating_sub(1), self.session.raw_cap, ); self.dspark = Some(dspark); let mut proposals = proposals?; proposals.truncate( max_tokens .saturating_sub(1) .min(self.session.context.saturating_sub(self.session.position)) as usize, ); if let Some(stop) = proposals .iter() .position(|token| self.model.is_stop_token_for_reasoning(*token, reasoning)) { proposals.truncate(stop + 1); } let no_draft = proposals.is_empty(); if proposals.len() < 2 { let dspark = self.dspark.as_mut().expect("DSpark disappeared"); dspark.commit_proposed_prefix(1, self.session.raw_cap); dspark.scheduler_note(0, no_draft); return Ok(emitted); } let (verified, accepted) = self.verify_target_suffix_stochastic( &proposals, temperature, top_p, min_p, top_k, rng, cancelled, )?; emitted.extend_from_slice(&verified); let dspark = self.dspark.as_mut().expect("DSpark disappeared"); dspark.accepted += accepted as u64; dspark.scheduler_note(accepted as u32, no_draft); Ok(emitted) } pub(super) fn prefill( &mut self, tokens: &[i32], mut progress: impl FnMut(u32) -> bool, ) -> Result { if self.session.position == 0 && self.ssd.is_some() { unsafe { ds4_gpu_stream_expert_cache_reset_route_hotness() }; } let streaming_decode_cap = if matches!(self.model.shape.model, ModelChoice::DeepSeekV4Flash0731) && self.weights.layers.first().is_some_and(|layer| { matches!(layer.expert_gate.kind, Q4_K | MXFP4) && layer.expert_up.kind == layer.expert_gate.kind && layer.expert_down.kind == layer.expert_gate.kind }) { 64 } else { 18 }; if self.ssd.is_some() && !self.quality && !tokens.is_empty() && tokens.len() <= streaming_decode_cap { for (index, &token) in tokens.iter().enumerate() { if !progress(self.session.position) { return Ok(index); } self.eval(token)?; } return Ok(tokens.len()); } if tokens.len() < 4 && self.session.position != 0 { for (index, &token) in tokens.iter().enumerate() { self.eval(token)?; if !progress(self.session.position) { return Ok(index + 1); } } return Ok(tokens.len()); } let end = self .session .position .checked_add(u32::try_from(tokens.len()).map_err(|_| "prefill is too large")?) .ok_or("prefill position overflow")?; if end > self.session.context { return Err(format!( "the Rust Metal executor currently supports {} tokens per session", self.session.context )); } let mut consumed = 0_usize; while consumed < tokens.len() { if !progress(self.session.position) { return Ok(consumed); } let pos = self.session.position; let mut cap = self.session.prefill_cap; if pos != 0 { cap = cap.min(self.session.raw_cap); let offset = pos % self.session.prefill_cap; if offset != 0 { cap = cap.min(self.session.prefill_cap - offset); } } let rows = (tokens.len() - consumed).min(cap as usize); self.eval_batch(&tokens[consumed..consumed + rows])?; consumed += rows; if !progress(self.session.position) { return Ok(consumed); } } Ok(tokens.len()) } fn eval_batch(&mut self, tokens: &[i32]) -> Result<(), String> { self.eval_batch_inner(tokens, false).map(|_| ()) } fn eval_batch_tops(&mut self, tokens: &[i32]) -> Result { self.eval_batch_inner(tokens, true) } fn eval_batch_inner( &mut self, tokens: &[i32], collect_tops: bool, ) -> Result { let rows = u32::try_from(tokens.len()).map_err(|_| "prefill batch is too large")?; if rows == 0 || rows > self.session.prefill_cap { return Err("prefill batch exceeds the configured prefill workspace".into()); } if tokens .iter() .any(|token| *token < 0 || *token as u64 >= self.model.shape.vocab) { return Err("prefill contains a token outside the vocabulary".into()); } let mut batch = BatchScratch::allocate(&self.model, self.session.position, rows, collect_tops)?; if let Some(dspark) = &mut self.dspark { dspark.begin_capture(); } batch.tokens.write_i32(tokens)?; let map = self.model.main.map_ptr().cast(); let size = self.model.main.len(); let shape = self.model.shape; let pos = self.session.position; let mut prefixes = (0..if collect_tops { rows } else { 0 }) .map(|_| SpecPrefixFrontier { layers: (0..shape.layers) .map(|_| LayerFrontier { compression: None, indexer: None, }) .collect(), }) .collect::>(); if let Some(ssd) = &self.ssd { ssd.static_decode_map_current .store(false, std::sync::atomic::Ordering::Release); install_deepseek_model_spans( &self.model, &deepseek_token_model_spans(&self.model)?, "DeepSeek prefill token mapping", )?; } let mut pread = if prefill_pread_enabled() { self.ssd .as_ref() .zip(self.weights.layers.first()) .map(|(ssd, weights)| start_prefill_pread(&self.model, ssd, weights, 0, rows)) .transpose()? } else { None }; let pipelined_verifier = collect_tops && self.ssd.is_none() && self.profile.is_none(); let mut commands = Some(Commands::begin()?); call( unsafe { ds4_gpu_embed_tokens_hc_tensor( batch.current_hc.raw(), batch.tokens.raw(), map, size, self.weights.token_embedding.offset, shape.vocab as u32, rows, shape.embd as u32, shape.hc as u32, ) }, "batch token embedding", )?; if !pipelined_verifier { commands .take() .expect("batch commands are active") .finish()?; } for (index, (weights, state)) in self .weights .layers .iter() .zip(&mut self.session.layers) .enumerate() { let started = Instant::now(); if let Some(job) = pread.take() { job.finish()?; } let layer_selected_addr = self.ssd.is_some() && prefill_selected_addr(rows, shape, weights); if let Some(ssd) = &self.ssd { install_deepseek_model_spans( &self.model, &deepseek_layer_model_spans( &self.model, weights, index as u32, layer_selected_addr, ssd.per_expert_bytes, )?, "DeepSeek prefill layer mapping", )?; if prefill_pread_enabled() && let Some(next) = self.weights.layers.get(index + 1) { pread = Some(start_prefill_pread( &self.model, ssd, next, index as u32 + 1, rows, )?); } } if !pipelined_verifier { commands = Some(Commands::begin()?); } encode_batch_layer( &batch, state, weights, shape, map, size, index as u32, pos, rows, self.session.raw_cap, self.steering.as_ref(), collect_tops.then_some(prefixes.as_mut_slice()), )?; if let Some(profile) = &mut self.profile { profile.record( index, pos, &batch.router_selected, &batch.router_weights, rows, index < shape.hash_layers as usize, )?; } if let Some(dspark) = &mut self.dspark { dspark.capture_batch( index as u32, &batch.next_hc, rows, self.session.prefill_cap, shape, )?; } let seeded_from_map = match &self.ssd { Some(ssd) => ssd.seed_mapped_layer(&self.model, weights, index, true)?, None => true, }; if pipelined_verifier { if (index + 1) % 4 == 0 { commands .as_mut() .expect("batch commands are active") .flush()?; } } else { commands .take() .expect("batch commands are active") .finish()?; } if !seeded_from_map { let seeded = self .ssd .as_ref() .expect("SSD preload fallback lost its plan") .seed_mapped_layer(&self.model, weights, index, false)?; if !seeded { return Err(format!( "Metal could not preload SSD experts for layer {index}" )); } } throttle( &mut self.prefill_layer_average[index], started.elapsed(), self.power_percent, ); std::mem::swap(&mut batch.current_hc, &mut batch.next_hc); } if pipelined_verifier { commands .take() .expect("batch commands are active") .finish()?; } if self.dspark.is_some() && self.ssd.is_some() { install_speculative_model_maps(&self.model, "DSpark prefill support mapping")?; } if let (Some(dspark), Some(support)) = (&mut self.dspark, self.model.support.as_ref()) { dspark.seed_batch_cache( support, pos, rows, self.session.prefill_cap, self.session.raw_cap, shape, )?; } if self.ssd.is_some() { install_deepseek_model_spans( &self.model, &deepseek_output_model_spans(&self.model)?, "DeepSeek prefill output mapping", )?; } let (tops, output_logits) = if collect_tops { commands = Some(Commands::begin()?); encode_batch_output(&batch, &self.weights, shape, map, size, rows)?; commands .take() .expect("batch commands are active") .finish()?; let logits = batch .output_logits .as_ref() .expect("batch output logits are allocated"); let mut all_logits = vec![0.0; (u64::from(rows) * shape.vocab) as usize]; logits.read_f32(&mut all_logits)?; let output_logits = all_logits .chunks_exact(shape.vocab as usize) .map(<[f32]>::to_vec) .collect::>(); let tops = output_logits.iter().map(|logits| argmax(logits)).collect(); self.logits .clone_from(output_logits.last().expect("batch has an output row")); (tops, output_logits) } else { let row = rows - 1; let commands = Commands::begin()?; self.session.scratch.current_hc.copy_from( 0, &batch.current_hc, u64::from(row) * shape.hc * shape.embd * 4, shape.hc * shape.embd * 4, "selecting a prefill output row", )?; encode_output(&self.session.scratch, &self.weights, shape, map, size)?; commands.finish()?; self.session.scratch.logits.read_f32(&mut self.logits)?; (vec![argmax(&self.logits)], Vec::new()) }; self.session.position += rows; self.tokens.extend_from_slice(tokens); if let Some(profile) = &self.profile { profile.write()?; } Ok(BatchVerification { tops, logits: output_logits, prefixes, }) } pub(super) fn logits(&self) -> &[f32] { &self.logits } fn execution_stats(&self) -> ExecutionStats { let (speculative_mode, drafted_tokens, accepted_draft_tokens) = self .dspark .as_ref() .map_or((0, 0, 0), |dspark| (2, dspark.drafted, dspark.accepted)); let mut stats = ExecutionStats { speculative_mode, speculative_cycles: self.speculative_cycles, drafted_tokens, accepted_draft_tokens, verifier_passes: self.verifier_passes, verifier_ms: self.verifier_ns / 1_000_000, ..ExecutionStats::default() }; if let Some(ssd) = &self.ssd { let mut native = StreamExpertCacheStats::default(); unsafe { ds4_gpu_stream_expert_cache_get_stats(&mut native) }; let selected = ssd.selected_experts.get(); stats.ssd_enabled = true; stats.ssd_resident_bytes = ssd.resident_bytes; stats.ssd_cache_experts = u64::from(ssd.cache_experts); stats.ssd_cache_entries = u64::from(native.current_count); stats.ssd_cache_bytes = ssd .per_expert_bytes .saturating_mul(u64::from(ssd.cache_experts)); stats.ssd_preloaded_experts = u64::from(ssd.preload_experts); stats.ssd_cache_hits = native.hits; stats.ssd_cache_misses = native.misses; stats.ssd_cache_evictions = native.evictions; stats.ssd_cache_wraps = native.wraps; stats.ssd_buffer_allocs = native.buffer_allocs; stats.ssd_buffer_reuses = native.buffer_reuses; stats.ssd_pread_bytes = native.pread_bytes; stats.ssd_pread_ms = native.pread_ms.max(0.0) as u64; stats.ssd_evict_advise_bytes = native.evict_advise_bytes; stats.ssd_willneed_advise_bytes = native.willneed_advise_bytes; stats.ssd_selected_requests = ssd.selected_requests.get(); stats.ssd_requested_bytes = ssd.per_expert_bytes.saturating_mul(selected); stats.ssd_wait_ms = ssd.selected_wait_ns.get() / 1_000_000; } stats } pub(super) fn model(&self) -> &Model { &self.model } pub(super) fn context(&self) -> u32 { self.session.context } pub(super) fn position(&self) -> u32 { self.session.position } pub(super) fn reset(&mut self) -> Result<(), String> { self.session = Session::new(&self.model, self.session.context, self.session.prefill_cap)?; if let Some(dspark) = &mut self.dspark { dspark.capture_mask = 0; dspark.cache_start = 0; dspark.cache_len = 0; } self.tokens.clear(); self.checkpoint_tag = [0; 32]; Ok(()) } fn blank_resident_state(&self) -> Result { let session = Session::new(&self.model, self.session.context, self.session.prefill_cap)?; let dspark = match ( self.model.support_kind, self.model.support.as_ref(), self.speculative.dspark, ) { (Some(SupportKind::DSpark), Some(support), true) => Some(Dspark::new( &self.model, support, &session, self.speculative, self.quality, )?), _ => None, }; Ok(DeepSeekResidentState { session, dspark, logits: vec![0.0; self.model.shape.vocab as usize], tokens: Vec::new(), 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_state(), Ok)?; std::mem::swap(&mut self.session, &mut incoming.session); std::mem::swap(&mut self.dspark, &mut incoming.dspark); std::mem::swap(&mut self.logits, &mut incoming.logits); std::mem::swap(&mut self.tokens, &mut incoming.tokens); std::mem::swap(&mut self.checkpoint_tag, &mut incoming.checkpoint_tag); *state = Some(incoming); Ok(()) } pub(super) fn align_prompt(&mut self, tokens: &[i32]) -> Result { if !tokens.starts_with(&self.tokens) { self.reset()?; } Ok(self.tokens.len()) } pub(super) fn tokens(&self) -> &[i32] { &self.tokens } pub(super) fn checkpoint_tag(&self) -> [u8; 32] { self.checkpoint_tag } /// Marks what the live KV state represents when no checkpoint file is /// written for it, so the next turn can still continue in memory. pub(super) fn note_checkpoint_tag(&mut self, tag: [u8; 32]) { self.checkpoint_tag = tag; } fn encode_token(&mut self, token: u32) -> Result<(), String> { self.encode_token_embedding(token)?; for index in 0..self.weights.layers.len() { self.encode_token_layer(index, token)?; if index == 3 && index + 1 < self.weights.layers.len() { call( unsafe { ds4_gpu_flush_commands() }, "flushing the DeepSeek decode graph", )?; } } encode_output( &self.session.scratch, &self.weights, self.model.shape, self.model.main.map_ptr().cast(), self.model.main.len(), ) } fn encode_streaming_token(&mut self, token: u32) -> Result<(), String> { let ssd = self .ssd .as_ref() .expect("streaming token lost its SSD plan"); if !ssd .static_decode_map_current .load(std::sync::atomic::Ordering::Acquire) { install_deepseek_model_spans( &self.model, &DeepSeekModelSpans { ranges: ssd.model_spans.clone(), max_tensor_bytes: self.model.main.max_tensor_bytes(), }, "DeepSeek static decode mapping", )?; ssd.static_decode_map_current .store(true, std::sync::atomic::Ordering::Release); } let layer_batch = self.profile.is_none(); let mut commands = Some(Commands::begin()?); self.encode_token_embedding(token)?; if !layer_batch { commands .take() .expect("decode command batch disappeared") .finish()?; } for index in 0..self.weights.layers.len() { if !layer_batch { commands = Some(Commands::begin()?); } self.encode_token_layer(index, token)?; if !layer_batch { commands .take() .expect("decode command batch disappeared") .finish()?; } } if !layer_batch { commands = Some(Commands::begin()?); } encode_output( &self.session.scratch, &self.weights, self.model.shape, self.model.main.map_ptr().cast(), self.model.main.len(), )?; commands .take() .expect("decode command batch disappeared") .finish() } fn encode_token_embedding(&mut self, token: u32) -> Result<(), String> { let shape = self.model.shape; let map = self.model.main.map_ptr().cast(); let size = self.model.main.len(); let scratch = &mut self.session.scratch; call( unsafe { ds4_gpu_embed_token_hc_tensor( scratch.current_hc.raw(), map, size, self.weights.token_embedding.offset, shape.vocab as u32, token, shape.embd as u32, shape.hc as u32, ) }, "token embedding", ) } fn encode_token_layer(&mut self, index: usize, token: u32) -> Result<(), String> { let shape = self.model.shape; let scratch = &mut self.session.scratch; encode_layer( scratch, &mut self.session.layers[index], &self.weights.layers[index], shape, self.model.main.map_ptr().cast(), self.model.main.len(), index as u32, self.session.position, token, self.session.raw_cap, self.steering.as_ref(), self.ssd.as_ref(), self.quality, self.profile.is_some(), )?; if let Some(profile) = &mut self.profile { profile.record( index, self.session.position, &scratch.router_selected, &scratch.router_weights, 1, index < shape.hash_layers as usize, )?; } if let Some(dspark) = &mut self.dspark { dspark.capture_decode(index as u32, &scratch.current_hc, shape)?; } Ok(()) } } /// Model-family dispatch over the two Rust-owned Metal graphs. pub(super) enum Executor { DeepSeek(Box), Glm(Box), } pub(super) enum ResidentState { DeepSeek(Box), Glm(Box), } impl Executor { #[allow(dead_code)] pub(super) fn open( model: Model, context: u32, quality: bool, prefill_chunk: u32, ) -> Result { Self::open_configured( model, context, quality, prefill_chunk, 100, EngineSpeculativeSettings { glm_mtp: false, glm_mtp_timing: false, dspark: false, dspark_confidence_threshold: 0.9, dspark_confidence_threshold_set: false, dspark_strict: false, dspark_exact_sampling: false, }, crate::settings::EngineSsdSettings { enabled: false, cold: false, cache_experts: 0, cache_bytes: 0, full_layers: 0, full_layers_set: false, preload_experts: 0, }, EngineSteeringSettings { file: None, ffn_scale: 0.0, attention_scale: 0.0, }, None, ) } #[allow(clippy::too_many_arguments)] pub(super) fn open_configured( model: Model, context: u32, quality: bool, prefill_chunk: u32, power_percent: u8, speculative: EngineSpeculativeSettings, ssd: EngineSsdSettings, steering: EngineSteeringSettings, expert_profile_path: Option<&str>, ) -> Result { match model.shape.family { ModelFamily::DeepSeek => DeepSeekExecutor::open_profile( model, context, quality, prefill_chunk, power_percent, speculative, ssd, steering, expert_profile_path, ) .map(Box::new) .map(Self::DeepSeek), ModelFamily::Glm => GlmExecutor::open_profile( model, context, quality, ssd, speculative, expert_profile_path, ) .map(Box::new) .map(Self::Glm), } } #[allow(clippy::too_many_arguments)] pub(super) fn eval_speculative_sampled( &mut self, token: i32, max_tokens: u32, reasoning: ReasoningMode, temperature: f32, top_p: f32, min_p: f32, top_k: i32, rng: &mut Rng, cancelled: &std::sync::atomic::AtomicBool, ) -> Result, String> { match self { Self::DeepSeek(executor) => executor.eval_speculative_sampled( token, max_tokens, reasoning, temperature, top_p, min_p, top_k, rng, cancelled, ), Self::Glm(executor) => { executor.eval(token)?; Ok(vec![token]) } } } pub(super) fn eval(&mut self, token: i32) -> Result<(), String> { match self { Self::DeepSeek(executor) => executor.eval(token), Self::Glm(executor) => executor.eval(token), } } pub(super) fn eval_speculative_greedy( &mut self, token: i32, max_tokens: u32, reasoning: ReasoningMode, cancelled: &std::sync::atomic::AtomicBool, ) -> Result, String> { match self { Self::DeepSeek(executor) => { executor.eval_speculative_greedy(token, max_tokens, reasoning, cancelled) } Self::Glm(executor) => { let _ = reasoning; executor.eval_speculative_greedy(token, max_tokens, cancelled) } } } pub(super) fn prefill( &mut self, tokens: &[i32], progress: impl FnMut(u32) -> bool, ) -> Result { match self { Self::DeepSeek(executor) => executor.prefill(tokens, progress), Self::Glm(executor) => executor.prefill(tokens, progress), } } pub(super) fn logits(&self) -> &[f32] { match self { Self::DeepSeek(executor) => executor.logits(), Self::Glm(executor) => executor.logits(), } } pub(super) fn execution_stats(&self) -> ExecutionStats { match self { Self::DeepSeek(executor) => executor.execution_stats(), Self::Glm(executor) => executor.execution_stats(), } } pub(super) fn model(&self) -> &Model { match self { Self::DeepSeek(executor) => executor.model(), Self::Glm(executor) => executor.model(), } } pub(super) fn context(&self) -> u32 { match self { Self::DeepSeek(executor) => executor.context(), Self::Glm(executor) => executor.context(), } } pub(super) fn position(&self) -> u32 { match self { Self::DeepSeek(executor) => executor.position(), Self::Glm(executor) => executor.position(), } } pub(super) fn reset(&mut self) -> Result<(), String> { match self { Self::DeepSeek(executor) => executor.reset(), Self::Glm(executor) => executor.reset(), } } pub(super) fn swap_resident_state( &mut self, state: &mut Option, ) -> Result<(), String> { match self { Self::DeepSeek(executor) => { let mut inner = match state.take() { Some(ResidentState::DeepSeek(state)) => Some(*state), Some(ResidentState::Glm(_)) => { return Err("resident session belongs to a different model family".into()); } None => None, }; executor.swap_resident_state(&mut inner)?; *state = inner.map(|state| ResidentState::DeepSeek(Box::new(state))); } Self::Glm(executor) => { let mut inner = match state.take() { Some(ResidentState::Glm(state)) => Some(*state), Some(ResidentState::DeepSeek(_)) => { return Err("resident session belongs to a different model family".into()); } None => None, }; executor.swap_resident_state(&mut inner)?; *state = inner.map(|state| ResidentState::Glm(Box::new(state))); } } Ok(()) } pub(super) fn align_prompt(&mut self, tokens: &[i32]) -> Result { match self { Self::DeepSeek(executor) => executor.align_prompt(tokens), Self::Glm(executor) => executor.align_prompt(tokens), } } pub(super) fn tokens(&self) -> &[i32] { match self { Self::DeepSeek(executor) => executor.tokens(), Self::Glm(executor) => executor.tokens(), } } pub(super) fn checkpoint_tag(&self) -> [u8; 32] { match self { Self::DeepSeek(executor) => executor.checkpoint_tag(), Self::Glm(executor) => executor.checkpoint_tag(), } } pub(super) fn note_checkpoint_tag(&mut self, tag: [u8; 32]) { match self { Self::DeepSeek(executor) => executor.note_checkpoint_tag(tag), Self::Glm(executor) => executor.note_checkpoint_tag(tag), } } } #[allow(clippy::too_many_arguments)] fn refresh_ratio4_compressor_state( s: &BatchScratch, state: &CompressionState, weights: CompressorWeights, shape: super::Shape, map: *const c_void, size: u64, pos: u32, rows: u32, head_dim: u32, ) -> Result<(), String> { if rows < 4 { return Ok(()); } let width = 2 * u64::from(head_dim); let tail = s .norm .view(u64::from(rows - 4) * shape.embd * 4, 4 * shape.embd * 4)?; f16_rows( &s.compressed_kv, weights.kv, shape.embd, width, &tail, 4, map, size, )?; f16_rows( &s.compressed_score, weights.gate, shape.embd, width, &tail, 4, map, size, )?; call( unsafe { ds4_gpu_compressor_prefill_state_ratio4_tensor( state.state_kv.raw(), state.state_score.raw(), s.compressed_kv.raw(), s.compressed_score.raw(), map, size, weights.ape.offset, weights.ape.kind, head_dim, pos + rows - 4, ) }, "refreshing ratio-4 compressor state", ) } #[allow(clippy::too_many_arguments)] fn compress_attention_batch( s: &BatchScratch, state: &mut CompressionState, weights: CompressorWeights, shape: super::Shape, map: *const c_void, size: u64, pos: u32, rows: u32, original: u32, freq_base: f32, freq_scale: f32, ext: f32, attn_factor: f32, mut prefixes: Option<&mut [SpecPrefixFrontier]>, layer: usize, ) -> Result { let ratio = state.ratio; let chunk = rows / ratio; if prefixes.is_none() && (pos == 0 || (pos.is_multiple_of(ratio) && rows.is_multiple_of(ratio))) { let before = if pos == 0 { 0 } else { state.rows }; let target = s .compressed_stage .view(0, u64::from(chunk) * shape.head_dim * 4)?; let result = if pos == 0 || ratio != 4 { unsafe { ds4_gpu_compressor_prefill_tensor( target.raw(), state.state_kv.raw(), state.state_score.raw(), s.compressed_kv.raw(), s.compressed_score.raw(), map, size, weights.ape.offset, weights.ape.kind, weights.norm.offset, weights.norm.kind, shape.head_dim as u32, ratio, pos, rows, shape.rot as u32, original, true, freq_base, freq_scale, ext, attn_factor, shape.rope_beta_fast, shape.rope_beta_slow, shape.rms_epsilon, ) } } else { unsafe { ds4_gpu_compressor_prefill_ratio4_replay_tensor( target.raw(), state.state_kv.raw(), state.state_score.raw(), s.compressed_kv.raw(), s.compressed_score.raw(), map, size, weights.ape.offset, weights.ape.kind, weights.norm.offset, weights.norm.kind, shape.head_dim as u32, pos, rows, shape.rot as u32, original, true, freq_base, freq_scale, ext, attn_factor, shape.rope_beta_fast, shape.rope_beta_slow, shape.rms_epsilon, ) } }; call(result, "batch attention compression")?; if chunk != 0 { call( unsafe { ds4_gpu_tensor_copy_f32_to_f16( state.cache.raw(), u64::from(before) * shape.head_dim * 2, target.raw(), 0, u64::from(chunk) * shape.head_dim, ) }, "committing compressed attention KV", )?; } state.rows = before + chunk; if ratio == 4 { refresh_ratio4_compressor_state( s, state, weights, shape, map, size, pos, rows, shape.head_dim as u32, )?; } } else { let width = if ratio == 4 { 2 } else { 1 } * shape.head_dim; let target = s.compressed_stage.view(0, shape.head_dim * 4)?; for row in 0..rows { let absolute = pos + row; let emit = (absolute + 1).is_multiple_of(ratio); let kv = s .compressed_kv .view(u64::from(row) * width * 4, width * 4)?; let score = s .compressed_score .view(u64::from(row) * width * 4, width * 4)?; call( unsafe { ds4_gpu_compressor_update_tensor( kv.raw(), score.raw(), state.state_kv.raw(), state.state_score.raw(), target.raw(), map, size, weights.ape.offset, weights.ape.kind, weights.norm.offset, weights.norm.kind, shape.head_dim as u32, ratio, absolute, 0, shape.rot as u32, original, freq_base, freq_scale, ext, attn_factor, shape.rope_beta_fast, shape.rope_beta_slow, shape.rms_epsilon, false, false, false, ) }, "updating batched attention compression", )?; if emit { call( unsafe { ds4_gpu_dsv4_fp8_kv_quantize_tensor( target.raw(), 1, shape.head_dim as u32, shape.rot as u32, ) }, "rounding compressed attention KV", )?; call( unsafe { ds4_gpu_tensor_copy_f32_to_f16( state.cache.raw(), u64::from(state.rows) * shape.head_dim * 2, target.raw(), 0, shape.head_dim, ) }, "committing compressed attention KV", )?; state.rows += 1; } if let Some(prefixes) = prefixes.as_deref_mut() { let coefficient = if ratio == 4 { 2 } else { 1 }; prefixes[row as usize].layers[layer].compression = Some(capture_compression_frontier( state, coefficient * coefficient * u64::from(ratio) * shape.head_dim * 4, "capturing speculative compressor prefix", )?); } } } Ok(state.rows) } #[allow(clippy::too_many_arguments)] fn compress_index_batch( s: &BatchScratch, state: &mut CompressionState, weights: CompressorWeights, shape: super::Shape, map: *const c_void, size: u64, pos: u32, rows: u32, original: u32, freq_base: f32, freq_scale: f32, ext: f32, attn_factor: f32, mut prefixes: Option<&mut [SpecPrefixFrontier]>, layer: usize, ) -> Result<(), String> { let ratio = state.ratio; if prefixes.is_none() && (pos == 0 || (pos.is_multiple_of(ratio) && rows.is_multiple_of(ratio))) { let before = if pos == 0 { 0 } else { state.rows }; let chunk = rows / ratio; let target = state.cache.view( u64::from(before) * shape.indexer_head_dim * 4, u64::from(chunk) * shape.indexer_head_dim * 4, )?; let result = if pos == 0 { unsafe { ds4_gpu_compressor_prefill_tensor( target.raw(), state.state_kv.raw(), state.state_score.raw(), s.compressed_kv.raw(), s.compressed_score.raw(), map, size, weights.ape.offset, weights.ape.kind, weights.norm.offset, weights.norm.kind, shape.indexer_head_dim as u32, ratio, pos, rows, shape.rot as u32, original, false, freq_base, freq_scale, ext, attn_factor, shape.rope_beta_fast, shape.rope_beta_slow, shape.rms_epsilon, ) } } else { unsafe { ds4_gpu_compressor_prefill_ratio4_replay_tensor( target.raw(), state.state_kv.raw(), state.state_score.raw(), s.compressed_kv.raw(), s.compressed_score.raw(), map, size, weights.ape.offset, weights.ape.kind, weights.norm.offset, weights.norm.kind, shape.indexer_head_dim as u32, pos, rows, shape.rot as u32, original, false, freq_base, freq_scale, ext, attn_factor, shape.rope_beta_fast, shape.rope_beta_slow, shape.rms_epsilon, ) } }; call(result, "batch indexer compression")?; if chunk != 0 { call( unsafe { ds4_gpu_dsv4_indexer_qat_tensor( target.raw(), chunk, shape.indexer_head_dim as u32, ) }, "batch compressed index rounding", )?; } state.rows = before + chunk; refresh_ratio4_compressor_state( s, state, weights, shape, map, size, pos, rows, shape.indexer_head_dim as u32, )?; } else { let width = 2 * shape.indexer_head_dim; for row in 0..rows { let absolute = pos + row; let emit = (absolute + 1).is_multiple_of(ratio); let kv = s .compressed_kv .view(u64::from(row) * width * 4, width * 4)?; let score = s .compressed_score .view(u64::from(row) * width * 4, width * 4)?; let target = state.cache.view( u64::from(state.rows) * shape.indexer_head_dim * 4, shape.indexer_head_dim * 4, )?; call( unsafe { ds4_gpu_compressor_update_tensor( kv.raw(), score.raw(), state.state_kv.raw(), state.state_score.raw(), target.raw(), map, size, weights.ape.offset, weights.ape.kind, weights.norm.offset, weights.norm.kind, shape.indexer_head_dim as u32, ratio, absolute, 0, shape.rot as u32, original, freq_base, freq_scale, ext, attn_factor, shape.rope_beta_fast, shape.rope_beta_slow, shape.rms_epsilon, false, false, false, ) }, "updating batched indexer compression", )?; if emit { call( unsafe { ds4_gpu_dsv4_indexer_qat_tensor( target.raw(), 1, shape.indexer_head_dim as u32, ) }, "rounding compressed index", )?; state.rows += 1; } if let Some(prefixes) = prefixes.as_deref_mut() { prefixes[row as usize].layers[layer].indexer = Some(capture_compression_frontier( state, 4 * u64::from(ratio) * shape.indexer_head_dim * 4, "capturing speculative indexer prefix", )?); } } } Ok(()) } #[allow(clippy::too_many_arguments)] fn encode_batch_layer( s: &BatchScratch, state: &mut LayerState, w: &Layer, shape: super::Shape, map: *const c_void, size: u64, layer: u32, pos: u32, rows: u32, raw_cap: u32, steering: Option<&Steering>, mut prefixes: Option<&mut [SpecPrefixFrontier]>, ) -> Result<(), String> { let hc_dim = shape.hc * shape.embd; let mix_hc = 2 * shape.hc + shape.hc * shape.hc; let q_dim = shape.heads * shape.head_dim; let ratio = compression_ratio(shape, layer); let compressed = ratio != 0; let freq_base = if compressed { shape.compress_rope_base } else { shape.rope_base }; let freq_scale = if compressed { 1.0 / shape.rope_scale } else { 1.0 }; let ext = if compressed && shape.rope_scale > 1.0 { 1.0 } else { 0.0 }; let attn_factor = if ext != 0.0 { 1.0 / (1.0 + 0.1 * (1.0 / freq_scale).ln()) } else { 1.0 }; let original = if compressed { shape.original_context as u32 } else { 0 }; call( unsafe { ds4_gpu_hc_rms_scale_project_f16_tensor( s.hc_mix.raw(), s.flat_hc.raw(), map, size, w.hc_attn_fn.offset, hc_dim as u32, mix_hc as u32, s.current_hc.raw(), rows, shape.rms_epsilon, ) }, "batch attention HC projection", )?; call( unsafe { ds4_gpu_hc_split_weighted_sum_norm_tensor( s.current.raw(), s.norm.raw(), s.hc_split.raw(), s.hc_mix.raw(), s.current_hc.raw(), map, size, w.hc_attn_scale.offset, w.hc_attn_base.offset, w.attn_norm.offset, shape.embd as u32, shape.hc as u32, shape.hc_sinkhorn as u32, shape.hc_epsilon, shape.rms_epsilon, ) }, "batch attention HC mix", )?; q8_rows( &s.q_rank, w.attn_q_a, shape.embd, shape.lora_q, &s.norm, rows, map, size, )?; q8_rows( &s.kv_raw, w.attn_kv, shape.embd, shape.head_dim, &s.norm, rows, map, size, )?; call( unsafe { ds4_gpu_dsv4_qkv_rms_norm_rows_tensor( s.q_rank_norm.raw(), s.q_rank.raw(), map, size, w.attn_q_a_norm.offset, shape.lora_q as u32, s.kv.raw(), s.kv_raw.raw(), w.attn_kv_norm.offset, shape.head_dim as u32, rows, shape.rms_epsilon, ) }, "batch Q/KV norm", )?; let fused_q = unsafe { ds4_gpu_attn_q_b_f16_head_rms_rope_tail_tensor( s.q.raw(), s.q_half.raw(), map, size, w.attn_q_b.offset, shape.lora_q, q_dim, s.q_rank_norm.raw(), rows, shape.heads as u32, shape.head_dim as u32, shape.rot as u32, pos, original, false, freq_base, freq_scale, ext, attn_factor, shape.rope_beta_fast, shape.rope_beta_slow, shape.rms_epsilon, ) } != 0; if !fused_q { q8_rows( &s.q, w.attn_q_b, shape.lora_q, q_dim, &s.q_rank_norm, rows, map, size, )?; call( unsafe { ds4_gpu_head_rms_norm_tensor( s.q.raw(), rows, shape.heads as u32, shape.head_dim as u32, shape.rms_epsilon, ) }, "batch Q norm", )?; call( unsafe { ds4_gpu_rope_tail_tensor( s.q.raw(), rows, shape.heads as u32, shape.head_dim as u32, shape.rot as u32, pos, original, false, freq_base, freq_scale, ext, attn_factor, shape.rope_beta_fast, shape.rope_beta_slow, ) }, "batch Q RoPE", )?; } call( unsafe { ds4_gpu_rope_tail_tensor( s.kv.raw(), rows, shape.head_kv as u32, shape.head_dim as u32, shape.rot as u32, pos, original, false, freq_base, freq_scale, ext, attn_factor, shape.rope_beta_fast, shape.rope_beta_slow, ) }, "batch KV RoPE", )?; call( unsafe { ds4_gpu_dsv4_fp8_kv_quantize_tensor( s.kv.raw(), rows, shape.head_dim as u32, shape.rot as u32, ) }, "batch KV rounding", )?; call( unsafe { ds4_gpu_store_raw_kv_batch_tensor( state.raw_cache.raw(), s.kv.raw(), raw_cap, pos, rows, shape.head_dim as u32, ) }, "batch raw KV store", )?; let mut compressed_rows = 0; if let (Some(weights), Some(compression)) = (w.attn_compressor.as_ref(), state.compression.as_mut()) { let width = if ratio == 4 { 2 } else { 1 } * shape.head_dim; f16_rows( &s.compressed_kv, weights.kv, shape.embd, width, &s.norm, rows, map, size, )?; f16_rows( &s.compressed_score, weights.gate, shape.embd, width, &s.norm, rows, map, size, )?; compressed_rows = compress_attention_batch( s, compression, *weights, shape, map, size, pos, rows, original, freq_base, freq_scale, ext, attn_factor, prefixes.as_deref_mut(), layer as usize, )?; } if ratio == 4 { let weights = w .indexer .ok_or("ratio-4 layer is missing indexer weights")?; let indexer = state .indexer .as_mut() .ok_or("ratio-4 layer is missing indexer state")?; let width = 2 * shape.indexer_head_dim; f16_rows( &s.compressed_kv, weights.compressor.kv, shape.embd, width, &s.norm, rows, map, size, )?; f16_rows( &s.compressed_score, weights.compressor.gate, shape.embd, width, &s.norm, rows, map, size, )?; compress_index_batch( s, indexer, weights.compressor, shape, map, size, pos, rows, original, freq_base, freq_scale, ext, attn_factor, prefixes, layer as usize, )?; matmul_rows( &s.indexer_q, weights.q, shape.lora_q, shape.indexer_heads * shape.indexer_head_dim, &s.q_rank_norm, rows, map, size, )?; call( unsafe { ds4_gpu_rope_tail_tensor( s.indexer_q.raw(), rows, shape.indexer_heads as u32, shape.indexer_head_dim as u32, shape.rot as u32, pos, original, false, freq_base, freq_scale, ext, attn_factor, shape.rope_beta_fast, shape.rope_beta_slow, ) }, "batch indexer Q RoPE", )?; call( unsafe { ds4_gpu_dsv4_indexer_qat_tensor( s.indexer_q.raw(), rows * shape.indexer_heads as u32, shape.indexer_head_dim as u32, ) }, "batch indexer Q rounding", )?; f16_rows( &s.indexer_weights, weights.proj, shape.embd, shape.indexer_heads, &s.norm, rows, map, size, )?; } if pos == 0 && compressed_rows == 0 { call( unsafe { ds4_gpu_attention_prefill_raw_heads_tensor( s.heads.raw(), map, size, w.attn_sinks.offset, s.q.raw(), s.kv.raw(), rows, shape.sliding_window as u32, shape.heads as u32, shape.head_dim as u32, ) }, "batch raw prefill attention", )?; } else if pos == 0 && ratio == 4 && compressed_rows > shape.indexer_top_k as u32 { let indexer = state .indexer .as_ref() .ok_or("ratio-4 layer is missing indexer state")?; call( unsafe { ds4_gpu_indexer_scores_prefill_tensor( s.indexer_scores.raw(), s.indexer_q.raw(), s.indexer_weights.raw(), indexer.cache.raw(), compressed_rows, rows, shape.indexer_heads as u32, shape.indexer_head_dim as u32, ratio, ((shape.indexer_heads * shape.indexer_head_dim) as f32) .sqrt() .recip(), ) }, "batch indexer scoring", )?; call( unsafe { ds4_gpu_indexer_topk_tensor( s.indexer_selected.raw(), s.indexer_scores.raw(), compressed_rows, rows, shape.indexer_top_k as u32, ) }, "batch indexer top-k", )?; let compression = state .compression .as_ref() .ok_or("compressed layer is missing attention state")?; call( unsafe { ds4_gpu_attention_indexed_mixed_batch_heads_tensor( s.heads.raw(), map, size, w.attn_sinks.offset, s.q.raw(), state.raw_cache.raw(), compression.cache.raw(), 1, s.indexer_selected.raw(), rows, pos, rows, raw_cap, 0, compressed_rows, shape.indexer_top_k as u32, shape.sliding_window as u32, ratio, shape.heads as u32, shape.head_dim as u32, ) }, "batch indexed prefill attention", )?; } else if pos == 0 { let compression = state .compression .as_ref() .ok_or("compressed layer is missing attention state")?; call( unsafe { ds4_gpu_attention_prefill_static_mixed_heads_tensor( s.heads.raw(), map, size, w.attn_sinks.offset, s.q.raw(), s.kv.raw(), compression.cache.raw(), 1, rows, compressed_rows, shape.sliding_window as u32, ratio, shape.heads as u32, shape.head_dim as u32, ) }, "batch mixed prefill attention", )?; } else { let (n_raw, raw_start) = raw_batch_span(pos, rows, raw_cap, shape.sliding_window as u32); if ratio == 0 { call( unsafe { ds4_gpu_attention_decode_raw_batch_heads_tensor( s.heads.raw(), map, size, w.attn_sinks.offset, s.q.raw(), state.raw_cache.raw(), rows, pos, n_raw, raw_cap, raw_start, shape.sliding_window as u32, shape.heads as u32, shape.head_dim as u32, ) }, "batch resumed raw attention", )?; } else if ratio == 4 && compressed_rows > shape.indexer_top_k as u32 { let indexer = state .indexer .as_ref() .ok_or("ratio-4 layer is missing indexer state")?; call( unsafe { ds4_gpu_indexer_scores_decode_batch_tensor( s.indexer_scores.raw(), s.indexer_q.raw(), s.indexer_weights.raw(), indexer.cache.raw(), compressed_rows, rows, pos, shape.indexer_heads as u32, shape.indexer_head_dim as u32, ratio, ((shape.indexer_heads * shape.indexer_head_dim) as f32) .sqrt() .recip(), ) }, "batch resumed indexer scoring", )?; call( unsafe { ds4_gpu_indexer_topk_tensor( s.indexer_selected.raw(), s.indexer_scores.raw(), compressed_rows, rows, shape.indexer_top_k as u32, ) }, "batch resumed indexer top-k", )?; let compression = state .compression .as_ref() .ok_or("compressed layer is missing attention state")?; call( unsafe { ds4_gpu_attention_indexed_mixed_batch_heads_tensor( s.heads.raw(), map, size, w.attn_sinks.offset, s.q.raw(), state.raw_cache.raw(), compression.cache.raw(), 1, s.indexer_selected.raw(), rows, pos, n_raw, raw_cap, raw_start, compressed_rows, shape.indexer_top_k as u32, shape.sliding_window as u32, ratio, shape.heads as u32, shape.head_dim as u32, ) }, "batch resumed indexed attention", )?; } else { let compression = state .compression .as_ref() .ok_or("compressed layer is missing attention state")?; call( unsafe { ds4_gpu_attention_decode_mixed_batch_heads_tensor( s.heads.raw(), map, size, w.attn_sinks.offset, s.q.raw(), state.raw_cache.raw(), compression.cache.raw(), 1, std::ptr::null(), 0, rows, pos, n_raw, raw_cap, raw_start, compressed_rows, shape.sliding_window as u32, ratio, shape.heads as u32, shape.head_dim as u32, ) }, "batch resumed mixed attention", )?; } } call( unsafe { ds4_gpu_rope_tail_tensor( s.heads.raw(), rows, shape.heads as u32, shape.head_dim as u32, shape.rot as u32, pos, original, true, freq_base, freq_scale, ext, attn_factor, shape.rope_beta_fast, shape.rope_beta_slow, ) }, "batch inverse attention RoPE", )?; let group_dim = shape.head_dim * (shape.heads / shape.out_groups); let attention_steering = steering.is_some_and(|value| value.attention_scale != 0.0); let half_output = !attention_steering && unsafe { ds4_gpu_attention_output_q8_batch_f16_tensor( s.q_half.raw(), s.attention_low.raw(), map, size, w.attn_output_a.offset, w.attn_output_b.offset, group_dim, shape.lora_o, shape.out_groups as u32, shape.embd, s.heads.raw(), rows, ) } != 0; if half_output { call( unsafe { ds4_gpu_hc_expand_split_half_tensor( s.after_attention_hc.raw(), s.q_half.raw(), s.current_hc.raw(), s.hc_split.raw(), shape.embd as u32, shape.hc as u32, ) }, "batch attention HC expansion", )?; } else { call( unsafe { ds4_gpu_attention_output_q8_batch_tensor( s.attention_out.raw(), s.attention_low.raw(), s.attention_group_tmp.raw(), s.attention_low_tmp.raw(), map, size, w.attn_output_a.offset, w.attn_output_b.offset, group_dim, shape.lora_o, shape.out_groups as u32, shape.embd, s.heads.raw(), rows, ) }, "batch attention output", )?; if let Some(steering) = steering { steering.apply(&s.attention_out, layer, rows, true)?; } call( unsafe { ds4_gpu_hc_expand_split_tensor( s.after_attention_hc.raw(), s.attention_out.raw(), s.current_hc.raw(), s.hc_split.raw(), shape.embd as u32, shape.hc as u32, ) }, "batch attention HC expansion", )?; } call( unsafe { ds4_gpu_hc_rms_scale_project_f16_tensor( s.hc_mix.raw(), s.flat_hc.raw(), map, size, w.hc_ffn_fn.offset, hc_dim as u32, mix_hc as u32, s.after_attention_hc.raw(), rows, shape.rms_epsilon, ) }, "batch FFN HC projection", )?; call( unsafe { ds4_gpu_hc_split_weighted_sum_norm_tensor( s.current.raw(), s.norm.raw(), s.hc_split.raw(), s.hc_mix.raw(), s.after_attention_hc.raw(), map, size, w.hc_ffn_scale.offset, w.hc_ffn_base.offset, w.ffn_norm.offset, shape.embd as u32, shape.hc as u32, shape.hc_sinkhorn as u32, shape.hc_epsilon, shape.rms_epsilon, ) }, "batch FFN HC mix", )?; f16_rows( &s.router_logits, w.router, shape.embd, shape.experts, &s.norm, rows, map, size, )?; call( unsafe { ds4_gpu_router_select_batch_tensor( s.router_selected.raw(), s.router_weights.raw(), s.router_probs.raw(), map, size, w.router_bias.map_or(0, |value| value.offset), w.router_hash.map_or(0, |value| value.offset), w.router_hash.map_or(0, |value| value.dims[1] as u32), 0, 0, w.router_bias.is_some(), w.router_hash.is_some(), s.router_logits.raw(), s.tokens.raw(), shape.experts as u32, shape.experts_used as u32, shape.expert_weight_scale, rows, ) }, "batch expert routing", )?; let gate_row = w.expert_gate.bytes / (w.expert_gate.dims[1] * w.expert_gate.dims[2]); let down_row = w.expert_down.bytes / (w.expert_down.dims[1] * w.expert_down.dims[2]); let mut mid_f16 = false; call( unsafe { ds4_gpu_routed_moe_batch_tensor( s.routed_out.raw(), s.routed_gate.raw(), s.routed_up.raw(), s.routed_mid.raw(), s.routed_experts.raw(), map, size, w.expert_gate.offset, w.expert_up.offset, w.expert_down.offset, w.expert_gate.kind, w.expert_down.kind, w.expert_gate.dims[1] * gate_row, gate_row, w.expert_down.dims[1] * down_row, down_row, w.expert_gate.dims[0] as u32, w.expert_down.dims[0] as u32, w.expert_down.dims[1] as u32, s.router_selected.raw(), s.router_weights.raw(), shape.experts as u32, shape.experts_used as u32, shape.swiglu_clamp, s.norm.raw(), layer, rows, &mut mid_f16, false, ) }, "batch routed experts", )?; q8_rows( &s.shared_gate, w.shared_gate, shape.embd, shape.ff_expert, &s.norm, rows, map, size, )?; q8_rows( &s.shared_up, w.shared_up, shape.embd, shape.ff_expert, &s.norm, rows, map, size, )?; call( unsafe { ds4_gpu_swiglu_tensor( s.shared_mid.raw(), s.shared_gate.raw(), s.shared_up.raw(), rows * shape.ff_expert as u32, shape.swiglu_clamp, 1.0, ) }, "batch shared expert activation", )?; let ffn_steering = steering.filter(|value| value.ffn_scale != 0.0); let shared_down_f16 = ffn_steering.is_none() && unsafe { ds4_gpu_matmul_q8_0_f16_out_tensor( s.q_half.raw(), map, size, w.shared_down.offset, shape.ff_expert, shape.embd, s.shared_mid.raw(), u64::from(rows), ) } != 0; if !shared_down_f16 { q8_rows( &s.shared_out, w.shared_down, shape.ff_expert, shape.embd, &s.shared_mid, rows, map, size, )?; } if let Some(steering) = ffn_steering { call( unsafe { ds4_gpu_add_tensor( s.attention_out.raw(), s.routed_out.raw(), s.shared_out.raw(), rows * shape.embd as u32, ) }, "combining batch FFN output", )?; steering.apply(&s.attention_out, layer, rows, false)?; call( unsafe { ds4_gpu_hc_expand_split_tensor( s.next_hc.raw(), s.attention_out.raw(), s.after_attention_hc.raw(), s.hc_split.raw(), shape.embd as u32, shape.hc as u32, ) }, "batch steered FFN HC expansion", ) } else if shared_down_f16 { call( unsafe { ds4_gpu_hc_expand_add_split_half_add_tensor( s.next_hc.raw(), s.routed_out.raw(), s.q_half.raw(), s.after_attention_hc.raw(), s.hc_split.raw(), shape.embd as u32, shape.hc as u32, ) }, "batch half shared FFN HC expansion", ) } else { call( unsafe { ds4_gpu_hc_expand_add_split_tensor( s.next_hc.raw(), s.routed_out.raw(), s.shared_out.raw(), s.after_attention_hc.raw(), s.hc_split.raw(), shape.embd as u32, shape.hc as u32, ) }, "batch FFN HC expansion", ) } } // Mirrors one fixed DS4 tape invocation; grouping these scalar dimensions would // only hide the Metal kernel contract behind another one-use type. #[allow(clippy::too_many_arguments)] fn encode_layer( s: &Scratch, state: &mut LayerState, w: &Layer, shape: super::Shape, map: *const c_void, size: u64, layer: u32, pos: u32, token: u32, raw_cap: u32, steering: Option<&Steering>, ssd: Option<&SsdPlan>, quality: bool, profile_active: bool, ) -> Result<(), String> { encode_layer_with_cache_rows( s, state, w, shape, map, size, layer, pos, token, raw_cap, steering, ssd, None, quality, profile_active, ) } fn decode_feature_allowed( pre_m5: bool, m5: bool, pre_m5_disabled: bool, m5_disabled: bool, ) -> bool { (pre_m5 && !pre_m5_disabled) || (m5 && !m5_disabled) } fn decode_feature_enabled(pre_m5_disable: &CStr, m5_disable: &CStr) -> bool { let pre_m5 = unsafe { ds4_gpu_device_is_pre_m5_apple_silicon() } != 0; let m5 = unsafe { ds4_gpu_device_is_m5_apple_silicon() } != 0; decode_feature_allowed( pre_m5, m5, environment_present(pre_m5_disable), environment_present(m5_disable), ) } #[allow(clippy::too_many_arguments)] fn decode_hc_mix( s: &Scratch, residual: &Buffer, mix_weight: Weight, scale: Weight, base: Weight, norm_weight: Weight, shape: super::Shape, map: *const c_void, size: u64, purpose: &'static str, ) -> Result<(), String> { let hc_dim = shape.hc * shape.embd; let mix_hc = 2 * shape.hc + shape.hc * shape.hc; let fused = hc_dim == 16_384 && mix_hc == 24 && mix_weight.kind == F16 && !environment_present(c"DS4_METAL_DISABLE_PRE_M5_HC_NORM_MIX_FUSE") && decode_feature_enabled( c"DS4_METAL_DISABLE_PRE_M5_HC_PRODUCER_PRE_NORM_FUSE", c"DS4_METAL_DISABLE_M5_HC_PRODUCER_PRE_NORM_FUSE", ) && unsafe { ds4_gpu_hc_rms_norm_mix_f16_available() } != 0; if fused { let result = unsafe { ds4_gpu_hc_rms_norm_mix_split_norm_f16_tensor( s.hc_mix.raw(), s.current.raw(), s.norm.raw(), s.hc_split.raw(), residual.raw(), map, size, mix_weight.offset, scale.offset, base.offset, norm_weight.offset, hc_dim as u32, mix_hc as u32, shape.embd as u32, shape.hc as u32, shape.hc_sinkhorn as u32, shape.rms_epsilon, shape.hc_epsilon, shape.rms_epsilon, ) }; if result < 0 { return Err(format!("Metal failed while {purpose}")); } if result > 0 { return Ok(()); } call( unsafe { ds4_gpu_hc_rms_norm_mix_f16_tensor( s.hc_mix.raw(), residual.raw(), map, size, mix_weight.offset, hc_dim as u32, mix_hc as u32, shape.rms_epsilon, ) }, purpose, )?; } else { call( unsafe { ds4_gpu_rms_norm_plain_tensor( s.flat_hc.raw(), residual.raw(), hc_dim as u32, shape.rms_epsilon, ) }, purpose, )?; matmul(&s.hc_mix, mix_weight, hc_dim, mix_hc, &s.flat_hc, map, size)?; } call( unsafe { ds4_gpu_hc_split_weighted_sum_norm_tensor( s.current.raw(), s.norm.raw(), s.hc_split.raw(), s.hc_mix.raw(), residual.raw(), map, size, scale.offset, base.offset, norm_weight.offset, shape.embd as u32, shape.hc as u32, shape.hc_sinkhorn as u32, shape.hc_epsilon, shape.rms_epsilon, ) }, purpose, ) } #[allow(clippy::too_many_arguments)] fn try_qkv_pair_compressor_fusion( s: &Scratch, state: &LayerState, w: &Layer, shape: super::Shape, map: *const c_void, size: u64, layer: u32, pos: u32, resident: bool, ) -> Result { let ratio = compression_ratio(shape, layer); let Some(attn_weights) = w.attn_compressor else { return Ok(false); }; let Some(attn_state) = state.compression.as_ref() else { return Ok(false); }; if !resident || !matches!(ratio, 4 | 128) || w.attn_q_a.kind != Q8_0 || w.attn_kv.kind != Q8_0 || attn_weights.kv.kind != F16 || attn_weights.gate.kind != F16 || attn_weights.kv.dims != [ shape.embd, (if ratio == 4 { 2 } else { 1 }) * shape.head_dim, 1, ] || attn_weights.gate.dims != attn_weights.kv.dims { return Ok(false); } let enabled = if ratio == 4 { decode_feature_enabled( c"DS4_METAL_DISABLE_PRE_M5_QKV_PAIR_QUAD_FUSE", c"DS4_METAL_DISABLE_M5_QKV_PAIR_QUAD_FUSE", ) } else { decode_feature_enabled( c"DS4_METAL_DISABLE_PRE_M5_QKV_PAIR_COMPRESSOR_FUSE", c"DS4_METAL_DISABLE_M5_QKV_PAIR_COMPRESSOR_FUSE", ) }; if !enabled { return Ok(false); } let (out1_kv, out1_score, state1_kv, state1_score, weights1, width1) = if ratio == 4 { let Some(index_weights) = w.indexer.map(|weights| weights.compressor) else { return Ok(false); }; let Some(index_state) = state.indexer.as_ref() else { return Ok(false); }; if index_weights.kv.kind != F16 || index_weights.gate.kind != F16 || index_weights.kv.dims != [shape.embd, 2 * shape.indexer_head_dim, 1] || index_weights.gate.dims != index_weights.kv.dims { return Ok(false); } ( &s.index_compressed_kv, &s.index_compressed_score, &index_state.state_kv, &index_state.state_score, index_weights, (2 * shape.indexer_head_dim) as u32, ) } else { ( &s.compressed_kv, &s.compressed_score, &attn_state.state_kv, &attn_state.state_score, attn_weights, 0, ) }; let result = unsafe { ds4_gpu_qkv_pair_quad_compressor_store_tensor( s.q_rank.raw(), s.kv_raw.raw(), s.compressed_kv.raw(), s.compressed_score.raw(), out1_kv.raw(), out1_score.raw(), attn_state.state_kv.raw(), attn_state.state_score.raw(), state1_kv.raw(), state1_score.raw(), map, size, w.attn_q_a.offset, w.attn_kv.offset, attn_weights.kv.offset, attn_weights.gate.offset, weights1.kv.offset, weights1.gate.offset, attn_weights.ape.offset, attn_weights.ape.kind, weights1.ape.offset, weights1.ape.kind, shape.embd as u32, shape.lora_q as u32, shape.head_dim as u32, (u64::from(if ratio == 4 { 2_u32 } else { 1 }) * shape.head_dim) as u32, width1, s.norm.raw(), ratio, pos, ) }; if result < 0 { Err("Metal failed while fusing Q/KV and compressor projections".into()) } else { Ok(result > 0) } } #[allow(clippy::too_many_arguments)] fn encode_layer_with_cache_rows( s: &Scratch, state: &mut LayerState, w: &Layer, shape: super::Shape, map: *const c_void, size: u64, layer: u32, pos: u32, token: u32, raw_cap: u32, steering: Option<&Steering>, ssd: Option<&SsdPlan>, cache_rows: Option, quality: bool, profile_active: bool, ) -> Result<(), String> { let compressed = layer >= 2; let freq_base = if compressed { shape.compress_rope_base } else { shape.rope_base }; let freq_scale = if compressed { 1.0 / shape.rope_scale } else { 1.0 }; let ext = if compressed && shape.rope_scale > 1.0 { 1.0 } else { 0.0 }; let attn_factor = if ext != 0.0 { 1.0 / (1.0 + 0.1 * (1.0 / freq_scale).ln()) } else { 1.0 }; let original = if compressed { shape.original_context as u32 } else { 0 }; let raw_cache = state.raw_cache.raw(); decode_hc_mix( s, &s.current_hc, w.hc_attn_fn, w.hc_attn_scale, w.hc_attn_base, w.attn_norm, shape, map, size, "attention HC mix", )?; let qkv_compressors_fused = try_qkv_pair_compressor_fusion( s, state, w, shape, map, size, layer, pos, ssd.is_none() && cache_rows.is_none(), )?; if !qkv_compressors_fused { call( unsafe { ds4_gpu_matmul_q8_0_pair_tensor( s.q_rank.raw(), s.kv_raw.raw(), map, size, w.attn_q_a.offset, w.attn_kv.offset, shape.embd, shape.lora_q, shape.head_dim, s.norm.raw(), 1, ) }, "Q/KV projection", )?; } let raw_row = pos % raw_cap; let kv_norm_store_fused = cache_rows.is_none() && shape.head_kv == 1 && shape.head_dim == 512 && shape.rot == 64 && raw_row < raw_cap && !environment_present(c"DS4_METAL_DISABLE_PRE_M5_QKV_NORM_KV_STORE_FUSE") && (unsafe { ds4_gpu_device_is_pre_m5_apple_silicon() } != 0 || unsafe { ds4_gpu_device_is_m5_apple_silicon() } != 0) && unsafe { ds4_gpu_kv_rope_fp8_fuse_available() } != 0; if kv_norm_store_fused { call( unsafe { ds4_gpu_dsv4_qkv_rms_norm_kv_rope_fp8_store_tensor( s.q_rank_norm.raw(), s.q_rank.raw(), map, size, w.attn_q_a_norm.offset, shape.lora_q as u32, s.kv.raw(), s.kv_raw.raw(), w.attn_kv_norm.offset, shape.head_dim as u32, raw_cache, raw_cap as u64, raw_row, shape.rot as u32, pos, original, freq_base, freq_scale, ext, attn_factor, shape.rope_beta_fast, shape.rope_beta_slow, shape.rms_epsilon, ) }, "Q/KV norm, RoPE, and cache write", )?; } else { call( unsafe { ds4_gpu_dsv4_qkv_rms_norm_rows_tensor( s.q_rank_norm.raw(), s.q_rank.raw(), map, size, w.attn_q_a_norm.offset, shape.lora_q as u32, s.kv.raw(), s.kv_raw.raw(), w.attn_kv_norm.offset, shape.head_dim as u32, 1, shape.rms_epsilon, ) }, "Q/KV norm", )?; } q8( &s.q, w.attn_q_b, shape.lora_q, shape.heads * shape.head_dim, &s.q_rank_norm, map, size, )?; let fused_q_rope = unsafe { ds4_gpu_head_rms_norm_rope_tail_tensor( s.q.raw(), 1, shape.heads as u32, shape.head_dim as u32, shape.rot as u32, pos, original, false, freq_base, freq_scale, ext, attn_factor, shape.rope_beta_fast, shape.rope_beta_slow, shape.rms_epsilon, ) } != 0; if !fused_q_rope { call( unsafe { ds4_gpu_head_rms_norm_tensor( s.q.raw(), 1, shape.heads as u32, shape.head_dim as u32, shape.rms_epsilon, ) }, "Q norm", )?; call( unsafe { ds4_gpu_rope_tail_tensor( s.q.raw(), 1, shape.heads as u32, shape.head_dim as u32, shape.rot as u32, pos, original, false, freq_base, freq_scale, ext, attn_factor, shape.rope_beta_fast, shape.rope_beta_slow, ) }, "Q RoPE", )?; } if !kv_norm_store_fused { call( unsafe { ds4_gpu_rope_tail_tensor( s.kv.raw(), 1, shape.head_kv as u32, shape.head_dim as u32, shape.rot as u32, pos, original, false, freq_base, freq_scale, ext, attn_factor, shape.rope_beta_fast, shape.rope_beta_slow, ) }, "KV RoPE", )?; call( unsafe { ds4_gpu_kv_fp8_store_raw_tensor( s.kv.raw(), raw_cache, raw_cap, raw_row, shape.head_dim as u32, shape.rot as u32, ) }, "KV cache write", )?; } let (n_raw, _) = raw_decode_span( cache_rows.unwrap_or(pos), raw_cap, shape.sliding_window as u32, ); let raw_start = (pos + 1 - n_raw) % raw_cap; if let (Some(attn_weights), Some(compression)) = (w.attn_compressor, state.compression.as_mut()) { if let (Some(index_weights), Some(indexer)) = (w.indexer, state.indexer.as_mut()) { update_compression_pair( s, compression, attn_weights, indexer, index_weights.compressor, shape, map, size, pos, original, freq_base, freq_scale, ext, attn_factor, qkv_compressors_fused, cache_rows.is_none(), ssd.is_none() && cache_rows.is_none(), )?; } else { update_compression( s, compression, attn_weights, shape, map, size, pos, original, freq_base, freq_scale, ext, attn_factor, qkv_compressors_fused, )?; } } let (comp_cache, n_comp) = state .compression .as_ref() .map_or((std::ptr::null(), 0), |compression| { (compression.cache.raw().cast_const(), compression.rows) }); let indexed = if n_comp > 1_024 { match (w.indexer, state.indexer.as_ref()) { (Some(weights), Some(indexer)) if indexer.rows > shape.indexer_top_k as u32 => { matmul( &s.indexer_q, weights.q, shape.lora_q, shape.indexer_heads * shape.indexer_head_dim, &s.q_rank_norm, map, size, )?; call( unsafe { ds4_gpu_rope_tail_tensor( s.indexer_q.raw(), 1, shape.indexer_heads as u32, shape.indexer_head_dim as u32, shape.rot as u32, pos, original, false, freq_base, freq_scale, ext, attn_factor, shape.rope_beta_fast, shape.rope_beta_slow, ) }, "indexer Q RoPE", )?; call( unsafe { ds4_gpu_dsv4_indexer_qat_tensor( s.indexer_q.raw(), shape.indexer_heads as u32, shape.indexer_head_dim as u32, ) }, "indexer Q quantization", )?; f16( &s.indexer_weights, weights.proj, shape.embd, shape.indexer_heads, &s.norm, map, size, )?; call( unsafe { ds4_gpu_indexer_score_one_tensor( s.indexer_scores.raw(), s.indexer_q.raw(), s.indexer_weights.raw(), indexer.cache.raw(), indexer.rows, shape.indexer_heads as u32, shape.indexer_head_dim as u32, ((shape.indexer_heads * shape.indexer_head_dim) as f32) .sqrt() .recip(), ) }, "indexer scoring", )?; call( unsafe { ds4_gpu_indexer_topk_tensor( s.indexer_selected.raw(), s.indexer_scores.raw(), indexer.rows, 1, shape.indexer_top_k as u32, ) }, "indexer top-k selection", )?; Some(indexer) } _ => None, } } else { None }; let fuse_inverse_rope = !environment_present(c"DS4_METAL_DISABLE_PRE_M5_ATTN_INV_ROPE_FUSE") && (unsafe { ds4_gpu_device_is_pre_m5_apple_silicon() } != 0 || unsafe { ds4_gpu_device_is_m5_apple_silicon() } != 0) && unsafe { ds4_gpu_decode_attn_rope_fuse_available() } != 0; let mut inverse_rope_armed = false; if indexed.is_some() { call( unsafe { ds4_gpu_attention_indexed_mixed_batch_heads_tensor( s.heads.raw(), map, size, w.attn_sinks.offset, s.q.raw(), raw_cache, comp_cache, 1, s.indexer_selected.raw(), 1, pos, n_raw, raw_cap, raw_start, n_comp, shape.indexer_top_k as u32, shape.sliding_window as u32, compression_ratio(shape, layer), shape.heads as u32, shape.head_dim as u32, ) }, "indexed attention", )?; } else { if fuse_inverse_rope { unsafe { ds4_gpu_set_decode_attn_rope_fuse( shape.head_dim as u32, shape.rot as u32, pos, original, true, freq_base, freq_scale, ext, attn_factor, shape.rope_beta_fast, shape.rope_beta_slow, ) }; inverse_rope_armed = true; } call( unsafe { ds4_gpu_attention_decode_heads_tensor( s.heads.raw(), map, size, w.attn_sinks.offset, s.q.raw(), raw_cache, n_raw, raw_cap, raw_start, comp_cache, 1, n_comp, std::ptr::null(), 0, shape.heads as u32, shape.head_dim as u32, ) }, "attention", )?; } if !(inverse_rope_armed && unsafe { ds4_gpu_decode_attn_rope_fuse_used() } != 0) { call( unsafe { ds4_gpu_rope_tail_tensor( s.heads.raw(), 1, shape.heads as u32, shape.head_dim as u32, shape.rot as u32, pos, original, true, freq_base, freq_scale, ext, attn_factor, shape.rope_beta_fast, shape.rope_beta_slow, ) }, "inverse attention RoPE", )?; } let group_dim = shape.head_dim * (shape.heads / shape.out_groups); call( unsafe { ds4_gpu_attention_output_low_q8_tensor( s.attention_low.raw(), map, size, w.attn_output_a.offset, group_dim, shape.lora_o, shape.out_groups as u32, s.heads.raw(), ) }, "attention low projection", )?; if let Some(steering) = steering.filter(|value| value.attention_scale != 0.0) { q8( &s.attention_out, w.attn_output_b, shape.out_groups * shape.lora_o, shape.embd, &s.attention_low, map, size, )?; steering.apply(&s.attention_out, layer, 1, true)?; call( unsafe { ds4_gpu_hc_expand_split_tensor( s.next_hc.raw(), s.attention_out.raw(), s.current_hc.raw(), s.hc_split.raw(), shape.embd as u32, shape.hc as u32, ) }, "steered attention HC expansion", )?; } else { call( unsafe { ds4_gpu_matmul_q8_0_hc_expand_tensor( s.next_hc.raw(), s.attention_out.raw(), map, size, w.attn_output_b.offset, shape.out_groups * shape.lora_o, shape.embd, s.attention_low.raw(), s.current_hc.raw(), s.hc_split.raw(), shape.embd as u32, shape.hc as u32, ) }, "attention output", )?; } decode_hc_mix( s, &s.next_hc, w.hc_ffn_fn, w.hc_ffn_scale, w.hc_ffn_base, w.ffn_norm, shape, map, size, "FFN HC mix", )?; let gate_row = w.expert_gate.bytes / (w.expert_gate.dims[1] * w.expert_gate.dims[2]); let down_row = w.expert_down.bytes / (w.expert_down.dims[1] * w.expert_down.dims[2]); let gate_expert_bytes = w.expert_gate.dims[1] * gate_row; let down_expert_bytes = w.expert_down.dims[1] * down_row; let parallel_eligible = !quality && ssd.is_none() && cache_rows.is_none() && !profile_active && steering.is_none_or(|value| value.ffn_scale == 0.0) && !environment_present(c"DS4_METAL_MOE_ONE_STAGE_PROFILE") && !environment_present(c"DS4_METAL_MOE_WRITE_CLAMPED_ACT") && !environment_present(c"DS4_METAL_DISABLE_ROUTED_PAIR_SWIGLU_FUSION") && !environment_present(c"DS4_METAL_Q8_MV_NSG") && w.expert_gate.kind == IQ2_XXS && w.expert_up.kind == IQ2_XXS && w.expert_down.kind == Q2_K && shape.experts == 256 && shape.experts_used == 6 && shape.embd == 4096 && w.expert_gate.dims[0] == 4096 && w.expert_down.dims[0] == 2048 && w.expert_down.dims[1] == 4096 && gate_row == 1056 && gate_expert_bytes == 2_162_688 && down_row == 672 && down_expert_bytes == 2_752_512 && w.shared_down.kind == Q8_0 && w.router.kind == F16 && w.router.dims[0] == shape.embd && w.router.dims[1] == shape.experts && decode_feature_enabled( c"DS4_METAL_DISABLE_PRE_M5_PARALLEL_FULL_FFN", c"DS4_METAL_DISABLE_M5_PARALLEL_FULL_FFN", ); let router_fusion_eligible = !quality && ssd.is_none() && cache_rows.is_none() && w.shared_gate.kind == Q8_0 && w.shared_up.kind == Q8_0 && w.router.kind == F16 && w.router.dims[0] == shape.embd && w.router.dims[1] == shape.experts && (!((unsafe { ds4_gpu_device_is_pre_m5_apple_silicon() } != 0) && environment_present(c"DS4_METAL_DISABLE_PRE_M5_ROUTER_SHARED_FUSE"))) && (!((unsafe { ds4_gpu_device_is_m5_apple_silicon() } != 0) && environment_present(c"DS4_METAL_DISABLE_M5_ROUTER_SHARED_FUSE"))) && (unsafe { ds4_gpu_device_is_pre_m5_apple_silicon() } != 0 || unsafe { ds4_gpu_device_is_m5_apple_silicon() } != 0); let mut router_shared_done = false; let router_project_select = router_fusion_eligible && parallel_eligible && w.router_hash.is_none() && !environment_present(c"DS4_METAL_DISABLE_M5_ROUTER_PROJECT_SELECT_FUSE") && unsafe { ds4_gpu_device_is_m5_apple_silicon() } != 0; let mut router_selected_done = false; let mut router_projected = false; let mut parallel_router_ready = false; if router_project_select { let fused = unsafe { ds4_gpu_router_project_select_fused_tensor( s.router_logits.raw(), s.router_probs.raw(), s.router_selected.raw(), s.router_weights.raw(), map, size, w.router.offset, w.router_bias.map_or(0, |value| value.offset), w.router_bias.is_some(), s.norm.raw(), ) }; if fused < 0 { return Err("Metal failed while fusing router projection and selection".into()); } router_selected_done = fused > 0; router_projected = router_selected_done; parallel_router_ready = router_selected_done; } if router_fusion_eligible && !router_selected_done { let fused = unsafe { ds4_gpu_router_shared_gate_up_q8_0_tensor( s.router_logits.raw(), s.shared_gate.raw(), s.shared_up.raw(), s.shared_mid.raw(), map, size, w.router.offset, w.shared_gate.offset, w.shared_up.offset, shape.embd, shape.experts, shape.ff_expert, s.norm.raw(), shape.swiglu_clamp, parallel_eligible, ) }; if fused < 0 { return Err("Metal failed while fusing router and shared-expert projections".into()); } if fused > 0 { router_projected = true; router_shared_done = !parallel_eligible; parallel_router_ready = parallel_eligible; } } if !router_selected_done { if !router_projected { matmul( &s.router_logits, w.router, shape.embd, shape.experts, &s.norm, map, size, )?; } call( unsafe { ds4_gpu_router_select_tensor( s.router_selected.raw(), s.router_weights.raw(), s.router_probs.raw(), map, size, w.router_bias.map_or(0, |v| v.offset), w.router_hash.map_or(0, |v| v.offset), w.router_hash.map_or(0, |v| v.dims[1] as u32), token, shape.experts as u32, shape.experts_used as u32, shape.expert_weight_scale, 0, 0, w.router_bias.is_some(), w.router_hash.is_some(), s.router_logits.raw(), ) }, "expert routing", )?; } if let Some(ssd) = ssd { let table = StreamExpertTable { model_map: map, model_size: size, layer, total_experts: shape.experts as u32, gate_offset: w.expert_gate.offset, up_offset: w.expert_up.offset, down_offset: w.expert_down.offset, gate_expert_bytes: w.expert_gate.dims[1] * gate_row, down_expert_bytes: w.expert_down.dims[1] * down_row, }; ssd.begin_selected(&s.router_selected, table, shape.experts_used as u32)?; } let parallel = (parallel_eligible && parallel_router_ready) .then(|| { ParallelFfn::start( &s.shared_gate, &s.shared_up, &s.shared_mid, &s.shared_out, map, size, w.shared_gate.offset, w.shared_up.offset, w.shared_down.offset, shape.embd as u32, shape.ff_expert as u32, &s.norm, shape.swiglu_clamp, ) }) .flatten(); if parallel.is_none() && !router_shared_done { call( unsafe { ds4_gpu_shared_gate_up_swiglu_q8_0_tensor( s.shared_gate.raw(), s.shared_up.raw(), s.shared_mid.raw(), map, size, w.shared_gate.offset, w.shared_up.offset, shape.embd, shape.ff_expert, s.norm.raw(), shape.swiglu_clamp, ) }, "shared expert gate/up", )?; } if steering.is_some_and(|value| value.ffn_scale != 0.0) { q8( &s.shared_out, w.shared_down, shape.ff_expert, shape.embd, &s.shared_mid, map, size, )?; } if let Some(ssd) = ssd { ssd.finish_selected(true)?; } call( unsafe { ds4_gpu_routed_moe_one_tensor( s.routed_out.raw(), s.routed_gate.raw(), s.routed_up.raw(), s.routed_mid.raw(), s.routed_experts.raw(), map, size, w.expert_gate.offset, w.expert_up.offset, w.expert_down.offset, w.expert_gate.kind, w.expert_down.kind, w.expert_gate.dims[1] * gate_row, gate_row, w.expert_down.dims[1] * down_row, down_row, w.expert_gate.dims[0] as u32, w.expert_down.dims[0] as u32, w.expert_down.dims[1] as u32, s.router_selected.raw(), s.router_weights.raw(), shape.experts as u32, shape.experts_used as u32, shape.swiglu_clamp, s.norm.raw(), std::ptr::null(), layer, false, ) }, "routed experts", )?; let parallel_done = if let Some(work) = parallel { work.finish()?; true } else { false }; if let Some(steering) = steering.filter(|value| value.ffn_scale != 0.0) { call( unsafe { ds4_gpu_add_tensor( s.attention_out.raw(), s.routed_out.raw(), s.shared_out.raw(), shape.embd as u32, ) }, "combining FFN output", )?; steering.apply(&s.attention_out, layer, 1, false)?; call( unsafe { ds4_gpu_hc_expand_split_tensor( s.current_hc.raw(), s.attention_out.raw(), s.next_hc.raw(), s.hc_split.raw(), shape.embd as u32, shape.hc as u32, ) }, "steered FFN HC expansion", ) } else if parallel_done { call( unsafe { ds4_gpu_hc_expand_add_split_tensor( s.current_hc.raw(), s.routed_out.raw(), s.shared_out.raw(), s.next_hc.raw(), s.hc_split.raw(), shape.embd as u32, shape.hc as u32, ) }, "parallel FFN HC expansion", ) } else { call( unsafe { ds4_gpu_shared_down_hc_expand_q8_0_tensor( s.current_hc.raw(), s.shared_out.raw(), map, size, w.shared_down.offset, shape.ff_expert, shape.embd, s.shared_mid.raw(), s.routed_out.raw(), s.next_hc.raw(), s.hc_split.raw(), shape.embd as u32, shape.hc as u32, ) }, "shared expert output", ) } } #[allow(clippy::too_many_arguments)] fn project_compressor_pair( kv: &Buffer, score: &Buffer, state: &CompressionState, weights: CompressorWeights, input: &Buffer, input_width: u64, width: u32, map: *const c_void, size: u64, pos: u32, ) -> Result { let fused = unsafe { ds4_gpu_matmul_f16_pair_compressor_store_tensor( kv.raw(), score.raw(), state.state_kv.raw(), state.state_score.raw(), map, size, weights.kv.offset, weights.gate.offset, weights.ape.offset, weights.ape.kind, input_width, width, input.raw(), state.ratio, pos, ) }; if fused < 0 { return Err("Metal failed while storing compressor projections".into()); } if fused == 0 { call( unsafe { ds4_gpu_matmul_f16_pair_tensor( kv.raw(), score.raw(), map, size, weights.kv.offset, weights.gate.offset, input_width, width as u64, input.raw(), 1, ) }, "compressor projection", )?; } Ok(fused > 0) } #[allow(clippy::too_many_arguments)] fn update_compression_pair( s: &Scratch, attn: &mut CompressionState, attn_weights: CompressorWeights, index: &mut CompressionState, index_weights: CompressorWeights, shape: super::Shape, map: *const c_void, size: u64, pos: u32, original: u32, freq_base: f32, freq_scale: f32, ext: f32, attn_factor: f32, projected: bool, full_phase: bool, resident: bool, ) -> Result<(), String> { let attn_width = (2 * shape.head_dim) as u32; let index_width = (2 * shape.indexer_head_dim) as u32; let mut stored = projected; if !stored && full_phase && attn.ratio == 4 && attn_weights.kv.kind == F16 && attn_weights.gate.kind == F16 && index_weights.kv.kind == F16 && index_weights.gate.kind == F16 && !environment_present(c"DS4_METAL_DISABLE_PRE_M5_COMPRESSOR_QUAD_STORE") && (unsafe { ds4_gpu_device_is_pre_m5_apple_silicon() } != 0 || unsafe { ds4_gpu_device_is_m5_apple_silicon() } != 0) { let result = unsafe { ds4_gpu_matmul_f16_quad_compressor_store_tensor( s.compressed_kv.raw(), s.compressed_score.raw(), s.index_compressed_kv.raw(), s.index_compressed_score.raw(), attn.state_kv.raw(), attn.state_score.raw(), index.state_kv.raw(), index.state_score.raw(), map, size, attn_weights.kv.offset, attn_weights.gate.offset, index_weights.kv.offset, index_weights.gate.offset, attn_weights.ape.offset, attn_weights.ape.kind, index_weights.ape.offset, index_weights.ape.kind, shape.embd, attn_width, index_width, s.norm.raw(), 4, pos, ) }; if result < 0 { return Err("Metal failed while fusing compressor projections".into()); } stored = result > 0; } let (attn_stored, index_stored) = if stored { (true, true) } else { ( project_compressor_pair( &s.compressed_kv, &s.compressed_score, attn, attn_weights, &s.norm, shape.embd, attn_width, map, size, pos, )?, project_compressor_pair( &s.index_compressed_kv, &s.index_compressed_score, index, index_weights, &s.norm, shape.embd, index_width, map, size, pos, )?, ) }; let emit = (pos + 1).is_multiple_of(4); let fused_finalize = emit && resident && shape.head_dim == 512 && shape.indexer_head_dim == 128 && attn_weights.norm.kind == F32 && index_weights.norm.kind == F32 && unsafe { ds4_gpu_kv_rope_fp8_fuse_available() } != 0 && decode_feature_enabled( c"DS4_METAL_DISABLE_PRE_M5_COMP_FINALIZE_FUSE", c"DS4_METAL_DISABLE_M5_COMP_FINALIZE_FUSE", ); call( unsafe { ds4_gpu_compressor_update_tensor( s.compressed_kv.raw(), s.compressed_score.raw(), attn.state_kv.raw(), attn.state_score.raw(), s.compressed_stage.raw(), map, size, attn_weights.ape.offset, attn_weights.ape.kind, attn_weights.norm.offset, attn_weights.norm.kind, shape.head_dim as u32, 4, pos, 0, shape.rot as u32, original, freq_base, freq_scale, ext, attn_factor, shape.rope_beta_fast, shape.rope_beta_slow, shape.rms_epsilon, attn_stored, true, fused_finalize, ) }, "attention compressor update", )?; call( unsafe { ds4_gpu_compressor_update_tensor( s.index_compressed_kv.raw(), s.index_compressed_score.raw(), index.state_kv.raw(), index.state_score.raw(), index.cache.raw(), map, size, index_weights.ape.offset, index_weights.ape.kind, index_weights.norm.offset, index_weights.norm.kind, shape.indexer_head_dim as u32, 4, pos, index.rows, shape.rot as u32, original, freq_base, freq_scale, ext, attn_factor, shape.rope_beta_fast, shape.rope_beta_slow, shape.rms_epsilon, index_stored, true, fused_finalize, ) }, "index compressor update", )?; if !emit { return Ok(()); } if fused_finalize { let fused = unsafe { ds4_gpu_dsv4_comp_row_finalize_tensor( s.compressed_stage.raw(), attn.cache.raw(), attn.rows, attn_weights.norm.offset, index.cache.raw(), index.rows, index_weights.norm.offset, attn.state_kv.raw(), attn.state_score.raw(), index.state_kv.raw(), index.state_score.raw(), map, size, pos + 1 - 4, shape.rot as u32, original, freq_base, freq_scale, ext, attn_factor, shape.rope_beta_fast, shape.rope_beta_slow, shape.rms_epsilon, ) }; if fused != 1 { return Err("Metal failed while finalizing compressor rows".into()); } } else { call( unsafe { ds4_gpu_dsv4_fp8_kv_quantize_tensor( s.compressed_stage.raw(), 1, shape.head_dim as u32, shape.rot as u32, ) }, "compressed KV quantization", )?; call( unsafe { ds4_gpu_tensor_copy_f32_to_f16( attn.cache.raw(), attn.rows as u64 * shape.head_dim * 2, s.compressed_stage.raw(), 0, shape.head_dim, ) }, "compressed KV cache write", )?; let row = index.cache.view( index.rows as u64 * shape.indexer_head_dim * 4, shape.indexer_head_dim * 4, )?; call( unsafe { ds4_gpu_dsv4_indexer_qat_tensor(row.raw(), 1, shape.indexer_head_dim as u32) }, "compressed index quantization", )?; } attn.rows += 1; index.rows += 1; Ok(()) } #[allow(clippy::too_many_arguments)] fn update_compression( s: &Scratch, state: &mut CompressionState, weights: CompressorWeights, shape: super::Shape, map: *const c_void, size: u64, pos: u32, original: u32, freq_base: f32, freq_scale: f32, ext: f32, attn_factor: f32, projected: bool, ) -> Result<(), String> { let emit = update_compressor_stage( s, state, weights, shape, map, size, pos, original, freq_base, freq_scale, ext, attn_factor, shape.head_dim as u32, projected, )?; if emit { call( unsafe { ds4_gpu_dsv4_fp8_kv_quantize_tensor( s.compressed_stage.raw(), 1, shape.head_dim as u32, shape.rot as u32, ) }, "compressed KV quantization", )?; call( unsafe { ds4_gpu_tensor_copy_f32_to_f16( state.cache.raw(), state.rows as u64 * shape.head_dim * 2, s.compressed_stage.raw(), 0, shape.head_dim, ) }, "compressed KV cache write", )?; state.rows += 1; } Ok(()) } #[allow(clippy::too_many_arguments)] fn update_compressor_stage( s: &Scratch, state: &CompressionState, weights: CompressorWeights, shape: super::Shape, map: *const c_void, size: u64, pos: u32, original: u32, freq_base: f32, freq_scale: f32, ext: f32, attn_factor: f32, head_dim: u32, projected: bool, ) -> Result { let width = if state.ratio == 4 { 2 } else { 1 } * head_dim; let fused = if projected { 1 } else { unsafe { ds4_gpu_matmul_f16_pair_compressor_store_tensor( s.compressed_kv.raw(), s.compressed_score.raw(), state.state_kv.raw(), state.state_score.raw(), map, size, weights.kv.offset, weights.gate.offset, weights.ape.offset, weights.ape.kind, shape.embd, width, s.norm.raw(), state.ratio, pos, ) } }; if fused < 0 { return Err("Metal failed while storing compressor projections".into()); } if fused == 0 { call( unsafe { ds4_gpu_matmul_f16_pair_tensor( s.compressed_kv.raw(), s.compressed_score.raw(), map, size, weights.kv.offset, weights.gate.offset, shape.embd, width as u64, s.norm.raw(), 1, ) }, "compressor projection", )?; } call( unsafe { ds4_gpu_compressor_update_tensor( s.compressed_kv.raw(), s.compressed_score.raw(), state.state_kv.raw(), state.state_score.raw(), s.compressed_stage.raw(), map, size, weights.ape.offset, weights.ape.kind, weights.norm.offset, weights.norm.kind, head_dim, state.ratio, pos, 0, shape.rot as u32, original, freq_base, freq_scale, ext, attn_factor, shape.rope_beta_fast, shape.rope_beta_slow, shape.rms_epsilon, fused > 0, true, false, ) }, "compressor update", )?; Ok((pos + 1).is_multiple_of(state.ratio)) } fn compression_ratio(shape: super::Shape, layer: u32) -> u32 { match shape.model { crate::model::ModelChoice::DeepSeekV4Flash0731 if layer < 2 => 0, crate::model::ModelChoice::DeepSeekV4Pro if layer < 2 => 128, crate::model::ModelChoice::DeepSeekV4Flash0731 | crate::model::ModelChoice::DeepSeekV4Pro if layer.is_multiple_of(2) => { 4 } crate::model::ModelChoice::DeepSeekV4Flash0731 | crate::model::ModelChoice::DeepSeekV4Pro => 128, crate::model::ModelChoice::Glm52 => 0, } } fn raw_decode_span(pos: u32, raw_cap: u32, window: u32) -> (u32, u32) { let n_raw = (pos + 1) .min(raw_cap) .min(if window == 0 { raw_cap } else { window }); (n_raw, (pos + 1 - n_raw) % raw_cap) } fn raw_batch_span(pos: u32, rows: u32, raw_cap: u32, window: u32) -> (u32, u32) { if rows == 0 || raw_cap == 0 { return (0, 0); } let last = pos + rows - 1; let needed = rows .saturating_add(if rows == 1 { window.saturating_sub(1) } else { window }) .min(last + 1) .min(raw_cap); (needed, (last + 1 - needed) % raw_cap) } fn encode_output( s: &Scratch, w: &Weights, shape: super::Shape, map: *const c_void, size: u64, ) -> Result<(), String> { let hc_dim = shape.hc * shape.embd; call( unsafe { ds4_gpu_rms_norm_plain_tensor( s.flat_hc.raw(), s.current_hc.raw(), hc_dim as u32, shape.rms_epsilon, ) }, "output HC norm", )?; f16( &s.output_pre, w.output_hc_fn, hc_dim, shape.hc, &s.flat_hc, map, size, )?; call( unsafe { ds4_gpu_output_hc_weights_tensor( s.output_weights.raw(), s.output_pre.raw(), map, size, w.output_hc_scale.offset, w.output_hc_base.offset, shape.hc as u32, shape.hc_epsilon, ) }, "output HC weights", )?; call( unsafe { ds4_gpu_hc_weighted_sum_tensor( s.output_embedding.raw(), s.current_hc.raw(), s.output_weights.raw(), shape.embd as u32, shape.hc as u32, ) }, "output HC collapse", )?; call( unsafe { ds4_gpu_rms_norm_weight_tensor( s.output_norm.raw(), s.output_embedding.raw(), map, size, w.output_norm.offset, shape.embd as u32, shape.rms_epsilon, ) }, "output norm", )?; q8( &s.logits, w.output, shape.embd, shape.vocab, &s.output_norm, map, size, ) } fn encode_batch_output( s: &BatchScratch, w: &Weights, shape: super::Shape, map: *const c_void, size: u64, rows: u32, ) -> Result<(), String> { let head_rows = if rows > 1 && rows < 8 { 8 } else { rows }; let hc_dim = shape.hc * shape.embd; let output_pre = s.hc_mix.view(0, u64::from(rows) * shape.hc * 4)?; let output_weights = s.hc_split.view(0, u64::from(rows) * shape.hc * 4)?; let output_embedding = s.current.view(0, u64::from(rows) * shape.embd * 4)?; let output_norm = s.norm.view(0, u64::from(head_rows) * shape.embd * 4)?; let logits = s .output_logits .as_ref() .ok_or("batch output logits are not allocated")?; call( unsafe { ds4_gpu_rms_norm_plain_rows_tensor( s.flat_hc.raw(), s.current_hc.raw(), hc_dim as u32, rows, shape.rms_epsilon, ) }, "batch output HC norm", )?; f16_rows( &output_pre, w.output_hc_fn, hc_dim, shape.hc, &s.flat_hc, rows, map, size, )?; call( unsafe { ds4_gpu_output_hc_weights_tensor( output_weights.raw(), output_pre.raw(), map, size, w.output_hc_scale.offset, w.output_hc_base.offset, shape.hc as u32, shape.hc_epsilon, ) }, "batch output HC weights", )?; call( unsafe { ds4_gpu_hc_weighted_sum_tensor( output_embedding.raw(), s.current_hc.raw(), output_weights.raw(), shape.embd as u32, shape.hc as u32, ) }, "batch output HC collapse", )?; call( unsafe { ds4_gpu_rms_norm_weight_rows_tensor( output_norm.raw(), output_embedding.raw(), map, size, w.output_norm.offset, shape.embd as u32, rows, shape.rms_epsilon, ) }, "batch output norm", )?; if head_rows > rows { s.norm .view( u64::from(rows) * shape.embd * 4, u64::from(head_rows - rows) * shape.embd * 4, )? .fill(0.0, u64::from(head_rows - rows) * shape.embd)?; } matmul_rows( logits, w.output, shape.embd, shape.vocab, &output_norm, head_rows, map, size, ) } fn argmax(values: &[f32]) -> i32 { values .iter() .enumerate() .max_by(|left, right| left.1.total_cmp(right.1)) .map_or(-1, |(index, _)| index as i32) } fn dense_row(model: &Gguf, weight: Weight, row: u32) -> Result, String> { if u64::from(row) >= weight.dims[1] { return Err("dense row is outside the tensor".into()); } let width = weight.dims[0] as usize; let row_bytes = match weight.kind { F32 => width.checked_mul(4), F16 => width.checked_mul(2), Q8_0 => width.div_ceil(32).checked_mul(34), _ => None, } .ok_or("unsupported DSpark dense tensor layout")?; let offset = weight .offset .checked_add(u64::from(row) * row_bytes as u64) .ok_or("DSpark dense row offset overflow")?; if offset > model.len() || row_bytes as u64 > model.len() - offset { return Err("DSpark dense row is outside the GGUF mapping".into()); } let bytes = unsafe { std::slice::from_raw_parts(model.map_ptr().add(offset as usize), row_bytes) }; let mut out = vec![0.0; width]; match weight.kind { F32 => { for (value, bytes) in out.iter_mut().zip(bytes.chunks_exact(4)) { *value = f32::from_le_bytes(bytes.try_into().expect("four bytes")); } } F16 => { for (value, bytes) in out.iter_mut().zip(bytes.chunks_exact(2)) { *value = half_to_f32(u16::from_le_bytes(bytes.try_into().expect("two bytes"))); } } Q8_0 => { for (block, bytes) in bytes.chunks_exact(34).enumerate() { let scale = half_to_f32(u16::from_le_bytes([bytes[0], bytes[1]])); for (index, quantized) in bytes[2..].iter().enumerate() { let output = block * 32 + index; if output == width { break; } out[output] = scale * f32::from(*quantized as i8); } } } _ => unreachable!(), } Ok(out) } fn dense_dot(model: &Gguf, weight: Weight, row: u32, values: &[f32]) -> Result { if weight.dims[0] as usize != values.len() || u64::from(row) >= weight.dims[1] { return Err("DSpark dense dot has mismatched dimensions".into()); } let width = values.len(); let row_bytes = match weight.kind { F32 => width.checked_mul(4), F16 => width.checked_mul(2), Q8_0 => width.div_ceil(32).checked_mul(34), _ => None, } .ok_or("unsupported DSpark dense tensor layout")?; let offset = weight .offset .checked_add(u64::from(row) * row_bytes as u64) .ok_or("DSpark dense dot offset overflow")?; if offset > model.len() || row_bytes as u64 > model.len() - offset { return Err("DSpark dense dot is outside the GGUF mapping".into()); } let bytes = unsafe { std::slice::from_raw_parts(model.map_ptr().add(offset as usize), row_bytes) }; Ok(dense_dot_bytes(weight.kind, width, bytes, values)) } fn dense_dot_bytes(kind: u32, width: usize, bytes: &[u8], values: &[f32]) -> f32 { match kind { F32 => bytes .chunks_exact(4) .zip(values) .map(|(bytes, value)| f32::from_le_bytes(bytes.try_into().expect("four bytes")) * value) .sum(), F16 => bytes .chunks_exact(2) .zip(values) .map(|(bytes, value)| { half_to_f32(u16::from_le_bytes(bytes.try_into().expect("two bytes"))) * value }) .sum(), Q8_0 => { let mut sum = 0.0; for (block, bytes) in bytes.chunks_exact(34).enumerate() { let scale = half_to_f32(u16::from_le_bytes([bytes[0], bytes[1]])); for (index, quantized) in bytes[2..].iter().enumerate() { let input = block * 32 + index; if input == width { break; } sum += scale * f32::from(*quantized as i8) * values[input]; } } sum } _ => unreachable!(), } } fn dense_argmax( model: &Gguf, weight: Weight, values: &[f32], logits: &[f32], ) -> Result { if weight.dims[0] as usize != values.len() || weight.dims[1] as usize != logits.len() { return Err("DSpark dense argmax has mismatched dimensions".into()); } let width = values.len(); let row_bytes = match weight.kind { F32 => width.checked_mul(4), F16 => width.checked_mul(2), Q8_0 => width.div_ceil(32).checked_mul(34), _ => None, } .ok_or("unsupported DSpark dense tensor layout")?; let bytes_len = row_bytes .checked_mul(logits.len()) .ok_or("DSpark dense argmax size overflow")?; if weight.offset > model.len() || bytes_len as u64 > model.len() - weight.offset { return Err("DSpark dense argmax is outside the GGUF mapping".into()); } let bytes = unsafe { std::slice::from_raw_parts(model.map_ptr().add(weight.offset as usize), bytes_len) }; let quantized = (weight.kind == Q8_0).then(|| quantize_q8_activation(values)); let workers = std::thread::available_parallelism() .map_or(1, std::num::NonZero::get) .min(logits.len()); let chunk = logits.len().div_ceil(workers); let best = std::thread::scope(|scope| { let mut handles = Vec::with_capacity(workers); let quantized = quantized.as_ref(); for start in (0..logits.len()).step_by(chunk) { let end = (start + chunk).min(logits.len()); handles.push(scope.spawn(move || { let mut best = (start, f32::NEG_INFINITY); for token in start..end { let row = &bytes[token * row_bytes..(token + 1) * row_bytes]; let dot = quantized.as_ref().map_or_else( || dense_dot_bytes(weight.kind, width, row, values), |(values, scales)| dense_dot_q8(row, values, scales, width), ); let score = logits[token] + dot; if score > best.1 { best = (token, score); } } best })); } handles .into_iter() .map(|handle| handle.join().expect("DSpark argmax worker panicked")) .fold((0, f32::NEG_INFINITY), |best, candidate| { if candidate.1 > best.1 { candidate } else { best } }) }); Ok(best.0 as i32) } fn quantize_q8_activation(values: &[f32]) -> (Vec, Vec) { let blocks = values.len().div_ceil(32); let mut quantized = vec![0; blocks * 32]; let mut scales = Vec::with_capacity(blocks); for (block, values) in values.chunks(32).enumerate() { let scale = values .iter() .fold(0.0_f32, |max, value| max.max(value.abs())) / 127.0; let inverse = if scale == 0.0 { 0.0 } else { scale.recip() }; scales.push(scale); for (target, value) in quantized[block * 32..].iter_mut().zip(values) { *target = (value * inverse).round_ties_even().clamp(-128.0, 127.0) as i8; } } (quantized, scales) } fn dense_dot_q8(bytes: &[u8], values: &[i8], scales: &[f32], width: usize) -> f32 { bytes .chunks_exact(34) .zip(values.chunks_exact(32)) .zip(scales) .enumerate() .map(|(block, ((bytes, values), &scale))| { let count = (width - block * 32).min(32); let weight_scale = half_to_f32(u16::from_le_bytes([bytes[0], bytes[1]])); let dot = bytes[2..] .iter() .zip(values) .take(count) .map(|(&weight, &value)| i32::from(weight as i8) * i32::from(value)) .sum::(); weight_scale * scale * dot as f32 }) .sum() } fn half_to_f32(value: u16) -> f32 { let sign = u32::from(value & 0x8000) << 16; let exponent = u32::from((value >> 10) & 0x1f); let fraction = u32::from(value & 0x03ff); let bits = match exponent { 0 if fraction == 0 => sign, 0 => { let shift = fraction.leading_zeros() - 21; sign | ((127 - 14 - shift) << 23) | ((fraction << (shift + 1) & 0x03ff) << 13) } 31 => sign | 0x7f80_0000 | (fraction << 13), _ => sign | ((exponent + 127 - 15) << 23) | (fraction << 13), }; f32::from_bits(bits) } fn f16( out: &Buffer, weight: Weight, input: u64, output: u64, x: &Buffer, map: *const c_void, size: u64, ) -> Result<(), String> { call( unsafe { ds4_gpu_matmul_f16_tensor( out.raw(), map, size, weight.offset, input, output, x.raw(), 1, ) }, "F16 projection", ) } #[allow(clippy::too_many_arguments)] fn f16_rows( out: &Buffer, weight: Weight, input: u64, output: u64, x: &Buffer, rows: u32, map: *const c_void, size: u64, ) -> Result<(), String> { call( unsafe { ds4_gpu_matmul_f16_tensor( out.raw(), map, size, weight.offset, input, output, x.raw(), u64::from(rows), ) }, "batch F16 projection", ) } #[allow(clippy::too_many_arguments)] fn matmul_rows( out: &Buffer, weight: Weight, input: u64, output: u64, x: &Buffer, rows: u32, map: *const c_void, size: u64, ) -> Result<(), String> { match weight.kind { F16 => f16_rows(out, weight, input, output, x, rows, map, size), F32 => call( unsafe { ds4_gpu_matmul_f32_tensor( out.raw(), map, size, weight.offset, input, output, x.raw(), u64::from(rows), ) }, "batch F32 projection", ), Q8_0 => q8_rows(out, weight, input, output, x, rows, map, size), kind => Err(format!("unsupported Metal batch projection type {kind}")), } } fn matmul( out: &Buffer, weight: Weight, input: u64, output: u64, x: &Buffer, map: *const c_void, size: u64, ) -> Result<(), String> { match weight.kind { F16 => f16(out, weight, input, output, x, map, size), F32 => call( unsafe { ds4_gpu_matmul_f32_tensor( out.raw(), map, size, weight.offset, input, output, x.raw(), 1, ) }, "F32 projection", ), Q8_0 => q8(out, weight, input, output, x, map, size), kind => Err(format!("unsupported Metal projection type {kind}")), } } fn q8( out: &Buffer, weight: Weight, input: u64, output: u64, x: &Buffer, map: *const c_void, size: u64, ) -> Result<(), String> { call( unsafe { ds4_gpu_matmul_q8_0_tensor( out.raw(), map, size, weight.offset, input, output, x.raw(), 1, ) }, "Q8 projection", ) } #[allow(clippy::too_many_arguments)] fn q8_rows( out: &Buffer, weight: Weight, input: u64, output: u64, x: &Buffer, rows: u32, map: *const c_void, size: u64, ) -> Result<(), String> { call( unsafe { ds4_gpu_matmul_q8_0_tensor( out.raw(), map, size, weight.offset, input, output, x.raw(), u64::from(rows), ) }, "batch Q8 projection", ) } fn throttle(average: &mut f64, elapsed: Duration, power_percent: u8) { if power_percent >= 100 { return; } let sample = elapsed.as_secs_f64(); *average = if *average <= 0.0 || !average.is_finite() { sample } else { *average * 0.875 + sample * 0.125 }; let sleep = *average * (100.0 - f64::from(power_percent)) / f64::from(power_percent); std::thread::sleep(Duration::from_secs_f64(sleep)); } fn call(result: i32, operation: &str) -> Result<(), String> { if result == 0 { Err(format!("Metal failed while {operation}")) } else { Ok(()) } } fn check(result: i32, operation: &str) -> Result<(), String> { call(result, operation) } #[cfg(test)] mod tests { use super::{ compression_ratio, dspark_scheduler_pause, effective_prefill_cap, effective_raw_cap, estimated_deepseek_runtime_bytes, finish_deepseek_model_spans, gpu::ds4_gpu_print_memory_report, quantize_q8_activation, raw_batch_span, raw_decode_span, }; use crate::engine::{FLASH_0731 as FLASH, MXFP4, PRO}; fn installed_artifacts( model: crate::model::ModelChoice, dspark: bool, ) -> crate::model::EngineArtifacts { crate::model::engine_artifacts(model, dspark, &crate::app::models_path()) } #[test] fn compression_schedule_tracks_the_deepseek_model_shape() { assert_eq!(compression_ratio(FLASH, 0), 0); assert_eq!(compression_ratio(FLASH, 2), 4); assert_eq!(compression_ratio(PRO, 0), 128); assert_eq!(compression_ratio(PRO, 1), 128); } #[test] fn dspark_scheduler_matches_the_ds4_default_window() { assert_eq!(dspark_scheduler_pause(4, 8, 0), 0); assert_eq!(dspark_scheduler_pause(4, 5, 0), 2); assert_eq!(dspark_scheduler_pause(4, 8, 2), 4); } #[test] fn dspark_markov_activation_matches_q8_rounding() { let (values, scales) = quantize_q8_activation(&[0.0, 1.0, -1.0, 0.5]); assert_eq!(&values[..4], &[0, 127, -127, 64]); assert_eq!(scales, [1.0 / 127.0]); } #[test] fn pro_q4_model_spans_remain_isolated() { assert_eq!( finish_deepseek_model_spans( vec![(0, 10, false), (10, 20, true), (30, 10, false)], "test", ) .unwrap(), vec![(0, 10), (10, 20), (30, 10)], ); } #[test] fn decode_raw_span_tracks_the_logical_sliding_window() { assert_eq!(raw_decode_span(0, 4_352, 128), (1, 0)); assert_eq!(raw_decode_span(127, 4_352, 128), (128, 0)); assert_eq!(raw_decode_span(128, 4_352, 128), (128, 1)); assert_eq!(raw_decode_span(362, 4_352, 128), (128, 235)); assert_eq!(raw_decode_span(4_500, 4_352, 128), (128, 4_373 % 4_352)); } #[test] fn batch_raw_span_includes_the_chunk_and_previous_window() { assert_eq!(raw_batch_span(0, 4_096, 4_352, 128), (4_096, 0)); assert_eq!(raw_batch_span(4_096, 1_925, 4_352, 128), (2_053, 3_968)); assert_eq!(raw_batch_span(4_500, 1, 4_352, 128), (128, 21)); } #[test] fn half_conversion_handles_normal_and_subnormal_values() { assert_eq!(super::half_to_f32(0x0000), 0.0); assert_eq!(super::half_to_f32(0x3c00), 1.0); assert_eq!(super::half_to_f32(0xc000), -2.0); assert_eq!(super::half_to_f32(0x0400), 2.0_f32.powi(-14)); assert_eq!(super::half_to_f32(0x0001), 2.0_f32.powi(-24)); } #[test] fn decode_feature_policy_keeps_device_generations_and_rollbacks_separate() { use super::decode_feature_allowed; assert!(decode_feature_allowed(true, false, false, false)); assert!(!decode_feature_allowed(true, false, true, false)); assert!(decode_feature_allowed(false, true, false, false)); assert!(!decode_feature_allowed(false, true, false, true)); assert!(!decode_feature_allowed(false, false, false, false)); } #[test] fn streaming_memory_plan_matches_ds4_graph_formulas() { assert_eq!(effective_prefill_cap(16_384, 0), 4_096); assert_eq!(effective_prefill_cap(16_384, 2_048), 2_048); assert_eq!(effective_raw_cap(FLASH, 16_384, 4_096), 4_352); assert_eq!(effective_raw_cap(FLASH, 1_000, 1_000), 1_024); let raw = 43_u64 * 4_352 * 512 * 4; let ratio4 = 21_u64 * (16_384 / 4 + 2) * (512 * 2 + 128 * 4); let ratio128 = 20_u64 * (16_384 / 128 + 2) * 512 * 2; let scratch = 2_u64 * (16_384 / 4 + 2) * 4_096 * 4 + (4_096 / 4 + 2) as u64 * 512 * 4; assert_eq!( estimated_deepseek_runtime_bytes(FLASH, 16_384, 0), raw + ratio4 + ratio128 + scratch ); } #[test] fn pro_resident_admission_includes_weights_context_and_scratch() { let weights = 464_627_334_560_u64; let runtime = estimated_deepseek_runtime_bytes(PRO, 32_768, 4_096); assert_eq!( super::resident_deepseek_admission_for_weights(weights, PRO, 32_768, 4_096).unwrap(), weights + runtime + 512 * 1024 * 1024 ); } #[test] fn generated_hotlists_match_each_model_shape() { for (hotlist, shape) in [ (super::hotlist::FLASH, FLASH), (super::hotlist::PRO, PRO), (super::hotlist::GLM52, crate::engine::GLM), ] { assert!(hotlist.len() >= 4_096); assert!(hotlist.iter().all(|&(layer, expert)| { u32::from(layer) < shape.layers && u64::from(expert) < shape.experts })); } } #[test] #[ignore = "requires a 0731 Flash GGUF and an Apple M5 device"] fn flash_0731_m5_decode_performance_gate() { use super::{DeepSeekExecutor, Digest, Sha256, configure_sources}; use crate::engine::gguf::Gguf; use crate::engine::validation::validate_support; use crate::engine::{ChatTurn, Model, Rng, sample}; use crate::model::ModelChoice; use crate::settings::{ EngineSpeculativeSettings, EngineSsdSettings, EngineSteeringSettings, ReasoningMode, }; use std::sync::atomic::AtomicBool; use std::time::Instant; configure_sources().unwrap(); let path = std::env::var_os("DS4SERVER_BENCH_MODEL") .map(std::path::PathBuf::from) .unwrap_or_else(|| installed_artifacts(ModelChoice::DeepSeekV4Flash0731, false).model); let prompt_content = std::env::var_os("DS4SERVER_BENCH_CHAT_PROMPT_FILE") .map(|path| std::fs::read_to_string(path).unwrap()) .unwrap_or_else(|| "Count from one to two hundred, spelling out every number.".into()); let frontier = std::env::var("DS4SERVER_BENCH_FRONTIER") .ok() .map(|value| value.parse::().unwrap()); let dspark = std::env::var_os("DS4SERVER_BENCH_DSPARK").is_some(); let temperature = std::env::var("DS4SERVER_BENCH_TEMPERATURE") .ok() .map_or(0.0, |value| value.parse::().unwrap()); let confidence = std::env::var("DS4SERVER_BENCH_DSPARK_CONFIDENCE") .ok() .map(|value| value.parse::().unwrap()); let measured = std::env::var("DS4SERVER_BENCH_MEASURED") .ok() .map_or(128, |value| value.parse::().unwrap()); let run = || { let mut model = Model::open_main(&path, ModelChoice::DeepSeekV4Flash0731).unwrap(); if dspark { let support_path = installed_artifacts(ModelChoice::DeepSeekV4Flash0731, true) .support .unwrap(); let support = Gguf::open(&support_path).unwrap(); model.support_kind = Some(validate_support(&support, &model.shape).unwrap()); model.support = Some(support); } let expert_kind = model .main .tensor("blk.4.ffn_gate_exps.weight") .unwrap() .kind; if std::env::var_os("DS4SERVER_BENCH_EXPECT_MXFP4").is_some() { assert_eq!(expert_kind, MXFP4); } let mut prompt = model.render_conversation( "", &[ChatTurn { user: true, tool: false, system: false, skip_previous_eos: false, reasoning: None, reasoning_complete: true, content: prompt_content.clone(), }], ReasoningMode::Direct, ); if let Some(frontier) = frontier { assert!(prompt.len() >= frontier); prompt.truncate(frontier); } let eos = model.eos_token(); let streaming = std::env::var_os("DS4SERVER_BENCH_SSD").is_some(); let context = frontier .map(|frontier| u32::try_from(frontier + 129).unwrap()) .unwrap_or(4_096); let mut executor = DeepSeekExecutor::open( model, context, false, context, 100, EngineSpeculativeSettings { glm_mtp: false, glm_mtp_timing: false, dspark, dspark_confidence_threshold: confidence.unwrap_or(0.8), dspark_confidence_threshold_set: confidence.is_some(), dspark_strict: false, dspark_exact_sampling: false, }, EngineSsdSettings { enabled: streaming, cold: false, cache_experts: 0, cache_bytes: 0, full_layers: 0, full_layers_set: false, preload_experts: 0, }, EngineSteeringSettings { file: None, ffn_scale: 0.0, attention_scale: 0.0, }, ) .unwrap(); let prefill_started = Instant::now(); executor.prefill(&prompt, |_| true).unwrap(); let prefill_seconds = prefill_started.elapsed().as_secs_f64(); let mut generated = Vec::new(); let mut latencies_ms = Vec::new(); let mut rng = Rng::new(12_345); let cancelled = AtomicBool::new(false); let started = Instant::now(); while generated.len() < measured as usize { let token = sample(executor.logits(), temperature, 1.0, 0.05, 0, &mut rng); assert_ne!(token, eos, "benchmark reached EOS before {measured} tokens"); let token_started = Instant::now(); let remaining = measured - generated.len() as u32; let cycle = if temperature <= 0.0 { executor.eval_speculative_greedy( token, remaining, ReasoningMode::Direct, &cancelled, ) } else { executor.eval_speculative_sampled( token, remaining, ReasoningMode::Direct, temperature, 1.0, 0.05, 0, &mut rng, &cancelled, ) } .unwrap(); let per_token_ms = token_started.elapsed().as_secs_f64() * 1_000.0 / cycle.len() as f64; latencies_ms.extend(std::iter::repeat_n(per_token_ms, cycle.len())); generated.extend(cycle); } let seconds = started.elapsed().as_secs_f64(); let tokens_per_second = f64::from(measured) / seconds; let steady_seconds = latencies_ms[1..].iter().sum::() / 1_000.0; let steady_tokens_per_second = f64::from(measured - 1) / steady_seconds; let mut sorted = latencies_ms[1..].to_vec(); sorted.sort_by(f64::total_cmp); let p50 = sorted[sorted.len() / 2]; let p95 = sorted[(sorted.len() * 95).div_ceil(100) - 1]; let token_bytes = generated .iter() .flat_map(|token| token.to_le_bytes()) .collect::>(); let token_hash = Sha256::digest(token_bytes); let token_hash = token_hash .iter() .map(|byte| format!("{byte:02x}")) .collect::(); unsafe { ds4_gpu_print_memory_report(c"benchmark".as_ptr()) }; let stats = executor.execution_stats(); eprintln!( "DS4SERVER_METAL_PERF mode={} speculative={} temperature={temperature} model=flash-0731 expert_kind={expert_kind} context={context} prompt={} prefill_tps={:.6} measured={measured} seconds={seconds:.6} tps={tokens_per_second:.6} first_ms={:.6} steady_tps={steady_tokens_per_second:.6} p50_ms={p50:.6} p95_ms={p95:.6} cycles={} drafted={} accepted={} verifier_passes={} verifier_ms={} cache_entries={} cache_hits={} cache_misses={} pread_bytes={}", if streaming { "ssd" } else { "resident" }, if dspark { "dspark" } else { "plain" }, prompt.len(), prompt.len() as f64 / prefill_seconds, latencies_ms[0], stats.speculative_cycles, stats.drafted_tokens, stats.accepted_draft_tokens, stats.verifier_passes, stats.verifier_ms, stats.ssd_cache_entries, stats.ssd_cache_hits, stats.ssd_cache_misses, stats.ssd_pread_bytes, ); eprintln!("DS4SERVER_METAL_TOKEN_SHA256 {token_hash}"); if std::env::var_os("DS4SERVER_BENCH_TOKENS").is_some() { eprintln!("DS4SERVER_METAL_PROMPT_TOKENS {prompt:?}"); eprintln!("DS4SERVER_METAL_GENERATED_TOKENS {generated:?}"); let output = generated .iter() .filter_map(|&token| executor.model.token_bytes(token)) .flatten() .collect::>(); eprintln!( "DS4SERVER_METAL_GENERATED_TEXT {:?}", String::from_utf8_lossy(&output) ); } assert!(steady_tokens_per_second.is_finite() && steady_tokens_per_second > 0.0); steady_tokens_per_second }; let reference = std::env::var("DS4_REFERENCE_TPS") .ok() .map(|value| value.parse::().unwrap()); let baseline = std::env::var("DS4SERVER_BASELINE_TPS") .ok() .map(|value| value.parse::().unwrap()); let runs = if reference.is_some() || baseline.is_some() { 3 } else { 1 }; let mut results = Vec::with_capacity(runs); for run_index in 0..runs { if run_index > 0 { std::thread::sleep(std::time::Duration::from_secs(5)); } results.push(run()); } results.sort_by(f64::total_cmp); let median = results[results.len() / 2]; eprintln!("DS4SERVER_METAL_PERF_MEDIAN runs={runs} steady_tps={median:.6}"); if let Some(reference) = reference { assert!( median >= reference * 0.95, "DS4Server median {median:.3} tok/s is more than 5% below DS4 {reference:.3} tok/s" ); } if let Some(baseline) = baseline { assert!( median > baseline, "DS4Server median {median:.3} tok/s did not improve on {baseline:.3} tok/s" ); } } #[test] #[ignore = "requires a 0731 Flash GGUF and an Apple M5 device"] fn flash_0731_long_context_crosses_indexed_prefill_boundary() { use super::{DeepSeekExecutor, argmax, configure_sources}; use crate::engine::{ChatTurn, Model}; use crate::model::ModelChoice; use crate::settings::{ EngineSpeculativeSettings, EngineSsdSettings, EngineSteeringSettings, ReasoningMode, }; configure_sources().unwrap(); let path = std::env::var_os("DS4SERVER_BENCH_MODEL") .map(std::path::PathBuf::from) .unwrap_or_else(|| installed_artifacts(ModelChoice::DeepSeekV4Flash0731, false).model); let model = Model::open_main(&path, ModelChoice::DeepSeekV4Flash0731).unwrap(); let prompt = model.render_conversation( "", &[ChatTurn { user: true, tool: false, system: false, skip_previous_eos: false, reasoning: None, reasoning_complete: true, content: "hi ".repeat(4_100), }], ReasoningMode::Direct, ); assert!(prompt.len() > 4_096 && prompt.len() < 8_192); let streaming = std::env::var_os("DS4SERVER_BENCH_SSD").is_some(); let mut executor = DeepSeekExecutor::open( model, 8_192, false, 4_096, 100, EngineSpeculativeSettings { glm_mtp: false, glm_mtp_timing: false, dspark: false, dspark_confidence_threshold: 0.9, dspark_confidence_threshold_set: false, dspark_strict: false, dspark_exact_sampling: false, }, EngineSsdSettings { enabled: streaming, cold: false, cache_experts: if streaming { 4_096 } else { 0 }, cache_bytes: 0, full_layers: 0, full_layers_set: false, preload_experts: if streaming { 4_096 } else { 0 }, }, EngineSteeringSettings { file: None, ffn_scale: 0.0, attention_scale: 0.0, }, ) .unwrap(); let cancelled = executor.prefill(&prompt, |_| false).unwrap(); assert!(cancelled < prompt.len()); executor.reset().unwrap(); assert_eq!(executor.prefill(&prompt, |_| true).unwrap(), prompt.len()); executor.eval(argmax(executor.logits())).unwrap(); assert_eq!(executor.position(), prompt.len() as u32 + 1); } #[test] #[ignore = "requires the installed 0731 Flash and checkpoint-specific DSpark GGUF fixtures"] fn dspark_runs_a_target_owned_greedy_cycle() { run_dspark_target_owned_greedy_cycle(false); } #[test] #[ignore = "requires the installed 0731 Flash and checkpoint-specific DSpark GGUF fixtures"] fn ssd_streaming_supports_dspark() { run_dspark_target_owned_greedy_cycle(true); } fn run_dspark_target_owned_greedy_cycle(streaming: bool) { use super::{DeepSeekExecutor, argmax, configure_sources}; use crate::engine::Model; use crate::engine::gguf::Gguf; use crate::engine::validation::validate_support; use crate::model::ModelChoice; use crate::settings::{ EngineSpeculativeSettings, EngineSsdSettings, EngineSteeringSettings, }; use std::sync::atomic::AtomicBool; configure_sources().unwrap(); let artifacts = installed_artifacts(ModelChoice::DeepSeekV4Flash0731, true); let main_path = std::env::var_os("DS4SERVER_BENCH_MODEL") .map(std::path::PathBuf::from) .unwrap_or(artifacts.model); let support_path = artifacts.support.unwrap(); let mut model = Model::open_main(&main_path, ModelChoice::DeepSeekV4Flash0731).unwrap(); let support = Gguf::open(&support_path).unwrap(); model.support_kind = Some(validate_support(&support, &model.shape).unwrap()); model.support = Some(support); let prompt = model.render_conversation( "", &[crate::engine::ChatTurn { user: true, tool: false, system: false, skip_previous_eos: false, reasoning: None, reasoning_complete: true, content: "hi".into(), }], crate::settings::ReasoningMode::Direct, ); let mut executor = DeepSeekExecutor::open( model, 64, false, 64, 100, EngineSpeculativeSettings { glm_mtp: false, glm_mtp_timing: false, dspark: true, dspark_confidence_threshold: 0.9, dspark_confidence_threshold_set: false, dspark_strict: false, dspark_exact_sampling: false, }, EngineSsdSettings { enabled: streaming, cold: false, cache_experts: if streaming { 4_096 } else { 0 }, cache_bytes: 0, full_layers: 0, full_layers_set: false, preload_experts: if streaming { 4_096 } else { 0 }, }, EngineSteeringSettings { file: None, ffn_scale: 0.0, attention_scale: 0.0, }, ) .unwrap(); executor.prefill(&prompt, |_| true).unwrap(); let mut generated = Vec::new(); while generated.len() < 8 { let first = argmax(executor.logits()); let cycle = executor .eval_speculative_greedy( first, (8 - generated.len()) as u32, crate::settings::ReasoningMode::Direct, &AtomicBool::new(false), ) .unwrap(); generated.extend(cycle); } let dspark = executor.dspark.as_ref().unwrap(); assert!(dspark.drafted > 0); if std::env::var_os("DS4SERVER_BENCH_MODEL").is_none() { assert!(dspark.accepted > 0); } executor.reset().unwrap(); executor.dspark.as_mut().unwrap().strict = true; executor.prefill(&prompt, |_| true).unwrap(); let mut target_only = Vec::new(); while target_only.len() < 8 { let first = argmax(executor.logits()); let cycle = executor .eval_speculative_greedy( first, (8 - target_only.len()) as u32, crate::settings::ReasoningMode::Direct, &AtomicBool::new(false), ) .unwrap(); assert_eq!(cycle.len(), 1); target_only.extend(cycle); } assert_eq!(target_only, generated); let mut turns = Vec::new(); for _ in 0..4 { for (user, content) in [(true, "ping"), (false, "pong")] { turns.push(crate::engine::ChatTurn { user, tool: false, system: false, skip_previous_eos: false, reasoning: None, reasoning_complete: true, content: content.into(), }); } } let ping_pong = executor .model .render_conversation("", &turns, crate::settings::ReasoningMode::Direct); assert!(ping_pong.len() > executor.dspark.as_ref().unwrap().config.block_size as usize + 1); executor.reset().unwrap(); executor.dspark.as_mut().unwrap().strict = false; executor.prefill(&ping_pong, |_| true).unwrap(); let first = argmax(executor.logits()); assert!( !executor .eval_speculative_greedy( first, 4, crate::settings::ReasoningMode::Direct, &AtomicBool::new(false), ) .unwrap() .is_empty() ); } #[test] #[ignore = "requires the installed 0731 target and checkpoint-specific DSpark GGUF fixtures"] fn flash_0731_runs_exact_sampled_dspark() { use super::{DeepSeekExecutor, argmax, configure_sources}; use crate::engine::gguf::Gguf; use crate::engine::validation::validate_support; use crate::engine::{Model, Rng}; use crate::model::{ModelChoice, validate_engine_artifacts}; use crate::settings::{ EngineSpeculativeSettings, EngineSsdSettings, EngineSteeringSettings, ReasoningMode, }; use std::sync::atomic::AtomicBool; configure_sources().unwrap(); let artifacts = installed_artifacts(ModelChoice::DeepSeekV4Flash0731, true); let main_path = std::env::var_os("DS4SERVER_BENCH_MODEL") .map(std::path::PathBuf::from) .unwrap_or_else(|| { validate_engine_artifacts(ModelChoice::DeepSeekV4Flash0731, true, &artifacts) .unwrap(); artifacts.model.clone() }); let mut model = Model::open_main(&main_path, ModelChoice::DeepSeekV4Flash0731).unwrap(); let support = Gguf::open(artifacts.support.as_ref().unwrap()).unwrap(); model.support_kind = Some(validate_support(&support, &model.shape).unwrap()); model.support = Some(support); let prompt = model.render_conversation( "", &[crate::engine::ChatTurn { user: true, tool: false, system: false, skip_previous_eos: false, reasoning: None, reasoning_complete: true, content: "hi".into(), }], ReasoningMode::Direct, ); let streaming = std::env::var_os("DS4SERVER_BENCH_SSD").is_some(); let mut executor = DeepSeekExecutor::open( model, 64, false, 64, 100, EngineSpeculativeSettings { glm_mtp: false, glm_mtp_timing: false, dspark: true, dspark_confidence_threshold: 0.6, dspark_confidence_threshold_set: false, dspark_strict: false, dspark_exact_sampling: true, }, EngineSsdSettings { enabled: streaming, cold: false, cache_experts: if streaming { 4_096 } else { 0 }, cache_bytes: 0, full_layers: 0, full_layers_set: false, preload_experts: if streaming { 4_096 } else { 0 }, }, EngineSteeringSettings { file: None, ffn_scale: 0.0, attention_scale: 0.0, }, ) .unwrap(); executor.prefill(&prompt, |_| true).unwrap(); let first = argmax(executor.logits()); let cycle = executor .eval_speculative_sampled( first, 4, ReasoningMode::Direct, 0.8, 0.95, 0.0, 0, &mut Rng::new(7), &AtomicBool::new(false), ) .unwrap(); assert!(!cycle.is_empty()); assert!(executor.dspark.as_ref().unwrap().drafted > 0); assert!(executor.logits().iter().all(|logit| logit.is_finite())); assert!(executor.session.position >= prompt.len() as u32 + cycle.len() as u32); } #[test] #[ignore = "requires the installed 81 GiB Flash GGUF fixture and a Metal device"] fn flash_0731_resident_and_ssd_streaming_choose_the_same_tokens() { use super::{DeepSeekExecutor, argmax, configure_sources}; use crate::engine::Model; use crate::model::ModelChoice; use crate::settings::{ EngineSpeculativeSettings, EngineSsdSettings, EngineSteeringSettings, }; use std::path::Path; struct FusionRollbacks; impl FusionRollbacks { const NAMES: &'static [&'static str] = &[ "DS4_METAL_DISABLE_PRE_M5_HC_NORM_MIX_FUSE", "DS4_METAL_DISABLE_PRE_M5_HC_PRODUCER_PRE_NORM_FUSE", "DS4_METAL_DISABLE_M5_HC_PRODUCER_PRE_NORM_FUSE", "DS4_METAL_DISABLE_PRE_M5_QKV_PAIR_QUAD_FUSE", "DS4_METAL_DISABLE_M5_QKV_PAIR_QUAD_FUSE", "DS4_METAL_DISABLE_PRE_M5_QKV_PAIR_COMPRESSOR_FUSE", "DS4_METAL_DISABLE_M5_QKV_PAIR_COMPRESSOR_FUSE", "DS4_METAL_DISABLE_PRE_M5_QKV_NORM_KV_STORE_FUSE", "DS4_METAL_DISABLE_PRE_M5_ATTN_INV_ROPE_FUSE", "DS4_METAL_DISABLE_PRE_M5_PARALLEL_FULL_FFN", "DS4_METAL_DISABLE_M5_PARALLEL_FULL_FFN", "DS4_METAL_DISABLE_PRE_M5_ROUTER_SHARED_FUSE", "DS4_METAL_DISABLE_M5_ROUTER_SHARED_FUSE", "DS4_METAL_DISABLE_M5_ROUTER_PROJECT_SELECT_FUSE", "DS4_METAL_DISABLE_PRE_M5_COMPRESSOR_QUAD_STORE", "DS4_METAL_DISABLE_PRE_M5_COMP_FINALIZE_FUSE", "DS4_METAL_DISABLE_M5_COMP_FINALIZE_FUSE", ]; fn activate() -> Self { for name in Self::NAMES { unsafe { std::env::set_var(name, "1") }; } Self } } impl Drop for FusionRollbacks { fn drop(&mut self) { for name in Self::NAMES { unsafe { std::env::remove_var(name) }; } } } fn run(path: &Path, streaming: bool) -> Vec { let model = Model::open_main(path, ModelChoice::DeepSeekV4Flash0731).unwrap(); let prompt = model.render_conversation( "", &[crate::engine::ChatTurn { user: true, tool: false, system: false, skip_previous_eos: false, reasoning: None, reasoning_complete: true, content: "hi".into(), }], crate::settings::ReasoningMode::Direct, ); assert_eq!(prompt.len(), 5); let speculative = EngineSpeculativeSettings { glm_mtp: false, glm_mtp_timing: false, dspark: false, dspark_confidence_threshold: 0.9, dspark_confidence_threshold_set: false, dspark_strict: false, dspark_exact_sampling: false, }; let mut executor = DeepSeekExecutor::open( model, 32, false, 32, 100, speculative, EngineSsdSettings { enabled: streaming, cold: true, cache_experts: if streaming { 16 } else { 0 }, cache_bytes: 0, full_layers: 0, full_layers_set: false, preload_experts: 0, }, EngineSteeringSettings { file: None, ffn_scale: 0.0, attention_scale: 0.0, }, ) .unwrap(); executor.prefill(&prompt, |_| true).unwrap(); let mut tokens = Vec::new(); for _ in 0..4 { let token = argmax(executor.logits()); tokens.push(token); executor.eval(token).unwrap(); } tokens } configure_sources().unwrap(); let path = installed_artifacts(ModelChoice::DeepSeekV4Flash0731, false).model; let resident = run(&path, false); let resident_rolled_back = { let _rollbacks = FusionRollbacks::activate(); run(&path, false) }; assert_eq!(resident, resident_rolled_back); let streaming = run(&path, true); let streaming_rolled_back = { let _rollbacks = FusionRollbacks::activate(); run(&path, true) }; assert_eq!(streaming, streaming_rolled_back); assert_eq!(resident, streaming); } #[test] #[ignore = "requires the installed 81 GiB Flash GGUF fixture and a Metal device"] fn flash_0731_ssd_streaming_maps_and_seeds_batched_prefill_layers() { use super::{DeepSeekExecutor, argmax, configure_sources}; use crate::engine::Model; use crate::model::ModelChoice; use crate::settings::{ EngineSpeculativeSettings, EngineSsdSettings, EngineSteeringSettings, }; configure_sources().unwrap(); let path = installed_artifacts(ModelChoice::DeepSeekV4Flash0731, false).model; let model = Model::open_main(&path, ModelChoice::DeepSeekV4Flash0731).unwrap(); let prompt = model.render_conversation( "", &[crate::engine::ChatTurn { user: true, tool: false, system: false, skip_previous_eos: false, reasoning: None, reasoning_complete: true, content: "hi ".repeat(80), }], crate::settings::ReasoningMode::Direct, ); assert!(prompt.len() > 64); let mut executor = DeepSeekExecutor::open( model, 512, false, 256, 100, EngineSpeculativeSettings { glm_mtp: false, glm_mtp_timing: false, dspark: false, dspark_confidence_threshold: 0.9, dspark_confidence_threshold_set: false, dspark_strict: false, dspark_exact_sampling: false, }, EngineSsdSettings { enabled: true, cold: false, cache_experts: 16, cache_bytes: 0, full_layers: 0, full_layers_set: false, preload_experts: 16, }, EngineSteeringSettings { file: None, ffn_scale: 0.0, attention_scale: 0.0, }, ) .unwrap(); let fallback_layer = executor .ssd .as_ref() .unwrap() .preload_by_layer .iter() .position(|entries| !entries.is_empty()) .unwrap(); let ssd = executor.ssd.as_ref().unwrap(); assert!( !ssd.seed_mapped_layer( &executor.model, &executor.weights.layers[fallback_layer], fallback_layer, true, ) .unwrap() ); assert!( ssd.seed_mapped_layer( &executor.model, &executor.weights.layers[fallback_layer], fallback_layer, false, ) .unwrap() ); assert_eq!(executor.prefill(&prompt, |_| true).unwrap(), prompt.len()); let stats = executor.execution_stats(); assert_eq!(stats.ssd_preloaded_experts, 16); assert!(stats.ssd_cache_entries > 0); executor.eval(argmax(executor.logits())).unwrap(); } #[test] #[ignore = "requires the installed 81 GiB Flash GGUF fixture and a Metal device"] fn resident_multi_session_switching_preserves_each_kv_frontier() { use super::{DeepSeekExecutor, argmax, configure_sources}; use crate::engine::Model; use crate::model::ModelChoice; use crate::settings::{ EngineSpeculativeSettings, EngineSsdSettings, EngineSteeringSettings, ReasoningMode, }; configure_sources().unwrap(); let path = installed_artifacts(ModelChoice::DeepSeekV4Flash0731, false).model; let model = Model::open_main(&path, ModelChoice::DeepSeekV4Flash0731).unwrap(); let prompts = ["Reply with A.", "Reply with B."].map(|content| { model.render_conversation( "", &[crate::engine::ChatTurn { user: true, tool: false, system: false, skip_previous_eos: false, reasoning: None, reasoning_complete: true, content: content.into(), }], ReasoningMode::Direct, ) }); let mut executor = DeepSeekExecutor::open( model, 64, false, 64, 100, EngineSpeculativeSettings { glm_mtp: false, glm_mtp_timing: false, dspark: false, dspark_confidence_threshold: 0.9, dspark_confidence_threshold_set: false, dspark_strict: false, dspark_exact_sampling: false, }, EngineSsdSettings { enabled: false, cold: false, cache_experts: 0, cache_bytes: 0, full_layers: 0, full_layers_set: false, preload_experts: 0, }, EngineSteeringSettings { file: None, ffn_scale: 0.0, attention_scale: 0.0, }, ) .unwrap(); executor.prefill(&prompts[0], |_| true).unwrap(); let first = (executor.tokens().to_vec(), argmax(executor.logits())); let mut inactive = None; executor.swap_resident_state(&mut inactive).unwrap(); executor.prefill(&prompts[1], |_| true).unwrap(); let second = (executor.tokens().to_vec(), argmax(executor.logits())); executor.swap_resident_state(&mut inactive).unwrap(); assert_eq!( (executor.tokens(), argmax(executor.logits())), (&*first.0, first.1) ); executor.swap_resident_state(&mut inactive).unwrap(); assert_eq!( (executor.tokens(), argmax(executor.logits())), (&*second.0, second.1) ); } #[test] #[ignore = "requires the installed Flash GGUF, DS4 steering fixture, and a Metal device"] fn directional_steering_matches_the_ds4_token_oracle() { use super::{DeepSeekExecutor, argmax, configure_sources}; use crate::engine::Model; use crate::model::ModelChoice; use crate::settings::{ EngineSpeculativeSettings, EngineSsdSettings, EngineSteeringSettings, }; configure_sources().unwrap(); let model_path = installed_artifacts(ModelChoice::DeepSeekV4Flash0731, false).model; let steering_path = std::env::var("DS4_STEERING_FILE") .expect("set DS4_STEERING_FILE to the DS4 verbosity direction fixture"); let cases = [ (1.0, 0.0, [19_923, 3, 1_730, 588, 342, 8_233, 440, 4_316]), (0.0, 1.0, [19_923, 3, 1_730, 588, 342, 1_694, 440, 4_316]), ]; for (ffn_scale, attention_scale, expected) in cases { let model = Model::open_main(&model_path, ModelChoice::DeepSeekV4Flash0731).unwrap(); let prompt = model.render_conversation( "", &[crate::engine::ChatTurn { user: true, tool: false, system: false, skip_previous_eos: false, reasoning: None, reasoning_complete: true, content: "hi".into(), }], crate::settings::ReasoningMode::Direct, ); let mut executor = DeepSeekExecutor::open( model, 32, false, 32, 100, EngineSpeculativeSettings { glm_mtp: false, glm_mtp_timing: false, dspark: false, dspark_confidence_threshold: 0.9, dspark_confidence_threshold_set: false, dspark_strict: false, dspark_exact_sampling: false, }, EngineSsdSettings { enabled: false, cold: false, cache_experts: 0, cache_bytes: 0, full_layers: 0, full_layers_set: false, preload_experts: 0, }, EngineSteeringSettings { file: Some(steering_path.clone()), ffn_scale, attention_scale, }, ) .unwrap(); executor.prefill(&prompt, |_| true).unwrap(); let mut tokens = Vec::new(); for _ in 0..8 { let token = argmax(executor.logits()); tokens.push(token); executor.eval(token).unwrap(); } assert_eq!(tokens, expected); } } #[test] #[ignore = "requires the installed DeepSeek V4 Pro GGUF and Apple Metal"] fn pro_ssd_streaming_full_layer_and_selected_batch_agree() { use super::{DeepSeekExecutor, argmax, configure_sources}; use crate::engine::Model; use crate::model::ModelChoice; use crate::settings::{ EngineSpeculativeSettings, EngineSsdSettings, EngineSteeringSettings, }; configure_sources().unwrap(); let path = installed_artifacts(ModelChoice::DeepSeekV4Pro, false).model; if !path.is_file() { eprintln!("skipping unavailable Pro fixture: {}", path.display()); return; } let run = |cache_experts| { let model = Model::open_main(&path, ModelChoice::DeepSeekV4Pro).unwrap(); let prompt = model.render_conversation( "", &[crate::engine::ChatTurn { user: true, tool: false, system: false, skip_previous_eos: false, reasoning: None, reasoning_complete: true, content: "hi ".repeat(32), }], crate::settings::ReasoningMode::Direct, ); assert!(prompt.len() > 18); let mut executor = DeepSeekExecutor::open( model, 256, false, 256, 100, EngineSpeculativeSettings { glm_mtp: false, glm_mtp_timing: false, dspark: false, dspark_confidence_threshold: 0.9, dspark_confidence_threshold_set: false, dspark_strict: false, dspark_exact_sampling: false, }, EngineSsdSettings { enabled: true, cold: true, cache_experts, cache_bytes: 0, full_layers: 0, full_layers_set: false, preload_experts: 0, }, EngineSteeringSettings { file: None, ffn_scale: 0.0, attention_scale: 0.0, }, ) .unwrap(); executor.prefill(&prompt, |_| true).unwrap(); (0..4) .map(|_| { let token = argmax(executor.logits()); executor.eval(token).unwrap(); token }) .collect::>() }; assert_eq!(run(32), run(384)); } }