Support DeepSeek V4 Flash 0731
This commit is contained in:
185
src/engine.rs
185
src/engine.rs
@@ -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 {
|
||||
|
||||
Reference in New Issue
Block a user