Accelerate DeepSeek DSpark inference
This commit is contained in:
@@ -365,6 +365,28 @@ struct Dspark {
|
||||
strict: bool,
|
||||
drafted: u64,
|
||||
accepted: u64,
|
||||
scheduler_cycles: u32,
|
||||
scheduler_accepted: u32,
|
||||
scheduler_no_draft: u32,
|
||||
scheduler_skip: u32,
|
||||
scheduler_lifetime_accepted: u32,
|
||||
scheduler_long_accept_seen: bool,
|
||||
last_confidence: Option<f32>,
|
||||
}
|
||||
|
||||
fn dspark_scheduler_pause(cycles: u32, accepted: u32, no_draft: u32) -> u32 {
|
||||
if cycles == 0 {
|
||||
return 0;
|
||||
}
|
||||
let low_acceptance = u64::from(accepted) * 1_000 < u64::from(cycles) * 1_500;
|
||||
let many_no_draft = no_draft * 2 >= cycles;
|
||||
if many_no_draft {
|
||||
4
|
||||
} else if low_acceptance {
|
||||
2
|
||||
} else {
|
||||
0
|
||||
}
|
||||
}
|
||||
|
||||
impl Dspark {
|
||||
@@ -411,7 +433,7 @@ impl Dspark {
|
||||
raw_caches: (0..config.stages)
|
||||
.map(|_| Buffer::floats(u64::from(session.raw_cap) * shape.head_dim))
|
||||
.collect::<Result<_, _>>()?,
|
||||
scratch: BatchScratch::allocate(model, session.context, rows)?,
|
||||
scratch: BatchScratch::allocate(model, session.context, rows, false)?,
|
||||
logits: Buffer::floats(u64::from(config.block_size) * shape.vocab)?,
|
||||
config,
|
||||
weights,
|
||||
@@ -430,9 +452,57 @@ impl Dspark {
|
||||
strict: settings.dspark_strict || quality,
|
||||
drafted: 0,
|
||||
accepted: 0,
|
||||
scheduler_cycles: 0,
|
||||
scheduler_accepted: 0,
|
||||
scheduler_no_draft: 0,
|
||||
scheduler_skip: 0,
|
||||
scheduler_lifetime_accepted: 0,
|
||||
scheduler_long_accept_seen: false,
|
||||
last_confidence: None,
|
||||
})
|
||||
}
|
||||
|
||||
fn scheduler_should_skip(&mut self) -> bool {
|
||||
if self.scheduler_skip == 0 {
|
||||
return false;
|
||||
}
|
||||
self.scheduler_skip -= 1;
|
||||
true
|
||||
}
|
||||
|
||||
fn scheduler_note(&mut self, accepted: u32, no_draft: bool) {
|
||||
self.scheduler_cycles += 1;
|
||||
self.scheduler_accepted = self.scheduler_accepted.saturating_add(accepted);
|
||||
self.scheduler_lifetime_accepted =
|
||||
self.scheduler_lifetime_accepted.saturating_add(accepted);
|
||||
self.scheduler_long_accept_seen |= accepted > 2;
|
||||
self.scheduler_no_draft += u32::from(no_draft);
|
||||
if no_draft {
|
||||
let skip = if self.scheduler_lifetime_accepted == 0
|
||||
&& self
|
||||
.last_confidence
|
||||
.is_some_and(|confidence| confidence <= 0.5)
|
||||
{
|
||||
7
|
||||
} else if self.scheduler_lifetime_accepted != 0 && !self.scheduler_long_accept_seen {
|
||||
4
|
||||
} else {
|
||||
3
|
||||
};
|
||||
self.scheduler_skip = self.scheduler_skip.max(skip);
|
||||
}
|
||||
if self.scheduler_cycles >= 4 {
|
||||
self.scheduler_skip = self.scheduler_skip.max(dspark_scheduler_pause(
|
||||
self.scheduler_cycles,
|
||||
self.scheduler_accepted,
|
||||
self.scheduler_no_draft,
|
||||
));
|
||||
self.scheduler_cycles = 0;
|
||||
self.scheduler_accepted = 0;
|
||||
self.scheduler_no_draft = 0;
|
||||
}
|
||||
}
|
||||
|
||||
fn target_slot(&self, layer: u32) -> Option<u32> {
|
||||
self.config
|
||||
.target_layers
|
||||
@@ -889,12 +959,9 @@ impl Dspark {
|
||||
},
|
||||
"embedding DSpark draft block",
|
||||
)?;
|
||||
commands.finish()?;
|
||||
|
||||
for stage in 0..self.weights.len() {
|
||||
self.eval_stage(support, stage, pos, raw_cap, shape)?;
|
||||
if stage + 1 < self.weights.len() {
|
||||
let commands = Commands::begin()?;
|
||||
self.stage_input_hc.copy_from(
|
||||
shape.hc * shape.embd * 4,
|
||||
&self.scratch.next_hc,
|
||||
@@ -902,10 +969,8 @@ impl Dspark {
|
||||
u64::from(self.config.block_size) * shape.hc * shape.embd * 4,
|
||||
"feeding the next DSpark stage",
|
||||
)?;
|
||||
commands.finish()?;
|
||||
}
|
||||
}
|
||||
let commands = Commands::begin()?;
|
||||
self.stage_output_hc.copy_from(
|
||||
0,
|
||||
&self.scratch.next_hc,
|
||||
@@ -913,8 +978,7 @@ impl Dspark {
|
||||
u64::from(self.config.block_size) * shape.hc * shape.embd * 4,
|
||||
"capturing DSpark stage output",
|
||||
)?;
|
||||
commands.finish()?;
|
||||
self.eval_output_heads(base, support, base_weights, token)
|
||||
self.eval_output_heads(commands, base, support, base_weights, token)
|
||||
}
|
||||
|
||||
fn eval_stage(
|
||||
@@ -955,7 +1019,6 @@ impl Dspark {
|
||||
.scratch
|
||||
.kv
|
||||
.view(shape.head_dim * 4, u64::from(draft) * shape.head_dim * 4)?;
|
||||
let commands = Commands::begin()?;
|
||||
call(
|
||||
unsafe {
|
||||
ds4_gpu_rms_norm_plain_rows_tensor(
|
||||
@@ -1216,7 +1279,7 @@ impl Dspark {
|
||||
"expanding DSpark attention HC",
|
||||
)?;
|
||||
self.eval_stage_ffn(stage_index, shape, support)?;
|
||||
commands.finish()
|
||||
Ok(())
|
||||
}
|
||||
|
||||
fn eval_stage_ffn(
|
||||
@@ -1415,6 +1478,7 @@ impl Dspark {
|
||||
|
||||
fn eval_output_heads(
|
||||
&mut self,
|
||||
commands: Commands,
|
||||
base: &Model,
|
||||
support: &Gguf,
|
||||
base_weights: &Weights,
|
||||
@@ -1454,7 +1518,6 @@ impl Dspark {
|
||||
.scratch
|
||||
.norm
|
||||
.view(0, u64::from(draft) * shape.embd * 4)?;
|
||||
let commands = Commands::begin()?;
|
||||
call(
|
||||
unsafe {
|
||||
ds4_gpu_rms_norm_plain_rows_tensor(
|
||||
@@ -1540,6 +1603,7 @@ impl Dspark {
|
||||
output_norm.read_f32(&mut hidden)?;
|
||||
let mut proposals = Vec::with_capacity(draft as usize);
|
||||
let mut previous = first_token as u32;
|
||||
self.last_confidence = None;
|
||||
for row in 0..draft as usize {
|
||||
let state = dense_row(support, markov_w1, previous)?;
|
||||
let mut features = Vec::with_capacity(shape.embd as usize + state.len());
|
||||
@@ -1554,19 +1618,16 @@ impl Dspark {
|
||||
let value = confidence_logit.exp();
|
||||
value / (1.0 + value)
|
||||
};
|
||||
if row == 0 {
|
||||
self.last_confidence = Some(confidence_logit);
|
||||
}
|
||||
if self.confidence_threshold > 0.0 && confidence_value < self.confidence_threshold {
|
||||
break;
|
||||
}
|
||||
let row_logits = &logits[row * shape.vocab as usize..(row + 1) * shape.vocab as usize];
|
||||
let mut best = (0_i32, f32::NEG_INFINITY);
|
||||
for (token, &logit) in row_logits.iter().enumerate() {
|
||||
let value = logit + dense_dot(support, markov_w2, token as u32, &state)?;
|
||||
if value > best.1 {
|
||||
best = (token as i32, value);
|
||||
}
|
||||
}
|
||||
proposals.push(best.0);
|
||||
previous = best.0 as u32;
|
||||
let best = dense_argmax(support, markov_w2, &state, row_logits)?;
|
||||
proposals.push(best);
|
||||
previous = best as u32;
|
||||
}
|
||||
self.drafted += proposals.len() as u64;
|
||||
Ok(proposals)
|
||||
@@ -2625,12 +2686,18 @@ struct BatchScratch {
|
||||
shared_up: Buffer,
|
||||
shared_mid: Buffer,
|
||||
shared_out: Buffer,
|
||||
output_logits: Option<Buffer>,
|
||||
}
|
||||
|
||||
impl BatchScratch {
|
||||
fn allocate(model: &Model, pos: u32, rows: u32) -> Result<Self, String> {
|
||||
fn allocate(model: &Model, pos: u32, rows: u32, output_logits: bool) -> Result<Self, String> {
|
||||
let shape = model.shape;
|
||||
let rows = u64::from(rows);
|
||||
let output_rows = if output_logits && rows > 1 && rows < 8 {
|
||||
8
|
||||
} else {
|
||||
rows
|
||||
};
|
||||
let rows = u64::from(output_rows);
|
||||
let hc_dim = shape.hc * shape.embd;
|
||||
let mix_hc = 2 * shape.hc + shape.hc * shape.hc;
|
||||
let q_dim = shape.heads * shape.head_dim;
|
||||
@@ -2680,6 +2747,9 @@ impl BatchScratch {
|
||||
shared_up: Buffer::floats(rows * shape.ff_expert)?,
|
||||
shared_mid: Buffer::floats(rows * shape.ff_expert)?,
|
||||
shared_out: Buffer::floats(rows * shape.embd)?,
|
||||
output_logits: output_logits
|
||||
.then(|| Buffer::floats(rows * shape.vocab))
|
||||
.transpose()?,
|
||||
})
|
||||
}
|
||||
}
|
||||
@@ -3123,6 +3193,10 @@ impl DeepSeekExecutor {
|
||||
|
||||
pub(super) fn eval(&mut self, token: i32) -> Result<(), String> {
|
||||
self.eval_target(token)?;
|
||||
self.seed_dspark_current_cache()
|
||||
}
|
||||
|
||||
fn seed_dspark_current_cache(&mut self) -> Result<(), String> {
|
||||
if let (Some(_), Some(ssd)) = (&self.dspark, &self.ssd) {
|
||||
install_speculative_model_maps(&self.model, "DSpark support mapping")?;
|
||||
ssd.static_decode_map_current
|
||||
@@ -3586,6 +3660,14 @@ impl DeepSeekExecutor {
|
||||
return Ok(accepted);
|
||||
}
|
||||
if self.dspark.is_some() {
|
||||
let scheduler_skip = self
|
||||
.dspark
|
||||
.as_mut()
|
||||
.is_some_and(Dspark::scheduler_should_skip);
|
||||
if max_tokens < 10 || scheduler_skip {
|
||||
self.seed_dspark_current_cache()?;
|
||||
return Ok(accepted);
|
||||
}
|
||||
if self.ssd.is_some() {
|
||||
install_speculative_model_maps(&self.model, "DSpark support mapping")?;
|
||||
}
|
||||
@@ -3611,6 +3693,7 @@ impl DeepSeekExecutor {
|
||||
{
|
||||
proposals.truncate(stop + 1);
|
||||
}
|
||||
let no_draft = proposals.is_empty();
|
||||
let verified = self.verify_target_suffix(&proposals, cancelled)?;
|
||||
if verified.is_empty() {
|
||||
self.dspark
|
||||
@@ -3619,7 +3702,9 @@ impl DeepSeekExecutor {
|
||||
.commit_proposed_prefix(1, self.session.raw_cap);
|
||||
}
|
||||
accepted.extend_from_slice(&verified);
|
||||
self.dspark.as_mut().expect("DSpark disappeared").accepted += verified.len() as u64;
|
||||
let dspark = self.dspark.as_mut().expect("DSpark disappeared");
|
||||
dspark.accepted += verified.len() as u64;
|
||||
dspark.scheduler_note(verified.len() as u32, no_draft);
|
||||
return Ok(accepted);
|
||||
}
|
||||
let Some(draft_cap) = self.legacy_mtp.as_ref().map(|mtp| {
|
||||
@@ -3711,6 +3796,14 @@ impl DeepSeekExecutor {
|
||||
{
|
||||
return Ok(emitted);
|
||||
}
|
||||
let scheduler_skip = self
|
||||
.dspark
|
||||
.as_mut()
|
||||
.is_some_and(Dspark::scheduler_should_skip);
|
||||
if max_tokens < 10 || scheduler_skip {
|
||||
self.seed_dspark_current_cache()?;
|
||||
return Ok(emitted);
|
||||
}
|
||||
if self.ssd.is_some() {
|
||||
install_speculative_model_maps(&self.model, "DSpark support mapping")?;
|
||||
}
|
||||
@@ -3736,11 +3829,11 @@ impl DeepSeekExecutor {
|
||||
{
|
||||
proposals.truncate(stop + 1);
|
||||
}
|
||||
let no_draft = proposals.is_empty();
|
||||
if proposals.len() < 2 {
|
||||
self.dspark
|
||||
.as_mut()
|
||||
.expect("DSpark disappeared")
|
||||
.commit_proposed_prefix(1, self.session.raw_cap);
|
||||
let dspark = self.dspark.as_mut().expect("DSpark disappeared");
|
||||
dspark.commit_proposed_prefix(1, self.session.raw_cap);
|
||||
dspark.scheduler_note(0, no_draft);
|
||||
return Ok(emitted);
|
||||
}
|
||||
let (verified, accepted) = self.verify_target_suffix_stochastic(
|
||||
@@ -3753,7 +3846,9 @@ impl DeepSeekExecutor {
|
||||
cancelled,
|
||||
)?;
|
||||
emitted.extend_from_slice(&verified);
|
||||
self.dspark.as_mut().expect("DSpark disappeared").accepted += accepted as u64;
|
||||
let dspark = self.dspark.as_mut().expect("DSpark disappeared");
|
||||
dspark.accepted += accepted as u64;
|
||||
dspark.scheduler_note(accepted as u32, no_draft);
|
||||
Ok(emitted)
|
||||
}
|
||||
|
||||
@@ -4020,7 +4115,8 @@ impl DeepSeekExecutor {
|
||||
{
|
||||
return Err("prefill contains a token outside the vocabulary".into());
|
||||
}
|
||||
let mut batch = BatchScratch::allocate(&self.model, self.session.position, rows)?;
|
||||
let mut batch =
|
||||
BatchScratch::allocate(&self.model, self.session.position, rows, collect_tops)?;
|
||||
if let Some(dspark) = &mut self.dspark {
|
||||
dspark.begin_capture();
|
||||
}
|
||||
@@ -4057,7 +4153,8 @@ impl DeepSeekExecutor {
|
||||
} else {
|
||||
None
|
||||
};
|
||||
let commands = Commands::begin()?;
|
||||
let pipelined_verifier = collect_tops && self.ssd.is_none() && self.profile.is_none();
|
||||
let mut commands = Some(Commands::begin()?);
|
||||
call(
|
||||
unsafe {
|
||||
ds4_gpu_embed_tokens_hc_tensor(
|
||||
@@ -4074,7 +4171,12 @@ impl DeepSeekExecutor {
|
||||
},
|
||||
"batch token embedding",
|
||||
)?;
|
||||
commands.finish()?;
|
||||
if !pipelined_verifier {
|
||||
commands
|
||||
.take()
|
||||
.expect("batch commands are active")
|
||||
.finish()?;
|
||||
}
|
||||
|
||||
for (index, (weights, state)) in self
|
||||
.weights
|
||||
@@ -4113,7 +4215,9 @@ impl DeepSeekExecutor {
|
||||
)?);
|
||||
}
|
||||
}
|
||||
let commands = Commands::begin()?;
|
||||
if !pipelined_verifier {
|
||||
commands = Some(Commands::begin()?);
|
||||
}
|
||||
encode_batch_layer(
|
||||
&batch,
|
||||
state,
|
||||
@@ -4151,7 +4255,19 @@ impl DeepSeekExecutor {
|
||||
Some(ssd) => ssd.seed_mapped_layer(&self.model, weights, index, true)?,
|
||||
None => true,
|
||||
};
|
||||
commands.finish()?;
|
||||
if pipelined_verifier {
|
||||
if (index + 1) % 4 == 0 {
|
||||
commands
|
||||
.as_mut()
|
||||
.expect("batch commands are active")
|
||||
.flush()?;
|
||||
}
|
||||
} else {
|
||||
commands
|
||||
.take()
|
||||
.expect("batch commands are active")
|
||||
.finish()?;
|
||||
}
|
||||
if !seeded_from_map {
|
||||
let seeded = self
|
||||
.ssd
|
||||
@@ -4171,6 +4287,12 @@ impl DeepSeekExecutor {
|
||||
);
|
||||
std::mem::swap(&mut batch.current_hc, &mut batch.next_hc);
|
||||
}
|
||||
if pipelined_verifier {
|
||||
commands
|
||||
.take()
|
||||
.expect("batch commands are active")
|
||||
.finish()?;
|
||||
}
|
||||
|
||||
if self.dspark.is_some() && self.ssd.is_some() {
|
||||
install_speculative_model_maps(&self.model, "DSpark prefill support mapping")?;
|
||||
@@ -4193,11 +4315,29 @@ impl DeepSeekExecutor {
|
||||
)?;
|
||||
}
|
||||
|
||||
let output_rows = if collect_tops { rows } else { 1 };
|
||||
let first_output = rows - output_rows;
|
||||
let mut tops = Vec::with_capacity(output_rows as usize);
|
||||
let mut output_logits = Vec::with_capacity(output_rows as usize);
|
||||
for row in first_output..rows {
|
||||
let (tops, output_logits) = if collect_tops {
|
||||
commands = Some(Commands::begin()?);
|
||||
encode_batch_output(&batch, &self.weights, shape, map, size, rows)?;
|
||||
commands
|
||||
.take()
|
||||
.expect("batch commands are active")
|
||||
.finish()?;
|
||||
let logits = batch
|
||||
.output_logits
|
||||
.as_ref()
|
||||
.expect("batch output logits are allocated");
|
||||
let mut all_logits = vec![0.0; (u64::from(rows) * shape.vocab) as usize];
|
||||
logits.read_f32(&mut all_logits)?;
|
||||
let output_logits = all_logits
|
||||
.chunks_exact(shape.vocab as usize)
|
||||
.map(<[f32]>::to_vec)
|
||||
.collect::<Vec<_>>();
|
||||
let tops = output_logits.iter().map(|logits| argmax(logits)).collect();
|
||||
self.logits
|
||||
.clone_from(output_logits.last().expect("batch has an output row"));
|
||||
(tops, output_logits)
|
||||
} else {
|
||||
let row = rows - 1;
|
||||
let commands = Commands::begin()?;
|
||||
self.session.scratch.current_hc.copy_from(
|
||||
0,
|
||||
@@ -4209,11 +4349,8 @@ impl DeepSeekExecutor {
|
||||
encode_output(&self.session.scratch, &self.weights, shape, map, size)?;
|
||||
commands.finish()?;
|
||||
self.session.scratch.logits.read_f32(&mut self.logits)?;
|
||||
tops.push(argmax(&self.logits));
|
||||
if collect_tops {
|
||||
output_logits.push(self.logits.clone());
|
||||
}
|
||||
}
|
||||
(vec![argmax(&self.logits)], Vec::new())
|
||||
};
|
||||
self.session.position += rows;
|
||||
self.tokens.extend_from_slice(tokens);
|
||||
if let Some(profile) = &self.profile {
|
||||
@@ -7938,6 +8075,108 @@ fn encode_output(
|
||||
)
|
||||
}
|
||||
|
||||
fn encode_batch_output(
|
||||
s: &BatchScratch,
|
||||
w: &Weights,
|
||||
shape: super::Shape,
|
||||
map: *const c_void,
|
||||
size: u64,
|
||||
rows: u32,
|
||||
) -> Result<(), String> {
|
||||
let head_rows = if rows > 1 && rows < 8 { 8 } else { rows };
|
||||
let hc_dim = shape.hc * shape.embd;
|
||||
let output_pre = s.hc_mix.view(0, u64::from(rows) * shape.hc * 4)?;
|
||||
let output_weights = s.hc_split.view(0, u64::from(rows) * shape.hc * 4)?;
|
||||
let output_embedding = s.current.view(0, u64::from(rows) * shape.embd * 4)?;
|
||||
let output_norm = s.norm.view(0, u64::from(head_rows) * shape.embd * 4)?;
|
||||
let logits = s
|
||||
.output_logits
|
||||
.as_ref()
|
||||
.ok_or("batch output logits are not allocated")?;
|
||||
call(
|
||||
unsafe {
|
||||
ds4_gpu_rms_norm_plain_rows_tensor(
|
||||
s.flat_hc.raw(),
|
||||
s.current_hc.raw(),
|
||||
hc_dim as u32,
|
||||
rows,
|
||||
shape.rms_epsilon,
|
||||
)
|
||||
},
|
||||
"batch output HC norm",
|
||||
)?;
|
||||
f16_rows(
|
||||
&output_pre,
|
||||
w.output_hc_fn,
|
||||
hc_dim,
|
||||
shape.hc,
|
||||
&s.flat_hc,
|
||||
rows,
|
||||
map,
|
||||
size,
|
||||
)?;
|
||||
call(
|
||||
unsafe {
|
||||
ds4_gpu_output_hc_weights_tensor(
|
||||
output_weights.raw(),
|
||||
output_pre.raw(),
|
||||
map,
|
||||
size,
|
||||
w.output_hc_scale.offset,
|
||||
w.output_hc_base.offset,
|
||||
shape.hc as u32,
|
||||
shape.hc_epsilon,
|
||||
)
|
||||
},
|
||||
"batch output HC weights",
|
||||
)?;
|
||||
call(
|
||||
unsafe {
|
||||
ds4_gpu_hc_weighted_sum_tensor(
|
||||
output_embedding.raw(),
|
||||
s.current_hc.raw(),
|
||||
output_weights.raw(),
|
||||
shape.embd as u32,
|
||||
shape.hc as u32,
|
||||
)
|
||||
},
|
||||
"batch output HC collapse",
|
||||
)?;
|
||||
call(
|
||||
unsafe {
|
||||
ds4_gpu_rms_norm_weight_rows_tensor(
|
||||
output_norm.raw(),
|
||||
output_embedding.raw(),
|
||||
map,
|
||||
size,
|
||||
w.output_norm.offset,
|
||||
shape.embd as u32,
|
||||
rows,
|
||||
shape.rms_epsilon,
|
||||
)
|
||||
},
|
||||
"batch output norm",
|
||||
)?;
|
||||
if head_rows > rows {
|
||||
s.norm
|
||||
.view(
|
||||
u64::from(rows) * shape.embd * 4,
|
||||
u64::from(head_rows - rows) * shape.embd * 4,
|
||||
)?
|
||||
.fill(0.0, u64::from(head_rows - rows) * shape.embd)?;
|
||||
}
|
||||
matmul_rows(
|
||||
logits,
|
||||
w.output,
|
||||
shape.embd,
|
||||
shape.vocab,
|
||||
&output_norm,
|
||||
head_rows,
|
||||
map,
|
||||
size,
|
||||
)
|
||||
}
|
||||
|
||||
#[allow(clippy::too_many_arguments)]
|
||||
fn encode_mtp_output(
|
||||
s: &Scratch,
|
||||
@@ -8115,7 +8354,11 @@ fn dense_dot(model: &Gguf, weight: Weight, row: u32, values: &[f32]) -> Result<f
|
||||
}
|
||||
let bytes =
|
||||
unsafe { std::slice::from_raw_parts(model.map_ptr().add(offset as usize), row_bytes) };
|
||||
let sum = match weight.kind {
|
||||
Ok(dense_dot_bytes(weight.kind, width, bytes, values))
|
||||
}
|
||||
|
||||
fn dense_dot_bytes(kind: u32, width: usize, bytes: &[u8], values: &[f32]) -> f32 {
|
||||
match kind {
|
||||
F32 => bytes
|
||||
.chunks_exact(4)
|
||||
.zip(values)
|
||||
@@ -8143,8 +8386,111 @@ fn dense_dot(model: &Gguf, weight: Weight, row: u32, values: &[f32]) -> Result<f
|
||||
sum
|
||||
}
|
||||
_ => unreachable!(),
|
||||
}
|
||||
}
|
||||
|
||||
fn dense_argmax(
|
||||
model: &Gguf,
|
||||
weight: Weight,
|
||||
values: &[f32],
|
||||
logits: &[f32],
|
||||
) -> Result<i32, String> {
|
||||
if weight.dims[0] as usize != values.len() || weight.dims[1] as usize != logits.len() {
|
||||
return Err("DSpark dense argmax has mismatched dimensions".into());
|
||||
}
|
||||
let width = values.len();
|
||||
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_len = row_bytes
|
||||
.checked_mul(logits.len())
|
||||
.ok_or("DSpark dense argmax size overflow")?;
|
||||
if weight.offset > model.len() || bytes_len as u64 > model.len() - weight.offset {
|
||||
return Err("DSpark dense argmax is outside the GGUF mapping".into());
|
||||
}
|
||||
let bytes = unsafe {
|
||||
std::slice::from_raw_parts(model.map_ptr().add(weight.offset as usize), bytes_len)
|
||||
};
|
||||
Ok(sum)
|
||||
let quantized = (weight.kind == Q8_0).then(|| quantize_q8_activation(values));
|
||||
let workers = std::thread::available_parallelism()
|
||||
.map_or(1, std::num::NonZero::get)
|
||||
.min(logits.len());
|
||||
let chunk = logits.len().div_ceil(workers);
|
||||
let best = std::thread::scope(|scope| {
|
||||
let mut handles = Vec::with_capacity(workers);
|
||||
let quantized = quantized.as_ref();
|
||||
for start in (0..logits.len()).step_by(chunk) {
|
||||
let end = (start + chunk).min(logits.len());
|
||||
handles.push(scope.spawn(move || {
|
||||
let mut best = (start, f32::NEG_INFINITY);
|
||||
for token in start..end {
|
||||
let row = &bytes[token * row_bytes..(token + 1) * row_bytes];
|
||||
let dot = quantized.as_ref().map_or_else(
|
||||
|| dense_dot_bytes(weight.kind, width, row, values),
|
||||
|(values, scales)| dense_dot_q8(row, values, scales, width),
|
||||
);
|
||||
let score = logits[token] + dot;
|
||||
if score > best.1 {
|
||||
best = (token, score);
|
||||
}
|
||||
}
|
||||
best
|
||||
}));
|
||||
}
|
||||
handles
|
||||
.into_iter()
|
||||
.map(|handle| handle.join().expect("DSpark argmax worker panicked"))
|
||||
.fold((0, f32::NEG_INFINITY), |best, candidate| {
|
||||
if candidate.1 > best.1 {
|
||||
candidate
|
||||
} else {
|
||||
best
|
||||
}
|
||||
})
|
||||
});
|
||||
Ok(best.0 as i32)
|
||||
}
|
||||
|
||||
fn quantize_q8_activation(values: &[f32]) -> (Vec<i8>, Vec<f32>) {
|
||||
let blocks = values.len().div_ceil(32);
|
||||
let mut quantized = vec![0; blocks * 32];
|
||||
let mut scales = Vec::with_capacity(blocks);
|
||||
for (block, values) in values.chunks(32).enumerate() {
|
||||
let scale = values
|
||||
.iter()
|
||||
.fold(0.0_f32, |max, value| max.max(value.abs()))
|
||||
/ 127.0;
|
||||
let inverse = if scale == 0.0 { 0.0 } else { scale.recip() };
|
||||
scales.push(scale);
|
||||
for (target, value) in quantized[block * 32..].iter_mut().zip(values) {
|
||||
*target = (value * inverse).round_ties_even().clamp(-128.0, 127.0) as i8;
|
||||
}
|
||||
}
|
||||
(quantized, scales)
|
||||
}
|
||||
|
||||
fn dense_dot_q8(bytes: &[u8], values: &[i8], scales: &[f32], width: usize) -> f32 {
|
||||
bytes
|
||||
.chunks_exact(34)
|
||||
.zip(values.chunks_exact(32))
|
||||
.zip(scales)
|
||||
.enumerate()
|
||||
.map(|(block, ((bytes, values), &scale))| {
|
||||
let count = (width - block * 32).min(32);
|
||||
let weight_scale = half_to_f32(u16::from_le_bytes([bytes[0], bytes[1]]));
|
||||
let dot = bytes[2..]
|
||||
.iter()
|
||||
.zip(values)
|
||||
.take(count)
|
||||
.map(|(&weight, &value)| i32::from(weight as i8) * i32::from(value))
|
||||
.sum::<i32>();
|
||||
weight_scale * scale * dot as f32
|
||||
})
|
||||
.sum()
|
||||
}
|
||||
|
||||
fn half_to_f32(value: u16) -> f32 {
|
||||
@@ -8364,9 +8710,9 @@ fn check(result: i32, operation: &str) -> Result<(), String> {
|
||||
#[cfg(test)]
|
||||
mod tests {
|
||||
use super::{
|
||||
compression_ratio, effective_prefill_cap, effective_raw_cap,
|
||||
compression_ratio, dspark_scheduler_pause, effective_prefill_cap, effective_raw_cap,
|
||||
estimated_deepseek_runtime_bytes, finish_deepseek_model_spans,
|
||||
gpu::ds4_gpu_print_memory_report, raw_batch_span, raw_decode_span,
|
||||
gpu::ds4_gpu_print_memory_report, quantize_q8_activation, raw_batch_span, raw_decode_span,
|
||||
};
|
||||
use crate::engine::{FLASH, MXFP4, PRO};
|
||||
|
||||
@@ -8386,6 +8732,20 @@ mod tests {
|
||||
assert_eq!(compression_ratio(PRO, 1), 128);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn dspark_scheduler_matches_the_ds4_default_window() {
|
||||
assert_eq!(dspark_scheduler_pause(4, 8, 0), 0);
|
||||
assert_eq!(dspark_scheduler_pause(4, 5, 0), 2);
|
||||
assert_eq!(dspark_scheduler_pause(4, 8, 2), 4);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn dspark_markov_activation_matches_q8_rounding() {
|
||||
let (values, scales) = quantize_q8_activation(&[0.0, 1.0, -1.0, 0.5]);
|
||||
assert_eq!(&values[..4], &[0, 127, -127, 64]);
|
||||
assert_eq!(scales, [1.0 / 127.0]);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn pro_q4_model_spans_remain_isolated() {
|
||||
assert_eq!(
|
||||
@@ -8479,11 +8839,14 @@ mod tests {
|
||||
#[ignore = "requires a 0731 Flash GGUF and an Apple M5 device"]
|
||||
fn flash_0731_m5_decode_performance_gate() {
|
||||
use super::{DeepSeekExecutor, Digest, Sha256, configure_sources};
|
||||
use crate::engine::{ChatTurn, Model};
|
||||
use crate::engine::gguf::Gguf;
|
||||
use crate::engine::validation::validate_support;
|
||||
use crate::engine::{ChatTurn, Model, Rng, sample};
|
||||
use crate::model::ModelChoice;
|
||||
use crate::settings::{
|
||||
EngineSpeculativeSettings, EngineSsdSettings, EngineSteeringSettings, ReasoningMode,
|
||||
};
|
||||
use std::sync::atomic::AtomicBool;
|
||||
use std::time::Instant;
|
||||
|
||||
configure_sources().unwrap();
|
||||
@@ -8498,8 +8861,27 @@ mod tests {
|
||||
let frontier = std::env::var("DS4SERVER_BENCH_FRONTIER")
|
||||
.ok()
|
||||
.map(|value| value.parse::<usize>().unwrap());
|
||||
let dspark = std::env::var_os("DS4SERVER_BENCH_DSPARK").is_some();
|
||||
let temperature = std::env::var("DS4SERVER_BENCH_TEMPERATURE")
|
||||
.ok()
|
||||
.map_or(0.0, |value| value.parse::<f32>().unwrap());
|
||||
let confidence = std::env::var("DS4SERVER_BENCH_DSPARK_CONFIDENCE")
|
||||
.ok()
|
||||
.map(|value| value.parse::<f32>().unwrap());
|
||||
let measured = std::env::var("DS4SERVER_BENCH_MEASURED")
|
||||
.ok()
|
||||
.map_or(128, |value| value.parse::<u32>().unwrap());
|
||||
let run = || {
|
||||
let model = Model::open_main(&path, ModelChoice::DeepSeekV4Flash0731).unwrap();
|
||||
let mut model = Model::open_main(&path, ModelChoice::DeepSeekV4Flash0731).unwrap();
|
||||
if dspark {
|
||||
let support_path =
|
||||
installed_artifacts(ModelChoice::DeepSeekV4Flash0731, false, true)
|
||||
.mtp
|
||||
.unwrap();
|
||||
let support = Gguf::open(&support_path).unwrap();
|
||||
model.support_kind = Some(validate_support(&support, &model.shape).unwrap());
|
||||
model.support = Some(support);
|
||||
}
|
||||
let expert_kind = model
|
||||
.main
|
||||
.tensor("blk.4.ffn_gate_exps.weight")
|
||||
@@ -8541,9 +8923,9 @@ mod tests {
|
||||
mtp_margin: 3.0,
|
||||
glm_mtp: false,
|
||||
glm_mtp_timing: false,
|
||||
dspark: false,
|
||||
dspark_confidence_threshold: 0.9,
|
||||
dspark_confidence_threshold_set: false,
|
||||
dspark,
|
||||
dspark_confidence_threshold: confidence.unwrap_or(0.8),
|
||||
dspark_confidence_threshold_set: confidence.is_some(),
|
||||
dspark_strict: false,
|
||||
dspark_exact_sampling: false,
|
||||
},
|
||||
@@ -8568,20 +8950,39 @@ mod tests {
|
||||
let prefill_seconds = prefill_started.elapsed().as_secs_f64();
|
||||
let mut generated = Vec::new();
|
||||
let mut latencies_ms = Vec::new();
|
||||
let measured = 128_u32;
|
||||
let mut rng = Rng::new(12_345);
|
||||
let cancelled = AtomicBool::new(false);
|
||||
let started = Instant::now();
|
||||
for _ in 0..measured {
|
||||
let token = executor
|
||||
.logits()
|
||||
.iter()
|
||||
.enumerate()
|
||||
.filter(|(token, _)| *token as i32 != eos)
|
||||
.max_by(|left, right| left.1.total_cmp(right.1))
|
||||
.map_or(-1, |(token, _)| token as i32);
|
||||
generated.push(token);
|
||||
while generated.len() < measured as usize {
|
||||
let token = sample(executor.logits(), temperature, 1.0, 0.05, 0, &mut rng);
|
||||
assert_ne!(token, eos, "benchmark reached EOS before {measured} tokens");
|
||||
let token_started = Instant::now();
|
||||
executor.eval(token).unwrap();
|
||||
latencies_ms.push(token_started.elapsed().as_secs_f64() * 1_000.0);
|
||||
let remaining = measured - generated.len() as u32;
|
||||
let cycle = if temperature <= 0.0 {
|
||||
executor.eval_speculative_greedy(
|
||||
token,
|
||||
remaining,
|
||||
ReasoningMode::Direct,
|
||||
&cancelled,
|
||||
)
|
||||
} else {
|
||||
executor.eval_speculative_sampled(
|
||||
token,
|
||||
remaining,
|
||||
ReasoningMode::Direct,
|
||||
temperature,
|
||||
1.0,
|
||||
0.05,
|
||||
0,
|
||||
&mut rng,
|
||||
&cancelled,
|
||||
)
|
||||
}
|
||||
.unwrap();
|
||||
let per_token_ms =
|
||||
token_started.elapsed().as_secs_f64() * 1_000.0 / cycle.len() as f64;
|
||||
latencies_ms.extend(std::iter::repeat_n(per_token_ms, cycle.len()));
|
||||
generated.extend(cycle);
|
||||
}
|
||||
let seconds = started.elapsed().as_secs_f64();
|
||||
let tokens_per_second = f64::from(measured) / seconds;
|
||||
@@ -8603,11 +9004,17 @@ mod tests {
|
||||
unsafe { ds4_gpu_print_memory_report(c"benchmark".as_ptr()) };
|
||||
let stats = executor.execution_stats();
|
||||
eprintln!(
|
||||
"DS4SERVER_METAL_PERF mode={} model=flash-0731 expert_kind={expert_kind} context={context} prompt={} prefill_tps={:.6} measured={measured} seconds={seconds:.6} tps={tokens_per_second:.6} first_ms={:.6} steady_tps={steady_tokens_per_second:.6} p50_ms={p50:.6} p95_ms={p95:.6} cache_entries={} cache_hits={} cache_misses={} pread_bytes={}",
|
||||
"DS4SERVER_METAL_PERF mode={} speculative={} temperature={temperature} model=flash-0731 expert_kind={expert_kind} context={context} prompt={} prefill_tps={:.6} measured={measured} seconds={seconds:.6} tps={tokens_per_second:.6} first_ms={:.6} steady_tps={steady_tokens_per_second:.6} p50_ms={p50:.6} p95_ms={p95:.6} cycles={} drafted={} accepted={} verifier_passes={} verifier_ms={} cache_entries={} cache_hits={} cache_misses={} pread_bytes={}",
|
||||
if streaming { "ssd" } else { "resident" },
|
||||
if dspark { "dspark" } else { "plain" },
|
||||
prompt.len(),
|
||||
prompt.len() as f64 / prefill_seconds,
|
||||
latencies_ms[0],
|
||||
stats.speculative_cycles,
|
||||
stats.drafted_tokens,
|
||||
stats.accepted_draft_tokens,
|
||||
stats.verifier_passes,
|
||||
stats.verifier_ms,
|
||||
stats.ssd_cache_entries,
|
||||
stats.ssd_cache_hits,
|
||||
stats.ssd_cache_misses,
|
||||
|
||||
@@ -1644,6 +1644,13 @@ impl Commands {
|
||||
"executing Metal commands",
|
||||
)
|
||||
}
|
||||
|
||||
pub(super) fn flush(&mut self) -> Result<(), String> {
|
||||
check(
|
||||
unsafe { ds4_gpu_flush_commands() },
|
||||
"flushing Metal commands",
|
||||
)
|
||||
}
|
||||
}
|
||||
|
||||
pub(super) struct ParallelFfn(bool);
|
||||
|
||||
Reference in New Issue
Block a user