3467 lines
116 KiB
Rust
3467 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 {
|
|
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, &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);
|
|
}
|
|
}
|