use super::checkpoint::{read_buffer, read_u32, read_u64, write_buffer, write_u32, write_u64}; use super::*; use crate::settings::EngineSsdSettings; const CHECKPOINT_MAGIC: &[u8; 8] = b"DS4GLM01"; const CHECKPOINT_VERSION: u32 = 1; const CACHE_F16: bool = true; const DECODE_FLUSH_LAYERS: usize = 4; const AUTO_CACHE_BYTES: u64 = 12 * 1024 * 1024 * 1024; const STREAMING_TOKEN_PREFILL_MAX: u32 = 64; fn live_prefix_rewind_target(live: &[i32], incoming: &[i32]) -> Option { (incoming.len() > 1 && incoming.len() < live.len() && live.starts_with(incoming)) .then_some(incoming.len() - 1) } const STREAMING_FULL_ATTN_CONTEXT: u32 = 8192; const LONG_CONTEXT_THRESHOLD: u32 = 65_536; const LONG_CONTEXT_FULL_ATTN_CONTEXT: u32 = 4096; const INDEXED_PREFILL_CHUNK: u32 = 4096; const INDEXED_PREFILL_SCORE_BYTES: u64 = 256 * 1024 * 1024; const INDEXED_PREFILL_ATTN_SLICE: u32 = 2048; #[derive(Clone, Copy)] struct SparseWeights { router: Weight, bias: Weight, gate: Weight, up: Weight, down: Weight, shared_gate: Weight, shared_up: Weight, shared_down: Weight, } #[derive(Clone, Copy)] struct DenseWeights { gate: Weight, up: Weight, down: Weight, } #[derive(Clone, Copy)] struct NextnWeights { eh_proj: Weight, enorm: Weight, hnorm: Weight, shared_head_norm: Weight, } struct GlmLayer { attn_norm: Weight, q_a: Weight, q_a_norm: Weight, q_b: Weight, kv_a: Weight, kv_norm: Weight, k_b: Weight, v_b: Weight, output: Weight, indexer_k: Weight, indexer_q: Weight, indexer_k_norm: Weight, indexer_k_bias: Weight, indexer_proj: Weight, ffn_norm: Weight, dense: Option, sparse: Option, nextn: Option, } struct GlmWeights { embedding: Weight, output_norm: Weight, output: Weight, layers: Vec, nextn: Option, } impl GlmWeights { fn bind(model: &Model) -> Result { let main = &model.main; let normal_layers = model.shape.layers - model.shape.nextn; let mut layers: Vec = (0..model.shape.layers) .map(|index| { let required = |suffix: &str| Weight::bind(main, &format!("blk.{index}.{suffix}")); let dense = (index < model.shape.leading_dense) .then(|| { Ok::<_, String>(DenseWeights { gate: required("ffn_gate.weight")?, up: required("ffn_up.weight")?, down: required("ffn_down.weight")?, }) }) .transpose()?; let sparse = (index >= model.shape.leading_dense) .then(|| { Ok::<_, String>(SparseWeights { router: required("ffn_gate_inp.weight")?, bias: required("exp_probs_b.bias")?, gate: required("ffn_gate_exps.weight")?, up: required("ffn_up_exps.weight")?, 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")?, }) }) .transpose()?; let nextn = (index >= normal_layers) .then(|| { Ok::<_, String>(NextnWeights { eh_proj: required("nextn.eh_proj.weight")?, enorm: required("nextn.enorm.weight")?, hnorm: required("nextn.hnorm.weight")?, shared_head_norm: required("nextn.shared_head_norm.weight")?, }) }) .transpose()?; Ok(GlmLayer { attn_norm: required("attn_norm.weight")?, q_a: required("attn_q_a.weight")?, q_a_norm: required("attn_q_a_norm.weight")?, q_b: required("attn_q_b.weight")?, kv_a: required("attn_kv_a_mqa.weight")?, kv_norm: required("attn_kv_a_norm.weight")?, k_b: required("attn_k_b.weight")?, v_b: required("attn_v_b.weight")?, output: required("attn_output.weight")?, indexer_k: required("indexer.attn_k.weight")?, indexer_q: required("indexer.attn_q_b.weight")?, indexer_k_norm: required("indexer.k_norm.weight")?, indexer_k_bias: required("indexer.k_norm.bias")?, indexer_proj: required("indexer.proj.weight")?, ffn_norm: required("ffn_norm.weight")?, dense, sparse, nextn, }) }) .collect::>()?; let nextn = (model.shape.nextn != 0) .then(|| layers.pop().expect("validated GLM nextn layer disappeared")); Ok(Self { embedding: Weight::bind(main, "token_embd.weight")?, output_norm: Weight::bind(main, "output_norm.weight")?, output: Weight::bind(main, "output.weight")?, layers, nextn, }) } } struct LayerCache { kv: Buffer, rope: Buffer, indexer: Option, } impl LayerCache { fn allocate(shape: super::super::Shape, layer: usize, context: u32) -> Result { Ok(Self { kv: Buffer::bytes(u64::from(context) * shape.kv_lora * 2)?, rope: Buffer::bytes(u64::from(context) * shape.rot * 2)?, indexer: full_indexer_layer(shape, layer) .then(|| Buffer::bytes(u64::from(context) * shape.indexer_head_dim * 2)) .transpose()?, }) } } struct GlmScratch { current: Buffer, next: Buffer, attn_norm: Buffer, q_rank: Buffer, q_rank_norm: Buffer, q: Buffer, kv_raw: Buffer, indexer_k: Buffer, indexer_q: Buffer, indexer_weights: Buffer, indexer_scores: Buffer, indexer_selected: Buffer, qk_low: Buffer, heads: Buffer, attn_out: Buffer, after_attn: Buffer, ffn_norm: Buffer, ffn_gate: Buffer, ffn_up: Buffer, ffn_mid: Buffer, routed_down: Buffer, ffn_out: Buffer, ffn_sum: Buffer, router_logits: Buffer, router_probs: Buffer, router_selected: Buffer, router_weights: Buffer, output_norm: Buffer, logits: Buffer, } struct GlmBatchScratch { tokens: Buffer, current: Buffer, next: Buffer, attn_norm: Buffer, q_rank: Buffer, q_rank_norm: Buffer, q: Buffer, kv_raw: Buffer, kv_norm: Buffer, indexer_k: Buffer, indexer_q: Buffer, indexer_weights: Buffer, indexer_scores: Buffer, indexer_selected: Buffer, qk_low: Buffer, attn_lora: Buffer, heads: Buffer, attn_out: Buffer, after_attn: Buffer, ffn_norm: Buffer, ffn_gate: Buffer, ffn_up: Buffer, ffn_mid: Buffer, routed_gate: Buffer, routed_up: Buffer, routed_down: Buffer, ffn_out: Buffer, router_logits: Buffer, router_probs: Buffer, router_selected: Buffer, router_weights: Buffer, } struct GlmMtp { cache: LayerCache, selected: Buffer, concat: Buffer, pending: Option, min_pos: Option, cycles: u64, drafted: u64, accepted: u64, verifier_passes: u64, verifier_ns: u64, timing: bool, } #[derive(Clone, Copy)] struct GlmStreamingPlan { settings: EngineSsdSettings, per_expert: u64, cache_experts: u32, planned_expert_bytes: u64, } impl GlmScratch { fn allocate(model: &Model, context: u32) -> Result { let shape = model.shape; let q = shape.heads * shape.key_mla; let heads = shape.heads * shape.value_mla; Ok(Self { current: Buffer::floats(shape.embd)?, next: Buffer::floats(shape.embd)?, attn_norm: Buffer::floats(shape.embd)?, q_rank: Buffer::floats(shape.lora_q)?, q_rank_norm: Buffer::floats(shape.lora_q)?, q: Buffer::floats(q)?, kv_raw: Buffer::floats(shape.head_dim)?, indexer_k: Buffer::floats(shape.indexer_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))?, indexer_selected: Buffer::bytes(shape.indexer_top_k * 4)?, qk_low: Buffer::floats(shape.heads * shape.kv_lora)?, heads: Buffer::floats(heads)?, attn_out: Buffer::floats(shape.embd)?, after_attn: Buffer::floats(shape.embd)?, ffn_norm: Buffer::floats(shape.embd)?, ffn_gate: Buffer::floats(shape.ff_dense.max(shape.experts_used * shape.ff_expert))?, ffn_up: Buffer::floats(shape.ff_dense.max(shape.experts_used * shape.ff_expert))?, ffn_mid: Buffer::floats(shape.ff_dense.max(shape.experts_used * shape.ff_expert))?, routed_down: Buffer::floats(shape.experts_used * shape.embd)?, ffn_out: Buffer::floats(shape.embd)?, ffn_sum: 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)?, output_norm: Buffer::floats(shape.embd)?, logits: Buffer::floats(shape.vocab)?, }) } } impl GlmBatchScratch { fn allocate(model: &Model, rows: u32, context: u32) -> Result { let shape = model.shape; let rows = u64::from(rows); let q = shape.heads * shape.key_mla; let qk_low = shape.heads * shape.kv_lora; let heads = shape.heads * shape.value_mla; let hidden = shape.ff_dense.max(shape.ff_expert); let routed_mid = shape.experts_used * shape.ff_expert; let score_rows = (INDEXED_PREFILL_SCORE_BYTES / (u64::from(context) * 4)) .max(1) .min(rows); Ok(Self { tokens: Buffer::bytes(rows * 4)?, current: Buffer::floats(rows * shape.embd)?, next: Buffer::floats(rows * shape.embd)?, attn_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)?, kv_raw: Buffer::floats(rows * shape.head_dim)?, kv_norm: Buffer::floats(rows * shape.kv_lora)?, indexer_k: Buffer::floats(rows * shape.indexer_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(score_rows * u64::from(context))?, indexer_selected: Buffer::bytes(rows * shape.indexer_top_k * 4)?, qk_low: Buffer::floats(rows * qk_low)?, attn_lora: Buffer::floats(rows * qk_low)?, heads: Buffer::floats(rows * heads)?, attn_out: Buffer::floats(rows * shape.embd)?, after_attn: Buffer::floats(rows * shape.embd)?, ffn_norm: Buffer::floats(rows * shape.embd)?, ffn_gate: Buffer::floats(rows * hidden)?, ffn_up: Buffer::floats(rows * hidden)?, ffn_mid: Buffer::floats(rows * hidden.max(routed_mid))?, routed_gate: Buffer::floats(rows * routed_mid)?, routed_up: Buffer::floats(rows * routed_mid)?, routed_down: Buffer::floats(rows * shape.experts_used * shape.embd)?, ffn_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)?, }) } fn score_rows(&self, context: u32, rows: u32) -> u32 { u32::try_from( (INDEXED_PREFILL_SCORE_BYTES / (u64::from(context) * 4)) .max(1) .min(u64::from(rows)), ) .expect("GLM score row cap is bounded by the batch") } } pub(in crate::engine) struct GlmExecutor { weights: GlmWeights, scratch: GlmScratch, caches: Vec, logits: Vec, tokens: Vec, context: u32, quality: bool, ssd: EngineSsdSettings, profile: Option, mtp: Option, ssd_resident_bytes: u64, ssd_cache_bytes: u64, ssd_cache_experts: u64, ssd_preloaded_experts: u64, checkpoint_tag: [u8; 32], model_modified: (u64, u32), model_identity: [u8; 32], streaming_spans: Option>, _context: Context, model: Model, } pub(in crate::engine) struct GlmResidentState { scratch: GlmScratch, caches: Vec, logits: Vec, tokens: Vec, checkpoint_tag: [u8; 32], mtp: Option, } impl GlmExecutor { #[allow(dead_code)] pub(super) fn open( model: Model, context: u32, quality: bool, ssd: EngineSsdSettings, ) -> Result { Self::open_profile( model, context, quality, ssd, 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, }, None, ) } pub(super) fn open_profile( model: Model, context: u32, quality: bool, ssd: EngineSsdSettings, speculative: EngineSpeculativeSettings, expert_profile_path: Option<&str>, ) -> Result { if context == 0 || u64::from(context) > model.shape.original_context { return Err(format!( "GLM context must be between 1 and {} tokens", model.shape.original_context )); } let weights = GlmWeights::bind(&model)?; let streaming = glm_streaming_plan(&model, &weights, ssd)?; let effective_ssd = streaming.map_or(ssd, |plan| plan.settings); let admission = admission_bytes(&model, &weights, context, streaming.as_ref())?; let model_spans = streaming .as_ref() .map(|plan| glm_streaming_model_spans(&model, &weights, plan)) .transpose()?; let context_spans = model_spans .as_ref() .map(|(spans, max_tensor_bytes)| (spans.as_slice(), *max_tensor_bytes)); let context_handle = Context::open( &model, quality, effective_ssd.enabled, admission, context_spans, )?; let model_spans = model_spans.map(|(spans, _)| spans); configure_streaming(&model, &weights, streaming.as_ref())?; let scratch = GlmScratch::allocate(&model, context)?; let caches = (0..weights.layers.len()) .map(|layer| LayerCache::allocate(model.shape, layer, context)) .collect::>()?; 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 = model.checkpoint_identity(); let profile = ExpertProfile::new( expert_profile_path, model.shape.model, model.shape.layers, model.shape.experts, model.shape.experts_used, )?; let mtp = speculative .glm_mtp .then(|| { Ok::<_, String>(GlmMtp { cache: LayerCache::allocate( model.shape, (model.shape.layers - 1) as usize, context, )?, selected: Buffer::bytes(u64::from(context) * 4)?, concat: Buffer::floats(2 * model.shape.embd)?, pending: None, min_pos: None, cycles: 0, drafted: 0, accepted: 0, verifier_passes: 0, verifier_ns: 0, timing: speculative.glm_mtp_timing, }) }) .transpose()?; let ssd_resident_bytes = if let Some(spans) = &model_spans { spans.iter().map(|span| span.1).sum() } else { 0 }; let (ssd_cache_bytes, ssd_cache_experts, ssd_preloaded_experts) = streaming .map(|plan| { let preload = preload_count(plan.settings, plan.cache_experts); ( plan.per_expert .saturating_mul(u64::from(plan.cache_experts)), u64::from(plan.cache_experts), u64::from(preload), ) }) .unwrap_or_default(); Ok(Self { weights, scratch, caches, logits: vec![0.0; model.shape.vocab as usize], tokens: Vec::new(), context, quality, ssd: effective_ssd, profile, mtp, ssd_resident_bytes, ssd_cache_bytes, ssd_cache_experts, ssd_preloaded_experts, checkpoint_tag: [0; 32], model_modified, model_identity, streaming_spans: model_spans, _context: context_handle, model, }) } pub(super) fn eval(&mut self, token: i32) -> Result<(), String> { let shape = self.model.shape; if token < 0 || token as u64 >= shape.vocab { return Err(format!("token {token} is outside the vocabulary")); } let pos = u32::try_from(self.tokens.len()).map_err(|_| "GLM position overflow")?; if pos >= self.context { return Err(format!( "the GLM Metal executor supports {} tokens", self.context )); } let map = self.model.main.map_ptr().cast(); let size = self.model.main.len(); let commands = Commands::begin()?; call( unsafe { ds4_gpu_embed_token_quant_tensor( self.scratch.current.raw(), map, size, self.weights.embedding.offset, self.weights.embedding.kind, shape.vocab as u32, token as u32, shape.embd as u32, ) }, "GLM token embedding", )?; let mut selected_count = 0; for (layer_index, ((layer, cache), ordinal)) in self .weights .layers .iter() .zip(&self.caches) .zip(0_u32..) .enumerate() { self.encode_layer( layer, cache, layer_index, ordinal, pos, &mut selected_count, None, )?; if layer.sparse.is_some() && let Some(profile) = &mut self.profile { profile.record( ordinal as usize, pos, &self.scratch.router_selected, &self.scratch.router_weights, 1, false, )?; } std::mem::swap(&mut self.scratch.current, &mut self.scratch.next); if (layer_index + 1).is_multiple_of(DECODE_FLUSH_LAYERS) && layer_index + 1 < self.weights.layers.len() { call( unsafe { ds4_gpu_flush_commands() }, "flushing the GLM decode graph", )?; } } norm( &self.scratch.output_norm, &self.scratch.current, self.weights.output_norm, shape.embd as u32, shape.rms_epsilon, map, size, )?; project( &self.scratch.logits, self.weights.output, shape.embd, shape.vocab, &self.scratch.output_norm, map, size, self.ssd.enabled, )?; commands.finish()?; self.scratch.logits.read_f32(&mut self.logits)?; self.tokens.push(token); if let Some(profile) = &self.profile { profile.write()?; } Ok(()) } #[allow(clippy::too_many_arguments)] fn encode_batch_layer( &self, batch: &GlmBatchScratch, layer: &GlmLayer, cache: &LayerCache, ordinal: u32, pos: u32, rows: u32, selected_ready: &mut bool, selected_count: &mut u32, ) -> Result<(), String> { let shape = self.model.shape; let map = self.model.main.map_ptr().cast(); let size = self.model.main.len(); let q_dim = shape.heads * shape.key_mla; let q_nope = shape.key_mla - shape.rot; norm_rows( &batch.attn_norm, &batch.current, layer.attn_norm, shape.embd as u32, rows, shape.rms_epsilon, map, size, )?; project_rows( &batch.q_rank, layer.q_a, shape.embd, shape.lora_q, &batch.attn_norm, rows, map, size, )?; norm_rows( &batch.q_rank_norm, &batch.q_rank, layer.q_a_norm, shape.lora_q as u32, rows, shape.rms_epsilon, map, size, )?; project_rows( &batch.q, layer.q_b, shape.lora_q, q_dim, &batch.q_rank_norm, rows, map, size, )?; call( unsafe { ds4_gpu_glm_rope_tail_tensor( batch.q.raw(), rows, shape.heads as u32, shape.key_mla as u32, shape.rot as u32, pos, 0, shape.rope_base, 1.0, 0.0, 1.0, 0.0, 0.0, ) }, "applying batched GLM query RoPE", )?; if full_indexer_layer(shape, ordinal as usize) { let indexer_cache = cache .indexer .as_ref() .ok_or("GLM full-indexer layer is missing its key cache")?; project_rows( &batch.indexer_k, layer.indexer_k, shape.embd, shape.indexer_head_dim, &batch.current, rows, map, size, )?; call( unsafe { ds4_gpu_glm_store_indexer_k_tensor( indexer_cache.raw(), batch.indexer_k.raw(), map, size, layer.indexer_k_norm.offset, layer.indexer_k_bias.offset, pos, rows, self.context, shape.indexer_head_dim as u32, shape.rot as u32, 0, 1.0e-6, shape.rope_base, 1.0, 0.0, 1.0, 0.0, 0.0, CACHE_F16, ) }, "storing batched GLM indexer keys", )?; } project_rows( &batch.kv_raw, layer.kv_a, shape.embd, shape.head_dim, &batch.attn_norm, rows, map, size, )?; call( unsafe { ds4_gpu_glm_kv_lora_rms_norm_tensor( batch.kv_norm.raw(), batch.kv_raw.raw(), map, size, layer.kv_norm.offset, rows, shape.head_dim as u32, shape.kv_lora as u32, shape.rms_epsilon, ) }, "normalizing batched GLM compact KV", )?; call( unsafe { ds4_gpu_glm_store_compact_kv_tensor( cache.kv.raw(), cache.rope.raw(), batch.kv_norm.raw(), batch.kv_raw.raw(), pos, rows, self.context, shape.head_dim as u32, shape.kv_lora as u32, shape.rot as u32, CACHE_F16, ) }, "storing batched GLM compact KV", )?; let visible = pos + rows; let top_k = shape.indexer_top_k as u32; *selected_count = visible.min(top_k); if full_indexer_layer(shape, ordinal as usize) { if visible <= top_k { call( unsafe { ds4_gpu_glm_fill_selected_range_batch_tensor( batch.indexer_selected.raw(), rows, pos, *selected_count, self.context, ) }, "selecting the causal GLM prefill range", )?; } else { self.select_batch_indexer_rows(batch, layer, cache, pos, rows, visible)?; } *selected_ready = true; } else if !*selected_ready { return Err("GLM indexed prefill reached a non-indexer layer before selection".into()); } call( unsafe { ds4_gpu_glm_qk_lowrank_typed_batch_tensor( batch.qk_low.raw(), batch.q.raw(), map, size, layer.k_b.offset, layer.k_b.kind, rows, shape.heads as u32, shape.kv_lora as u32, q_nope as u32, shape.key_mla as u32, ) }, "projecting batched GLM low-rank queries", )?; self.encode_batch_attention( batch, layer, cache, pos, rows, *selected_count, visible <= top_k, )?; project_rows( &batch.attn_out, layer.output, shape.heads * shape.value_mla, shape.embd, &batch.heads, rows, map, size, )?; let residual = rows .checked_mul(shape.embd as u32) .ok_or("GLM prefill residual size overflow")?; call( unsafe { ds4_gpu_add_tensor( batch.after_attn.raw(), batch.current.raw(), batch.attn_out.raw(), residual, ) }, "adding the batched GLM attention residual", )?; self.encode_batch_ffn(batch, layer, ordinal, rows, residual) } fn select_batch_indexer_rows( &self, batch: &GlmBatchScratch, layer: &GlmLayer, cache: &LayerCache, pos: u32, rows: u32, visible: u32, ) -> Result<(), String> { let shape = self.model.shape; let map = self.model.main.map_ptr().cast(); let size = self.model.main.len(); project_rows( &batch.indexer_q, layer.indexer_q, shape.lora_q, shape.indexer_heads * shape.indexer_head_dim, &batch.q_rank_norm, rows, map, size, )?; call( unsafe { ds4_gpu_glm_indexer_rope_tail_tensor( batch.indexer_q.raw(), rows, shape.indexer_heads as u32, shape.indexer_head_dim as u32, shape.rot as u32, pos, 0, shape.rope_base, 1.0, 0.0, 1.0, 0.0, 0.0, ) }, "applying batched GLM indexer RoPE", )?; f32_project_rows( &batch.indexer_weights, layer.indexer_proj, shape.embd, shape.indexer_heads, &batch.current, rows, map, size, )?; let cache = cache .indexer .as_ref() .ok_or("GLM full-indexer layer is missing its key cache")?; let score_rows = batch.score_rows(self.context, rows); let q_row_bytes = shape.indexer_heads * shape.indexer_head_dim * 4; let weight_row_bytes = shape.indexer_heads * 4; let selected_row_bytes = shape.indexer_top_k * 4; let scale = 1.0 / ((shape.indexer_heads * shape.indexer_head_dim) as f32).sqrt(); let mut start = 0; while start < rows { let slice = (rows - start).min(score_rows); let q = batch.indexer_q.view( u64::from(start) * q_row_bytes, u64::from(slice) * q_row_bytes, )?; let weights = batch.indexer_weights.view( u64::from(start) * weight_row_bytes, u64::from(slice) * weight_row_bytes, )?; let selected = batch.indexer_selected.view( u64::from(start) * selected_row_bytes, u64::from(slice) * selected_row_bytes, )?; call( unsafe { ds4_gpu_glm_indexer_scores_batch_tensor( batch.indexer_scores.raw(), q.raw(), weights.raw(), cache.raw(), visible, slice, pos + start, shape.indexer_heads as u32, shape.indexer_head_dim as u32, scale, CACHE_F16, ) }, "scoring batched GLM indexer rows", )?; call( unsafe { ds4_gpu_indexer_topk_tensor( selected.raw(), batch.indexer_scores.raw(), visible, slice, shape.indexer_top_k as u32, ) }, "selecting batched GLM indexed-attention rows", )?; start += slice; } Ok(()) } #[allow(clippy::too_many_arguments)] fn encode_batch_attention( &self, batch: &GlmBatchScratch, layer: &GlmLayer, cache: &LayerCache, pos: u32, rows: u32, selected_count: u32, causal_range: bool, ) -> Result<(), String> { let shape = self.model.shape; let map = self.model.main.map_ptr().cast(); let size = self.model.main.len(); let q_dim = shape.heads * shape.key_mla; let qk_low_dim = shape.heads * shape.kv_lora; let heads_dim = shape.heads * shape.value_mla; let selected_row_bytes = u64::from(selected_count) * 4; let mut start = 0; while start < rows { let slice = (rows - start).min(INDEXED_PREFILL_ATTN_SLICE); let q = batch .q .view(u64::from(start) * q_dim * 4, u64::from(slice) * q_dim * 4)?; let qk_low = batch.qk_low.view( u64::from(start) * qk_low_dim * 4, u64::from(slice) * qk_low_dim * 4, )?; let lora = batch.attn_lora.view( u64::from(start) * qk_low_dim * 4, u64::from(slice) * qk_low_dim * 4, )?; let heads = batch.heads.view( u64::from(start) * heads_dim * 4, u64::from(slice) * heads_dim * 4, )?; if causal_range { call( unsafe { ds4_gpu_glm_attention_indexed_batch_lora_causal_tensor( lora.raw(), q.raw(), qk_low.raw(), cache.kv.raw(), cache.rope.raw(), slice, pos + start, selected_count, self.context, CACHE_F16, shape.heads as u32, shape.kv_lora as u32, (shape.key_mla - shape.rot) as u32, shape.rot as u32, 0, shape.rope_base, 1.0, 0.0, 1.0, 0.0, 0.0, ) }, "running causal batched GLM indexed attention", )?; } else { let selected = batch.indexer_selected.view( u64::from(start) * selected_row_bytes, u64::from(slice) * selected_row_bytes, )?; call( unsafe { ds4_gpu_glm_attention_indexed_batch_lora_valid_tensor( lora.raw(), q.raw(), qk_low.raw(), cache.kv.raw(), cache.rope.raw(), selected.raw(), slice, selected_count, self.context, CACHE_F16, shape.heads as u32, shape.kv_lora as u32, (shape.key_mla - shape.rot) as u32, shape.rot as u32, 0, shape.rope_base, 1.0, 0.0, 1.0, 0.0, 0.0, ) }, "running batched GLM indexed attention", )?; } call( unsafe { ds4_gpu_glm_value_project_typed_batch_heads_tensor( heads.raw(), lora.raw(), map, size, layer.v_b.offset, layer.v_b.kind, slice, shape.heads as u32, shape.kv_lora as u32, shape.value_mla as u32, ) }, "projecting batched GLM attention values", )?; start += slice; } Ok(()) } fn encode_batch_ffn( &self, batch: &GlmBatchScratch, layer: &GlmLayer, ordinal: u32, rows: u32, residual: u32, ) -> Result<(), String> { let shape = self.model.shape; let map = self.model.main.map_ptr().cast(); let size = self.model.main.len(); norm_rows( &batch.ffn_norm, &batch.after_attn, layer.ffn_norm, shape.embd as u32, rows, shape.rms_epsilon, map, size, )?; if let Some(dense) = layer.dense { project_rows( &batch.ffn_gate, dense.gate, shape.embd, shape.ff_dense, &batch.ffn_norm, rows, map, size, )?; project_rows( &batch.ffn_up, dense.up, shape.embd, shape.ff_dense, &batch.ffn_norm, rows, map, size, )?; call( unsafe { ds4_gpu_swiglu_tensor( batch.ffn_mid.raw(), batch.ffn_gate.raw(), batch.ffn_up.raw(), rows * shape.ff_dense as u32, 0.0, 1.0, ) }, "activating the batched dense GLM FFN", )?; project_rows( &batch.ffn_out, dense.down, shape.ff_dense, shape.embd, &batch.ffn_mid, rows, map, size, )?; return call( unsafe { ds4_gpu_add_tensor( batch.next.raw(), batch.after_attn.raw(), batch.ffn_out.raw(), residual, ) }, "adding the batched dense GLM residual", ); } let sparse = layer.sparse.ok_or("GLM layer has no FFN weights")?; f32_project_rows( &batch.router_logits, sparse.router, shape.embd, shape.experts, &batch.ffn_norm, rows, map, size, )?; call( unsafe { ds4_gpu_glm_router_select_batch_tensor( batch.router_selected.raw(), batch.router_weights.raw(), batch.router_probs.raw(), map, size, sparse.bias.offset, batch.router_logits.raw(), shape.experts as u32, shape.experts_used as u32, shape.expert_weight_scale, rows, ) }, "routing a batched GLM FFN", )?; let (gate_expert, gate_row) = expert_layout(sparse.gate, shape.experts); let (up_expert, up_row) = expert_layout(sparse.up, shape.experts); let (down_expert, down_row) = expert_layout(sparse.down, shape.experts); let force_resident = self.ssd.enabled && ordinal.saturating_sub(shape.leading_dense) < self.ssd.full_layers; let mut mid_is_f16 = false; let routed = unsafe { if sparse.gate.kind == IQ2_XXS { ds4_gpu_routed_moe_batch_tensor( batch.ffn_out.raw(), batch.routed_gate.raw(), batch.routed_up.raw(), batch.ffn_mid.raw(), batch.routed_down.raw(), map, size, sparse.gate.offset, sparse.up.offset, sparse.down.offset, sparse.gate.kind, sparse.down.kind, gate_expert, gate_row, down_expert, down_row, shape.embd as u32, shape.ff_expert as u32, shape.embd as u32, batch.router_selected.raw(), batch.router_weights.raw(), shape.experts as u32, shape.experts_used as u32, 0.0, batch.ffn_norm.raw(), ordinal, rows, &mut mid_is_f16, force_resident, ) } else { ds4_gpu_glm_routed_moe_batch_tensor( batch.ffn_out.raw(), batch.ffn_mid.raw(), map, size, sparse.gate.offset, sparse.up.offset, sparse.down.offset, sparse.gate.kind, sparse.up.kind, sparse.down.kind, gate_expert, gate_row, up_expert, up_row, down_expert, down_row, shape.embd as u32, shape.ff_expert as u32, shape.embd as u32, batch.router_selected.raw(), batch.router_weights.raw(), shape.experts as u32, shape.experts_used as u32, ordinal, batch.ffn_norm.raw(), rows, (shape.experts_used * shape.ff_expert) as u32, force_resident, ) } }; call(routed, "running batched routed GLM experts") .map_err(|error| format!("{error} in layer {ordinal}"))?; project_rows( &batch.ffn_gate, sparse.shared_gate, shape.embd, shape.ff_expert, &batch.ffn_norm, rows, map, size, )?; project_rows( &batch.ffn_up, sparse.shared_up, shape.embd, shape.ff_expert, &batch.ffn_norm, rows, map, size, )?; call( unsafe { ds4_gpu_swiglu_tensor( batch.ffn_mid.raw(), batch.ffn_gate.raw(), batch.ffn_up.raw(), rows * shape.ff_expert as u32, 0.0, 1.0, ) }, "activating the batched shared GLM expert", )?; project_rows( &batch.attn_out, sparse.shared_down, shape.ff_expert, shape.embd, &batch.ffn_mid, rows, map, size, )?; call( unsafe { ds4_gpu_add3_tensor( batch.next.raw(), batch.after_attn.raw(), batch.ffn_out.raw(), batch.attn_out.raw(), residual, ) }, "adding the batched sparse GLM residual", ) } pub(super) fn eval_speculative_greedy( &mut self, token: i32, max_tokens: u32, cancelled: &std::sync::atomic::AtomicBool, ) -> Result, String> { if self.mtp.is_none() { self.eval(token)?; return Ok(vec![token]); } if let Some(mtp) = &mut self.mtp { mtp.cycles += 1; } let pending = self.mtp.as_ref().and_then(|mtp| mtp.pending); let started = Instant::now(); self.eval(token)?; if let Some(mtp) = &mut self.mtp { mtp.verifier_passes += 1; } let pos = self.position() - 1; let next = argmax(&self.logits); let mut draft = self.mtp_step(next, pos)?; let mut accepted = vec![token]; if pending == Some(next) && max_tokens > 1 && self.position() < self.context && !cancelled.load(std::sync::atomic::Ordering::Relaxed) { self.eval(next)?; accepted.push(next); if let Some(mtp) = &mut self.mtp { mtp.accepted += 1; mtp.verifier_passes += 1; } draft = self.mtp_step(argmax(&self.logits), pos + 1)?; } if let Some(mtp) = &mut self.mtp { mtp.pending = Some(draft); mtp.verifier_ns = mtp .verifier_ns .saturating_add(u64::try_from(started.elapsed().as_nanos()).unwrap_or(u64::MAX)); } Ok(accepted) } fn mtp_step(&mut self, next_token: i32, pos: u32) -> Result { let mut mtp = self.mtp.take().ok_or("GLM MTP is not configured")?; let result = (|| { let layer = self .weights .nextn .as_ref() .ok_or("GLM nextn layer is missing")?; let nextn = layer.nextn.ok_or("GLM nextn input weights are missing")?; let min_pos = *mtp.min_pos.get_or_insert(pos); let selected_count = pos .checked_sub(min_pos) .and_then(|count| count.checked_add(1)) .ok_or("GLM MTP attention range overflow")?; if selected_count > self.context { return Err("GLM MTP attention range exceeds the context".into()); } let selected = (min_pos..=pos) .map(|position| position as i32) .collect::>(); mtp.selected.write_i32(&selected)?; let shape = self.model.shape; let map = self.model.main.map_ptr().cast(); let size = self.model.main.len(); let enorm = mtp.concat.view(0, shape.embd * 4)?; let hnorm = mtp.concat.view(shape.embd * 4, shape.embd * 4)?; let started = Instant::now(); let commands = Commands::begin()?; call( unsafe { ds4_gpu_embed_token_quant_tensor( self.scratch.next.raw(), map, size, self.weights.embedding.offset, self.weights.embedding.kind, shape.vocab as u32, next_token as u32, shape.embd as u32, ) }, "GLM MTP token embedding", )?; norm( &enorm, &self.scratch.next, nextn.enorm, shape.embd as u32, shape.rms_epsilon, map, size, )?; norm( &hnorm, &self.scratch.current, nextn.hnorm, shape.embd as u32, shape.rms_epsilon, map, size, )?; project( &self.scratch.current, nextn.eh_proj, 2 * shape.embd, shape.embd, &mtp.concat, map, size, self.ssd.enabled, )?; let mut count = selected_count; self.encode_layer( layer, &mtp.cache, (shape.layers - 1) as usize, shape.layers - 1, pos, &mut count, Some((&mtp.selected, selected_count)), )?; if let Some(profile) = &mut self.profile { profile.record( (shape.layers - 1) as usize, pos, &self.scratch.router_selected, &self.scratch.router_weights, 1, false, )?; } norm( &self.scratch.output_norm, &self.scratch.next, nextn.shared_head_norm, shape.embd as u32, shape.rms_epsilon, map, size, )?; project( &self.scratch.logits, self.weights.output, shape.embd, shape.vocab, &self.scratch.output_norm, map, size, self.ssd.enabled, )?; commands.finish()?; let mut logits = vec![0.0; shape.vocab as usize]; self.scratch.logits.read_f32(&mut logits)?; if let Some(profile) = &self.profile { profile.write()?; } let draft = argmax(&logits); mtp.drafted += 1; if mtp.timing { eprintln!( "ds4: GLM MTP draft {draft} at position {pos} in {:.1} ms", started.elapsed().as_secs_f64() * 1000.0 ); } Ok(draft) })(); self.mtp = Some(mtp); result } #[allow(clippy::too_many_arguments)] fn encode_layer( &self, layer: &GlmLayer, cache: &LayerCache, _layer_index: usize, ordinal: u32, pos: u32, selected_count: &mut u32, attention_selected: Option<(&Buffer, u32)>, ) -> Result<(), String> { let shape = self.model.shape; let map = self.model.main.map_ptr().cast(); let size = self.model.main.len(); let q_dim = shape.heads * shape.key_mla; let q_nope = shape.key_mla - shape.rot; norm( &self.scratch.attn_norm, &self.scratch.current, layer.attn_norm, shape.embd as u32, shape.rms_epsilon, map, size, )?; project( &self.scratch.q_rank, layer.q_a, shape.embd, shape.lora_q, &self.scratch.attn_norm, map, size, self.ssd.enabled, )?; project( &self.scratch.kv_raw, layer.kv_a, shape.embd, shape.head_dim, &self.scratch.attn_norm, map, size, self.ssd.enabled, )?; call( unsafe { ds4_gpu_glm_qkv_norm_store_compact_kv_tensor( self.scratch.q_rank_norm.raw(), self.scratch.q_rank.raw(), map, size, layer.q_a_norm.offset, shape.lora_q as u32, cache.kv.raw(), cache.rope.raw(), self.scratch.kv_raw.raw(), layer.kv_norm.offset, pos, 1, self.context, shape.head_dim as u32, shape.kv_lora as u32, shape.rot as u32, CACHE_F16, shape.rms_epsilon, ) }, "storing GLM compact KV", )?; project( &self.scratch.q, layer.q_b, shape.lora_q, q_dim, &self.scratch.q_rank_norm, map, size, self.ssd.enabled, )?; call( unsafe { ds4_gpu_glm_rope_tail_tensor( self.scratch.q.raw(), 1, shape.heads as u32, shape.key_mla as u32, shape.rot as u32, pos, 0, shape.rope_base, 1.0, 0.0, 1.0, 0.0, 0.0, ) }, "applying GLM query RoPE", )?; if let Some((_, count)) = attention_selected { *selected_count = count; } else if let Some(indexer_cache) = &cache.indexer { project( &self.scratch.indexer_k, layer.indexer_k, shape.embd, shape.indexer_head_dim, &self.scratch.current, map, size, self.ssd.enabled, )?; call( unsafe { ds4_gpu_glm_store_indexer_k_tensor( indexer_cache.raw(), self.scratch.indexer_k.raw(), map, size, layer.indexer_k_norm.offset, layer.indexer_k_bias.offset, pos, 1, self.context, shape.indexer_head_dim as u32, shape.rot as u32, 0, 1.0e-6, shape.rope_base, 1.0, 0.0, 1.0, 0.0, 0.0, CACHE_F16, ) }, "storing the GLM indexer key", )?; let visible = pos + 1; *selected_count = visible.min(shape.indexer_top_k as u32); if visible <= shape.indexer_top_k as u32 { call( unsafe { ds4_gpu_glm_fill_selected_range_tensor( self.scratch.indexer_selected.raw(), *selected_count, ) }, "selecting the visible GLM context", )?; } else { project( &self.scratch.indexer_q, layer.indexer_q, shape.lora_q, shape.indexer_heads * shape.indexer_head_dim, &self.scratch.q_rank_norm, map, size, self.ssd.enabled, )?; call( unsafe { ds4_gpu_glm_indexer_rope_tail_tensor( self.scratch.indexer_q.raw(), 1, shape.indexer_heads as u32, shape.indexer_head_dim as u32, shape.rot as u32, pos, 0, shape.rope_base, 1.0, 0.0, 1.0, 0.0, 0.0, ) }, "applying GLM indexer RoPE", )?; f32_project( &self.scratch.indexer_weights, layer.indexer_proj, shape.embd, shape.indexer_heads, &self.scratch.current, map, size, )?; let scale = 1.0 / ((shape.indexer_heads * shape.indexer_head_dim) as f32).sqrt(); call( unsafe { ds4_gpu_glm_indexer_score_one_tensor( self.scratch.indexer_scores.raw(), self.scratch.indexer_q.raw(), self.scratch.indexer_weights.raw(), indexer_cache.raw(), visible, shape.indexer_heads as u32, shape.indexer_head_dim as u32, scale, CACHE_F16, ) }, "scoring the GLM indexer", )?; call( unsafe { ds4_gpu_indexer_topk_tensor( self.scratch.indexer_selected.raw(), self.scratch.indexer_scores.raw(), visible, 1, *selected_count, ) }, "selecting GLM indexed attention rows", )?; } } if *selected_count == 0 { return Err("GLM indexer did not select an attention context".into()); } call( unsafe { ds4_gpu_glm_qk_lowrank_typed_tensor( self.scratch.qk_low.raw(), self.scratch.q.raw(), map, size, layer.k_b.offset, layer.k_b.kind, shape.heads as u32, shape.kv_lora as u32, q_nope as u32, shape.key_mla as u32, ) }, "projecting the GLM low-rank query", )?; call( unsafe { ds4_gpu_glm_attention_indexed_decode_typed_tensor( self.scratch.heads.raw(), self.scratch.q.raw(), self.scratch.qk_low.raw(), cache.kv.raw(), cache.rope.raw(), map, size, layer.v_b.offset, layer.v_b.kind, attention_selected.map_or_else( || self.scratch.indexer_selected.raw(), |value| value.0.raw(), ), *selected_count, self.context, CACHE_F16, shape.heads as u32, shape.kv_lora as u32, q_nope as u32, shape.rot as u32, shape.value_mla as u32, 0, shape.rope_base, 1.0, 0.0, 1.0, 0.0, 0.0, ) }, "running GLM indexed attention", )?; project( &self.scratch.attn_out, layer.output, shape.heads * shape.value_mla, shape.embd, &self.scratch.heads, map, size, self.ssd.enabled, )?; call( unsafe { ds4_gpu_add_rms_norm_weight_tensor( self.scratch.ffn_norm.raw(), self.scratch.after_attn.raw(), self.scratch.current.raw(), self.scratch.attn_out.raw(), map, size, layer.ffn_norm.offset, shape.embd as u32, shape.rms_epsilon, ) }, "normalizing the GLM FFN input", )?; if let Some(dense) = layer.dense { project( &self.scratch.ffn_gate, dense.gate, shape.embd, shape.ff_dense, &self.scratch.ffn_norm, map, size, self.ssd.enabled, )?; project( &self.scratch.ffn_up, dense.up, shape.embd, shape.ff_dense, &self.scratch.ffn_norm, map, size, self.ssd.enabled, )?; call( unsafe { ds4_gpu_swiglu_tensor( self.scratch.ffn_mid.raw(), self.scratch.ffn_gate.raw(), self.scratch.ffn_up.raw(), shape.ff_dense as u32, 0.0, 1.0, ) }, "activating the dense GLM FFN", )?; project( &self.scratch.ffn_out, dense.down, shape.ff_dense, shape.embd, &self.scratch.ffn_mid, map, size, self.ssd.enabled, )?; call( unsafe { ds4_gpu_add_tensor( self.scratch.next.raw(), self.scratch.after_attn.raw(), self.scratch.ffn_out.raw(), shape.embd as u32, ) }, "adding the dense GLM residual", )?; } else if let Some(sparse) = layer.sparse { f32_project( &self.scratch.router_logits, sparse.router, shape.embd, shape.experts, &self.scratch.ffn_norm, map, size, )?; call( unsafe { ds4_gpu_glm_router_select_tensor( self.scratch.router_selected.raw(), self.scratch.router_weights.raw(), self.scratch.router_probs.raw(), map, size, sparse.bias.offset, self.scratch.router_logits.raw(), shape.experts as u32, shape.experts_used as u32, shape.expert_weight_scale, ) }, "routing the GLM experts", )?; let force_resident = self.ssd.enabled && ordinal.saturating_sub(shape.leading_dense) < self.ssd.full_layers; if self.ssd.enabled && !force_resident { let table = expert_table(map, size, ordinal, shape, sparse); call( unsafe { ds4_gpu_glm_stream_expert_cache_begin_selected_load_tensor( &table, self.scratch.router_selected.raw(), shape.experts_used as u32, ) }, "loading selected GLM experts", )?; } let (gate_expert, gate_row) = expert_layout(sparse.gate, shape.experts); let (up_expert, up_row) = expert_layout(sparse.up, shape.experts); let (down_expert, down_row) = expert_layout(sparse.down, shape.experts); project( &self.scratch.ffn_gate, sparse.shared_gate, shape.embd, shape.ff_expert, &self.scratch.ffn_norm, map, size, self.ssd.enabled, )?; project( &self.scratch.ffn_up, sparse.shared_up, shape.embd, shape.ff_expert, &self.scratch.ffn_norm, map, size, self.ssd.enabled, )?; call( unsafe { ds4_gpu_swiglu_tensor( self.scratch.ffn_mid.raw(), self.scratch.ffn_gate.raw(), self.scratch.ffn_up.raw(), shape.ff_expert as u32, 0.0, 1.0, ) }, "activating the shared GLM expert", )?; project( &self.scratch.ffn_sum, sparse.shared_down, shape.ff_expert, shape.embd, &self.scratch.ffn_mid, map, size, self.ssd.enabled, )?; let routed = call( unsafe { if sparse.gate.kind == IQ2_XXS { ds4_gpu_routed_moe_one_tensor( self.scratch.ffn_out.raw(), self.scratch.ffn_gate.raw(), self.scratch.ffn_up.raw(), self.scratch.ffn_mid.raw(), self.scratch.routed_down.raw(), map, size, sparse.gate.offset, sparse.up.offset, sparse.down.offset, sparse.gate.kind, sparse.down.kind, gate_expert, gate_row, down_expert, down_row, shape.embd as u32, shape.ff_expert as u32, shape.embd as u32, self.scratch.router_selected.raw(), self.scratch.router_weights.raw(), shape.experts as u32, shape.experts_used as u32, 0.0, self.scratch.ffn_norm.raw(), std::ptr::null(), ordinal, force_resident, ) } else { ds4_gpu_glm_routed_moe_one_tensor( self.scratch.ffn_out.raw(), self.scratch.ffn_mid.raw(), map, size, sparse.gate.offset, sparse.up.offset, sparse.down.offset, sparse.gate.kind, sparse.up.kind, sparse.down.kind, gate_expert, gate_row, up_expert, up_row, down_expert, down_row, shape.embd as u32, shape.ff_expert as u32, shape.embd as u32, self.scratch.router_selected.raw(), self.scratch.router_weights.raw(), shape.experts as u32, shape.experts_used as u32, ordinal, self.scratch.ffn_norm.raw(), force_resident, ) } }, "running the routed GLM experts", ); routed.map_err(|error| format!("{error} in layer {ordinal}"))?; call( unsafe { ds4_gpu_add3_tensor( self.scratch.next.raw(), self.scratch.after_attn.raw(), self.scratch.ffn_out.raw(), self.scratch.ffn_sum.raw(), shape.embd as u32, ) }, "adding the sparse GLM residual", )?; } Ok(()) } pub(super) fn prefill( &mut self, tokens: &[i32], mut progress: impl FnMut(u32) -> bool, ) -> Result { let end = self .position() .checked_add(u32::try_from(tokens.len()).map_err(|_| "GLM prefill is too large")?) .ok_or("GLM prefill position overflow")?; if end > self.context { return Err(format!( "the GLM Metal executor supports {} tokens", self.context )); } let mut completed = 0; while completed < tokens.len() { if !progress(self.position()) { break; } let remaining = &tokens[completed..]; if self.use_streaming_token_prefill(remaining.len()) { for &token in remaining { self.eval(token)?; completed += 1; if !progress(self.position()) { return Ok(completed); } } continue; } let pos = self.position(); let rows = indexed_prefill_rows(pos, remaining.len(), self.model.shape.indexer_top_k as u32); let cancelled = self.eval_batch(&remaining[..rows], &mut progress)?; completed += rows; if cancelled { break; } } Ok(completed) } fn use_streaming_token_prefill(&self, remaining: usize) -> bool { if !self.ssd.enabled || self.quality || remaining == 0 { return false; } if std::env::var_os("DS4_GLM_DISABLE_STREAMING_TOKEN_PREFILL").is_some() || std::env::var_os("DS4_METAL_GLM_DISABLE_STREAMING_TOKEN_PREFILL").is_some() { return false; } let max_tokens = std::env::var("DS4_METAL_GLM_STREAMING_TOKEN_PREFILL_MAX") .ok() .or_else(|| std::env::var("DS4_GLM_STREAMING_TOKEN_PREFILL_MAX").ok()) .and_then(|value| value.parse::().ok()) .unwrap_or(STREAMING_TOKEN_PREFILL_MAX); let Ok(rows) = u32::try_from(remaining) else { return false; }; streaming_token_prefill_eligible(self.context, self.position(), rows, max_tokens) } fn eval_batch( &mut self, tokens: &[i32], progress: &mut impl FnMut(u32) -> bool, ) -> Result { let rows = u32::try_from(tokens.len()).map_err(|_| "GLM prefill batch is too large")?; if rows == 0 || rows > INDEXED_PREFILL_CHUNK { return Err("GLM indexed prefill batch exceeds DS4's workspace".into()); } if tokens .iter() .any(|token| *token < 0 || *token as u64 >= self.model.shape.vocab) { return Err("GLM prefill contains a token outside the vocabulary".into()); } let pos = self.position(); let mut batch = GlmBatchScratch::allocate(&self.model, rows, self.context)?; batch.tokens.write_i32(tokens)?; let shape = self.model.shape; let map = self.model.main.map_ptr().cast(); let size = self.model.main.len(); let commands = Commands::begin()?; call( unsafe { ds4_gpu_embed_tokens_quant_tensor( batch.current.raw(), batch.tokens.raw(), map, size, self.weights.embedding.offset, self.weights.embedding.kind, shape.vocab as u32, rows, shape.embd as u32, ) }, "embedding a GLM prefill batch", )?; commands.finish()?; let mut selected_ready = false; let mut selected_count = 0; let mut cancelled = false; for index in 0..self.weights.layers.len() { if self.ssd.enabled && self.weights.layers[index] .sparse .is_some_and(|sparse| sparse.gate.kind == IQ2_XXS) && (index as u32).saturating_sub(shape.leading_dense) >= self.ssd.full_layers { install_glm_model_spans( &self.model, &glm_layer_model_spans(&self.model, index as u32)?, "GLM indexed-prefill layer", )?; } let commands = Commands::begin()?; self.encode_batch_layer( &batch, &self.weights.layers[index], &self.caches[index], index as u32, pos, rows, &mut selected_ready, &mut selected_count, )?; commands.finish()?; if self.weights.layers[index].sparse.is_some() && let Some(profile) = &mut self.profile { profile.record( index, pos, &batch.router_selected, &batch.router_weights, rows, false, )?; } std::mem::swap(&mut batch.current, &mut batch.next); let layer_done = (index + 1) as u64; let token_equivalent = u32::try_from(u64::from(rows) * layer_done / self.weights.layers.len() as u64) .expect("GLM prefill progress is bounded by the batch"); cancelled |= !progress(pos + token_equivalent); } if let Some(spans) = &self.streaming_spans { install_glm_model_spans(&self.model, spans, "GLM decode/output")?; } let commands = Commands::begin()?; let last = batch .current .view(u64::from(rows - 1) * shape.embd * 4, shape.embd * 4)?; norm( &self.scratch.output_norm, &last, self.weights.output_norm, shape.embd as u32, shape.rms_epsilon, map, size, )?; project( &self.scratch.logits, self.weights.output, shape.embd, shape.vocab, &self.scratch.output_norm, map, size, self.ssd.enabled, )?; commands.finish()?; self.scratch.logits.read_f32(&mut self.logits)?; self.tokens.extend_from_slice(tokens); if let Some(profile) = &self.profile { profile.write()?; } Ok(cancelled) } pub(super) fn logits(&self) -> &[f32] { &self.logits } pub(super) fn model(&self) -> &Model { &self.model } pub(super) fn context(&self) -> u32 { self.context } pub(super) fn position(&self) -> u32 { self.tokens.len() as u32 } pub(super) fn tokens(&self) -> &[i32] { &self.tokens } pub(super) fn checkpoint_tag(&self) -> [u8; 32] { self.checkpoint_tag } pub(super) fn execution_stats(&self) -> ExecutionStats { let mut stats = self .mtp .as_ref() .map_or_else(ExecutionStats::default, |mtp| ExecutionStats { speculative_mode: 3, speculative_cycles: mtp.cycles, drafted_tokens: mtp.drafted, accepted_draft_tokens: mtp.accepted, verifier_passes: mtp.verifier_passes, verifier_ms: mtp.verifier_ns / 1_000_000, ..ExecutionStats::default() }); if self.ssd.enabled { let mut native = StreamExpertCacheStats::default(); unsafe { ds4_gpu_stream_expert_cache_get_stats(&mut native) }; stats.ssd_enabled = true; stats.ssd_resident_bytes = self.ssd_resident_bytes; stats.ssd_cache_bytes = self.ssd_cache_bytes; stats.ssd_cache_experts = self.ssd_cache_experts; stats.ssd_preloaded_experts = self.ssd_preloaded_experts; stats.ssd_cache_entries = u64::from(native.current_count); 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 = native.hits.saturating_add(native.misses); stats.ssd_requested_bytes = native.misses.saturating_mul( self.ssd_cache_bytes .checked_div(self.ssd_cache_experts) .unwrap_or(0), ); stats.ssd_wait_ms = stats.ssd_pread_ms; } stats } pub(super) fn note_checkpoint_tag(&mut self, tag: [u8; 32]) { self.checkpoint_tag = tag; } pub(super) fn reset(&mut self) -> Result<(), String> { self.scratch = GlmScratch::allocate(&self.model, self.context)?; self.caches = (0..self.weights.layers.len()) .map(|layer| LayerCache::allocate(self.model.shape, layer, self.context)) .collect::>()?; self.tokens.clear(); self.checkpoint_tag = [0; 32]; if let Some(mtp) = &mut self.mtp { mtp.cache = LayerCache::allocate( self.model.shape, (self.model.shape.layers - 1) as usize, self.context, )?; mtp.pending = None; mtp.min_pos = None; } Ok(()) } fn blank_resident_state(&self) -> Result { let mtp = self .mtp .as_ref() .map(|current| { Ok::<_, String>(GlmMtp { cache: LayerCache::allocate( self.model.shape, (self.model.shape.layers - 1) as usize, self.context, )?, selected: Buffer::bytes(u64::from(self.context) * 4)?, concat: Buffer::floats(2 * self.model.shape.embd)?, pending: None, min_pos: None, cycles: 0, drafted: 0, accepted: 0, verifier_passes: 0, verifier_ns: 0, timing: current.timing, }) }) .transpose()?; Ok(GlmResidentState { scratch: GlmScratch::allocate(&self.model, self.context)?, caches: (0..self.weights.layers.len()) .map(|layer| LayerCache::allocate(self.model.shape, layer, self.context)) .collect::>()?, logits: vec![0.0; self.model.shape.vocab as usize], tokens: Vec::new(), checkpoint_tag: [0; 32], mtp, }) } 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.scratch, &mut incoming.scratch); std::mem::swap(&mut self.caches, &mut incoming.caches); 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); std::mem::swap(&mut self.mtp, &mut incoming.mtp); *state = Some(incoming); Ok(()) } pub(super) fn align_prompt(&mut self, tokens: &[i32]) -> Result { if let Some(rewind) = live_prefix_rewind_target(&self.tokens, tokens) { self.tokens.truncate(rewind); if let Some(mtp) = &mut self.mtp { mtp.pending = None; mtp.min_pos = None; } return Ok(rewind); } if !tokens.starts_with(&self.tokens) { self.reset()?; } Ok(self.tokens.len()) } pub(super) fn save_checkpoint( &mut self, path: &Path, tag: [u8; 32], progress: &mut impl FnMut(u64), ) -> Result<(), String> { if let Some(parent) = path.parent() { fs::create_dir_all(parent).map_err(|e| e.to_string())?; } let temporary = path.with_extension("tmp"); let mut file = File::create(&temporary).map_err(|e| e.to_string())?; file.write_all(CHECKPOINT_MAGIC) .map_err(|e| e.to_string())?; let shape = self.model.shape; for value in [ CHECKPOINT_VERSION, self.context, self.position(), self.quality as u32, self.weights.layers.len() as u32, shape.kv_lora as u32, shape.rot as u32, shape.indexer_head_dim as u32, shape.vocab as u32, ] { write_u32(&mut file, value)?; } write_u64(&mut file, self.model.main.len())?; write_u64(&mut file, self.model_modified.0)?; write_u32(&mut file, self.model_modified.1)?; file.write_all(&self.model_identity) .map_err(|e| e.to_string())?; file.write_all(&tag).map_err(|e| e.to_string())?; for &token in &self.tokens { write_u32(&mut file, token as u32)?; } for &logit in &self.logits { write_u32(&mut file, logit.to_bits())?; } let mut chunk = vec![0; CHECKPOINT_IO_CHUNK]; let rows = u64::from(self.position()); for cache in &self.caches { write_buffer( &mut file, &cache.kv, 0, rows * shape.kv_lora * 2, &mut chunk, progress, )?; write_buffer( &mut file, &cache.rope, 0, rows * shape.rot * 2, &mut chunk, progress, )?; if let Some(indexer) = &cache.indexer { write_buffer( &mut file, indexer, 0, rows * shape.indexer_head_dim * 2, &mut chunk, progress, )?; } } file.sync_all().map_err(|e| e.to_string())?; fs::rename(&temporary, path).map_err(|e| e.to_string())?; self.checkpoint_tag = tag; Ok(()) } pub(super) fn load_checkpoint( &mut self, path: &Path, progress: &mut impl FnMut(u64), ) -> Result { let mut file = match File::open(path) { Ok(file) => file, Err(error) if error.kind() == std::io::ErrorKind::NotFound => return Ok(false), Err(error) => return Err(error.to_string()), }; let mut magic = [0; 8]; file.read_exact(&mut magic).map_err(|e| e.to_string())?; if &magic != CHECKPOINT_MAGIC || read_u32(&mut file)? != CHECKPOINT_VERSION || read_u32(&mut file)? != self.context { return Err("GLM checkpoint does not match the current executor".into()); } let token_count = read_u32(&mut file)?; let shape = self.model.shape; if token_count > self.context || read_u32(&mut file)? != self.quality as u32 || read_u32(&mut file)? != self.weights.layers.len() as u32 || read_u32(&mut file)? != shape.kv_lora as u32 || read_u32(&mut file)? != shape.rot as u32 || read_u32(&mut file)? != shape.indexer_head_dim as u32 || read_u32(&mut file)? != shape.vocab as u32 || read_u64(&mut file)? != self.model.main.len() || read_u64(&mut file)? != self.model_modified.0 || read_u32(&mut file)? != self.model_modified.1 { return Err("GLM checkpoint was written for a different model or configuration".into()); } let mut identity = [0; 32]; file.read_exact(&mut identity).map_err(|e| e.to_string())?; if identity != self.model_identity { return Err("GLM checkpoint model identity changed".into()); } let mut tag = [0; 32]; file.read_exact(&mut tag).map_err(|e| e.to_string())?; let mut tokens = Vec::with_capacity(token_count as usize); for _ in 0..token_count { let token = read_u32(&mut file)?; if u64::from(token) >= self.model.shape.vocab { return Err("GLM checkpoint token is invalid".into()); } tokens.push(token as i32); } let mut logits = Vec::with_capacity(shape.vocab as usize); for _ in 0..shape.vocab { logits.push(f32::from_bits(read_u32(&mut file)?)); } self.reset()?; let mut chunk = vec![0; CHECKPOINT_IO_CHUNK]; let rows = u64::from(token_count); for cache in &self.caches { read_buffer( &mut file, &cache.kv, 0, rows * shape.kv_lora * 2, &mut chunk, progress, )?; read_buffer( &mut file, &cache.rope, 0, rows * shape.rot * 2, &mut chunk, progress, )?; if let Some(indexer) = &cache.indexer { read_buffer( &mut file, indexer, 0, rows * shape.indexer_head_dim * 2, &mut chunk, progress, )?; } } let mut trailing = [0]; if file.read(&mut trailing).map_err(|e| e.to_string())? != 0 { return Err("GLM checkpoint has trailing data".into()); } self.tokens = tokens; self.logits = logits; self.checkpoint_tag = tag; Ok(true) } } fn streaming_token_prefill_eligible(context: u32, pos: u32, rows: u32, max: u32) -> bool { let full_cap = if context >= LONG_CONTEXT_THRESHOLD { LONG_CONTEXT_FULL_ATTN_CONTEXT } else { STREAMING_FULL_ATTN_CONTEXT } .min(context); max != 0 && rows != 0 && rows <= max && pos < full_cap && rows <= full_cap - pos } fn indexed_prefill_rows(pos: u32, remaining: usize, top_k: u32) -> usize { let mut rows = remaining.min(INDEXED_PREFILL_CHUNK as usize); if pos < top_k { rows = rows.min((top_k - pos) as usize); } rows } fn full_indexer_layer(shape: super::super::Shape, layer: usize) -> bool { layer < (shape.layers - shape.nextn) as usize && (layer < shape.leading_dense as usize || (layer >= 6 && (layer - 6).is_multiple_of(4))) } fn expert_layout(weight: Weight, experts: u64) -> (u64, u64) { let expert = weight.bytes / experts; (expert, expert / weight.dims[1]) } fn expert_table( map: *const c_void, size: u64, layer: u32, shape: super::super::Shape, weights: SparseWeights, ) -> StreamExpertTable { StreamExpertTable { model_map: map, model_size: size, layer, total_experts: shape.experts as u32, gate_offset: weights.gate.offset, up_offset: weights.up.offset, down_offset: weights.down.offset, gate_expert_bytes: weights.gate.bytes / shape.experts, down_expert_bytes: weights.down.bytes / shape.experts, } } fn configure_streaming( model: &Model, weights: &GlmWeights, plan: Option<&GlmStreamingPlan>, ) -> Result<(), String> { let Some(plan) = plan else { return Ok(()) }; unsafe { ds4_gpu_set_streaming_expert_cache_expert_bytes(plan.per_expert); ds4_gpu_set_streaming_expert_cache_budget(plan.cache_experts); } let preload = preload_count(plan.settings, plan.cache_experts); let mut by_layer = vec![Vec::<(i32, u32)>::new(); weights.layers.len()]; let mut loaded = 0_u32; for &(layer, expert) in hotlist::GLM52 { if loaded == preload { break; } let layer = usize::from(layer); if layer < model .shape .leading_dense .saturating_add(plan.settings.full_layers) as usize || layer >= weights.layers.len() || u64::from(expert) >= model.shape.experts { continue; } by_layer[layer].push((i32::from(expert), preload - loaded)); loaded += 1; } for (layer_index, entries) in by_layer.iter().enumerate() { let Some(sparse) = weights.layers[layer_index].sparse else { continue; }; if entries.is_empty() { continue; } let ids = entries.iter().map(|entry| entry.0).collect::>(); let priorities = entries.iter().map(|entry| entry.1).collect::>(); let table = expert_table( model.main.map_ptr().cast(), model.main.len(), layer_index as u32, model.shape, sparse, ); call( unsafe { ds4_gpu_stream_expert_cache_seed_experts( &table, ids.as_ptr(), priorities.as_ptr(), ids.len() as u32, ) }, "preloading GLM SSD experts", )?; } Ok(()) } fn preload_count(ssd: EngineSsdSettings, budget: u32) -> u32 { if ssd.cold { return 0; } if ssd.preload_experts != 0 { return ssd.preload_experts.min(budget); } let cap = env::var("DS4_METAL_STREAMING_EXPERT_AUTO_PRELOAD_CAP") .ok() .and_then(|value| value.parse::().ok()) .unwrap_or(4096); if cap == 0 { budget } else { budget.min(cap) } } fn sparse_expert_bytes(weights: SparseWeights, experts: u64) -> Result { (weights.gate.bytes / experts) .checked_add(weights.up.bytes / experts) .and_then(|bytes| bytes.checked_add(weights.down.bytes / experts)) .ok_or_else(|| "GLM expert size overflow".into()) } fn glm_streaming_plan( model: &Model, weights: &GlmWeights, mut ssd: EngineSsdSettings, ) -> Result, String> { if !ssd.enabled { return Ok(None); } let shape = model.shape; let first_sparse = weights .layers .iter() .find_map(|layer| layer.sparse) .ok_or("GLM model has no routed expert layers")?; let per_expert = sparse_expert_bytes(first_sparse, shape.experts)?; let max_experts = shape.layers.saturating_mul(shape.experts as u32); if ssd.cache_experts != 0 { let cache_experts = ssd.cache_experts.min(max_experts); let full_layer_bytes = glm_full_layer_bytes(weights, shape, ssd.full_layers)?; return Ok(Some(GlmStreamingPlan { settings: ssd, per_expert, cache_experts, planned_expert_bytes: per_expert .saturating_mul(u64::from(cache_experts)) .saturating_add(full_layer_bytes), })); } let mut total = if ssd.cache_bytes != 0 { ssd.cache_bytes } else { let recommended = unsafe { ds4_gpu_recommended_working_set_size() }; if recommended == 0 { return Err( "Metal did not report a working-set size; set an explicit GLM SSD cache budget" .into(), ); } 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); let target = recommended.saturating_mul(percent) / 100; target.saturating_sub(glm_resident_weight_bytes(model)?) }; total = total.min(AUTO_CACHE_BYTES); let total_experts = ((total / per_expert).max(1)).min(u64::from(max_experts)); total = total_experts.saturating_mul(per_expert); let max_layer = weights .layers .iter() .chain(weights.nextn.iter()) .filter_map(|layer| layer.sparse) .map(|weights| { sparse_expert_bytes(weights, shape.experts) .map(|bytes| bytes.saturating_mul(shape.experts)) }) .collect::, _>>()? .into_iter() .max() .ok_or("GLM model has no routed expert layers")?; let prefill_headroom = max_layer.saturating_mul(2); if prefill_headroom >= total { return Err(format!( "GLM SSD cache budget is too small: two routed prefill layers need {:.2} GiB", prefill_headroom as f64 / 1_073_741_824.0 )); } let after_prefill = total - prefill_headroom; let supported = weights.layers.len() as u32 - shape.leading_dense; let requested_full = if ssd.full_layers_set { ssd.full_layers.min(supported) } else { let target = (after_prefill / 7).min(10 * 1024 * 1024 * 1024); (1..=supported) .take_while(|layers| { glm_full_layer_bytes(weights, shape, *layers).is_ok_and(|bytes| bytes <= target) }) .last() .unwrap_or(0) }; let minimum_dynamic = per_expert.saturating_mul(shape.experts); let mut full_layers = requested_full; let full_layer_bytes = loop { let bytes = glm_full_layer_bytes(weights, shape, full_layers)?; if full_layers == 0 || bytes.saturating_add(minimum_dynamic) < after_prefill { break bytes; } full_layers -= 1; }; ssd.full_layers = full_layers; let cache_experts = dynamic_expert_budget( total, prefill_headroom, full_layer_bytes, per_expert, max_experts, )?; Ok(Some(GlmStreamingPlan { settings: ssd, per_expert, cache_experts, planned_expert_bytes: total, })) } fn dynamic_expert_budget( total: u64, prefill_headroom: u64, full_layers: u64, per_expert: u64, max_experts: u32, ) -> Result { let reserved = prefill_headroom .checked_add(full_layers) .ok_or("GLM SSD cache reserve overflow")?; let experts = total .checked_sub(reserved) .filter(|bytes| *bytes >= per_expert) .map(|bytes| bytes / per_expert) .ok_or("GLM SSD streaming has no memory for an expert cache")?; Ok(u32::try_from(experts).unwrap_or(u32::MAX).min(max_experts)) } fn glm_full_layer_bytes( weights: &GlmWeights, shape: super::super::Shape, layers: u32, ) -> Result { weights .layers .iter() .skip(shape.leading_dense as usize) .take(layers as usize) .try_fold(0_u64, |total, layer| { let sparse = layer.sparse.ok_or("GLM resident prefix is not sparse")?; let bytes = sparse_expert_bytes(sparse, shape.experts)? .checked_mul(shape.experts) .ok_or("GLM resident layer size overflow")?; total .checked_add(bytes) .ok_or_else(|| "GLM resident layer size overflow".into()) }) } fn glm_resident_weight_bytes(model: &Model) -> Result { model .main .tensors .iter() .filter(|(name, _)| { !name.ends_with("ffn_gate_exps.weight") && !name.ends_with("ffn_up_exps.weight") && !name.ends_with("ffn_down_exps.weight") }) .try_fold(0_u64, |total, (_, tensor)| total.checked_add(tensor.bytes)) .ok_or_else(|| "GLM resident tensor size overflow".into()) } fn glm_streaming_model_spans( model: &Model, weights: &GlmWeights, plan: &GlmStreamingPlan, ) -> Result<(Vec<(u64, u64)>, u64), String> { let full_before = model .shape .leading_dense .saturating_add(plan.settings.full_layers); let mut spans = model .main .tensors .iter() .filter(|(name, _)| { let routed = name.ends_with("ffn_gate_exps.weight") || name.ends_with("ffn_up_exps.weight") || name.ends_with("ffn_down_exps.weight"); if !routed { return true; } let layer = name .strip_prefix("blk.") .and_then(|name| name.split('.').next()) .and_then(|layer| layer.parse::().ok()); layer.is_some_and(|layer| { layer < full_before as usize || weights .layers .get(layer) .or_else(|| { (layer == weights.layers.len()) .then_some(weights.nextn.as_ref()) .flatten() }) .and_then(|layer| layer.sparse) .is_some_and(|sparse| { sparse_expert_bytes(sparse, model.shape.experts) .is_ok_and(|bytes| bytes != plan.per_expert) }) }) }) .map(|(_, tensor)| (tensor.offset, tensor.bytes)) .collect::>(); let max_tensor_bytes = spans .iter() .map(|(_, bytes)| *bytes) .max() .ok_or("GLM SSD streaming found no resident model tensors")?; 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("GLM 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)); } Ok((merged, max_tensor_bytes)) } fn glm_layer_model_spans(model: &Model, layer: u32) -> Result, String> { let prefix = format!("blk.{layer}."); let mut spans = model .main .tensors .iter() .filter(|(name, _)| name.starts_with(&prefix)) .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("GLM layer 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(format!("GLM layer {layer} has no model tensors")); } Ok(merged) } fn install_glm_model_spans( model: &Model, spans: &[(u64, u64)], purpose: &str, ) -> Result<(), String> { let (offsets, sizes): (Vec<_>, Vec<_>) = spans.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.len() as u32, model.main.max_tensor_bytes(), ) }, purpose, ) } fn admission_bytes( model: &Model, weights: &GlmWeights, context: u32, streaming: Option<&GlmStreamingPlan>, ) -> Result { let shape = model.shape; let normal_layers = u64::from(shape.layers - shape.nextn); let indexer_layers = (0..normal_layers as usize) .filter(|layer| full_indexer_layer(shape, *layer)) .count() as u64; let per_layer = u64::from(context) .checked_mul((shape.kv_lora + shape.rot) * 2) .ok_or("GLM compact-cache size overflow")?; let kv = normal_layers .checked_mul(per_layer) .and_then(|bytes| { bytes.checked_add(indexer_layers * u64::from(context) * shape.indexer_head_dim * 2) }) .ok_or("GLM compact-cache size overflow")?; let resident = if let Some(plan) = streaming { let off_slab = weights .layers .iter() .chain(weights.nextn.iter()) .filter_map(|layer| layer.sparse) .try_fold(0_u64, |total, sparse| { let bytes = sparse_expert_bytes(sparse, shape.experts)?; if bytes == plan.per_expert { return Ok::<_, String>(total); } total .checked_add(bytes.saturating_mul(shape.experts)) .ok_or_else(|| "GLM off-slab resident size overflow".into()) })?; glm_resident_weight_bytes(model)? .checked_add(off_slab) .ok_or("GLM resident weight size overflow")? } else { model.main.len() - model.main.data_offset() }; let cache = streaming.map_or(0, |plan| plan.planned_expert_bytes); let scratch = glm_batch_scratch_bytes(shape, context.min(INDEXED_PREFILL_CHUNK), context)? .max(512 * 1024 * 1024); resident .checked_add(kv) .and_then(|bytes| bytes.checked_add(cache)) .and_then(|bytes| bytes.checked_add(scratch)) .ok_or_else(|| "GLM runtime memory size overflow".into()) } fn glm_batch_scratch_bytes( shape: super::super::Shape, rows: u32, context: u32, ) -> Result { let rows = u64::from(rows); let score_rows = (INDEXED_PREFILL_SCORE_BYTES / (u64::from(context) * 4)) .max(1) .min(rows); let q = shape.heads * shape.key_mla; let qk_low = shape.heads * shape.kv_lora; let heads = shape.heads * shape.value_mla; let hidden = shape.ff_dense.max(shape.ff_expert); let routed = shape.experts_used * shape.ff_expert; let row_floats = 7 * shape.embd + 2 * shape.lora_q + q + shape.head_dim + shape.kv_lora + shape.indexer_head_dim + shape.indexer_heads * shape.indexer_head_dim + shape.indexer_heads + 2 * qk_low + heads + 2 * hidden + hidden.max(routed) + 2 * routed + shape.experts_used * shape.embd + 2 * shape.experts + shape.experts_used; rows.checked_mul(row_floats) .and_then(|floats| floats.checked_mul(4)) .and_then(|bytes| bytes.checked_add(rows * 4)) .and_then(|bytes| bytes.checked_add(rows * shape.indexer_top_k * 4)) .and_then(|bytes| bytes.checked_add(rows * shape.experts_used * 4)) .and_then(|bytes| bytes.checked_add(score_rows * u64::from(context) * 4)) .ok_or_else(|| "GLM indexed-prefill scratch size overflow".into()) } fn norm( out: &Buffer, input: &Buffer, weight: Weight, width: u32, epsilon: f32, map: *const c_void, size: u64, ) -> Result<(), String> { call( unsafe { ds4_gpu_rms_norm_weight_tensor( out.raw(), input.raw(), map, size, weight.offset, width, epsilon, ) }, "normalizing GLM activations", ) } #[allow(clippy::too_many_arguments)] fn norm_rows( out: &Buffer, input: &Buffer, weight: Weight, width: u32, rows: u32, epsilon: f32, map: *const c_void, size: u64, ) -> Result<(), String> { call( unsafe { ds4_gpu_rms_norm_weight_rows_tensor( out.raw(), input.raw(), map, size, weight.offset, width, rows, epsilon, ) }, "normalizing batched GLM activations", ) } #[allow(clippy::too_many_arguments)] fn project( out: &Buffer, weight: Weight, input: u64, output: u64, x: &Buffer, map: *const c_void, size: u64, streaming: bool, ) -> Result<(), String> { let result = unsafe { if streaming { ds4_gpu_matmul_quant_tensor( out.raw(), map, size, weight.offset, weight.kind, input, output, x.raw(), 1, ) } else { ds4_gpu_matmul_quant_decode_mpp_model_view_tensor( out.raw(), map, size, weight.offset, weight.kind, input, output, x.raw(), 1, ) } }; call(result, "projecting GLM activations") } #[allow(clippy::too_many_arguments)] fn project_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_quant_tensor( out.raw(), map, size, weight.offset, weight.kind, input, output, x.raw(), u64::from(rows), ) }, "projecting batched GLM activations", ) } #[allow(clippy::too_many_arguments)] fn f32_project( out: &Buffer, weight: Weight, input: u64, output: u64, x: &Buffer, map: *const c_void, size: u64, ) -> Result<(), String> { call( unsafe { ds4_gpu_matmul_f32_tensor( out.raw(), map, size, weight.offset, input, output, x.raw(), 1, ) }, "projecting GLM F32 activations", ) } #[allow(clippy::too_many_arguments)] fn f32_project_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_f32_tensor( out.raw(), map, size, weight.offset, input, output, x.raw(), u64::from(rows), ) }, "projecting batched GLM F32 activations", ) } #[cfg(test)] mod tests { use super::{ GlmExecutor, argmax, dynamic_expert_budget, full_indexer_layer, indexed_prefill_rows, live_prefix_rewind_target, streaming_token_prefill_eligible, }; use crate::engine::{GLM, Model, ReasoningMode}; use crate::model::ModelChoice; use crate::settings::EngineSsdSettings; fn installed_glm_path() -> std::path::PathBuf { crate::model::engine_artifacts(ModelChoice::Glm52, false, &crate::app::models_path()).model } #[test] fn dsa_indexer_schedule_matches_the_reference() { assert!(full_indexer_layer(GLM, 0)); assert!(full_indexer_layer(GLM, 2)); assert!(!full_indexer_layer(GLM, 3)); assert!(full_indexer_layer(GLM, 6)); assert!(full_indexer_layer(GLM, 10)); assert!(!full_indexer_layer(GLM, 11)); assert!(!full_indexer_layer(GLM, 78)); } #[test] fn glm_prefill_schedule_matches_ds4_boundaries() { assert!(streaming_token_prefill_eligible(8192, 0, 64, 64)); assert!(!streaming_token_prefill_eligible(8192, 0, 65, 64)); assert!(!streaming_token_prefill_eligible(8192, 8180, 13, 64)); assert!(!streaming_token_prefill_eligible(65_536, 4096, 1, 64)); assert_eq!(indexed_prefill_rows(0, 5000, 2048), 2048); assert_eq!(indexed_prefill_rows(2048, 5000, 2048), 4096); assert_eq!(indexed_prefill_rows(0, 17, 2048), 17); } #[test] fn repeated_glm_prompt_rewinds_one_token_for_logits() { let live = [1, 2, 3, 4, 5, 6]; assert_eq!(live_prefix_rewind_target(&live, &[1, 2, 3, 4]), Some(3)); assert_eq!(live_prefix_rewind_target(&live, &[1]), None); assert_eq!(live_prefix_rewind_target(&live, &live), None); assert_eq!(live_prefix_rewind_target(&live, &[1, 2, 9]), None); } #[test] fn glm_byte_budget_reserves_prefill_before_dynamic_experts() { let mib = 1024 * 1024; assert_eq!( dynamic_expert_budget(12 * 1024 * mib, 6 * 1024 * mib, 0, 10 * mib, 20_000), Ok(614) ); assert!(dynamic_expert_budget(1024, 1024, 0, 1, 20_000).is_err()); } #[test] #[ignore = "requires the 197 GiB GLM 5.2 checkpoint and Apple Metal"] fn resident_and_streamed_glm_match_ds4_decode_oracles() { let path = installed_glm_path(); if !path.is_file() { eprintln!("skipping unavailable GLM fixture: {}", path.display()); return; } let prompt = b"Complete the C statement with the next exact token only:\nreturn snprintf(buf, sizeof(buf), \"%d\", value"; for streamed in [false, true] { let model = Model::open_main(&path, ModelChoice::Glm52).unwrap(); let tokens = model.tokenize(std::str::from_utf8(prompt).unwrap()); let executor = GlmExecutor::open( model, 4096, false, EngineSsdSettings { enabled: streamed, cold: false, cache_experts: 0, cache_bytes: 0, full_layers: 0, full_layers_set: false, preload_experts: 0, }, ); let mut executor = match executor { Ok(executor) => executor, Err(error) if !streamed && error.contains("Metal recommends at most") => { eprintln!("skipping resident GLM on this memory size: {error}"); continue; } Err(error) => panic!("GLM executor failed to open: {error}"), }; if streamed { assert_eq!(executor.ssd_cache_experts, 671); } assert_eq!(executor.prefill(&tokens, |_| true).unwrap(), tokens.len()); let token = executor .logits() .iter() .enumerate() .max_by(|a, b| a.1.total_cmp(b.1)) .unwrap() .0 as i32; assert_eq!( executor.model().token_bytes(token).as_deref(), Some(b");\n".as_slice()) ); let greeting = executor.model().render_prompt( "You are a helpful assistant", "Write one short greeting.", ReasoningMode::Direct, ); executor.reset().unwrap(); assert_eq!( executor.prefill(&greeting, |_| true).unwrap(), greeting.len() ); let mut generated = Vec::new(); for _ in 0..4 { let token = argmax(executor.logits()); generated.push(token); executor.eval(token).unwrap(); } assert_eq!(generated, [9703, 0, 2585, 646]); } } #[test] #[ignore = "requires the 197 GiB GLM 5.2 checkpoint and Apple Metal"] fn streamed_glm_uses_ds4_indexed_prefill_for_long_prompts() { let path = installed_glm_path(); if !path.is_file() { eprintln!("skipping unavailable GLM fixture: {}", path.display()); return; } let prompt = "Complete each C statement. Example: return snprintf(buf, sizeof(buf), \"%d\", value); Example: return snprintf(buf, sizeof(buf), \"%d\", value); Example: return snprintf(buf, sizeof(buf), \"%d\", value); Example: return snprintf(buf, sizeof(buf), \"%d\", value); Example: return snprintf(buf, sizeof(buf), \"%d\", value); Now complete exactly: return snprintf(buf, sizeof(buf), \"%d\", value"; let model = Model::open_main(&path, ModelChoice::Glm52).unwrap(); let tokens = model.render_prompt("You are a helpful assistant", prompt, ReasoningMode::Direct); assert_eq!(tokens.len(), 102); let mut executor = GlmExecutor::open( model, 256, false, EngineSsdSettings { enabled: true, cold: true, cache_experts: 16, cache_bytes: 0, full_layers: 0, full_layers_set: true, preload_experts: 0, }, ) .unwrap(); let mut updates = Vec::new(); assert_eq!( executor .prefill(&tokens, |position| { updates.push(position); true }) .unwrap(), tokens.len() ); assert_eq!(executor.position(), tokens.len() as u32); assert!(updates.len() < tokens.len()); assert!(executor.logits().iter().all(|logit| logit.is_finite())); assert_eq!(argmax(executor.logits()), 1215); assert!((executor.logits()[1215] - 24.178_133).abs() < 1.0e-3); executor.eval(1215).unwrap(); assert_eq!(executor.position(), 103); assert!(executor.logits().iter().all(|logit| logit.is_finite())); } #[test] #[ignore = "requires the 197 GiB GLM 5.2 checkpoint and Apple Metal"] fn glm_mtp_preserves_target_tokens_and_drafts() { use crate::settings::EngineSpeculativeSettings; use std::sync::atomic::AtomicBool; let path = installed_glm_path(); if !path.is_file() { eprintln!("skipping unavailable GLM fixture: {}", path.display()); return; } let run = |enabled| { let model = Model::open_main(&path, ModelChoice::Glm52).unwrap(); let tokens = model.tokenize("Write one short greeting."); let mut executor = GlmExecutor::open_profile( model, 128, false, EngineSsdSettings { enabled: true, cold: false, cache_experts: 0, cache_bytes: 0, full_layers: 0, full_layers_set: false, preload_experts: 0, }, EngineSpeculativeSettings { glm_mtp: enabled, glm_mtp_timing: false, dspark: false, dspark_confidence_threshold: 0.9, dspark_confidence_threshold_set: false, dspark_strict: false, dspark_exact_sampling: false, }, None, ) .unwrap(); executor.prefill(&tokens, |_| true).unwrap(); let mut generated = Vec::new(); while generated.len() < 8 { let token = argmax(executor.logits()); let mut cycle = executor .eval_speculative_greedy( token, (8 - generated.len()) as u32, &AtomicBool::new(false), ) .unwrap(); cycle.truncate(8 - generated.len()); generated.extend(cycle); } let drafted = executor.mtp.as_ref().map_or(0, |mtp| mtp.drafted); (generated, drafted) }; let baseline = run(false).0; let (mtp, drafted) = run(true); assert_eq!(baseline, mtp); assert!(drafted > 0); } }