Skip to main content

rig_candle/
generation.rs

1//! Generation configuration, sampling, incremental decoding, and inference sessions.
2
3use candle_transformers::generation::Sampling;
4use rig_core::completion::CompletionRequest;
5use serde::Deserialize;
6
7use crate::CandleError;
8
9/// Sampling and length defaults used when a completion request does not override them.
10#[derive(Debug, Clone, PartialEq)]
11pub struct GenerationConfig {
12    /// Maximum number of tokens generated for a request.
13    pub max_tokens: u64,
14    /// Sampling temperature. Zero selects greedy decoding.
15    pub temperature: f64,
16    /// Optional number of highest-probability tokens retained during sampling.
17    pub top_k: Option<usize>,
18    /// Optional nucleus-sampling probability threshold.
19    pub top_p: Option<f64>,
20    /// Deterministic random seed used by the sampler.
21    pub seed: u64,
22    /// Penalty applied to tokens repeated in the recent context. `1.0` disables it.
23    pub repeat_penalty: f32,
24    /// Number of recent tokens considered by the repeat penalty.
25    pub repeat_last_n: usize,
26}
27
28impl Default for GenerationConfig {
29    fn default() -> Self {
30        Self {
31            max_tokens: 256,
32            temperature: 0.8,
33            top_k: None,
34            top_p: Some(0.95),
35            seed: 299_792_458,
36            repeat_penalty: 1.1,
37            repeat_last_n: 64,
38        }
39    }
40}
41
42#[derive(Debug, Clone, Default, PartialEq)]
43enum OptionalGenerationOverride<T> {
44    #[default]
45    Inherit,
46    Set(T),
47    Disable,
48}
49
50impl<'de, T> Deserialize<'de> for OptionalGenerationOverride<T>
51where
52    T: Deserialize<'de>,
53{
54    fn deserialize<D>(deserializer: D) -> Result<Self, D::Error>
55    where
56        D: serde::Deserializer<'de>,
57    {
58        Option::<T>::deserialize(deserializer).map(|value| match value {
59            Some(value) => Self::Set(value),
60            None => Self::Disable,
61        })
62    }
63}
64
65impl<T> OptionalGenerationOverride<T> {
66    fn resolve(self, default: Option<T>) -> Option<T> {
67        match self {
68            Self::Inherit => default,
69            Self::Set(value) => Some(value),
70            Self::Disable => None,
71        }
72    }
73}
74
75#[derive(Debug, Default, Deserialize)]
76#[serde(default, deny_unknown_fields)]
77struct RequestGenerationOverrides {
78    top_k: OptionalGenerationOverride<usize>,
79    top_p: OptionalGenerationOverride<f64>,
80    seed: Option<u64>,
81    repeat_penalty: Option<f32>,
82    repeat_last_n: Option<usize>,
83}
84
85fn override_or<T>(value: Option<T>, default: T) -> T {
86    match value {
87        Some(value) => value,
88        None => default,
89    }
90}
91
92pub(crate) fn effective_generation(
93    request: &CompletionRequest,
94    defaults: &GenerationConfig,
95    vocab_size: usize,
96) -> Result<GenerationConfig, CandleError> {
97    let overrides = match &request.additional_params {
98        Some(value) => serde_json::from_value::<RequestGenerationOverrides>(value.clone())
99            .map_err(|error| CandleError::InvalidGeneration(error.to_string()))?,
100        None => RequestGenerationOverrides::default(),
101    };
102    let generation = GenerationConfig {
103        max_tokens: override_or(request.max_tokens, defaults.max_tokens),
104        temperature: override_or(request.temperature, defaults.temperature),
105        top_k: overrides.top_k.resolve(defaults.top_k),
106        top_p: overrides.top_p.resolve(defaults.top_p),
107        seed: override_or(overrides.seed, defaults.seed),
108        repeat_penalty: override_or(overrides.repeat_penalty, defaults.repeat_penalty),
109        repeat_last_n: override_or(overrides.repeat_last_n, defaults.repeat_last_n),
110    };
111    validate_generation(&generation, Some(vocab_size))?;
112    Ok(generation)
113}
114
115pub(crate) fn validate_generation(
116    generation: &GenerationConfig,
117    vocab_size: Option<usize>,
118) -> Result<(), CandleError> {
119    if generation.max_tokens == 0 {
120        return Err(CandleError::InvalidGeneration(
121            "max_tokens must be greater than zero".to_string(),
122        ));
123    }
124    if !generation.temperature.is_finite() || generation.temperature < 0.0 {
125        return Err(CandleError::InvalidGeneration(
126            "temperature must be finite and non-negative".to_string(),
127        ));
128    }
129    if let Some(top_k) = generation.top_k
130        && (top_k == 0 || vocab_size.is_some_and(|size| top_k > size))
131    {
132        return Err(CandleError::InvalidGeneration(
133            "top_k must be greater than zero and no larger than the vocabulary".to_string(),
134        ));
135    }
136    if let Some(top_p) = generation.top_p
137        && !(top_p.is_finite() && 0.0 < top_p && top_p <= 1.0)
138    {
139        return Err(CandleError::InvalidGeneration(
140            "top_p must be finite and in (0, 1]".to_string(),
141        ));
142    }
143    if !generation.repeat_penalty.is_finite() || generation.repeat_penalty <= 0.0 {
144        return Err(CandleError::InvalidGeneration(
145            "repeat_penalty must be finite and greater than zero".to_string(),
146        ));
147    }
148    Ok(())
149}
150
151pub(crate) fn sampling(config: &GenerationConfig) -> Sampling {
152    if config.temperature == 0.0 {
153        Sampling::ArgMax
154    } else {
155        match (config.top_k, config.top_p) {
156            (Some(k), Some(p)) => Sampling::TopKThenTopP {
157                k,
158                p,
159                temperature: config.temperature,
160            },
161            (Some(k), None) => Sampling::TopK {
162                k,
163                temperature: config.temperature,
164            },
165            (None, Some(p)) => Sampling::TopP {
166                p,
167                temperature: config.temperature,
168            },
169            (None, None) => Sampling::All {
170                temperature: config.temperature,
171            },
172        }
173    }
174}
175
176pub(crate) fn effective_output_limit(
177    prompt_tokens: usize,
178    requested_max_tokens: u64,
179    context_limit: usize,
180) -> Result<usize, CandleError> {
181    if prompt_tokens > context_limit {
182        return Err(CandleError::PromptTooLong {
183            prompt_tokens,
184            context_limit,
185        });
186    }
187    let remaining = context_limit - prompt_tokens;
188    if remaining == 0 {
189        return Err(CandleError::NoGenerationCapacity {
190            prompt_tokens,
191            context_limit,
192        });
193    }
194    let remaining = u64::try_from(remaining).map_err(|_| CandleError::NumericConversion {
195        field: "remaining_context_tokens",
196        value: u64::MAX,
197    })?;
198    max_tokens_to_usize(requested_max_tokens.min(remaining), usize::MAX as u64)
199}
200
201pub(crate) fn max_tokens_to_usize(value: u64, platform_max: u64) -> Result<usize, CandleError> {
202    if value > platform_max {
203        return Err(CandleError::NumericConversion {
204            field: "max_tokens",
205            value,
206        });
207    }
208    usize::try_from(value).map_err(|_| CandleError::NumericConversion {
209        field: "max_tokens",
210        value,
211    })
212}
213
214pub(crate) fn recent_tokens(tokens: &[u32], repeat_last_n: usize) -> &[u32] {
215    tokens
216        .get(tokens.len().saturating_sub(repeat_last_n)..)
217        .map_or(&[], |recent| recent)
218}
219
220pub(crate) fn next_cache_position(
221    prompt_tokens: usize,
222    generated_index: usize,
223) -> Result<usize, CandleError> {
224    prompt_tokens
225        .checked_add(generated_index)
226        .ok_or_else(|| CandleError::Inference("KV-cache position overflowed usize".to_string()))
227}
228
229use candle_core::Tensor;
230use candle_transformers::generation::LogitsProcessor;
231use candle_transformers::models::llama::{Cache, Llama};
232use candle_transformers::models::quantized_llama::ModelWeights as QuantizedLlama;
233use candle_transformers::models::quantized_qwen3::ModelWeights as QuantizedQwen3;
234use candle_transformers::utils::apply_repeat_penalty;
235use rig_core::OneOrMany;
236use rig_core::completion::{AssistantContent, CompletionResponse, GetTokenUsage};
237use rig_core::streaming::{RawStreamingChoice, RawStreamingToolCall};
238use tokenizers::tokenizer::DecodeStream;
239use tokenizers::{
240    DecoderWrapper, ModelWrapper, NormalizerWrapper, PostProcessorWrapper, PreTokenizerWrapper,
241    Tokenizer,
242};
243use web_time::{Duration, Instant};
244
245use crate::loader::{LoadedModel, LoadedWeights};
246use crate::profile::ModelFamily;
247use crate::runtime::{CancellationSignal, check_cancellation};
248use crate::types::{CandleCompletionResponse, FinishReason};
249
250type TokenDecodeStream<'a> = DecodeStream<
251    'a,
252    ModelWrapper,
253    NormalizerWrapper,
254    PreTokenizerWrapper,
255    PostProcessorWrapper,
256    DecoderWrapper,
257>;
258
259enum GenerationStep {
260    /// A token was sampled. Some token sequences need more IDs before they decode to valid UTF-8.
261    Token(Option<String>),
262    /// Generation and incremental decoding are complete.
263    Finished(CandleCompletionResponse),
264}
265
266pub(crate) struct IncrementalTextDecoder<'a> {
267    tokenizer: &'a Tokenizer,
268    stream: TokenDecodeStream<'a>,
269    token_ids: Vec<u32>,
270    text: String,
271    flushed: bool,
272}
273
274impl<'a> IncrementalTextDecoder<'a> {
275    pub(crate) fn new(tokenizer: &'a Tokenizer) -> Self {
276        Self {
277            tokenizer,
278            stream: tokenizer.decode_stream(true),
279            token_ids: Vec::new(),
280            text: String::new(),
281            flushed: false,
282        }
283    }
284
285    pub(crate) fn push(&mut self, token: u32) -> Result<Option<String>, CandleError> {
286        self.token_ids.push(token);
287        let fragment = self
288            .stream
289            .step(token)
290            .map_err(|error| CandleError::TokenizerDecoding(error.to_string()))?;
291        if let Some(fragment) = &fragment {
292            self.text.push_str(fragment);
293        }
294        Ok(fragment)
295    }
296
297    pub(crate) fn finish(&mut self) -> Result<Option<String>, CandleError> {
298        if self.flushed {
299            return Ok(None);
300        }
301        self.flushed = true;
302        let fully_decoded = self
303            .tokenizer
304            .decode(&self.token_ids, true)
305            .map_err(|error| CandleError::TokenizerDecoding(error.to_string()))?;
306        let suffix = fully_decoded.strip_prefix(&self.text).ok_or_else(|| {
307            CandleError::TokenizerDecoding(
308                "incremental decoding did not match complete decoding".to_string(),
309            )
310        })?;
311        if suffix.is_empty() {
312            Ok(None)
313        } else {
314            let suffix = suffix.to_string();
315            self.text.push_str(&suffix);
316            Ok(Some(suffix))
317        }
318    }
319
320    pub(crate) fn text(&self) -> &str {
321        &self.text
322    }
323}
324
325struct GenerationSession<'a> {
326    loaded: &'a LoadedModel,
327    generation: GenerationConfig,
328    cancellation: &'a CancellationSignal,
329    decoder: IncrementalTextDecoder<'a>,
330    weights: SessionWeights<'a>,
331    logits: Tensor,
332    processor: LogitsProcessor,
333    prompt_tokens: usize,
334    max_tokens: usize,
335    effective_max_tokens: u64,
336    all_tokens: Vec<u32>,
337    generated_tokens: usize,
338    finish_reason: Option<FinishReason>,
339    started: Instant,
340    prefill_duration: Duration,
341    time_to_first_token: Option<Duration>,
342    delivery_duration: Duration,
343}
344
345enum SessionWeights<'a> {
346    Safetensors { model: &'a Llama, cache: Cache },
347    QuantizedLlama(QuantizedLlama),
348    QuantizedQwen3(QuantizedQwen3),
349}
350
351impl SessionWeights<'_> {
352    fn forward(&mut self, input: &Tensor, position: usize) -> Result<Tensor, CandleError> {
353        match self {
354            Self::Safetensors { model, cache } => model
355                .forward(input, position, cache)
356                .and_then(|tensor| tensor.squeeze(0)),
357            Self::QuantizedLlama(model) => model
358                .forward(input, position)
359                .and_then(|tensor| tensor.squeeze(0)),
360            Self::QuantizedQwen3(model) => model
361                .forward(input, position)
362                .and_then(|tensor| tensor.squeeze(0)),
363        }
364        .map_err(|error| CandleError::Inference(error.to_string()))
365    }
366}
367
368impl<'a> GenerationSession<'a> {
369    pub(crate) fn new(
370        loaded: &'a LoadedModel,
371        request: CompletionRequest,
372        cancellation: &'a CancellationSignal,
373    ) -> Result<Self, CandleError> {
374        let prompt = crate::protocol::render_prompt(&request, loaded.profile.definition.protocol)?;
375        let generation =
376            effective_generation(&request, &loaded.generation, loaded.profile.vocab_size)?;
377        let encoding = loaded
378            .tokenizer
379            .encode(prompt, false)
380            .map_err(|error| CandleError::TokenizerEncoding(error.to_string()))?;
381        let prompt_ids = encoding.get_ids();
382        if prompt_ids.is_empty() {
383            return Err(CandleError::TokenizerEncoding(
384                "the rendered prompt produced no tokens".to_string(),
385            ));
386        }
387        let max_tokens = effective_output_limit(
388            prompt_ids.len(),
389            generation.max_tokens,
390            loaded.profile.context_limit,
391        )?;
392        let effective_max_tokens =
393            u64::try_from(max_tokens).map_err(|_| CandleError::NumericConversion {
394                field: "effective_max_tokens",
395                value: u64::MAX,
396            })?;
397
398        check_cancellation(cancellation)?;
399        let started = Instant::now();
400        let device = loaded.runtime.device();
401        let input = Tensor::new(prompt_ids, device)
402            .and_then(|tensor| tensor.unsqueeze(0))
403            .map_err(|error| CandleError::Inference(error.to_string()))?;
404        let mut weights = match &loaded.model {
405            LoadedWeights::Safetensors { model, config } => SessionWeights::Safetensors {
406                model,
407                cache: Cache::new(true, loaded.runtime.cache_dtype(), config, device)
408                    .map_err(|error| CandleError::Inference(error.to_string()))?,
409            },
410            LoadedWeights::QuantizedLlama(model) => SessionWeights::QuantizedLlama(model.clone()),
411            LoadedWeights::QuantizedQwen3(model) => SessionWeights::QuantizedQwen3(model.clone()),
412        };
413        check_cancellation(cancellation)?;
414        let logits = weights.forward(&input, 0)?;
415        let prefill_duration = started.elapsed();
416        let processor = LogitsProcessor::from_sampling(generation.seed, sampling(&generation));
417
418        Ok(Self {
419            loaded,
420            processor,
421            decoder: IncrementalTextDecoder::new(&loaded.tokenizer),
422            weights,
423            logits,
424            prompt_tokens: prompt_ids.len(),
425            max_tokens,
426            effective_max_tokens,
427            all_tokens: prompt_ids.to_vec(),
428            generated_tokens: 0,
429            finish_reason: None,
430            started,
431            prefill_duration,
432            time_to_first_token: None,
433            delivery_duration: Duration::ZERO,
434            generation,
435            cancellation,
436        })
437    }
438
439    fn next_token(&mut self) -> Result<GenerationStep, CandleError> {
440        check_cancellation(self.cancellation)?;
441
442        if self.finish_reason.is_some() {
443            return self.finish();
444        }
445
446        if self.generation.repeat_penalty != 1.0 && self.generation.repeat_last_n > 0 {
447            let recent = recent_tokens(&self.all_tokens, self.generation.repeat_last_n);
448            self.logits =
449                apply_repeat_penalty(&self.logits, self.generation.repeat_penalty, recent)
450                    .map_err(|error| CandleError::Inference(error.to_string()))?;
451        }
452        let token = self
453            .processor
454            .sample(&self.logits)
455            .map_err(|error| CandleError::Inference(error.to_string()))?;
456        self.generated_tokens = self.generated_tokens.checked_add(1).ok_or_else(|| {
457            CandleError::Inference("generated token count overflowed usize".to_string())
458        })?;
459        if self.time_to_first_token.is_none() {
460            self.time_to_first_token = Some(self.started.elapsed());
461        }
462
463        if self.loaded.profile.stop_tokens.contains(&token) {
464            self.finish_reason = Some(FinishReason::Eos);
465            return Ok(GenerationStep::Token(None));
466        }
467
468        self.all_tokens.push(token);
469        let fragment = self.decoder.push(token)?;
470
471        if self.generated_tokens >= self.max_tokens {
472            self.finish_reason = Some(FinishReason::MaxTokens);
473        } else {
474            check_cancellation(self.cancellation)?;
475            let generated_index = self.generated_tokens.saturating_sub(1);
476            let position = next_cache_position(self.prompt_tokens, generated_index)?;
477            let next = Tensor::new(&[token], self.loaded.runtime.device())
478                .and_then(|tensor| tensor.unsqueeze(0))
479                .map_err(|error| CandleError::Inference(error.to_string()))?;
480            self.logits = self.weights.forward(&next, position)?;
481        }
482
483        Ok(GenerationStep::Token(fragment))
484    }
485
486    pub(crate) fn finish(&mut self) -> Result<GenerationStep, CandleError> {
487        if let Some(suffix) = self.decoder.finish()? {
488            return Ok(GenerationStep::Token(Some(suffix)));
489        }
490
491        let finish_reason = self.finish_reason.ok_or_else(|| {
492            CandleError::Inference("generation finished without a finish reason".to_string())
493        })?;
494        let prompt_tokens = u64::try_from(self.prompt_tokens).map_err(|_| {
495            CandleError::Inference("prompt token count does not fit in u64".to_string())
496        })?;
497        let generated_tokens = u64::try_from(self.generated_tokens).map_err(|_| {
498            CandleError::Inference("generated token count does not fit in u64".to_string())
499        })?;
500        let generation_duration = self
501            .started
502            .elapsed()
503            .saturating_sub(self.delivery_duration);
504        let tokens_per_second = if generation_duration.is_zero() {
505            None
506        } else {
507            Some(generated_tokens as f64 / generation_duration.as_secs_f64())
508        };
509        Ok(GenerationStep::Finished(CandleCompletionResponse {
510            text: self.decoder.text().to_string(),
511            prompt_tokens,
512            generated_tokens,
513            requested_max_tokens: self.generation.max_tokens,
514            effective_max_tokens: self.effective_max_tokens,
515            finish_reason,
516            prefill_duration_ms: duration_millis(self.prefill_duration),
517            time_to_first_token_ms: self.time_to_first_token.map(duration_millis),
518            generation_duration_ms: duration_millis(generation_duration),
519            tokens_per_second,
520        }))
521    }
522
523    fn record_delivery_duration(&mut self, duration: Duration) {
524        self.delivery_duration = self.delivery_duration.saturating_add(duration);
525    }
526}
527
528fn duration_millis(duration: Duration) -> u64 {
529    u64::try_from(duration.as_millis()).map_or(u64::MAX, |value| value)
530}
531
532pub(crate) fn generate(
533    loaded: &LoadedModel,
534    request: CompletionRequest,
535    cancellation: &CancellationSignal,
536    mut emit: impl FnMut(String) -> Result<(), CandleError>,
537) -> Result<CandleCompletionResponse, CandleError> {
538    #[cfg(all(test, not(target_family = "wasm")))]
539    if let Some(control) = &loaded.test_control {
540        control.enter_generation()?;
541    }
542    let mut session = GenerationSession::new(loaded, request, cancellation)?;
543    loop {
544        match session.next_token()? {
545            GenerationStep::Token(Some(fragment)) if !fragment.is_empty() => {
546                let delivery_started = Instant::now();
547                let result = emit(fragment);
548                session.record_delivery_duration(delivery_started.elapsed());
549                result?;
550            }
551            GenerationStep::Token(_) => {}
552            GenerationStep::Finished(response) => return Ok(response),
553        }
554    }
555}
556
557pub(crate) fn infer(
558    loaded: &LoadedModel,
559    request: CompletionRequest,
560    cancellation: &CancellationSignal,
561) -> Result<CompletionResponse<CandleCompletionResponse>, CandleError> {
562    let parse_request = request.clone();
563    let mut raw_response = generate(loaded, request, cancellation, |_| Ok(()))?;
564    let parsed = crate::protocol::parse_assistant(
565        &raw_response.text,
566        &parse_request,
567        loaded.profile.definition.protocol,
568    )?;
569    raw_response.text = parsed.visible_text;
570    let choice = OneOrMany::many(parsed.items).map_err(|_| {
571        CandleError::Inference("output protocol produced no assistant content".to_string())
572    })?;
573    let usage = raw_response.token_usage();
574    Ok(CompletionResponse {
575        choice,
576        usage,
577        raw_response,
578        message_id: None,
579    })
580}
581
582pub(crate) fn stream_generate(
583    loaded: &LoadedModel,
584    request: CompletionRequest,
585    cancellation: &CancellationSignal,
586    mut emit: impl FnMut(RawStreamingChoice<CandleCompletionResponse>) -> Result<(), CandleError>,
587) -> Result<CandleCompletionResponse, CandleError> {
588    if loaded.profile.definition.protocol != ModelFamily::Qwen3 {
589        return generate(loaded, request, cancellation, |fragment| {
590            emit(RawStreamingChoice::Message(fragment))
591        });
592    }
593
594    let parse_request = request.clone();
595    // Qwen tool syntax can straddle arbitrary token boundaries. Buffer one
596    // model turn so control markup is never leaked as assistant text; complete
597    // tool calls are still delivered through Rig's streaming agent driver.
598    let mut response = generate(loaded, request, cancellation, |_| Ok(()))?;
599    let parsed = crate::protocol::parse_assistant(
600        &response.text,
601        &parse_request,
602        loaded.profile.definition.protocol,
603    )?;
604    response.text = parsed.visible_text;
605    for item in parsed.items {
606        match item {
607            AssistantContent::Text(text) => emit(RawStreamingChoice::Message(text.text))?,
608            AssistantContent::ToolCall(call) => {
609                let mut raw =
610                    RawStreamingToolCall::new(call.id, call.function.name, call.function.arguments);
611                raw.call_id = call.call_id;
612                raw.signature = call.signature;
613                raw.additional_params = call.additional_params;
614                emit(RawStreamingChoice::ToolCall(raw))?;
615            }
616            AssistantContent::Reasoning(reasoning) => {
617                for content in reasoning.content {
618                    emit(RawStreamingChoice::Reasoning {
619                        id: reasoning.id.clone(),
620                        content,
621                    })?;
622                }
623            }
624            AssistantContent::Image(_) => {
625                return Err(CandleError::Inference(
626                    "text-only Qwen output parser produced image content".to_string(),
627                ));
628            }
629        }
630    }
631    Ok(response)
632}