Support DeepSeek V4 Flash 0731

This commit is contained in:
Georg Bauer
2026-08-29 20:28:50 +02:00
parent ad855b321e
commit f1c177b754
23 changed files with 10510 additions and 2324 deletions

View File

@@ -8,7 +8,7 @@ mod validation;
#[cfg(target_os = "macos")]
use crate::metrics::{KvLookup, Metrics, SsdStats};
use crate::model::ModelChoice;
use crate::model::{ModelChoice, validate_engine_artifacts};
#[cfg(target_os = "macos")]
use crate::settings::TurnSettings;
use crate::settings::{EngineSettings, ReasoningMode};
@@ -232,6 +232,12 @@ pub(crate) struct ModelSummary {
impl Model {
#[allow(dead_code)]
pub(crate) fn open(settings: &EngineSettings) -> Result<Self, String> {
validate_engine_artifacts(
settings.model,
settings.artifacts.mtp.is_some() && !settings.speculative.dspark,
settings.speculative.dspark,
&settings.artifacts,
)?;
let mut model = Self::open_main(&settings.artifacts.model, settings.model)?;
if settings.execution.warm_weights {
model.main.warm()?;
@@ -1225,8 +1231,17 @@ impl Generator {
cancelled,
)?
} else {
self.executor.eval(token)?;
vec![token]
self.executor.eval_speculative_sampled(
token,
generation_limit - generated_tokens,
settings.reasoning_mode,
settings.temperature,
settings.top_p,
settings.min_p,
settings.top_k,
&mut rng,
cancelled,
)?
};
self.publish_execution_stats();
for token in cycle {
@@ -1545,12 +1560,30 @@ fn sample(
top_k: i32,
rng: &mut Rng,
) -> i32 {
let probabilities = sampling_probabilities(logits, temperature, top_p, min_p, top_k);
sample_probabilities(&probabilities, rng, None)
}
#[cfg(any(target_os = "macos", test))]
fn sampling_probabilities(
logits: &[f32],
temperature: f32,
top_p: f32,
min_p: f32,
top_k: i32,
) -> Vec<(usize, f32)> {
let greedy = || {
vec![(
logits
.iter()
.enumerate()
.max_by(|a, b| a.1.total_cmp(b.1))
.map_or(0, |(index, _)| index),
1.0,
)]
};
if temperature <= 0.0 {
return logits
.iter()
.enumerate()
.max_by(|a, b| a.1.total_cmp(b.1))
.map_or(0, |(index, _)| index as i32);
return greedy();
}
let maximum = logits
.iter()
@@ -1558,7 +1591,7 @@ fn sample(
.filter(|value| value.is_finite())
.fold(f32::NEG_INFINITY, f32::max);
if !maximum.is_finite() {
return 0;
return greedy();
}
let top_p = if top_p <= 0.0 || top_p > 1.0 {
1.0
@@ -1571,45 +1604,96 @@ fn sample(
.enumerate()
.filter(|(_, logit)| logit.is_finite())
.map(|(index, logit)| (index, ((*logit - maximum) / temperature).exp()))
.filter(|(_, probability)| *probability >= min_p)
.collect();
if probabilities.is_empty() {
return logits
.iter()
.enumerate()
.max_by(|a, b| a.1.total_cmp(b.1))
.map_or(0, |(index, _)| index as i32);
return greedy();
}
if top_p < 1.0 || top_k > 0 {
if top_p < 1.0 || top_k > 0 || min_p > 0.0 {
probabilities.sort_unstable_by(|a, b| b.1.total_cmp(&a.1).then_with(|| a.0.cmp(&b.0)));
if top_k > 0 {
probabilities.truncate(probabilities.len().min(top_k as usize));
probabilities.truncate(probabilities.len().min((top_k as usize).min(1024)));
}
}
if top_p < 1.0 {
let total: f32 = probabilities
.iter()
.map(|(_, probability)| probability)
.sum();
let mut kept = 0.0;
let count = probabilities
.iter()
.position(|(_, probability)| {
kept += *probability;
kept / total >= top_p
})
.map_or(probabilities.len(), |index| index + 1);
probabilities.truncate(count);
let total: f32 = probabilities
.iter()
.map(|(_, probability)| probability)
.sum();
let mut kept = 0.0;
let mut count = 0;
for (_, probability) in &probabilities {
if count > 0 && *probability < min_p {
break;
}
kept += *probability;
count += 1;
if kept / total >= top_p {
break;
}
}
let kept_total: f32 = probabilities.iter().map(|(_, p)| p).sum();
let mut choice = rng.unit() * kept_total;
for (token, probability) in &probabilities {
probabilities.truncate(count);
if probabilities.is_empty() || !kept.is_finite() || kept <= 0.0 {
return greedy();
}
for (_, probability) in &mut probabilities {
*probability /= kept;
}
if top_p >= 1.0 && top_k <= 0 {
probabilities.sort_unstable_by_key(|(token, _)| *token);
}
probabilities
}
#[cfg(any(target_os = "macos", test))]
fn sample_probabilities(
probabilities: &[(usize, f32)],
rng: &mut Rng,
excluded: Option<usize>,
) -> i32 {
let total: f32 = probabilities
.iter()
.filter(|(token, _)| Some(*token) != excluded)
.map(|(_, probability)| probability)
.sum();
let mut choice = rng.unit() * total;
for (token, probability) in probabilities {
if Some(*token) == excluded {
continue;
}
choice -= probability;
if choice <= 0.0 {
return *token as i32;
}
}
probabilities.last().map_or(0, |(token, _)| *token as i32)
probabilities
.iter()
.rev()
.find(|(token, _)| Some(*token) != excluded)
.map_or(0, |(token, _)| *token as i32)
}
#[cfg(any(target_os = "macos", test))]
fn exact_delta_sample(
logits: &[f32],
draft: i32,
temperature: f32,
top_p: f32,
min_p: f32,
top_k: i32,
rng: &mut Rng,
) -> (i32, bool) {
let mut probabilities = sampling_probabilities(logits, temperature, top_p, min_p, top_k);
let draft_probability = probabilities
.iter()
.find(|(token, _)| *token == draft as usize)
.map_or(0.0, |(_, probability)| *probability);
if rng.unit() <= draft_probability {
return (draft, true);
}
probabilities.sort_unstable_by_key(|(token, _)| *token);
(
sample_probabilities(&probabilities, rng, Some(draft as usize)),
false,
)
}
#[cfg(any(target_os = "macos", test))]
@@ -1665,6 +1749,37 @@ mod sampling_tests {
assert_eq!(chunks, ["hello "]);
}
#[test]
fn exact_delta_sampling_accepts_or_corrects_the_draft() {
let mut accept_rng = Rng::new(2);
let (accepted, was_draft) =
exact_delta_sample(&[10.0, 0.0], 0, 1.0, 1.0, 0.0, 0, &mut accept_rng);
assert_eq!((accepted, was_draft), (0, true));
let mut reject_rng = Rng::new(1);
let (replacement, was_draft) =
exact_delta_sample(&[0.0, 10.0], 0, 1.0, 1.0, 0.0, 0, &mut reject_rng);
assert_eq!((replacement, was_draft), (1, false));
}
#[test]
fn sampling_probabilities_match_ds4_filter_order() {
let probabilities = sampling_probabilities(&[0.0, 2.0, 1.0], 1.0, 0.8, 0.2, 0);
assert_eq!(probabilities.len(), 2);
assert_eq!(probabilities[0].0, 1);
assert_eq!(probabilities[1].0, 2);
assert!((probabilities.iter().map(|(_, value)| value).sum::<f32>() - 1.0).abs() < 1e-6);
let min_p_only = sampling_probabilities(&[0.0, 2.0, 1.0], 1.0, 1.0, 0.2, 0);
assert_eq!(
min_p_only
.iter()
.map(|(token, _)| *token)
.collect::<Vec<_>>(),
[1, 2]
);
}
#[test]
fn split_utf8_token_bytes_are_joined_before_decoding() {
let mut generated = ChatTurn {