Files
DS4Server/src/engine.rs
2026-07-25 11:22:59 +02:00

1093 lines
34 KiB
Rust

mod gguf;
#[cfg(target_os = "macos")]
mod metal;
mod tokenizer;
mod validation;
#[cfg(target_os = "macos")]
use crate::metrics::{KvLookup, Metrics};
use crate::model::ModelChoice;
use crate::settings::TurnSettings;
use crate::settings::{EngineSettings, ReasoningMode};
use gguf::{F16, F32, Gguf, I32, IQ2_XXS, Q2_K, Q4_0, Q4_K, Q5_K, Q6_K, Q8_0, Tensor, Value};
use sha2::{Digest, Sha256};
use std::path::{Path, PathBuf};
#[cfg(target_os = "macos")]
use std::sync::Arc;
use std::sync::atomic::{AtomicBool, Ordering};
#[cfg(target_os = "macos")]
use std::time::Instant;
use tokenizer::Tokenizer;
use validation::{validate_dspark, validate_main};
pub(crate) use validation::validate_model_artifact;
#[cfg(target_os = "macos")]
pub(crate) use metal::configure_sources as configure_metal_sources;
const DENSE: &[u32] = &[Q8_0, Q4_K, Q4_0];
const ROUTED: &[u32] = &[Q8_0, IQ2_XXS, Q2_K, Q4_K, Q5_K, Q6_K];
const PLAIN: &[u32] = &[F16, F32];
const DSPARK_DENSE: &[u32] = &[F16, F32, Q8_0];
#[derive(Clone, Copy, Debug, Eq, PartialEq)]
enum ModelFamily {
DeepSeek,
Glm,
}
#[derive(Clone, Copy)]
struct Shape {
model: ModelChoice,
family: ModelFamily,
layers: u32,
embd: u64,
vocab: u64,
heads: u64,
head_kv: u64,
head_dim: u64,
value_dim: u64,
rot: u64,
out_groups: u64,
lora_q: u64,
lora_o: u64,
experts: u64,
experts_used: u64,
expert_shared: u64,
ff_expert: u64,
ff_dense: u64,
hash_layers: u32,
sliding_window: u64,
indexer_heads: u64,
indexer_head_dim: u64,
indexer_top_k: u64,
hc: u64,
hc_sinkhorn: u64,
nextn: u32,
leading_dense: u32,
kv_lora: u64,
key_mla: u64,
value_mla: u64,
rms_epsilon: f32,
hc_epsilon: f32,
expert_weight_scale: f32,
swiglu_clamp: f32,
rope_base: f32,
rope_scale: f32,
rope_beta_fast: f32,
rope_beta_slow: f32,
compress_rope_base: f32,
original_context: u64,
}
const FLASH: Shape = Shape {
model: ModelChoice::DeepSeekV4Flash,
family: ModelFamily::DeepSeek,
layers: 43,
embd: 4096,
vocab: 129_280,
heads: 64,
head_kv: 1,
head_dim: 512,
value_dim: 512,
rot: 64,
out_groups: 8,
lora_q: 1024,
lora_o: 1024,
experts: 256,
experts_used: 6,
expert_shared: 1,
ff_expert: 2048,
ff_dense: 0,
hash_layers: 3,
sliding_window: 128,
indexer_heads: 64,
indexer_head_dim: 128,
indexer_top_k: 512,
hc: 4,
hc_sinkhorn: 20,
nextn: 0,
leading_dense: 0,
kv_lora: 0,
key_mla: 0,
value_mla: 0,
rms_epsilon: 1.0e-6,
hc_epsilon: 1.0e-6,
expert_weight_scale: 1.5,
swiglu_clamp: 10.0,
rope_base: 10_000.0,
rope_scale: 16.0,
rope_beta_fast: 32.0,
rope_beta_slow: 1.0,
compress_rope_base: 160_000.0,
original_context: 65_536,
};
const PRO: Shape = Shape {
model: ModelChoice::DeepSeekV4Pro,
layers: 61,
embd: 7168,
heads: 128,
out_groups: 16,
lora_q: 1536,
experts: 384,
ff_expert: 3072,
indexer_top_k: 1024,
expert_weight_scale: 2.5,
..FLASH
};
const GLM: Shape = Shape {
model: ModelChoice::Glm52,
family: ModelFamily::Glm,
layers: 79,
embd: 6144,
vocab: 154_880,
heads: 64,
head_kv: 1,
head_dim: 576,
value_dim: 512,
rot: 64,
out_groups: 0,
lora_q: 2048,
lora_o: 0,
experts: 256,
experts_used: 8,
expert_shared: 1,
ff_expert: 2048,
ff_dense: 12_288,
hash_layers: 0,
sliding_window: 0,
indexer_heads: 32,
indexer_head_dim: 128,
indexer_top_k: 2048,
hc: 0,
hc_sinkhorn: 0,
nextn: 1,
leading_dense: 3,
kv_lora: 512,
key_mla: 256,
value_mla: 256,
rms_epsilon: 1.0e-5,
hc_epsilon: 0.0,
expert_weight_scale: 2.5,
swiglu_clamp: 0.0,
rope_base: 8_000_000.0,
rope_scale: 1.0,
rope_beta_fast: 0.0,
rope_beta_slow: 0.0,
compress_rope_base: 0.0,
original_context: 1_048_576,
};
pub(crate) struct Model {
main: Gguf,
support: Option<Gguf>,
shape: Shape,
tokenizer: Tokenizer,
}
#[derive(Clone, Debug, Eq, PartialEq)]
pub(crate) struct ModelSummary {
pub(crate) model: ModelChoice,
pub(crate) mapped_bytes: u64,
pub(crate) tensor_count: usize,
pub(crate) vocabulary_size: usize,
pub(crate) support_loaded: bool,
}
impl Model {
#[allow(dead_code)]
pub(crate) fn open(settings: &EngineSettings) -> Result<Self, String> {
let mut model = Self::open_main(&settings.artifacts.model, settings.model)?;
if let Some(path) = &settings.artifacts.mtp {
let support = Gguf::open(path)?;
validate_dspark(&support, &model.shape)?;
model.support = Some(support);
}
Ok(model)
}
fn open_main(path: &Path, expected: ModelChoice) -> Result<Self, String> {
let main = Gguf::open(path)?;
let shape = validate_main(&main, expected)?;
let tokenizer = Tokenizer::load(&main, shape.family)?;
if tokenizer.vocab_size() != shape.vocab as usize {
return Err(format!(
"tokenizer has {} entries, expected {}",
tokenizer.vocab_size(),
shape.vocab
));
}
Ok(Self {
main,
support: None,
shape,
tokenizer,
})
}
pub(crate) fn summary(&self) -> ModelSummary {
ModelSummary {
model: self.shape.model,
mapped_bytes: self.main.len() + self.support.as_ref().map_or(0, Gguf::len),
tensor_count: self.main.tensors.len()
+ self
.support
.as_ref()
.map_or(0, |support| support.tensors.len()),
vocabulary_size: self.tokenizer.vocab_size(),
support_loaded: self.support.is_some(),
}
}
pub(crate) fn tokenize(&self, text: &str) -> Vec<i32> {
self.tokenizer.tokenize(text)
}
pub(crate) fn render_prompt(
&self,
system: &str,
prompt: &str,
reasoning: ReasoningMode,
) -> Vec<i32> {
self.tokenizer.encode_chat(system, prompt, reasoning)
}
fn render_conversation(
&self,
system: &str,
messages: &[ChatTurn],
reasoning: ReasoningMode,
) -> Vec<i32> {
self.tokenizer
.encode_conversation(system, messages, reasoning)
}
fn render_continuation(
&self,
prompt: &str,
reasoning: ReasoningMode,
skip_previous_eos: bool,
) -> Vec<i32> {
self.tokenizer
.encode_continuation(prompt, reasoning, skip_previous_eos)
}
pub(crate) fn token_bytes(&self, token: i32) -> Option<Vec<u8>> {
self.tokenizer.token_bytes(token)
}
pub(crate) fn eos_token(&self) -> i32 {
self.tokenizer.eos()
}
pub(crate) fn is_stop_token(&self, token: i32) -> bool {
self.tokenizer.is_stop(token)
}
pub(crate) fn is_think_start_token(&self, token: i32) -> bool {
self.tokenizer.is_think_start(token)
}
pub(crate) fn is_think_end_token(&self, token: i32) -> bool {
self.tokenizer.is_think_end(token)
}
pub(crate) fn is_stop_token_for_reasoning(&self, token: i32, reasoning: ReasoningMode) -> bool {
self.is_stop_token(token)
|| (reasoning == ReasoningMode::Direct
&& (self.is_think_start_token(token) || self.is_think_end_token(token)))
}
pub(crate) fn tensor_data(&self, name: &str) -> Result<&[u8], String> {
self.main.tensor_data(name)
}
}
#[cfg(target_os = "macos")]
pub(crate) struct Generator {
executor: metal::Executor,
checkpoint: Option<PathBuf>,
metrics: Arc<Metrics>,
}
#[derive(Clone)]
pub(crate) struct ChatTurn {
pub(crate) user: bool,
pub(crate) skip_previous_eos: bool,
pub(crate) reasoning: Option<String>,
pub(crate) reasoning_complete: bool,
pub(crate) content: String,
}
pub(crate) struct GenerationOutput {
pub(crate) message: ChatTurn,
pub(crate) prompt_tokens: u32,
pub(crate) cached_tokens: u32,
pub(crate) completion_tokens: u32,
pub(crate) finish_reason: &'static str,
pub(crate) previous_checkpoint_bytes: Option<u64>,
pub(crate) checkpoint_bytes: u64,
}
#[cfg(target_os = "macos")]
impl Generator {
pub(crate) fn open(settings: &EngineSettings, metrics: Arc<Metrics>) -> Result<Self, String> {
if settings.model != ModelChoice::DeepSeekV4Flash {
return Err("local generation currently supports DeepSeek V4 Flash only".into());
}
if settings.speculative.dspark || settings.ssd.enabled || settings.steering.file.is_some() {
return Err(
"DSpark, SSD streaming, and steering are not yet available in the Rust executor"
.into(),
);
}
let model = Model::open(settings)?;
let executor = metal::Executor::open(
model,
settings.context_tokens.max(1) as u32,
settings.execution.quality,
settings.execution.prefill_chunk,
)?;
Ok(Self {
executor,
checkpoint: None,
metrics,
})
}
pub(crate) fn summary(&self) -> ModelSummary {
self.executor.model().summary()
}
pub(crate) fn generate(
&mut self,
checkpoint: &Path,
messages: &[ChatTurn],
settings: &TurnSettings,
cancelled: &AtomicBool,
mut emit: impl FnMut(bool, String),
mut progress: impl FnMut(u32, u32, Option<f32>),
) -> Result<GenerationOutput, String> {
let history = messages
.split_last()
.map_or(messages, |(_, history)| history);
self.select_checkpoint(
checkpoint,
conversation_tag(&settings.system_prompt, settings.reasoning_mode, history),
)?;
let (mut output, prompt_complete) =
self.generate_inner(messages, settings, cancelled, &mut emit, &mut progress)?;
let mut completed = messages.to_vec();
completed.push(output.message.clone());
output.previous_checkpoint_bytes =
std::fs::metadata(checkpoint).ok().map(|item| item.len());
self.save_checkpoint(
checkpoint,
if prompt_complete {
conversation_tag(&settings.system_prompt, settings.reasoning_mode, &completed)
} else {
[0; 32]
},
)?;
output.checkpoint_bytes = std::fs::metadata(checkpoint)
.map(|item| item.len())
.unwrap_or(0);
Ok(output)
}
pub(crate) fn generate_transient(
&mut self,
directory: &Path,
messages: &[ChatTurn],
settings: &TurnSettings,
cancelled: &AtomicBool,
mut emit: impl FnMut(bool, String),
mut progress: impl FnMut(u32, u32, Option<f32>),
) -> Result<GenerationOutput, String> {
let history = messages
.split_last()
.map_or(messages, |(_, history)| history);
let history_tag =
conversation_tag(&settings.system_prompt, settings.reasoning_mode, history);
let checkpoint = directory.join(format!("{}.bin", hex_tag(history_tag)));
if self.executor.checkpoint_tag() != history_tag {
self.select_checkpoint(&checkpoint, history_tag)?;
if self.executor.checkpoint_tag() != history_tag {
self.executor.reset()?;
self.checkpoint = None;
let _ = std::fs::remove_file(&checkpoint);
}
} else {
self.metrics.kv_lookup(KvLookup::MemoryHit);
}
let (mut output, prompt_complete) =
self.generate_inner(messages, settings, cancelled, &mut emit, &mut progress)?;
let mut completed = messages.to_vec();
completed.push(output.message.clone());
let completed_tag =
conversation_tag(&settings.system_prompt, settings.reasoning_mode, &completed);
std::fs::create_dir_all(directory).map_err(|error| {
format!(
"Could not create transient KV cache directory {}: {error}",
directory.display()
)
})?;
let completed_checkpoint = directory.join(format!("{}.bin", hex_tag(completed_tag)));
output.previous_checkpoint_bytes = std::fs::metadata(&completed_checkpoint)
.ok()
.map(|item| item.len());
self.save_checkpoint(
&completed_checkpoint,
if prompt_complete {
completed_tag
} else {
[0; 32]
},
)?;
output.checkpoint_bytes = std::fs::metadata(&completed_checkpoint)
.map(|item| item.len())
.unwrap_or(0);
self.checkpoint = Some(completed_checkpoint);
Ok(output)
}
fn select_checkpoint(
&mut self,
checkpoint: &Path,
expected_tag: [u8; 32],
) -> Result<(), String> {
if self.checkpoint.as_deref() == Some(checkpoint) {
self.metrics
.kv_lookup(if self.executor.checkpoint_tag() == expected_tag {
KvLookup::MemoryHit
} else {
KvLookup::Miss
});
return Ok(());
}
self.executor.reset()?;
self.metrics.kv_read_started();
let started = Instant::now();
let loaded = self
.executor
.load_checkpoint(checkpoint, &mut |bytes| self.metrics.kv_read_bytes(bytes));
self.metrics
.kv_read_finished(started.elapsed(), loaded.is_err());
let lookup = match loaded {
Ok(true) if self.executor.checkpoint_tag() == expected_tag => KvLookup::DiskHit,
Ok(true) | Ok(false) => KvLookup::Miss,
Err(_) => {
self.executor.reset()?;
let _ = std::fs::remove_file(checkpoint);
KvLookup::Invalid
}
};
self.metrics.kv_lookup(lookup);
self.checkpoint = Some(checkpoint.to_owned());
Ok(())
}
fn save_checkpoint(&mut self, checkpoint: &Path, tag: [u8; 32]) -> Result<(), String> {
self.metrics.kv_write_started();
let started = Instant::now();
let result = self
.executor
.save_checkpoint(checkpoint, tag, &mut |bytes| {
self.metrics.kv_write_bytes(bytes);
});
self.metrics
.kv_write_finished(started.elapsed(), result.is_err());
result
}
fn generate_inner(
&mut self,
messages: &[ChatTurn],
settings: &TurnSettings,
cancelled: &AtomicBool,
emit: &mut impl FnMut(bool, String),
progress: &mut impl FnMut(u32, u32, Option<f32>),
) -> Result<(GenerationOutput, bool), String> {
let tokens = match messages.split_last() {
Some((latest, history))
if latest.user
&& self.executor.checkpoint_tag()
== conversation_tag(
&settings.system_prompt,
settings.reasoning_mode,
history,
) =>
{
let mut tokens = self.executor.tokens().to_vec();
tokens.extend(self.executor.model().render_continuation(
&latest.content,
settings.reasoning_mode,
latest.skip_previous_eos,
));
tokens
}
_ => self.executor.model().render_conversation(
&settings.system_prompt,
messages,
settings.reasoning_mode,
),
};
if tokens.is_empty() {
return Err("the rendered prompt is empty".into());
}
let max_context = self.executor.context() as usize;
if tokens.len() >= max_context {
return Err(format!(
"the conversation uses {} tokens; the configured context holds fewer than {max_context}",
tokens.len()
));
}
let reused = self.executor.align_prompt(&tokens)?;
self.metrics.kv_prefix_reused(reused);
progress(self.executor.position(), self.executor.context(), None);
let mut rng = Rng::new(settings.seed.unwrap_or(0x4453_3453_4552_5645));
let mut reasoning = settings.reasoning_mode != ReasoningMode::Direct;
let mut generated = ChatTurn {
user: false,
skip_previous_eos: false,
reasoning: reasoning.then(String::new),
reasoning_complete: !reasoning,
content: String::new(),
};
let mut emitted_reasoning = 0;
let mut emitted_content = 0;
let prompt_tokens = tokens.len();
let suffix = &tokens[reused..];
let completed = if (reused == 0 && tokens.len() > 1) || suffix.len() >= 4 {
let context = self.executor.context();
self.executor.prefill(suffix, |used| {
progress(used, context, None);
!cancelled.load(Ordering::Relaxed)
})?
} else {
let mut completed = 0;
for &token in suffix {
if cancelled.load(Ordering::Relaxed) {
break;
}
self.executor.eval(token)?;
completed += 1;
progress(self.executor.position(), self.executor.context(), None);
}
completed
};
if completed != suffix.len() {
return Ok((
GenerationOutput {
message: generated,
prompt_tokens: prompt_tokens as u32,
cached_tokens: reused as u32,
completion_tokens: 0,
finish_reason: "stop",
previous_checkpoint_bytes: None,
checkpoint_bytes: 0,
},
false,
));
}
progress(self.executor.position(), self.executor.context(), Some(0.0));
let generation_started = Instant::now();
let mut generated_tokens = 0_u32;
for _ in 0..settings
.max_generated_tokens
.max(0)
.min((max_context - self.executor.position() as usize) as i32)
{
if cancelled.load(Ordering::Relaxed) {
flush_generated(
&mut generated,
&mut emitted_reasoning,
&mut emitted_content,
&settings.stops,
emit,
);
return Ok((
GenerationOutput {
message: generated,
prompt_tokens: prompt_tokens as u32,
cached_tokens: reused as u32,
completion_tokens: generated_tokens,
finish_reason: "stop",
previous_checkpoint_bytes: None,
checkpoint_bytes: 0,
},
true,
));
}
let token = sample(
self.executor.logits(),
settings.temperature,
settings.top_p,
settings.min_p,
settings.top_k,
&mut rng,
);
if self
.executor
.model()
.is_stop_token_for_reasoning(token, settings.reasoning_mode)
{
flush_generated(
&mut generated,
&mut emitted_reasoning,
&mut emitted_content,
&settings.stops,
emit,
);
return Ok((
GenerationOutput {
message: generated,
prompt_tokens: prompt_tokens as u32,
cached_tokens: reused as u32,
completion_tokens: generated_tokens,
finish_reason: "stop",
previous_checkpoint_bytes: None,
checkpoint_bytes: 0,
},
true,
));
}
if self.executor.model().is_think_start_token(token) {
reasoning = true;
generated.reasoning.get_or_insert_default();
} else if self.executor.model().is_think_end_token(token) {
if let Some(reasoning_text) = &mut generated.reasoning
&& emit_safe_text(
reasoning_text,
&mut emitted_reasoning,
&settings.stops,
true,
true,
emit,
)
{
return Ok((
GenerationOutput {
message: generated,
prompt_tokens: prompt_tokens as u32,
cached_tokens: reused as u32,
completion_tokens: generated_tokens + 1,
finish_reason: "stop",
previous_checkpoint_bytes: None,
checkpoint_bytes: 0,
},
false,
));
}
reasoning = false;
generated.reasoning_complete = true;
emit(false, String::new());
} else if let Some(bytes) = self.executor.model().token_bytes(token) {
let content = String::from_utf8_lossy(&bytes).into_owned();
let stopped = if reasoning {
let text = generated.reasoning.get_or_insert_default();
text.push_str(&content);
emit_safe_text(
text,
&mut emitted_reasoning,
&settings.stops,
false,
true,
emit,
)
} else {
generated.reasoning_complete = true;
generated.content.push_str(&content);
emit_safe_text(
&mut generated.content,
&mut emitted_content,
&settings.stops,
false,
false,
emit,
)
};
if stopped {
return Ok((
GenerationOutput {
message: generated,
prompt_tokens: prompt_tokens as u32,
cached_tokens: reused as u32,
completion_tokens: generated_tokens + 1,
finish_reason: "stop",
previous_checkpoint_bytes: None,
checkpoint_bytes: 0,
},
false,
));
}
}
self.executor.eval(token)?;
generated_tokens += 1;
progress(
self.executor.position(),
self.executor.context(),
Some(
generated_tokens as f32
/ generation_started.elapsed().as_secs_f32().max(1.0e-6),
),
);
}
flush_generated(
&mut generated,
&mut emitted_reasoning,
&mut emitted_content,
&settings.stops,
emit,
);
Ok((
GenerationOutput {
message: generated,
prompt_tokens: prompt_tokens as u32,
cached_tokens: reused as u32,
completion_tokens: generated_tokens,
finish_reason: "length",
previous_checkpoint_bytes: None,
checkpoint_bytes: 0,
},
true,
))
}
}
#[cfg(target_os = "macos")]
fn flush_generated(
generated: &mut ChatTurn,
emitted_reasoning: &mut usize,
emitted_content: &mut usize,
stops: &[String],
emit: &mut impl FnMut(bool, String),
) {
if let Some(reasoning) = &mut generated.reasoning {
let _ = emit_safe_text(reasoning, emitted_reasoning, stops, true, true, emit);
}
let _ = emit_safe_text(
&mut generated.content,
emitted_content,
stops,
true,
false,
emit,
);
}
#[cfg(target_os = "macos")]
fn emit_safe_text(
text: &mut String,
emitted: &mut usize,
stops: &[String],
final_flush: bool,
reasoning: bool,
emit: &mut impl FnMut(bool, String),
) -> bool {
let stop = stops
.iter()
.filter_map(|stop| {
text[*emitted..]
.find(stop)
.map(|position| *emitted + position)
})
.min();
if let Some(stop) = stop {
if stop > *emitted {
emit(reasoning, text[*emitted..stop].to_owned());
}
text.truncate(stop);
*emitted = stop;
return true;
}
let hold = if final_flush {
0
} else {
stops
.iter()
.map(|stop| stop.len().saturating_sub(1))
.max()
.unwrap_or(0)
};
let mut safe = text.len().saturating_sub(hold);
while safe > *emitted && !text.is_char_boundary(safe) {
safe -= 1;
}
if safe > *emitted {
emit(reasoning, text[*emitted..safe].to_owned());
*emitted = safe;
}
false
}
#[cfg(target_os = "macos")]
fn hex_tag(tag: [u8; 32]) -> String {
const HEX: &[u8; 16] = b"0123456789abcdef";
let mut output = String::with_capacity(64);
for byte in tag {
output.push(HEX[(byte >> 4) as usize] as char);
output.push(HEX[(byte & 0x0f) as usize] as char);
}
output
}
#[cfg(target_os = "macos")]
fn conversation_tag(system: &str, reasoning: ReasoningMode, messages: &[ChatTurn]) -> [u8; 32] {
fn text(hasher: &mut Sha256, value: &str) {
hasher.update((value.len() as u64).to_le_bytes());
hasher.update(value.as_bytes());
}
let mut hasher = Sha256::new();
hasher.update(b"DS4Server chat checkpoint v1");
text(&mut hasher, system);
hasher.update([match reasoning {
ReasoningMode::Direct => 0,
ReasoningMode::High => 1,
ReasoningMode::Max => 2,
}]);
for message in messages {
hasher.update([u8::from(message.user)]);
match &message.reasoning {
Some(reasoning) => {
hasher.update([1]);
text(&mut hasher, reasoning);
}
None => hasher.update([0]),
}
hasher.update([u8::from(message.reasoning_complete)]);
text(&mut hasher, &message.content);
}
hasher.finalize().into()
}
#[cfg(target_os = "macos")]
fn sample(
logits: &[f32],
temperature: f32,
top_p: f32,
min_p: f32,
top_k: i32,
rng: &mut Rng,
) -> i32 {
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);
}
let maximum = logits
.iter()
.copied()
.filter(|value| value.is_finite())
.fold(f32::NEG_INFINITY, f32::max);
if !maximum.is_finite() {
return 0;
}
let top_p = if top_p <= 0.0 || top_p > 1.0 {
1.0
} else {
top_p
};
let min_p = min_p.max(0.0);
let mut probabilities: Vec<(usize, f32)> = logits
.iter()
.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);
}
if top_p < 1.0 || top_k > 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));
}
}
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 kept_total: f32 = probabilities.iter().map(|(_, p)| p).sum();
let mut choice = rng.unit() * kept_total;
for (token, probability) in &probabilities {
choice -= probability;
if choice <= 0.0 {
return *token as i32;
}
}
probabilities.last().map_or(0, |(token, _)| *token as i32)
}
#[cfg(target_os = "macos")]
struct Rng(u64);
#[cfg(target_os = "macos")]
impl Rng {
fn new(seed: u64) -> Self {
Self(seed.max(1))
}
fn unit(&mut self) -> f32 {
let mut value = self.0;
if value == 0 {
value = 0x9e37_79b9_7f4a_7c15;
}
value ^= value >> 12;
value ^= value << 25;
value ^= value >> 27;
self.0 = value;
let value = value.wrapping_mul(0x2545_f491_4f6c_dd1d);
((value >> 40) & 0xff_ffff) as f32 / 16_777_216.0
}
}
#[cfg(all(test, target_os = "macos"))]
mod sampling_tests {
use super::*;
#[test]
fn zero_temperature_is_greedy() {
let mut rng = Rng::new(1);
assert_eq!(sample(&[1.0, 4.0, 2.0], 0.0, 1.0, 0.0, 0, &mut rng), 1);
}
#[test]
fn top_k_and_stream_stops_are_applied_before_output() {
let mut rng = Rng::new(1);
assert_eq!(sample(&[1.0, 4.0, 2.0], 1.0, 1.0, 0.0, 1, &mut rng), 1);
let mut text = "hello STOP hidden".to_owned();
let mut emitted = 0;
let mut chunks = Vec::new();
assert!(emit_safe_text(
&mut text,
&mut emitted,
&["STOP".into()],
false,
false,
&mut |_, chunk| chunks.push(chunk),
));
assert_eq!(text, "hello ");
assert_eq!(chunks, ["hello "]);
}
#[test]
fn checkpoint_tag_covers_the_canonical_chat_state() {
let messages = [ChatTurn {
user: true,
skip_previous_eos: false,
reasoning: None,
reasoning_complete: true,
content: "Hello".into(),
}];
let tag = conversation_tag("System", ReasoningMode::High, &messages);
assert_eq!(
tag,
conversation_tag("System", ReasoningMode::High, &messages)
);
assert_ne!(
tag,
conversation_tag("Changed", ReasoningMode::High, &messages)
);
assert_ne!(
tag,
conversation_tag("System", ReasoningMode::Direct, &messages)
);
}
#[test]
#[ignore = "requires the 80 GiB Flash checkpoint and Apple Metal"]
fn metal_executes_real_flash_token() {
configure_metal_sources().unwrap();
let path = Path::new(env!("CARGO_MANIFEST_DIR")).join(
"../ds4/gguf/DeepSeek-V4-Flash-IQ2XXS-w2Q2K-AProjQ8-SExpQ8-OutQ8-chat-v2-imatrix.gguf",
);
let model = Model::open_main(&path, ModelChoice::DeepSeekV4Flash).unwrap();
let tokens = model.render_prompt(
"You are a helpful assistant",
"Hello",
ReasoningMode::Direct,
);
assert_eq!(tokens.len(), 10);
let mut executor = metal::Executor::open(model, 32_768, false, 0).unwrap();
assert_eq!(executor.context(), 32_768);
for &token in &tokens {
executor.eval(token).unwrap();
}
assert!(executor.logits().iter().all(|logit| logit.is_finite()));
let argmax = executor
.logits()
.iter()
.enumerate()
.max_by(|a, b| a.1.total_cmp(b.1))
.unwrap();
eprintln!(
"Rust logits: argmax={} value={} logit0={}",
argmax.0,
argmax.1,
executor.logits()[0]
);
assert_eq!(argmax.0, 19_923);
assert!((executor.logits()[0] - -7.675_424).abs() < 0.1);
let next = argmax.0 as i32;
let checkpoint =
std::env::temp_dir().join(format!("ds4-rust-kv-{}.bin", std::process::id()));
executor
.save_checkpoint(&checkpoint, [7; 32], &mut |_| {})
.unwrap();
executor.eval(next).unwrap();
let continued_logit = executor.logits()[0];
let continued_argmax = executor
.logits()
.iter()
.enumerate()
.max_by(|a, b| a.1.total_cmp(b.1))
.unwrap()
.0;
drop(executor);
let model = Model::open_main(&path, ModelChoice::DeepSeekV4Flash).unwrap();
let mut restored = metal::Executor::open(model, 32_768, false, 0).unwrap();
assert!(restored.load_checkpoint(&checkpoint, &mut |_| {}).unwrap());
assert_eq!(restored.position(), tokens.len() as u32);
assert_eq!(restored.tokens(), tokens);
assert_eq!(restored.checkpoint_tag(), [7; 32]);
restored.eval(next).unwrap();
assert_eq!(
restored
.logits()
.iter()
.enumerate()
.max_by(|a, b| a.1.total_cmp(b.1))
.unwrap()
.0,
continued_argmax
);
assert!((restored.logits()[0] - continued_logit).abs() < 1.0e-5);
std::fs::remove_file(checkpoint).unwrap();
}
}