1use ferrum_interfaces::{
20 model_executor::{DecodeInput, DecodeOutput},
21 tensor::TensorFactory,
22 KvCacheHandle, ModelExecutor,
23};
24use ferrum_types::{Result, TokenId};
25use rand::RngCore;
26use std::sync::Arc;
27
28fn softmax(logits: &[f32], temperature: f32) -> Vec<f32> {
32 if temperature == 0.0 {
33 let (argmax, _) =
35 logits
36 .iter()
37 .enumerate()
38 .fold((0usize, f32::NEG_INFINITY), |(bi, bv), (i, &v)| {
39 if v > bv {
40 (i, v)
41 } else {
42 (bi, bv)
43 }
44 });
45 let mut p = vec![0.0f32; logits.len()];
46 p[argmax] = 1.0;
47 return p;
48 }
49 let inv_t = 1.0 / temperature;
50 let mut max = f32::NEG_INFINITY;
51 for &l in logits {
52 if l > max {
53 max = l;
54 }
55 }
56 if !max.is_finite() {
57 max = 0.0;
58 }
59 let mut sum = 0.0f64;
60 let mut out = Vec::with_capacity(logits.len());
61 for &l in logits {
62 let e = ((l - max) * inv_t).exp();
63 out.push(e);
64 sum += e as f64;
65 }
66 let inv_sum = (1.0 / sum) as f32;
67 for x in out.iter_mut() {
68 *x *= inv_sum;
69 }
70 out
71}
72
73fn sample_categorical(probs: &[f32], u: f32) -> TokenId {
75 let mut acc = 0.0f32;
76 for (i, &p) in probs.iter().enumerate() {
77 acc += p;
78 if u <= acc {
79 return TokenId::new(i as u32);
80 }
81 }
82 TokenId::new((probs.len().saturating_sub(1)) as u32)
83}
84
85fn next_u(rng: &mut dyn RngCore) -> f32 {
87 (rng.next_u32() as f64 / (u32::MAX as f64 + 1.0)) as f32
89}
90
91fn residual(p_target: &[f32], p_draft: &[f32]) -> Vec<f32> {
94 debug_assert_eq!(p_target.len(), p_draft.len());
95 let mut r = Vec::with_capacity(p_target.len());
96 let mut sum = 0.0f64;
97 for (&pt, &pd) in p_target.iter().zip(p_draft.iter()) {
98 let d = (pt - pd).max(0.0);
99 r.push(d);
100 sum += d as f64;
101 }
102 if sum <= 0.0 {
103 return p_target.to_vec();
107 }
108 let inv = (1.0 / sum) as f32;
109 for x in r.iter_mut() {
110 *x *= inv;
111 }
112 r
113}
114
115pub struct Speculation<'a> {
125 pub draft_tokens: &'a [TokenId],
126 pub draft_logits: &'a [Vec<f32>],
127 pub target_logits: &'a [Vec<f32>],
128 pub temperature: f32,
129}
130
131#[derive(Debug, Clone, PartialEq)]
133pub struct SpeculationOutcome {
134 pub tokens: Vec<TokenId>,
137 pub rejected_at: usize,
140}
141
142pub fn verify_speculation(
152 spec: Speculation<'_>,
153 rng: &mut dyn RngCore,
154) -> Result<SpeculationOutcome> {
155 let n = spec.draft_tokens.len();
156 assert_eq!(spec.draft_logits.len(), n, "draft_logits count mismatch");
157 assert_eq!(
158 spec.target_logits.len(),
159 n + 1,
160 "target_logits must have N+1 rows (positions 0..N)"
161 );
162
163 let mut accepted: Vec<TokenId> = Vec::with_capacity(n + 1);
164
165 for i in 0..n {
166 let draft_token = spec.draft_tokens[i];
167 let idx = draft_token.get() as usize;
168
169 let p_target = softmax(&spec.target_logits[i], spec.temperature);
170 let p_draft = softmax(&spec.draft_logits[i], spec.temperature);
171
172 if idx >= p_target.len() || idx >= p_draft.len() {
173 let t = TokenId::new(
175 p_target
176 .iter()
177 .enumerate()
178 .fold((0, f32::NEG_INFINITY), |(bi, bv), (j, &v)| {
179 if v > bv {
180 (j, v)
181 } else {
182 (bi, bv)
183 }
184 })
185 .0 as u32,
186 );
187 accepted.push(t);
188 return Ok(SpeculationOutcome {
189 tokens: accepted,
190 rejected_at: i,
191 });
192 }
193
194 let pt = p_target[idx];
195 let pd = p_draft[idx].max(1e-20); let ratio = (pt / pd).min(1.0);
197 let u = next_u(rng);
198 if u < ratio {
199 accepted.push(draft_token);
201 } else {
202 let res = residual(&p_target, &p_draft);
204 let replacement = sample_categorical(&res, next_u(rng));
205 accepted.push(replacement);
206 return Ok(SpeculationOutcome {
207 tokens: accepted,
208 rejected_at: i,
209 });
210 }
211 }
212
213 let p_bonus = softmax(&spec.target_logits[n], spec.temperature);
217 let bonus = sample_categorical(&p_bonus, next_u(rng));
218 accepted.push(bonus);
219 Ok(SpeculationOutcome {
220 tokens: accepted,
221 rejected_at: n,
222 })
223}
224
225#[derive(Debug, Clone)]
227pub struct SpeculativeDecodingConfig {
228 pub num_speculative_tokens: usize,
232 pub temperature: f32,
235}
236
237impl Default for SpeculativeDecodingConfig {
238 fn default() -> Self {
239 Self {
240 num_speculative_tokens: 4,
241 temperature: 1.0,
242 }
243 }
244}
245
246pub struct SpeculativeRunner<'a> {
260 pub draft: &'a dyn ModelExecutor,
261 pub target: &'a dyn ModelExecutor,
262 pub tensor_factory: Arc<dyn TensorFactory>,
263 pub cfg: SpeculativeDecodingConfig,
264}
265
266pub struct SpeculativeStepOutcome {
268 pub tokens: Vec<TokenId>,
269 pub draft_kv: Arc<dyn KvCacheHandle>,
270 pub target_kv: Arc<dyn KvCacheHandle>,
271 pub rejected: bool,
274 pub rejected_at: usize,
275 pub draft_catchup_token: Option<TokenId>,
281}
282
283impl<'a> SpeculativeRunner<'a> {
284 pub async fn step(
288 &self,
289 last_token: TokenId,
290 draft_kv: Arc<dyn KvCacheHandle>,
291 target_kv: Arc<dyn KvCacheHandle>,
292 rng: &mut (dyn RngCore + Send),
293 ) -> Result<SpeculativeStepOutcome> {
294 let n = self.cfg.num_speculative_tokens.max(1);
295
296 let mut draft_tokens: Vec<TokenId> = Vec::with_capacity(n);
298 let mut draft_logits: Vec<Vec<f32>> = Vec::with_capacity(n);
299 let mut draft_kv_cur = draft_kv;
300 let mut draft_prev_token = last_token;
301 for _ in 0..n {
302 let input_tensor = tokens_to_tensor(&self.tensor_factory, &[draft_prev_token.get()])?;
303 let input = DecodeInput::new(input_tensor, draft_kv_cur.clone());
304 let output = self.draft.decode(&input).await?;
305 let logits = output.logits.to_vec_f32()?;
306 let next_token = argmax_token(&logits);
307 draft_tokens.push(next_token);
308 draft_logits.push(logits);
309 draft_kv_cur = output.kv_cache.clone();
310 draft_prev_token = next_token;
311 }
312
313 let mut verify_tokens = Vec::with_capacity(n + 1);
320 verify_tokens.push(last_token);
321 for i in 0..n {
322 verify_tokens.push(draft_tokens[i]);
323 }
324 let mut verify_inputs = Vec::with_capacity(verify_tokens.len());
325 let mut kv_for_verify = target_kv.clone();
326 for tok in &verify_tokens {
327 let input_tensor = tokens_to_tensor(&self.tensor_factory, &[tok.get()])?;
328 verify_inputs.push(DecodeInput::new(input_tensor, kv_for_verify.clone()));
329 kv_for_verify = kv_for_verify.clone();
332 }
333 let verify_outputs: Vec<DecodeOutput> = self.target.forward_verify(&verify_inputs).await?;
334 assert_eq!(verify_outputs.len(), n + 1);
335 let mut target_logits: Vec<Vec<f32>> = Vec::with_capacity(n + 1);
336 for out in &verify_outputs {
337 target_logits.push(out.logits.to_vec_f32()?);
338 }
339 let target_kv_cur = verify_outputs
340 .last()
341 .map(|o| o.kv_cache.clone())
342 .unwrap_or(target_kv);
343
344 let spec = Speculation {
345 draft_tokens: &draft_tokens,
346 draft_logits: &draft_logits,
347 target_logits: &target_logits,
348 temperature: self.cfg.temperature,
349 };
350 let outcome = verify_speculation(spec, rng)?;
351 let rejected = outcome.rejected_at < n;
352 let draft_catchup_token = if !rejected {
355 draft_tokens.last().copied()
356 } else {
357 None
358 };
359 Ok(SpeculativeStepOutcome {
360 tokens: outcome.tokens,
361 draft_kv: draft_kv_cur,
362 target_kv: target_kv_cur,
363 rejected,
364 rejected_at: outcome.rejected_at,
365 draft_catchup_token,
366 })
367 }
368}
369
370fn tokens_to_tensor(
371 factory: &Arc<dyn TensorFactory>,
372 token_ids: &[u32],
373) -> Result<ferrum_interfaces::tensor::TensorRef> {
374 use ferrum_types::{DataType, Device};
375 let f32_data: Vec<f32> = token_ids.iter().map(|&v| v as f32).collect();
376 let len = f32_data.len();
377 factory.from_slice(&f32_data, &[1, len], DataType::FP32, Device::CPU)
378}
379
380fn argmax_token(logits: &[f32]) -> TokenId {
381 let (idx, _) =
382 logits
383 .iter()
384 .enumerate()
385 .fold((0usize, f32::NEG_INFINITY), |(bi, bv), (i, &v)| {
386 if v > bv {
387 (i, v)
388 } else {
389 (bi, bv)
390 }
391 });
392 TokenId::new(idx as u32)
393}
394
395#[cfg(test)]
396mod tests {
397 use super::*;
398 use ferrum_interfaces::KvCacheHandle;
399 use ferrum_testkit::{ConfigurableModelExecutor, MockKvCacheHandle, MockTensorFactory};
400 use ferrum_types::RequestId;
401 use rand::{rngs::StdRng, SeedableRng};
402
403 fn biased_logits(vocab_size: usize, favored: u32, strength: f32) -> Vec<f32> {
406 let mut v = vec![0.0f32; vocab_size];
407 if (favored as usize) < vocab_size {
408 v[favored as usize] = strength;
409 }
410 v
411 }
412
413 #[test]
416 fn full_accept_when_draft_matches_target() {
417 let vocab = 32;
418 let drafts = vec![TokenId::new(3), TokenId::new(7), TokenId::new(11)];
419 let dl: Vec<Vec<f32>> = drafts
420 .iter()
421 .map(|t| biased_logits(vocab, t.get(), 20.0))
422 .collect();
423 let tl: Vec<Vec<f32>> = drafts
425 .iter()
426 .map(|t| biased_logits(vocab, t.get(), 20.0))
427 .chain(std::iter::once(biased_logits(vocab, 19, 20.0)))
428 .collect();
429
430 let mut rng = StdRng::seed_from_u64(1);
431 let out = verify_speculation(
432 Speculation {
433 draft_tokens: &drafts,
434 draft_logits: &dl,
435 target_logits: &tl,
436 temperature: 1.0,
437 },
438 &mut rng,
439 )
440 .unwrap();
441 assert_eq!(out.rejected_at, 3, "no rejections → rejected_at == N");
442 assert_eq!(
443 out.tokens,
444 vec![
445 TokenId::new(3),
446 TokenId::new(7),
447 TokenId::new(11),
448 TokenId::new(19)
449 ],
450 "should accept all three drafts + sample the bonus (token 19)"
451 );
452 }
453
454 #[test]
457 fn first_draft_rejected_residual_prefers_target() {
458 let vocab = 16;
459 let drafts = vec![TokenId::new(2)];
460 let dl = vec![biased_logits(vocab, 2, 20.0)];
461 let tl = vec![
463 biased_logits(vocab, 5, 20.0),
464 biased_logits(vocab, 0, 0.0), ];
466
467 let mut rng = StdRng::seed_from_u64(7);
468 let out = verify_speculation(
469 Speculation {
470 draft_tokens: &drafts,
471 draft_logits: &dl,
472 target_logits: &tl,
473 temperature: 1.0,
474 },
475 &mut rng,
476 )
477 .unwrap();
478
479 assert_eq!(out.rejected_at, 0);
480 assert_eq!(out.tokens.len(), 1);
481 assert_eq!(
482 out.tokens[0],
483 TokenId::new(5),
484 "residual should pick target's preferred token"
485 );
486 }
487
488 #[test]
491 fn partial_acceptance_second_draft_rejected() {
492 let vocab = 16;
493 let drafts = vec![TokenId::new(4), TokenId::new(9)];
494 let dl = vec![
495 biased_logits(vocab, 4, 20.0), biased_logits(vocab, 9, 20.0), ];
498 let tl = vec![
499 biased_logits(vocab, 4, 20.0), biased_logits(vocab, 1, 20.0), biased_logits(vocab, 0, 0.0), ];
503
504 let mut rng = StdRng::seed_from_u64(42);
505 let out = verify_speculation(
506 Speculation {
507 draft_tokens: &drafts,
508 draft_logits: &dl,
509 target_logits: &tl,
510 temperature: 1.0,
511 },
512 &mut rng,
513 )
514 .unwrap();
515
516 assert_eq!(out.rejected_at, 1);
517 assert_eq!(out.tokens.len(), 2);
518 assert_eq!(out.tokens[0], TokenId::new(4));
519 assert_eq!(
520 out.tokens[1],
521 TokenId::new(1),
522 "replacement should be the target's preferred token at position 1"
523 );
524 }
525
526 #[test]
528 fn zero_drafts_returns_bonus_only() {
529 let vocab = 8;
530 let tl = vec![biased_logits(vocab, 7, 20.0)];
531 let mut rng = StdRng::seed_from_u64(0);
532 let out = verify_speculation(
533 Speculation {
534 draft_tokens: &[],
535 draft_logits: &[],
536 target_logits: &tl,
537 temperature: 1.0,
538 },
539 &mut rng,
540 )
541 .unwrap();
542 assert_eq!(out.rejected_at, 0);
543 assert_eq!(out.tokens, vec![TokenId::new(7)]);
544 }
545
546 #[test]
549 fn greedy_temperature_full_accept_deterministic() {
550 let vocab = 16;
551 let drafts = vec![TokenId::new(2), TokenId::new(5)];
552 let dl: Vec<Vec<f32>> = drafts
553 .iter()
554 .map(|t| biased_logits(vocab, t.get(), 10.0))
555 .collect();
556 let tl: Vec<Vec<f32>> = drafts
557 .iter()
558 .map(|t| biased_logits(vocab, t.get(), 10.0))
559 .chain(std::iter::once(biased_logits(vocab, 13, 10.0)))
560 .collect();
561
562 let mut rng = StdRng::seed_from_u64(999);
563 let out = verify_speculation(
564 Speculation {
565 draft_tokens: &drafts,
566 draft_logits: &dl,
567 target_logits: &tl,
568 temperature: 0.0,
569 },
570 &mut rng,
571 )
572 .unwrap();
573
574 assert_eq!(out.rejected_at, 2);
575 assert_eq!(
576 out.tokens,
577 vec![TokenId::new(2), TokenId::new(5), TokenId::new(13)]
578 );
579 }
580
581 fn mock_kv(num_layers: usize) -> Arc<dyn KvCacheHandle> {
584 Arc::new(MockKvCacheHandle::new(RequestId::new(), num_layers, 0))
585 }
586
587 #[tokio::test]
590 async fn runner_full_accept_when_models_agree() {
591 let vocab = 64;
592 let draft: Arc<ConfigurableModelExecutor> = Arc::new(
593 ConfigurableModelExecutor::with_token_sequence(vocab, vec![13, 13, 13, 13, 13]),
594 );
595 let target: Arc<ConfigurableModelExecutor> = Arc::new(
596 ConfigurableModelExecutor::with_token_sequence(vocab, vec![13, 13, 13, 13, 13]),
597 );
598 let tf: Arc<dyn TensorFactory> = Arc::new(MockTensorFactory);
599 let runner = SpeculativeRunner {
600 draft: draft.as_ref(),
601 target: target.as_ref(),
602 tensor_factory: tf,
603 cfg: SpeculativeDecodingConfig {
604 num_speculative_tokens: 3,
605 temperature: 1.0,
606 },
607 };
608 let mut rng = StdRng::seed_from_u64(0);
609 let out = runner
610 .step(TokenId::new(5), mock_kv(12), mock_kv(12), &mut rng)
611 .await
612 .expect("step");
613
614 assert!(!out.rejected, "agreeing models should not reject");
615 assert_eq!(out.rejected_at, 3);
616 assert_eq!(out.tokens.len(), 4, "3 drafts + 1 bonus");
617 for &t in &out.tokens {
618 assert_eq!(
619 t.get(),
620 13,
621 "agreeing models should all emit the biased token 13"
622 );
623 }
624 }
625
626 #[tokio::test]
630 async fn runner_rejects_when_models_disagree() {
631 let vocab = 64;
632 let draft = Arc::new(ConfigurableModelExecutor::with_token_sequence(
633 vocab,
634 vec![7, 7, 7],
635 ));
636 let target = Arc::new(ConfigurableModelExecutor::with_token_sequence(
637 vocab,
638 vec![21, 21, 21, 21],
639 ));
640 let tf: Arc<dyn TensorFactory> = Arc::new(MockTensorFactory);
641 let runner = SpeculativeRunner {
642 draft: draft.as_ref(),
643 target: target.as_ref(),
644 tensor_factory: tf,
645 cfg: SpeculativeDecodingConfig {
646 num_speculative_tokens: 3,
647 temperature: 1.0,
648 },
649 };
650 let mut rng = StdRng::seed_from_u64(1);
651 let out = runner
652 .step(TokenId::new(0), mock_kv(12), mock_kv(12), &mut rng)
653 .await
654 .expect("step");
655
656 assert!(out.rejected);
657 assert_eq!(out.rejected_at, 0, "first draft should be rejected");
658 assert_eq!(out.tokens.len(), 1);
659 assert_eq!(
660 out.tokens[0].get(),
661 21,
662 "residual should sample target's preferred token"
663 );
664 }
665
666 #[test]
670 fn equal_distributions_always_accept() {
671 let vocab = 8;
672 let drafts = vec![TokenId::new(3)];
673 let dl = vec![vec![0.0f32; vocab]];
674 let tl = vec![vec![0.0f32; vocab], vec![0.0f32; vocab]];
675 for seed in 0..20u64 {
676 let mut rng = StdRng::seed_from_u64(seed);
677 let out = verify_speculation(
678 Speculation {
679 draft_tokens: &drafts,
680 draft_logits: &dl,
681 target_logits: &tl,
682 temperature: 1.0,
683 },
684 &mut rng,
685 )
686 .unwrap();
687 assert_eq!(out.rejected_at, 1, "seed {seed}: should accept");
688 assert_eq!(out.tokens.len(), 2);
689 assert_eq!(out.tokens[0], TokenId::new(3));
690 }
691 }
692}