spark_model/
engine.rs

1// SPDX-License-Identifier: AGPL-3.0-only
2
3//! Inference engine — generate loop for a single request.
4//!
5//! Orchestrates prefill → decode → sample loop using the [`Model`] trait.
6//! The engine is stateless — each call to [`generate`] creates a fresh
7//! sequence, runs inference, and returns output tokens.
8
9use anyhow::Result;
10use spark_runtime::sampler::SamplingParams;
11
12use crate::traits::Model;
13
14/// Result of a generate call.
15pub struct GenerateResult {
16    /// Output tokens (does not include prompt tokens).
17    pub output_tokens: Vec<u32>,
18    /// Why generation stopped: "stop" (EOS/stop token) or "length" (max_tokens).
19    pub finish_reason: String,
20}
21
22/// Generate response tokens from a prompt.
23///
24/// Runs prefill on the prompt tokens, then iteratively decodes up to
25/// `params.max_tokens` output tokens, stopping early on EOS or stop tokens.
26pub fn generate(
27    model: &dyn Model,
28    prompt_tokens: &[u32],
29    params: &SamplingParams,
30) -> Result<GenerateResult> {
31    let mut seq = model.alloc_sequence()?;
32    let stream = 0u64; // Default CUDA stream.
33
34    let result = generate_inner(model, prompt_tokens, params, &mut seq, stream);
35
36    // Always free GPU resources, even on error
37    model.free_sequence(&mut seq)?;
38
39    result
40}
41
42fn generate_inner(
43    model: &dyn Model,
44    prompt_tokens: &[u32],
45    params: &SamplingParams,
46    seq: &mut crate::traits::SequenceState,
47    stream: u64,
48) -> Result<GenerateResult> {
49    // ── Prefill ──
50    let logits_ptr = model.prefill(prompt_tokens, seq, stream)?;
51    let first_token = model.argmax_on_device(logits_ptr, stream)?;
52
53    let mut output_tokens = Vec::with_capacity(params.max_tokens);
54    output_tokens.push(first_token);
55
56    if params.stop_token_ids.contains(&first_token) {
57        return Ok(GenerateResult {
58            output_tokens,
59            finish_reason: "stop".to_string(),
60        });
61    }
62
63    // ── Decode loop ──
64    for _step in 1..params.max_tokens {
65        let last_token = *output_tokens.last().unwrap();
66        let logits_ptr = model.decode(last_token, seq, stream)?;
67        let token = model.argmax_on_device(logits_ptr, stream)?;
68
69        output_tokens.push(token);
70
71        if params.stop_token_ids.contains(&token) {
72            return Ok(GenerateResult {
73                output_tokens,
74                finish_reason: "stop".to_string(),
75            });
76        }
77    }
78
79    Ok(GenerateResult {
80        output_tokens,
81        finish_reason: "length".to_string(),
82    })
83}
84
85/// Generate response tokens with per-token callback.
86///
87/// Same as [`generate`] but calls `on_token(token_id)` after each token
88/// is produced (including the first token from prefill). The callback
89/// is synchronous — designed for the caller to send tokens through a
90/// channel without pulling in an async runtime dependency.
91pub fn generate_streaming<F>(
92    model: &dyn Model,
93    prompt_tokens: &[u32],
94    params: &SamplingParams,
95    mut on_token: F,
96) -> Result<GenerateResult>
97where
98    F: FnMut(u32),
99{
100    let mut seq = model.alloc_sequence()?;
101    let stream = 0u64;
102
103    let result = generate_streaming_inner(
104        model,
105        prompt_tokens,
106        params,
107        &mut on_token,
108        &mut seq,
109        stream,
110    );
111
112    model.free_sequence(&mut seq)?;
113
114    result
115}
116
117fn generate_streaming_inner<F>(
118    model: &dyn Model,
119    prompt_tokens: &[u32],
120    params: &SamplingParams,
121    on_token: &mut F,
122    seq: &mut crate::traits::SequenceState,
123    stream: u64,
124) -> Result<GenerateResult>
125where
126    F: FnMut(u32),
127{
128    let logits_ptr = model.prefill(prompt_tokens, seq, stream)?;
129    let first_token = model.argmax_on_device(logits_ptr, stream)?;
130
131    let mut output_tokens = Vec::with_capacity(params.max_tokens);
132    output_tokens.push(first_token);
133    on_token(first_token);
134
135    if params.stop_token_ids.contains(&first_token) {
136        return Ok(GenerateResult {
137            output_tokens,
138            finish_reason: "stop".to_string(),
139        });
140    }
141
142    for _step in 1..params.max_tokens {
143        let last_token = *output_tokens.last().unwrap();
144        let logits_ptr = model.decode(last_token, seq, stream)?;
145        let token = model.argmax_on_device(logits_ptr, stream)?;
146
147        output_tokens.push(token);
148        on_token(token);
149
150        if params.stop_token_ids.contains(&token) {
151            return Ok(GenerateResult {
152                output_tokens,
153                finish_reason: "stop".to_string(),
154            });
155        }
156    }
157
158    Ok(GenerateResult {
159        output_tokens,
160        finish_reason: "length".to_string(),
161    })
162}
163
164/// Generate with speculative decoding (MTP).
165///
166/// Delegates the speculative decode loop to `model.generate_speculative()`,
167/// which has access to GPU/buffers needed for the MTP proposer.
168pub fn generate_speculative(
169    model: &dyn Model,
170    prompt_tokens: &[u32],
171    params: &SamplingParams,
172    num_drafts: usize,
173) -> Result<GenerateResult> {
174    model.generate_speculative(prompt_tokens, params, num_drafts)
175}
176
177#[cfg(test)]
178mod tests;