Align DeepSeek and GLM execution with DS4
This commit is contained in:
@@ -0,0 +1,345 @@
|
||||
//! DS4's sampled-token path. Distribution materialization for speculative
|
||||
//! correction and MTPLX sampling are separate algorithms, not substitutes.
|
||||
use super::{Rng, ds4_sample_argmax};
|
||||
use std::cmp::{Ordering, Reverse};
|
||||
use std::collections::BinaryHeap;
|
||||
|
||||
#[derive(Clone, Copy)]
|
||||
struct Candidate {
|
||||
id: usize,
|
||||
logit: f32,
|
||||
probability: f32,
|
||||
}
|
||||
|
||||
impl PartialEq for Candidate {
|
||||
fn eq(&self, other: &Self) -> bool {
|
||||
self.logit == other.logit && self.id == other.id
|
||||
}
|
||||
}
|
||||
impl Eq for Candidate {}
|
||||
impl PartialOrd for Candidate {
|
||||
fn partial_cmp(&self, other: &Self) -> Option<Ordering> {
|
||||
Some(self.cmp(other))
|
||||
}
|
||||
}
|
||||
impl Ord for Candidate {
|
||||
fn cmp(&self, other: &Self) -> Ordering {
|
||||
// Candidates contain finite logits only. C treats signed zeros as
|
||||
// equal and breaks logit ties by the earlier vocabulary index.
|
||||
self.logit
|
||||
.partial_cmp(&other.logit)
|
||||
.unwrap()
|
||||
.then_with(|| other.id.cmp(&self.id))
|
||||
}
|
||||
}
|
||||
|
||||
fn ranked(heap: BinaryHeap<Reverse<Candidate>>) -> Vec<Candidate> {
|
||||
let mut candidates = heap.into_vec().into_iter().map(|c| c.0).collect::<Vec<_>>();
|
||||
candidates.sort_unstable_by(|a, b| b.cmp(a));
|
||||
candidates
|
||||
}
|
||||
|
||||
fn draw(candidates: &[Candidate], sum: f32, rng: &mut Rng) -> i32 {
|
||||
let mut remaining = rng.unit() * sum;
|
||||
for candidate in candidates {
|
||||
remaining -= candidate.probability;
|
||||
if remaining <= 0.0 {
|
||||
return candidate.id as i32;
|
||||
}
|
||||
}
|
||||
candidates.last().unwrap().id as i32
|
||||
}
|
||||
|
||||
fn filtered(candidates: &[Candidate], total: f32, top_p: f32, min_p: f32) -> (usize, f32, bool) {
|
||||
let minimum = (candidates[0].probability / total) * min_p;
|
||||
let mut sum = 0.0;
|
||||
let mut count = 0;
|
||||
let mut stopped_by_min_p = false;
|
||||
for (i, candidate) in candidates.iter().enumerate() {
|
||||
if i > 0 && candidate.probability / total < minimum {
|
||||
stopped_by_min_p = true;
|
||||
break;
|
||||
}
|
||||
sum += candidate.probability;
|
||||
count += 1;
|
||||
if sum / total >= top_p {
|
||||
break;
|
||||
}
|
||||
}
|
||||
(count, sum, stopped_by_min_p)
|
||||
}
|
||||
|
||||
#[allow(clippy::too_many_arguments)]
|
||||
fn fast_top_p(
|
||||
logits: &[f32],
|
||||
finite: usize,
|
||||
maximum: f32,
|
||||
best: i32,
|
||||
temperature: f32,
|
||||
top_p: f32,
|
||||
min_p: f32,
|
||||
rng: &mut Rng,
|
||||
) -> Option<i32> {
|
||||
const CAP: usize = 512;
|
||||
if finite > CAP && top_p >= 0.999 {
|
||||
return None;
|
||||
}
|
||||
let capacity = finite.min(CAP);
|
||||
let mut heap = BinaryHeap::<Reverse<Candidate>>::with_capacity(capacity);
|
||||
let mut total = 0.0;
|
||||
let mut heap_sum = 0.0;
|
||||
for (id, &logit) in logits.iter().enumerate() {
|
||||
if !logit.is_finite() {
|
||||
continue;
|
||||
}
|
||||
let probability = ((logit - maximum) / temperature).exp();
|
||||
total += probability;
|
||||
let candidate = Candidate {
|
||||
id,
|
||||
logit,
|
||||
probability,
|
||||
};
|
||||
if heap.len() < capacity {
|
||||
heap_sum += probability;
|
||||
heap.push(Reverse(candidate));
|
||||
} else if candidate > heap.peek().unwrap().0 {
|
||||
let mut worst = heap.peek_mut().unwrap();
|
||||
heap_sum -= worst.0.probability;
|
||||
worst.0 = candidate;
|
||||
heap_sum += probability;
|
||||
}
|
||||
}
|
||||
if total <= 0.0 || !total.is_finite() {
|
||||
return Some(best);
|
||||
}
|
||||
if heap.len() < finite && heap_sum < top_p * total {
|
||||
return None;
|
||||
}
|
||||
let candidates = ranked(heap);
|
||||
let min_p = if min_p > 0.0 { min_p } else { 0.0 };
|
||||
let (count, sum, stopped_by_min_p) = filtered(&candidates, total, top_p, min_p);
|
||||
// DS4 falls back if min-p stopped inside the heap but the unseen tail
|
||||
// might still pass its raw threshold. Do not consume RNG on fallback.
|
||||
if candidates.len() < finite
|
||||
&& stopped_by_min_p
|
||||
&& min_p > 0.0
|
||||
&& candidates.last().unwrap().probability >= candidates[0].probability * min_p
|
||||
{
|
||||
return None;
|
||||
}
|
||||
Some(if count == 0 {
|
||||
best
|
||||
} else {
|
||||
draw(&candidates[..count], sum, rng)
|
||||
})
|
||||
}
|
||||
|
||||
pub(super) fn sample(
|
||||
logits: &[f32],
|
||||
temperature: f32,
|
||||
top_p: f32,
|
||||
min_p: f32,
|
||||
top_k: i32,
|
||||
rng: &mut Rng,
|
||||
) -> i32 {
|
||||
if temperature <= 0.0 || logits.is_empty() {
|
||||
return ds4_sample_argmax(logits) as i32;
|
||||
}
|
||||
let top_p = if top_p <= 0.0 || top_p > 1.0 {
|
||||
1.0
|
||||
} else {
|
||||
top_p
|
||||
};
|
||||
let min_p = if min_p < 0.0 { 0.0 } else { min_p };
|
||||
if top_k > 0 {
|
||||
let capacity = (top_k as usize).min(1024).min(logits.len());
|
||||
let mut heap = BinaryHeap::<Reverse<Candidate>>::with_capacity(capacity);
|
||||
for (id, &logit) in logits.iter().enumerate() {
|
||||
if !logit.is_finite() {
|
||||
continue;
|
||||
}
|
||||
let candidate = Candidate {
|
||||
id,
|
||||
logit,
|
||||
probability: 0.0,
|
||||
};
|
||||
if heap.len() < capacity {
|
||||
heap.push(Reverse(candidate));
|
||||
} else if candidate > heap.peek().unwrap().0 {
|
||||
heap.peek_mut().unwrap().0 = candidate;
|
||||
}
|
||||
}
|
||||
if heap.is_empty() {
|
||||
return ds4_sample_argmax(logits) as i32;
|
||||
}
|
||||
let mut candidates = ranked(heap);
|
||||
let maximum = candidates[0].logit;
|
||||
let mut total = 0.0;
|
||||
for candidate in &mut candidates {
|
||||
candidate.probability = ((candidate.logit - maximum) / temperature).exp();
|
||||
total += candidate.probability;
|
||||
}
|
||||
if total <= 0.0 || !total.is_finite() {
|
||||
return candidates[0].id as i32;
|
||||
}
|
||||
let (count, sum, _) = filtered(&candidates, total, top_p, min_p);
|
||||
return if count == 0 {
|
||||
candidates[0].id as i32
|
||||
} else {
|
||||
draw(&candidates[..count], sum, rng)
|
||||
};
|
||||
}
|
||||
let mut maximum = -1.0e30_f32;
|
||||
let mut best = 0;
|
||||
let mut finite = 0;
|
||||
for (id, &logit) in logits.iter().enumerate() {
|
||||
if !logit.is_finite() {
|
||||
continue;
|
||||
}
|
||||
finite += 1;
|
||||
if logit > maximum {
|
||||
maximum = logit;
|
||||
best = id as i32;
|
||||
}
|
||||
}
|
||||
if finite == 0 {
|
||||
return ds4_sample_argmax(logits) as i32;
|
||||
}
|
||||
if top_p < 1.0
|
||||
&& let Some(token) = fast_top_p(
|
||||
logits,
|
||||
finite,
|
||||
maximum,
|
||||
best,
|
||||
temperature,
|
||||
top_p,
|
||||
min_p,
|
||||
rng,
|
||||
)
|
||||
{
|
||||
return token;
|
||||
}
|
||||
let mut candidates = Vec::with_capacity(finite);
|
||||
let mut total = 0.0;
|
||||
if top_p >= 1.0 {
|
||||
let min_relative = if min_p > 0.0 { min_p } else { 0.0 };
|
||||
if min_relative > 1.0 {
|
||||
return best;
|
||||
}
|
||||
// Same conservative expf-verified rejection boundary as DS4: skip
|
||||
// exponentials only when the value is guaranteed to fail min-p.
|
||||
let mut reject_scaled = None;
|
||||
if min_relative > 0.0 && min_relative.is_finite() {
|
||||
let mut cutoff = min_relative.ln();
|
||||
for _ in 0..8 {
|
||||
if !cutoff.is_finite() {
|
||||
break;
|
||||
}
|
||||
cutoff = cutoff.next_down();
|
||||
if cutoff.exp() < min_relative {
|
||||
reject_scaled = Some(cutoff);
|
||||
break;
|
||||
}
|
||||
}
|
||||
}
|
||||
for (id, &logit) in logits.iter().enumerate() {
|
||||
if !logit.is_finite() {
|
||||
continue;
|
||||
}
|
||||
let scaled = (logit - maximum) / temperature;
|
||||
if reject_scaled.is_some_and(|cutoff| scaled <= cutoff) {
|
||||
continue;
|
||||
}
|
||||
let probability = scaled.exp();
|
||||
if probability < min_relative {
|
||||
continue;
|
||||
}
|
||||
total += probability;
|
||||
candidates.push(Candidate {
|
||||
id,
|
||||
logit,
|
||||
probability,
|
||||
});
|
||||
}
|
||||
if total <= 0.0 || !total.is_finite() {
|
||||
return best;
|
||||
}
|
||||
let mut remaining = rng.unit() * total;
|
||||
for candidate in candidates {
|
||||
remaining -= candidate.probability;
|
||||
if remaining <= 0.0 {
|
||||
return candidate.id as i32;
|
||||
}
|
||||
}
|
||||
return best;
|
||||
}
|
||||
for (id, &logit) in logits.iter().enumerate() {
|
||||
if !logit.is_finite() {
|
||||
continue;
|
||||
}
|
||||
let probability = ((logit - maximum) / temperature).exp();
|
||||
total += probability;
|
||||
candidates.push(Candidate {
|
||||
id,
|
||||
logit,
|
||||
probability,
|
||||
});
|
||||
}
|
||||
if total <= 0.0 || !total.is_finite() {
|
||||
return best;
|
||||
}
|
||||
if min_p > 0.0 && min_p <= 1.0 {
|
||||
let minimum = (1.0 / total) * min_p;
|
||||
candidates.retain(|c| c.probability / total >= minimum);
|
||||
}
|
||||
if candidates.is_empty() {
|
||||
return best;
|
||||
}
|
||||
candidates.sort_unstable_by(|a, b| b.cmp(a));
|
||||
let min_p = if min_p > 0.0 { min_p } else { 0.0 };
|
||||
let (count, sum, _) = filtered(&candidates, total, top_p, min_p);
|
||||
if count == 0 {
|
||||
best
|
||||
} else {
|
||||
draw(&candidates[..count], sum, rng)
|
||||
}
|
||||
}
|
||||
|
||||
#[cfg(test)]
|
||||
mod tests {
|
||||
use super::*;
|
||||
|
||||
#[test]
|
||||
fn heap_fallback_preserves_rng_and_signed_zero_ties() {
|
||||
let flat = vec![0.0; 1024];
|
||||
let mut rng = Rng::new(42);
|
||||
for top_p in [0.95, 0.999] {
|
||||
assert_eq!(
|
||||
fast_top_p(&flat, 1024, 0.0, 0, 0.6, top_p, 0.0, &mut rng),
|
||||
None
|
||||
);
|
||||
}
|
||||
assert_eq!(rng.unit(), Rng::new(42).unit());
|
||||
let mut concentrated = vec![-100.0; 1024];
|
||||
concentrated[100] = 0.0;
|
||||
assert_eq!(
|
||||
fast_top_p(&concentrated, 1024, 0.0, 100, 0.6, 0.95, 0.0, &mut rng),
|
||||
Some(100)
|
||||
);
|
||||
for seed in 0..16 {
|
||||
for top_k in [0, 1] {
|
||||
assert_eq!(
|
||||
sample(
|
||||
&[-0.0, 0.0, -0.0, 0.0],
|
||||
1.0,
|
||||
0.95,
|
||||
1.1,
|
||||
top_k,
|
||||
&mut Rng::new(seed)
|
||||
),
|
||||
0
|
||||
);
|
||||
}
|
||||
}
|
||||
}
|
||||
}
|
||||
+7
-2
@@ -3,6 +3,7 @@ use sha2::{Digest, Sha256};
|
||||
use std::collections::HashMap;
|
||||
use std::fs::File;
|
||||
use std::path::{Path, PathBuf};
|
||||
use std::sync::Arc;
|
||||
|
||||
const MAGIC: u32 = 0x4655_4747;
|
||||
const MAX_COUNT: u64 = 1_000_000;
|
||||
@@ -48,7 +49,7 @@ pub(super) struct Tensor {
|
||||
|
||||
pub(super) struct Gguf {
|
||||
path: PathBuf,
|
||||
map: Mmap,
|
||||
map: Arc<Mmap>,
|
||||
data_offset: u64,
|
||||
max_tensor_bytes: u64,
|
||||
pub(super) metadata: HashMap<String, Value>,
|
||||
@@ -169,7 +170,7 @@ impl Gguf {
|
||||
|
||||
Ok(Self {
|
||||
path: path.to_owned(),
|
||||
map,
|
||||
map: Arc::new(map),
|
||||
data_offset: data_start,
|
||||
max_tensor_bytes,
|
||||
metadata,
|
||||
@@ -215,6 +216,10 @@ impl Gguf {
|
||||
self.map.as_ptr()
|
||||
}
|
||||
|
||||
pub(super) fn shared_map(&self) -> Arc<Mmap> {
|
||||
Arc::clone(&self.map)
|
||||
}
|
||||
|
||||
pub(super) fn data_offset(&self) -> u64 {
|
||||
self.data_offset
|
||||
}
|
||||
|
||||
+750
-280
File diff suppressed because it is too large
Load Diff
@@ -302,9 +302,7 @@ impl DeepSeekExecutor {
|
||||
self.logits = logits;
|
||||
self.checkpoint_tag = checkpoint_tag;
|
||||
if let Some(dspark) = &mut self.dspark {
|
||||
dspark.capture_mask = 0;
|
||||
dspark.cache_start = 0;
|
||||
dspark.cache_len = 0;
|
||||
dspark.reset_cache();
|
||||
}
|
||||
Ok(())
|
||||
}
|
||||
|
||||
+437
-41
@@ -1008,7 +1008,8 @@ impl GlmExecutor {
|
||||
)?;
|
||||
}
|
||||
std::mem::swap(&mut self.scratch.current, &mut self.scratch.next);
|
||||
if layer_index + 1 == DECODE_FLUSH_LAYERS {
|
||||
if glm_decode_flush_layer(layer_index + 1, self.weights.layers.len(), self.ssd.enabled)
|
||||
{
|
||||
call(
|
||||
unsafe { ds4_gpu_flush_commands() },
|
||||
"flushing the GLM decode graph",
|
||||
@@ -1419,7 +1420,15 @@ impl GlmExecutor {
|
||||
causal_range: bool,
|
||||
) -> Result<(), String> {
|
||||
let dsa = layer.dsa()?;
|
||||
self.encode_batch_attention_raw(batch, dsa, cache, pos, rows, selected_count, causal_range)
|
||||
self.encode_batch_attention_raw(
|
||||
batch,
|
||||
dsa,
|
||||
cache,
|
||||
pos,
|
||||
rows,
|
||||
selected_count,
|
||||
if causal_range { pos + rows } else { 0 },
|
||||
)
|
||||
}
|
||||
|
||||
#[allow(clippy::too_many_arguments)]
|
||||
@@ -1431,7 +1440,7 @@ impl GlmExecutor {
|
||||
pos: u32,
|
||||
rows: u32,
|
||||
selected_count: u32,
|
||||
causal_range: bool,
|
||||
dense_limit: u32,
|
||||
) -> Result<(), String> {
|
||||
let shape = self.model.shape;
|
||||
let map = self.model.main.map_ptr().cast();
|
||||
@@ -1446,7 +1455,12 @@ impl GlmExecutor {
|
||||
)?;
|
||||
let mut start = 0;
|
||||
while start < rows {
|
||||
let slice = (rows - start).min(INDEXED_PREFILL_ATTN_SLICE);
|
||||
let (slice, causal_range) = glm_attention_slice(pos + start, rows - start, dense_limit);
|
||||
let attention_selected_count = if causal_range {
|
||||
(pos + rows).min(dense_limit)
|
||||
} else {
|
||||
selected_count
|
||||
};
|
||||
let q = batch
|
||||
.q
|
||||
.view(u64::from(start) * q_dim * 4, u64::from(slice) * q_dim * 4)?;
|
||||
@@ -1471,7 +1485,7 @@ impl GlmExecutor {
|
||||
cache.kv.raw(),
|
||||
pos + start,
|
||||
slice,
|
||||
selected_count,
|
||||
attention_selected_count,
|
||||
self.context,
|
||||
CACHE_F16,
|
||||
shape.heads as u32,
|
||||
@@ -1492,7 +1506,7 @@ impl GlmExecutor {
|
||||
cache.rope.raw(),
|
||||
slice,
|
||||
pos + start,
|
||||
selected_count,
|
||||
attention_selected_count,
|
||||
self.context,
|
||||
CACHE_F16,
|
||||
shape.heads as u32,
|
||||
@@ -2007,6 +2021,29 @@ impl GlmExecutor {
|
||||
result
|
||||
}
|
||||
|
||||
pub(crate) fn rewind_speculative_output(&mut self, pos: usize) -> Result<(), String> {
|
||||
if pos > self.tokens.len() {
|
||||
return Err("GLM emitted token frontier exceeds evaluated tokens".into());
|
||||
}
|
||||
if pos == self.tokens.len() {
|
||||
return Ok(());
|
||||
}
|
||||
if self.model.shape.model == ModelChoice::Glm53Flash {
|
||||
return if self.rewind_glm53_mtp(pos)? {
|
||||
Ok(())
|
||||
} else {
|
||||
Err("GLM could not restore the emitted token frontier".into())
|
||||
};
|
||||
}
|
||||
self.tokens.truncate(pos);
|
||||
if let Some(mtp) = &mut self.mtp {
|
||||
mtp.pending = None;
|
||||
mtp.parent = None;
|
||||
mtp.rollback = None;
|
||||
}
|
||||
Ok(())
|
||||
}
|
||||
|
||||
fn rewind_glm53_mtp(&mut self, pos: usize) -> Result<bool, String> {
|
||||
let Some((start, first)) = self.mtp.as_ref().and_then(|m| m.rollback) else {
|
||||
return Ok(false);
|
||||
@@ -3137,7 +3174,7 @@ impl GlmExecutor {
|
||||
)?;
|
||||
|
||||
let visible = pos + rows;
|
||||
let causal = pos < dense_limit;
|
||||
let causal = visible <= dense_limit;
|
||||
let selected_count = if causal {
|
||||
visible.min(dense_limit)
|
||||
} else {
|
||||
@@ -3250,7 +3287,15 @@ impl GlmExecutor {
|
||||
},
|
||||
"projecting batched GLM 5.3 low-rank queries",
|
||||
)?;
|
||||
self.encode_batch_attention_raw(batch, weights, cache, pos, rows, selected_count, causal)?;
|
||||
self.encode_batch_attention_raw(
|
||||
batch,
|
||||
weights,
|
||||
cache,
|
||||
pos,
|
||||
rows,
|
||||
selected_count,
|
||||
dense_limit,
|
||||
)?;
|
||||
glm53_project_rows(
|
||||
&batch.attn_out,
|
||||
weights.output,
|
||||
@@ -3467,7 +3512,7 @@ impl GlmExecutor {
|
||||
"expanding GLM 5.3 FFN hyperconnections",
|
||||
)?;
|
||||
std::mem::swap(&mut self.scratch.hc_current, &mut self.scratch.hc_next);
|
||||
if ordinal + 1 == DECODE_FLUSH_LAYERS {
|
||||
if glm_decode_flush_layer(ordinal + 1, self.weights.layers.len(), self.ssd.enabled) {
|
||||
call(
|
||||
unsafe { ds4_gpu_flush_commands() },
|
||||
"flushing GLM 5.3 decode",
|
||||
@@ -4555,14 +4600,12 @@ impl GlmExecutor {
|
||||
}
|
||||
|
||||
let pos = self.position();
|
||||
let mut rows =
|
||||
indexed_prefill_rows(pos, remaining.len(), self.model.shape.indexer_top_k as u32);
|
||||
if self.model.shape.model == ModelChoice::Glm53Flash {
|
||||
let dense = glm53_dense_limit(self.context, self.ssd.enabled);
|
||||
if pos < dense {
|
||||
rows = rows.min((dense - pos) as usize);
|
||||
}
|
||||
}
|
||||
let rows = indexed_prefill_rows(
|
||||
self.model.shape.model,
|
||||
pos,
|
||||
remaining.len(),
|
||||
self.model.shape.indexer_top_k as u32,
|
||||
);
|
||||
let cancelled = self.eval_batch(
|
||||
&remaining[..rows],
|
||||
&mut progress,
|
||||
@@ -4760,13 +4803,7 @@ impl GlmExecutor {
|
||||
let shape = self.model.shape;
|
||||
let map = self.model.main.map_ptr().cast();
|
||||
let size = self.model.main.len();
|
||||
let mut dense_limit = glm53_dense_limit(self.context, self.ssd.enabled);
|
||||
// DS4 routes the entire verification pair through indexed attention
|
||||
// when it no longer fits the full-attention span. Normal prefill
|
||||
// already splits at this boundary before entering the batch function.
|
||||
if pos < dense_limit && pos + rows > dense_limit {
|
||||
dense_limit = 0;
|
||||
}
|
||||
let dense_limit = glm53_dense_limit(self.context, self.ssd.enabled);
|
||||
// Both HC split/sum and expansion derive their row count from the
|
||||
// output view, not a rows argument. Match DS4's token-sized views even
|
||||
// when a short prefill reuses the full 2048-row workspace.
|
||||
@@ -4930,6 +4967,16 @@ impl GlmExecutor {
|
||||
},
|
||||
"expanding batched GLM 5.3 attention hyperconnections",
|
||||
)?;
|
||||
#[cfg(test)]
|
||||
{
|
||||
commands = compare_glm53_reference_stages(
|
||||
commands,
|
||||
index,
|
||||
pos,
|
||||
rows as usize * shape.embd as usize,
|
||||
&[("attn_out", &batch.attn_out)],
|
||||
)?;
|
||||
}
|
||||
glm53_hc_pre(
|
||||
after_attn,
|
||||
&batch.ffn_norm,
|
||||
@@ -4947,6 +4994,16 @@ impl GlmExecutor {
|
||||
size,
|
||||
)?;
|
||||
self.encode_glm53_ffn_batch(batch, layer, index as u32, rows)?;
|
||||
#[cfg(test)]
|
||||
{
|
||||
commands = compare_glm53_reference_stages(
|
||||
commands,
|
||||
index,
|
||||
pos,
|
||||
rows as usize * shape.embd as usize,
|
||||
&[("ffn_out", &batch.next)],
|
||||
)?;
|
||||
}
|
||||
if let Some(steering) = &self.steering {
|
||||
steering.apply(&batch.next, index as u32, rows, false)?;
|
||||
}
|
||||
@@ -4963,10 +5020,11 @@ impl GlmExecutor {
|
||||
},
|
||||
"expanding batched GLM 5.3 FFN hyperconnections",
|
||||
)?;
|
||||
let drain = glm_prefill_flush_layers(rows, read_logits)
|
||||
&& (index + 1).is_multiple_of(PREFILL_DRAIN_LAYERS)
|
||||
&& index + 1 < self.weights.layers.len();
|
||||
if glm_prefill_flush_layers(rows, read_logits) {
|
||||
if (index + 1).is_multiple_of(PREFILL_DRAIN_LAYERS)
|
||||
&& index + 1 < self.weights.layers.len()
|
||||
{
|
||||
if drain {
|
||||
commands.finish()?;
|
||||
commands = Commands::begin()?;
|
||||
} else {
|
||||
@@ -4989,11 +5047,15 @@ impl GlmExecutor {
|
||||
if let Some(views) = &mut hc_views {
|
||||
views.swap(2, 3);
|
||||
}
|
||||
let done = u32::try_from(
|
||||
u64::from(rows) * (index as u64 + 1) / self.weights.layers.len() as u64,
|
||||
)
|
||||
.expect("GLM 5.3 prefill progress is bounded by the batch");
|
||||
cancelled |= !progress(pos + done);
|
||||
// DS4 reports completed work at drains, not every encoded layer.
|
||||
// An asynchronous flush alone does not advance GPU completion.
|
||||
if drain {
|
||||
let done = u32::try_from(
|
||||
u64::from(rows) * (index as u64 + 1) / self.weights.layers.len() as u64,
|
||||
)
|
||||
.expect("GLM 5.3 prefill progress is bounded by the batch");
|
||||
cancelled |= !progress(pos + done);
|
||||
}
|
||||
}
|
||||
commands.finish()?;
|
||||
|
||||
@@ -5024,6 +5086,7 @@ impl GlmExecutor {
|
||||
self.logits = logits;
|
||||
}
|
||||
self.tokens.extend_from_slice(tokens);
|
||||
cancelled |= !progress(self.position());
|
||||
if let Some(profile) = &self.profile {
|
||||
profile.write()?;
|
||||
}
|
||||
@@ -5513,14 +5576,35 @@ fn glm_prefill_flush_layers(rows: u32, logits_requested: bool) -> bool {
|
||||
logits_requested && rows > 8
|
||||
}
|
||||
|
||||
fn indexed_prefill_rows(pos: u32, remaining: usize, top_k: u32) -> usize {
|
||||
fn glm_decode_flush_layer(completed: usize, layers: usize, ssd_streaming: bool) -> bool {
|
||||
// The resident indexed DS4 graph flushes every four layers, except the
|
||||
// final layer. Static SSD mappings use the expert loader's own boundaries.
|
||||
!ssd_streaming && completed < layers && completed.is_multiple_of(DECODE_FLUSH_LAYERS)
|
||||
}
|
||||
|
||||
fn indexed_prefill_rows(model: ModelChoice, pos: u32, remaining: usize, top_k: u32) -> usize {
|
||||
let mut rows = remaining.min(INDEXED_PREFILL_CHUNK as usize);
|
||||
if pos < top_k {
|
||||
// GLM 5.3's active DS4 graph keeps full chunks across both the old
|
||||
// top-k boundary and the dense/sparse boundary. Only attention is sliced.
|
||||
if model != ModelChoice::Glm53Flash && pos < top_k {
|
||||
rows = rows.min((top_k - pos) as usize);
|
||||
}
|
||||
rows
|
||||
}
|
||||
|
||||
fn glm_attention_slice(pos: u32, remaining: u32, dense_limit: u32) -> (u32, bool) {
|
||||
let causal = pos < dense_limit;
|
||||
let rows = remaining.min(INDEXED_PREFILL_ATTN_SLICE);
|
||||
(
|
||||
if causal {
|
||||
rows.min(dense_limit - pos)
|
||||
} else {
|
||||
rows
|
||||
},
|
||||
causal,
|
||||
)
|
||||
}
|
||||
|
||||
fn full_indexer_layer(shape: super::super::Shape, layer: usize) -> bool {
|
||||
if shape.model == ModelChoice::Glm53Flash {
|
||||
return layer < (shape.layers - shape.nextn) as usize && !glm53_kda_layer(shape, layer);
|
||||
@@ -6432,12 +6516,72 @@ fn f32_project_rows(
|
||||
)
|
||||
}
|
||||
|
||||
// Diagnostic only: no stage reads, extra drains or environment checks in the
|
||||
// release app. Match original DS4's existing layer0/bootstrap tensor dumps.
|
||||
#[cfg(test)]
|
||||
fn compare_glm53_reference_stages(
|
||||
commands: Commands,
|
||||
layer: usize,
|
||||
pos: u32,
|
||||
values: usize,
|
||||
tensors: &[(&str, &Buffer)],
|
||||
) -> Result<Commands, String> {
|
||||
let Some(prefix) = std::env::var_os("DS4SERVER_GLM_STAGE_REFERENCE") else {
|
||||
return Ok(commands);
|
||||
};
|
||||
let selected_pos = std::env::var("DS4SERVER_GLM_STAGE_POS")
|
||||
.map_or(Ok(0), |v| v.parse::<u32>())
|
||||
.map_err(|e| e.to_string())?;
|
||||
if layer != 0 || pos != selected_pos {
|
||||
return Ok(commands);
|
||||
}
|
||||
commands.finish()?;
|
||||
for &(name, tensor) in tensors {
|
||||
let path = format!("{}_{name}-{layer}_pos{pos}.bin", prefix.to_string_lossy());
|
||||
let bytes = std::fs::read(&path).map_err(|e| format!("{path}: {e}"))?;
|
||||
if values == 0 || bytes.len() != values * 4 {
|
||||
return Err(format!(
|
||||
"reference tensor geometry mismatch: {path}: {} bytes, expected {}",
|
||||
bytes.len(),
|
||||
values * 4
|
||||
));
|
||||
}
|
||||
let reference: Vec<f32> = bytes
|
||||
.chunks_exact(4)
|
||||
.map(|b| f32::from_le_bytes(b.try_into().unwrap()))
|
||||
.collect();
|
||||
let mut actual = vec![0.0; reference.len()];
|
||||
tensor.read_f32(&mut actual)?;
|
||||
if !actual.iter().chain(&reference).all(|v| v.is_finite()) {
|
||||
return Err(format!("nonfinite stage {name}"));
|
||||
}
|
||||
let max_abs = actual
|
||||
.iter()
|
||||
.zip(&reference)
|
||||
.map(|(a, b)| (a - b).abs())
|
||||
.fold(0.0_f32, f32::max);
|
||||
let rms = (actual
|
||||
.iter()
|
||||
.zip(&reference)
|
||||
.map(|(a, b)| f64::from(a - b).powi(2))
|
||||
.sum::<f64>()
|
||||
/ actual.len() as f64)
|
||||
.sqrt();
|
||||
eprintln!(
|
||||
"{}",
|
||||
serde_json::json!({"event":"glm_stage_compare", "name":name, "layer":layer, "pos":pos, "values":actual.len(), "max_abs":max_abs, "rms":rms})
|
||||
);
|
||||
}
|
||||
Commands::begin()
|
||||
}
|
||||
|
||||
#[cfg(test)]
|
||||
mod tests {
|
||||
use super::{
|
||||
GlmExecutor, argmax, dynamic_expert_budget, full_indexer_layer, glm_prefill_flush_layers,
|
||||
glm53_dense_compact_prefill, glm53_dense_limit, indexed_prefill_rows,
|
||||
live_prefix_rewind_target, streaming_token_prefill_eligible,
|
||||
GlmExecutor, argmax, dynamic_expert_budget, full_indexer_layer, glm_attention_slice,
|
||||
glm_decode_flush_layer, glm_prefill_flush_layers, glm53_dense_compact_prefill,
|
||||
glm53_dense_limit, indexed_prefill_rows, live_prefix_rewind_target,
|
||||
streaming_token_prefill_eligible,
|
||||
};
|
||||
use crate::engine::{GLM, Model, ReasoningMode};
|
||||
use crate::model::ModelChoice;
|
||||
@@ -6464,9 +6608,30 @@ mod tests {
|
||||
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), 2048);
|
||||
assert_eq!(indexed_prefill_rows(0, 17, 2048), 17);
|
||||
assert_eq!(
|
||||
indexed_prefill_rows(ModelChoice::Glm52, 0, 5000, 2048),
|
||||
2048
|
||||
);
|
||||
assert_eq!(
|
||||
indexed_prefill_rows(ModelChoice::Glm52, 2048, 5000, 2048),
|
||||
2048
|
||||
);
|
||||
assert_eq!(indexed_prefill_rows(ModelChoice::Glm52, 0, 17, 2048), 17);
|
||||
assert_eq!(
|
||||
indexed_prefill_rows(ModelChoice::Glm52, 9, 2626, 2048),
|
||||
2039
|
||||
);
|
||||
for pos in [9, 2047, 4095, 8191] {
|
||||
assert_eq!(
|
||||
indexed_prefill_rows(ModelChoice::Glm53Flash, pos, 2626, 2048),
|
||||
2048
|
||||
);
|
||||
}
|
||||
for boundary in [4096, 8192] {
|
||||
assert_eq!(glm_attention_slice(boundary - 1, 2, boundary), (1, true));
|
||||
assert_eq!(glm_attention_slice(boundary, 1, boundary), (1, false));
|
||||
assert_eq!(glm_attention_slice(9, 2048, boundary), (2048, true));
|
||||
}
|
||||
assert_eq!(glm53_dense_limit(32_768, false), 4096);
|
||||
assert_eq!(glm53_dense_limit(32_768, true), 8192);
|
||||
assert_eq!(glm53_dense_limit(65_536, true), 4096);
|
||||
@@ -6477,6 +6642,212 @@ mod tests {
|
||||
assert!(!glm_prefill_flush_layers(2048, false));
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn glm_resident_decode_flushes_periodically_without_a_final_or_ssd_flush() {
|
||||
for layers in [1, 4, 5, 76, 78] {
|
||||
let actual = (1..=layers)
|
||||
.filter(|&completed| glm_decode_flush_layer(completed, layers, false))
|
||||
.collect::<Vec<_>>();
|
||||
assert_eq!(actual, (4..layers).step_by(4).collect::<Vec<_>>());
|
||||
assert!((1..=layers).all(|completed| !glm_decode_flush_layer(completed, layers, true)));
|
||||
}
|
||||
}
|
||||
|
||||
#[test]
|
||||
#[ignore = "installed GLM and DS4SERVER_GLM_REFERENCE_LOG from the standalone reference"]
|
||||
fn glm53_reference_prompt_tokens_match_shared_runtime() {
|
||||
let path = std::env::var("DS4SERVER_GLM_REFERENCE_LOG").unwrap();
|
||||
let events = std::fs::read_to_string(path)
|
||||
.unwrap()
|
||||
.lines()
|
||||
.map(|line| serde_json::from_str::<serde_json::Value>(line).unwrap())
|
||||
.collect::<Vec<_>>();
|
||||
let start = events
|
||||
.iter()
|
||||
.find(|e| e["event"] == "reference_start")
|
||||
.unwrap();
|
||||
assert_eq!(start["family"], "glm");
|
||||
let model = Model::open_main(
|
||||
std::path::Path::new(start["model"].as_str().unwrap()),
|
||||
ModelChoice::Glm53Flash,
|
||||
)
|
||||
.unwrap();
|
||||
let readme = std::fs::read_to_string(start["readme"].as_str().unwrap()).unwrap();
|
||||
let prompts = [
|
||||
"Reply with exactly: OK".to_owned(),
|
||||
format!("Give a summary of the following text:\n\n{readme}"),
|
||||
"Tell me a complete short story about a lighthouse keeper. Do not ask questions.".into(),
|
||||
"Write a Python function is_prime(n: int) -> bool, followed by five assert examples. No tools.".into(),
|
||||
];
|
||||
let ids = |v: &serde_json::Value| {
|
||||
v.as_array()
|
||||
.unwrap()
|
||||
.iter()
|
||||
.map(|v| v.as_i64().unwrap() as i32)
|
||||
.collect::<Vec<_>>()
|
||||
};
|
||||
let mut frontier = Vec::new();
|
||||
for (index, prompt) in prompts.iter().enumerate() {
|
||||
let expected = if index <= 1 {
|
||||
model.render_prompt("", prompt, ReasoningMode::Low)
|
||||
} else {
|
||||
let mut expected = frontier.clone();
|
||||
expected.extend(model.tokenizer.encode_continuation(
|
||||
prompt,
|
||||
ReasoningMode::Low,
|
||||
false,
|
||||
));
|
||||
expected
|
||||
};
|
||||
let receipt = events
|
||||
.iter()
|
||||
.find(|e| e["event"] == "reference_prefill" && e["turn"] == index)
|
||||
.unwrap();
|
||||
assert_eq!(
|
||||
ids(&receipt["tokens"]),
|
||||
expected,
|
||||
"reference turn {index} prompt differs"
|
||||
);
|
||||
let cached = if index <= 1 { 9 } else { frontier.len() };
|
||||
assert_eq!(
|
||||
receipt["cached_tokens"], cached,
|
||||
"bootstrap/continuation differs"
|
||||
);
|
||||
let result = events
|
||||
.iter()
|
||||
.find(|e| e["event"] == "reference_result" && e["turn"] == index)
|
||||
.unwrap();
|
||||
if index != 0 {
|
||||
assert_eq!(result["finish_reason"], "stop");
|
||||
}
|
||||
frontier = expected;
|
||||
frontier.extend(ids(&result["token_ids"]));
|
||||
}
|
||||
}
|
||||
|
||||
#[test]
|
||||
#[ignore = "requires the installed GLM 5.3 Flash Q2 checkpoint and Apple Metal"]
|
||||
fn glm53_reference_logits_replay_separates_sampling_from_execution() {
|
||||
use crate::engine::{Rng, sample};
|
||||
let events: Vec<serde_json::Value> =
|
||||
std::fs::read_to_string(std::env::var("DS4SERVER_GLM_REFERENCE_LOG").unwrap())
|
||||
.unwrap()
|
||||
.lines()
|
||||
.map(|l| serde_json::from_str(l).unwrap())
|
||||
.collect();
|
||||
let event = |name: &str| {
|
||||
events
|
||||
.iter()
|
||||
.find(|e| e["event"] == name && (name == "reference_start" || e["turn"] == 1))
|
||||
.unwrap()
|
||||
};
|
||||
let start = event("reference_start");
|
||||
assert_eq!(start["family"], "glm");
|
||||
assert_eq!(start["acceleration"], false);
|
||||
assert_eq!(start["seed"], 42);
|
||||
let trace = event("reference_logits_trace");
|
||||
let path: std::ffi::OsString = serde_json::from_value(trace["path"].clone()).unwrap();
|
||||
let bytes = std::fs::read(std::path::PathBuf::from(path)).unwrap();
|
||||
let n = trace["vocab"].as_u64().unwrap() as usize;
|
||||
assert_eq!(bytes.len(), 32 * n * 4);
|
||||
let values: Vec<f32> = bytes
|
||||
.chunks_exact(4)
|
||||
.map(|b| f32::from_le_bytes(b.try_into().unwrap()))
|
||||
.collect();
|
||||
let token_ids = |v: &serde_json::Value| {
|
||||
v.as_array()
|
||||
.unwrap()
|
||||
.iter()
|
||||
.map(|v| v.as_i64().unwrap() as i32)
|
||||
.collect::<Vec<_>>()
|
||||
};
|
||||
let prompt = token_ids(&event("reference_prefill")["tokens"]);
|
||||
let expected = token_ids(&event("reference_result")["token_ids"]);
|
||||
assert!(expected.len() >= 32);
|
||||
// First feed the exact C-produced rows through the production Rust
|
||||
// sampler, so a model-graph difference cannot masquerade as RNG drift.
|
||||
let mut host_rng = Rng::new(42);
|
||||
let host: Vec<i32> = values
|
||||
.chunks_exact(n)
|
||||
.map(|row| sample(row, 0.6, 0.95, 0.0, 0, &mut host_rng))
|
||||
.collect();
|
||||
assert_eq!(
|
||||
host,
|
||||
expected[..32],
|
||||
"Rust sampler differs on original DS4 logits"
|
||||
);
|
||||
super::super::configure_sources().unwrap();
|
||||
let model = Model::open_main(
|
||||
std::path::Path::new(start["model"].as_str().unwrap()),
|
||||
ModelChoice::Glm53Flash,
|
||||
)
|
||||
.unwrap();
|
||||
assert_eq!(n, model.shape.vocab as usize);
|
||||
let mut executor = GlmExecutor::open(
|
||||
model,
|
||||
32768,
|
||||
false,
|
||||
EngineSsdSettings {
|
||||
enabled: false,
|
||||
cold: false,
|
||||
cache_experts: 0,
|
||||
cache_bytes: 0,
|
||||
full_layers: 0,
|
||||
full_layers_set: false,
|
||||
preload_experts: 0,
|
||||
},
|
||||
)
|
||||
.unwrap();
|
||||
assert_eq!(event("reference_prefill")["cached_tokens"], 9);
|
||||
for part in [&prompt[..9], &prompt[9..]] {
|
||||
assert_eq!(
|
||||
executor
|
||||
.prefill(part, |pos| {
|
||||
eprintln!("reference replay prefill {pos}");
|
||||
true
|
||||
})
|
||||
.unwrap(),
|
||||
part.len()
|
||||
);
|
||||
}
|
||||
let mut gpu_rng = Rng::new(42);
|
||||
let mut different = Vec::new();
|
||||
for (step, row) in values.chunks_exact(n).enumerate() {
|
||||
let actual = executor.logits();
|
||||
let chosen = sample(actual, 0.6, 0.95, 0.0, 0, &mut gpu_rng);
|
||||
let max_abs = actual
|
||||
.iter()
|
||||
.zip(row)
|
||||
.map(|(a, b)| (a - b).abs())
|
||||
.fold(0.0_f32, f32::max);
|
||||
let rms = (actual
|
||||
.iter()
|
||||
.zip(row)
|
||||
.map(|(a, b)| f64::from(a - b).powi(2))
|
||||
.sum::<f64>()
|
||||
/ n as f64)
|
||||
.sqrt();
|
||||
eprintln!(
|
||||
"{}",
|
||||
serde_json::json!({"event":"glm_logits_compare", "step":step, "expected":expected[step], "actual":chosen, "max_abs":max_abs, "rms":rms})
|
||||
);
|
||||
if chosen != expected[step]
|
||||
|| !actual
|
||||
.iter()
|
||||
.zip(row)
|
||||
.all(|(a, b)| a.to_bits() == b.to_bits())
|
||||
{
|
||||
different.push(step);
|
||||
}
|
||||
// Never diverge the context: both backends see the reference token.
|
||||
executor.eval(expected[step]).unwrap();
|
||||
}
|
||||
assert!(
|
||||
different.is_empty(),
|
||||
"logits or sampled outputs differ on identical history at steps {different:?}"
|
||||
);
|
||||
}
|
||||
|
||||
#[test]
|
||||
#[ignore = "requires the installed GLM 5.3 Flash Q2 checkpoint and Apple Metal"]
|
||||
fn glm53_m5_text_performance_gate() {
|
||||
@@ -7171,12 +7542,30 @@ mod tests {
|
||||
.unwrap();
|
||||
let prepare = |executor: &mut GlmExecutor| {
|
||||
executor.reset().unwrap();
|
||||
let mut reported = Vec::new();
|
||||
executor
|
||||
.prefill(&prompt, |pos| {
|
||||
reported.push(pos);
|
||||
eprintln!("verify2 prefill {pos}");
|
||||
true
|
||||
})
|
||||
.unwrap();
|
||||
if prompt.len() <= super::INDEXED_PREFILL_CHUNK as usize {
|
||||
let layers = executor.weights.layers.len();
|
||||
let mut expected = vec![0];
|
||||
if prompt.len() > 8 {
|
||||
expected.extend(
|
||||
(super::PREFILL_DRAIN_LAYERS..layers)
|
||||
.step_by(super::PREFILL_DRAIN_LAYERS)
|
||||
.map(|done| (prompt.len() * done / layers) as u32),
|
||||
);
|
||||
}
|
||||
expected.push(prompt.len() as u32);
|
||||
assert_eq!(
|
||||
reported, expected,
|
||||
"progress must follow completed GPU work"
|
||||
);
|
||||
}
|
||||
};
|
||||
let recurrent = |executor: &GlmExecutor| {
|
||||
executor
|
||||
@@ -7230,6 +7619,9 @@ mod tests {
|
||||
eprintln!("partial prefill HC workspace guards passed");
|
||||
}
|
||||
let before = recurrent(&executor);
|
||||
assert!(executor.rewind_speculative_output(frontier + 1).is_err());
|
||||
executor.rewind_speculative_output(frontier).unwrap();
|
||||
assert_eq!(recurrent(&executor), before);
|
||||
let first = argmax(executor.logits());
|
||||
executor.eval(first).unwrap();
|
||||
let after_first = recurrent(&executor);
|
||||
@@ -7259,7 +7651,11 @@ mod tests {
|
||||
executor.mtp.as_ref().unwrap().rollback,
|
||||
Some((start, first))
|
||||
);
|
||||
assert!(executor.rewind_glm53_mtp(start as usize + keep).unwrap());
|
||||
// The UI/headless consumer uses this same path to remove an
|
||||
// MTP-returned stop token before retaining the ongoing chat.
|
||||
executor
|
||||
.rewind_speculative_output(start as usize + keep)
|
||||
.unwrap();
|
||||
assert_eq!(executor.position(), start + keep as u32);
|
||||
assert!(
|
||||
recurrent(&executor) == *if keep == 0 { &before } else { &after_first },
|
||||
|
||||
@@ -45,6 +45,9 @@ pub(super) struct QwenKernelArgs {
|
||||
pub(super) struct GpuCanarySample {
|
||||
pub(super) scheduled_seconds: f64,
|
||||
pub(super) completed_seconds: f64,
|
||||
pub(super) gpu_wait_seconds: f64,
|
||||
pub(super) gpu_interval_seconds: f64,
|
||||
pub(super) host_return_seconds: f64,
|
||||
}
|
||||
|
||||
#[derive(Clone, Copy, Default)]
|
||||
|
||||
@@ -0,0 +1,327 @@
|
||||
//! DSpark's CPU Markov head: persistent workers, DS4 row partition/reduction.
|
||||
use super::{F16, F32, Gguf, Q8_0, Weight, dense_dot_bytes, dense_dot_q8, quantize_q8_activation};
|
||||
use memmap2::Mmap;
|
||||
use std::sync::{Arc, mpsc};
|
||||
use std::thread::{self, JoinHandle};
|
||||
|
||||
type Best = (usize, f32);
|
||||
|
||||
struct Input {
|
||||
values: Vec<f32>,
|
||||
quantized: Option<(Vec<i8>, Vec<f32>)>,
|
||||
logits: Vec<f32>,
|
||||
}
|
||||
|
||||
struct Matrix {
|
||||
map: Arc<Mmap>,
|
||||
weight: Weight,
|
||||
width: usize,
|
||||
rows: usize,
|
||||
row_bytes: usize,
|
||||
}
|
||||
|
||||
impl Matrix {
|
||||
fn best(&self, input: &Input, start: usize, end: usize) -> Best {
|
||||
let bytes = &self.map[self.weight.offset as usize..];
|
||||
let mut best = (start, -f32::MAX);
|
||||
for token in start..end {
|
||||
let row = &bytes[token * self.row_bytes..(token + 1) * self.row_bytes];
|
||||
let dot = input.quantized.as_ref().map_or_else(
|
||||
|| dense_dot_bytes(self.weight.kind, self.width, row, &input.values),
|
||||
|(values, scales)| dense_dot_q8(row, values, scales, self.width),
|
||||
);
|
||||
let score = input.logits[token] + dot;
|
||||
if score > best.1 {
|
||||
best = (token, score);
|
||||
}
|
||||
}
|
||||
best
|
||||
}
|
||||
}
|
||||
|
||||
struct Worker {
|
||||
input: Option<mpsc::SyncSender<Arc<Input>>>,
|
||||
result: mpsc::Receiver<Best>,
|
||||
thread: Option<JoinHandle<()>>,
|
||||
}
|
||||
|
||||
impl Drop for Worker {
|
||||
fn drop(&mut self) {
|
||||
self.input.take();
|
||||
if let Some(thread) = self.thread.take() {
|
||||
let _ = thread.join();
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
pub(super) struct MarkovPool {
|
||||
matrix: Arc<Matrix>,
|
||||
workers: Vec<Worker>,
|
||||
chunk: usize,
|
||||
}
|
||||
|
||||
fn worker_count(online: usize, requested: Option<&str>) -> usize {
|
||||
requested
|
||||
.and_then(|value| value.parse::<usize>().ok())
|
||||
.filter(|&value| value > 0)
|
||||
.unwrap_or(online.min(12))
|
||||
.clamp(1, 32)
|
||||
}
|
||||
|
||||
impl MarkovPool {
|
||||
pub(super) fn new(model: &Gguf, weight: Weight) -> Result<Self, String> {
|
||||
let threads = worker_count(
|
||||
thread::available_parallelism().map_or(1, std::num::NonZero::get),
|
||||
std::env::var("DS4_THREADS").ok().as_deref(),
|
||||
);
|
||||
Self::from_map(model.shared_map(), weight, threads)
|
||||
}
|
||||
|
||||
fn from_map(map: Arc<Mmap>, weight: Weight, threads: usize) -> Result<Self, String> {
|
||||
let width = usize::try_from(weight.dims[0]).map_err(|_| "Markov width overflow")?;
|
||||
let rows = usize::try_from(weight.dims[1]).map_err(|_| "Markov rows overflow")?;
|
||||
if width == 0 || rows == 0 || rows > i32::MAX as usize {
|
||||
return Err("DSpark dense argmax has invalid dimensions".into());
|
||||
}
|
||||
let row_bytes = match weight.kind {
|
||||
F32 => width.checked_mul(4),
|
||||
F16 => width.checked_mul(2),
|
||||
Q8_0 => width.div_ceil(32).checked_mul(34),
|
||||
_ => None,
|
||||
}
|
||||
.ok_or("unsupported DSpark dense tensor layout")?;
|
||||
let bytes = row_bytes
|
||||
.checked_mul(rows)
|
||||
.ok_or("DSpark dense argmax size overflow")?;
|
||||
if weight.offset > map.len() as u64 || bytes as u64 > map.len() as u64 - weight.offset {
|
||||
return Err("DSpark dense argmax is outside the GGUF mapping".into());
|
||||
}
|
||||
let matrix = Arc::new(Matrix {
|
||||
map,
|
||||
weight,
|
||||
width,
|
||||
rows,
|
||||
row_bytes,
|
||||
});
|
||||
let mut pool = Self {
|
||||
matrix,
|
||||
workers: Vec::new(),
|
||||
chunk: rows.div_ceil(threads),
|
||||
};
|
||||
for slot in 1..threads {
|
||||
let matrix = Arc::clone(&pool.matrix);
|
||||
let start = slot * pool.chunk;
|
||||
let end = (start + pool.chunk).min(rows);
|
||||
let (input, jobs) = mpsc::sync_channel::<Arc<Input>>(1);
|
||||
let (results, result) = mpsc::channel();
|
||||
let thread = thread::Builder::new()
|
||||
.name(format!("dspark-markov-{slot}"))
|
||||
.spawn(move || {
|
||||
while let Ok(input) = jobs.recv() {
|
||||
let best = matrix.best(&input, start, end);
|
||||
// Return buffer ownership before signalling completion.
|
||||
drop(input);
|
||||
if results.send(best).is_err() {
|
||||
break;
|
||||
}
|
||||
}
|
||||
})
|
||||
.map_err(|error| format!("Cannot start DSpark Markov worker: {error}"))?;
|
||||
pool.workers.push(Worker {
|
||||
input: Some(input),
|
||||
result,
|
||||
thread: Some(thread),
|
||||
});
|
||||
}
|
||||
Ok(pool)
|
||||
}
|
||||
|
||||
pub(super) fn argmax(&mut self, values: &[f32], logits: &mut Vec<f32>) -> Result<i32, String> {
|
||||
if values.len() != self.matrix.width || logits.len() != self.matrix.rows {
|
||||
return Err("DSpark dense argmax has mismatched dimensions".into());
|
||||
}
|
||||
let input = Arc::new(Input {
|
||||
values: values.to_vec(),
|
||||
quantized: (self.matrix.weight.kind == Q8_0).then(|| quantize_q8_activation(values)),
|
||||
logits: std::mem::take(logits),
|
||||
});
|
||||
let parallel = !self.workers.is_empty() && self.matrix.rows >= 512;
|
||||
let mut failed = false;
|
||||
if parallel {
|
||||
for worker in &self.workers {
|
||||
failed |= worker
|
||||
.input
|
||||
.as_ref()
|
||||
.expect("worker input open")
|
||||
.send(Arc::clone(&input))
|
||||
.is_err();
|
||||
}
|
||||
}
|
||||
let mut best = self.matrix.best(
|
||||
&input,
|
||||
0,
|
||||
if parallel {
|
||||
self.chunk
|
||||
} else {
|
||||
self.matrix.rows
|
||||
},
|
||||
);
|
||||
if parallel {
|
||||
// Join results in slot order: equal scores keep the first token.
|
||||
// Drain every dispatched job even when another worker failed.
|
||||
for worker in &self.workers {
|
||||
match worker.result.recv() {
|
||||
Ok(candidate) if candidate.1 > best.1 => best = candidate,
|
||||
Ok(_) => {}
|
||||
Err(_) => failed = true,
|
||||
}
|
||||
}
|
||||
}
|
||||
*logits = Arc::try_unwrap(input)
|
||||
.map_err(|_| "DSpark worker retained input")?
|
||||
.logits;
|
||||
if failed {
|
||||
return Err("DSpark Markov worker failed".into());
|
||||
}
|
||||
Ok(best.0 as i32)
|
||||
}
|
||||
}
|
||||
|
||||
#[cfg(test)]
|
||||
mod tests {
|
||||
use super::*;
|
||||
|
||||
#[test]
|
||||
#[ignore = "requires installed 0731 DSpark support; CPU dispatch diagnostic, not parity acceptance"]
|
||||
fn installed_markov_worker_dispatch() {
|
||||
use std::time::Instant;
|
||||
let artifacts = crate::model::engine_artifacts(
|
||||
crate::model::ModelChoice::DeepSeekV4Flash0731,
|
||||
true,
|
||||
&crate::app::models_path(),
|
||||
);
|
||||
let model = Gguf::open(artifacts.support.as_ref().unwrap()).unwrap();
|
||||
let weight = Weight::bind(&model, "mtp.2.markov_head.markov_w2.weight").unwrap();
|
||||
let state_weight = Weight::bind(&model, "mtp.2.markov_head.markov_w1.weight").unwrap();
|
||||
let values = super::super::dense_row(&model, state_weight, 671).unwrap();
|
||||
let mut pool12 = MarkovPool::from_map(model.shared_map(), weight, 12).unwrap();
|
||||
let mut pool18 = MarkovPool::from_map(model.shared_map(), weight, 18).unwrap();
|
||||
let mut logits = vec![0.0; weight.dims[1] as usize];
|
||||
let expected = pool12.argmax(&values, &mut logits).unwrap();
|
||||
for round in 0..4 {
|
||||
let order = if round % 2 == 0 { [0, 1, 2] } else { [2, 1, 0] };
|
||||
for mode in order {
|
||||
let started = Instant::now();
|
||||
for _ in 0..128 {
|
||||
let token = if mode == 0 {
|
||||
// Same row arithmetic/input ownership, old per-call
|
||||
// dispatch pattern: isolates worker lifetime/count.
|
||||
let matrix = &pool18.matrix;
|
||||
let input = Input {
|
||||
values: values.clone(),
|
||||
quantized: Some(quantize_q8_activation(&values)),
|
||||
logits: std::mem::take(&mut logits),
|
||||
};
|
||||
let chunk = matrix.rows.div_ceil(18);
|
||||
let best = thread::scope(|scope| {
|
||||
let handles = (0..matrix.rows)
|
||||
.step_by(chunk)
|
||||
.map(|start| {
|
||||
let input = &input;
|
||||
scope.spawn(move || {
|
||||
matrix.best(input, start, (start + chunk).min(matrix.rows))
|
||||
})
|
||||
})
|
||||
.collect::<Vec<_>>();
|
||||
handles.into_iter().map(|h| h.join().unwrap()).fold(
|
||||
(0, -f32::MAX),
|
||||
|best, candidate| {
|
||||
if candidate.1 > best.1 {
|
||||
candidate
|
||||
} else {
|
||||
best
|
||||
}
|
||||
},
|
||||
)
|
||||
});
|
||||
logits = input.logits;
|
||||
best.0 as i32
|
||||
} else if mode == 1 {
|
||||
pool12.argmax(&values, &mut logits).unwrap()
|
||||
} else {
|
||||
pool18.argmax(&values, &mut logits).unwrap()
|
||||
};
|
||||
assert_eq!(token, expected);
|
||||
}
|
||||
println!(
|
||||
"markov_dispatch round={round} mode={} calls=128 elapsed_ms={:.3}",
|
||||
["scoped18", "persistent12", "persistent18"][mode],
|
||||
started.elapsed().as_secs_f64() * 1000.0
|
||||
);
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn markov_workers_reuse_threads_and_buffers_with_ordered_ties() {
|
||||
assert_eq!(worker_count(18, None), 12);
|
||||
assert_eq!(worker_count(4, Some("0")), 4);
|
||||
assert_eq!(worker_count(18, Some("99")), 32);
|
||||
assert_eq!(worker_count(18, Some("1")), 1);
|
||||
let rows = 1025;
|
||||
let mut map = memmap2::MmapMut::map_anon(rows * 4).unwrap();
|
||||
for (row, bytes) in map.chunks_exact_mut(4).enumerate() {
|
||||
bytes.copy_from_slice(&((row % 17) as f32).to_le_bytes());
|
||||
}
|
||||
let map = Arc::new(map.make_read_only().unwrap());
|
||||
let weight = Weight {
|
||||
offset: 0,
|
||||
kind: F32,
|
||||
bytes: (rows * 4) as u64,
|
||||
dims: [1, rows as u64, 1],
|
||||
};
|
||||
for threads in [1, 4, 12] {
|
||||
let mut pool = MarkovPool::from_map(Arc::clone(&map), weight, threads).unwrap();
|
||||
let ids = pool
|
||||
.workers
|
||||
.iter()
|
||||
.map(|w| w.thread.as_ref().unwrap().thread().id())
|
||||
.collect::<Vec<_>>();
|
||||
let mut logits = vec![0.0; rows];
|
||||
let pointer = logits.as_ptr();
|
||||
assert_eq!(pool.argmax(&[2.0], &mut logits).unwrap(), 16);
|
||||
logits[rows - 1] = 1000.0;
|
||||
assert_eq!(pool.argmax(&[2.0], &mut logits).unwrap(), (rows - 1) as i32);
|
||||
assert_eq!(logits.as_ptr(), pointer);
|
||||
assert!(pool.argmax(&[], &mut logits).is_err());
|
||||
assert_eq!(logits.len(), rows);
|
||||
assert_eq!(
|
||||
ids,
|
||||
pool.workers
|
||||
.iter()
|
||||
.map(|w| w.thread.as_ref().unwrap().thread().id())
|
||||
.collect::<Vec<_>>()
|
||||
);
|
||||
if let Some(worker) = pool.workers.first_mut() {
|
||||
let (closed, receiver) = mpsc::sync_channel(1);
|
||||
drop(receiver);
|
||||
worker.input.replace(closed);
|
||||
worker.thread.take().unwrap().join().unwrap();
|
||||
assert!(pool.argmax(&[2.0], &mut logits).is_err());
|
||||
assert_eq!(logits.as_ptr(), pointer);
|
||||
assert_eq!(logits[rows - 1], 1000.0);
|
||||
}
|
||||
}
|
||||
assert!(
|
||||
MarkovPool::from_map(
|
||||
map,
|
||||
Weight {
|
||||
offset: 1,
|
||||
..weight
|
||||
},
|
||||
1
|
||||
)
|
||||
.is_err()
|
||||
);
|
||||
}
|
||||
}
|
||||
@@ -357,8 +357,8 @@ fn qwen_request_progress_uses_ui_frontiers_and_decode_timer() {
|
||||
let timing = PromptTiming {
|
||||
evaluated_tokens: 20,
|
||||
eval_seconds: 0.125,
|
||||
mtp_history_seconds: 0.025,
|
||||
restore_seconds: 0.010,
|
||||
mtp_history_seconds: Some(0.025),
|
||||
restore_seconds: Some(0.010),
|
||||
};
|
||||
assert_eq!(metrics.snapshot().prompt_timing, None);
|
||||
metrics.set_prompt_timing(Some(timing));
|
||||
@@ -399,7 +399,7 @@ fn qwen_request_progress_uses_ui_frontiers_and_decode_timer() {
|
||||
let exact_hit = PromptTiming {
|
||||
evaluated_tokens: 0,
|
||||
eval_seconds: 0.0,
|
||||
mtp_history_seconds: 0.0,
|
||||
mtp_history_seconds: Some(0.0),
|
||||
..timing
|
||||
};
|
||||
metrics.set_prompt_timing(Some(exact_hit));
|
||||
|
||||
@@ -1163,8 +1163,8 @@ impl Execution {
|
||||
observe(TurnProgress::PromptReady(crate::metrics::PromptTiming {
|
||||
evaluated_tokens: prompt.len() - cached,
|
||||
eval_seconds: prepared.eval_seconds + history_seconds,
|
||||
mtp_history_seconds: history_seconds,
|
||||
restore_seconds,
|
||||
mtp_history_seconds: Some(history_seconds),
|
||||
restore_seconds: Some(restore_seconds),
|
||||
}))?;
|
||||
// The bank owns this boundary before decode starts, including when a
|
||||
// subsequent callback aborts the request. Keeping it in this stack
|
||||
|
||||
@@ -78,7 +78,8 @@ fn mtplx_canonical_source_bodies_preserve_pinned_hashes() {
|
||||
.split("// Runtime unit SHA256: ")
|
||||
.skip(1)
|
||||
.collect::<Vec<_>>();
|
||||
assert_eq!(runtime_units.len(), 22);
|
||||
// tools/mtplx-kernel-source.py --check verifies this complete pinned export.
|
||||
assert_eq!(runtime_units.len(), 26);
|
||||
for unit in runtime_units {
|
||||
let expected = unit.lines().next().unwrap();
|
||||
let body = unit
|
||||
|
||||
@@ -1092,6 +1092,108 @@ fn unicode_punctuation(cp: u32) -> bool {
|
||||
mod tests {
|
||||
use super::*;
|
||||
|
||||
#[test]
|
||||
#[ignore = "CPU-only; requires DS4SERVER_CHAT_REFERENCE and its installed GGUF"]
|
||||
fn ds4_chat_matches_original_session_tokens() {
|
||||
use super::super::{Model, ModelRef, conversation_tag, render_text_prompt};
|
||||
use crate::model::ModelChoice;
|
||||
let reference = std::env::var_os("DS4SERVER_CHAT_REFERENCE").unwrap();
|
||||
let fixture: serde_json::Value =
|
||||
serde_json::from_slice(&fs::read(reference).unwrap()).unwrap();
|
||||
let choice = match fixture["family"].as_str().unwrap() {
|
||||
"deepseek" => ModelChoice::DeepSeekV4Flash0731,
|
||||
"glm" => ModelChoice::Glm53Flash,
|
||||
other => panic!("unsupported reference family: {other}"),
|
||||
};
|
||||
let model =
|
||||
Model::open_main(Path::new(fixture["model"].as_str().unwrap()), choice).unwrap();
|
||||
let tokenizer = &model.tokenizer;
|
||||
let cases = fixture["cases"].as_array().unwrap();
|
||||
assert_eq!(cases.len(), 3);
|
||||
let mut previous = Vec::new();
|
||||
for (index, case) in cases.iter().enumerate() {
|
||||
let prompt = case["prompt"].as_str().unwrap();
|
||||
let expected: Vec<i32> = serde_json::from_value(case["tokens"].clone()).unwrap();
|
||||
let actual = if index == 0 {
|
||||
tokenizer.encode_chat("", prompt, ReasoningMode::Low)
|
||||
} else {
|
||||
let mut tokens = previous;
|
||||
tokens.extend(tokenizer.encode_continuation(prompt, ReasoningMode::Low, false));
|
||||
tokens
|
||||
};
|
||||
if actual != expected {
|
||||
let first = actual
|
||||
.iter()
|
||||
.zip(&expected)
|
||||
.position(|(a, b)| a != b)
|
||||
.unwrap_or(actual.len().min(expected.len()));
|
||||
let context = |tokens: &[i32]| {
|
||||
tokens[first.saturating_sub(4)..tokens.len().min(first + 8)]
|
||||
.iter()
|
||||
.map(|&id| {
|
||||
(
|
||||
id,
|
||||
String::from_utf8_lossy(&tokenizer.token_bytes(id).unwrap())
|
||||
.into_owned(),
|
||||
)
|
||||
})
|
||||
.collect::<Vec<_>>()
|
||||
};
|
||||
panic!(
|
||||
"turn {} first mismatch {first}, lengths {}/{}; Rust {:?}; DS4 {:?}",
|
||||
index + 1,
|
||||
actual.len(),
|
||||
expected.len(),
|
||||
context(&actual),
|
||||
context(&expected)
|
||||
);
|
||||
}
|
||||
if index == 0 {
|
||||
let settings = crate::settings::TurnSettings {
|
||||
kv_cache: crate::settings::KvCachePreferences::default().settings(),
|
||||
context_tokens: 32768,
|
||||
max_generated_tokens: i32::MAX,
|
||||
system_prompt: String::new(),
|
||||
temperature: 0.6,
|
||||
top_p: 0.95,
|
||||
min_p: 0.0,
|
||||
top_k: 0,
|
||||
stops: vec![],
|
||||
seed: Some(42),
|
||||
reasoning_mode: ReasoningMode::Low,
|
||||
};
|
||||
let messages = [ChatTurn {
|
||||
user: true,
|
||||
tool: false,
|
||||
system: false,
|
||||
skip_previous_eos: false,
|
||||
reasoning: None,
|
||||
reasoning_complete: true,
|
||||
content: prompt.to_owned(),
|
||||
}];
|
||||
let bootstrap = model.render_history("", &[], settings.reasoning_mode);
|
||||
let rendered = render_text_prompt(
|
||||
ModelRef::Gguf(&model),
|
||||
&bootstrap,
|
||||
conversation_tag("", settings.reasoning_mode, &[]),
|
||||
&messages,
|
||||
&settings,
|
||||
);
|
||||
assert!(
|
||||
rendered == expected,
|
||||
"shared bootstrap renderer differs: {}/{} tokens; starts {:?}/{:?}",
|
||||
rendered.len(),
|
||||
expected.len(),
|
||||
&rendered[..4],
|
||||
&expected[..4]
|
||||
);
|
||||
}
|
||||
previous = expected;
|
||||
let reply: Vec<i32> = serde_json::from_value(case["reply_tokens"].clone()).unwrap();
|
||||
previous.extend(reply);
|
||||
}
|
||||
}
|
||||
|
||||
#[test]
|
||||
#[ignore = "requires local tokenizer and tools/qwen-chat-reference.py output; no inference"]
|
||||
fn qwen_plain_chat_matches_mtplx_server_tokens() {
|
||||
|
||||
Reference in New Issue
Block a user