Files
DS4Server/src/engine/metal/glm.rs
2026-08-30 20:42:23 +02:00

3472 lines
116 KiB
Rust

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<usize> {
(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<DenseWeights>,
sparse: Option<SparseWeights>,
nextn: Option<NextnWeights>,
}
struct GlmWeights {
embedding: Weight,
output_norm: Weight,
output: Weight,
layers: Vec<GlmLayer>,
nextn: Option<GlmLayer>,
}
impl GlmWeights {
fn bind(model: &Model) -> Result<Self, String> {
let main = &model.main;
let normal_layers = model.shape.layers - model.shape.nextn;
let mut layers: Vec<GlmLayer> = (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::<Result<_, String>>()?;
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<Buffer>,
}
impl LayerCache {
fn allocate(shape: super::super::Shape, layer: usize, context: u32) -> Result<Self, String> {
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<i32>,
min_pos: Option<u32>,
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<Self, String> {
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<Self, String> {
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<LayerCache>,
logits: Vec<f32>,
tokens: Vec<i32>,
context: u32,
quality: bool,
ssd: EngineSsdSettings,
profile: Option<ExpertProfile>,
mtp: Option<GlmMtp>,
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<Vec<(u64, u64)>>,
_context: Context,
model: Model,
}
pub(in crate::engine) struct GlmResidentState {
scratch: GlmScratch,
caches: Vec<LayerCache>,
logits: Vec<f32>,
tokens: Vec<i32>,
checkpoint_tag: [u8; 32],
mtp: Option<GlmMtp>,
}
impl GlmExecutor {
#[allow(dead_code)]
pub(super) fn open(
model: Model,
context: u32,
quality: bool,
ssd: EngineSsdSettings,
) -> Result<Self, String> {
Self::open_profile(
model,
context,
quality,
ssd,
EngineSpeculativeSettings {
mtp_draft_tokens: 1,
mtp_margin: 3.0,
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<Self, String> {
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::<Result<_, _>>()?;
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<Vec<i32>, 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<i32, String> {
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::<Vec<_>>();
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<usize, String> {
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::<u32>().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<bool, String> {
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::<Result<_, _>>()?;
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<GlmResidentState, String> {
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::<Result<_, _>>()?,
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<GlmResidentState>,
) -> 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<usize, String> {
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<bool, String> {
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::<Vec<_>>();
let priorities = entries.iter().map(|entry| entry.1).collect::<Vec<_>>();
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::<u32>().ok())
.unwrap_or(4096);
if cap == 0 { budget } else { budget.min(cap) }
}
fn sparse_expert_bytes(weights: SparseWeights, experts: u64) -> Result<u64, String> {
(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<Option<GlmStreamingPlan>, 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::<u64>().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::<Result<Vec<_>, _>>()?
.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<u32, String> {
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<u64, String> {
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<u64, String> {
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::<usize>().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::<Vec<_>>();
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<Vec<(u64, u64)>, 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::<Vec<_>>();
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<u64, String> {
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<u64, String> {
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, 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 {
mtp_draft_tokens: 2,
mtp_margin: 3.0,
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);
}
}