1use candle_transformers::generation::Sampling;
4use rig_core::completion::CompletionRequest;
5use serde::Deserialize;
6
7use crate::CandleError;
8
9#[derive(Debug, Clone, PartialEq)]
11pub struct GenerationConfig {
12 pub max_tokens: u64,
14 pub temperature: f64,
16 pub top_k: Option<usize>,
18 pub top_p: Option<f64>,
20 pub seed: u64,
22 pub repeat_penalty: f32,
24 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 Token(Option<String>),
262 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 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}