use crate::error::Result; use candle_core::Tensor; use candle_transformers::generation::{LogitsProcessor, Sampling}; pub const DEFAULT_TEMPERATURE: f64 = 0.9; pub const DEFAULT_TOPK: usize = 50; pub const DEFAULT_TOPP: f64 = 0.9; pub struct CsmSampler { inner: LogitsProcessor, } impl CsmSampler { /// Build a sampler with explicit knobs. `top_p == 0.0` (or `>= 1.0`) disables /// nucleus filtering and falls back to pure top-k. `temperature <= 0.0` selects /// argmax (deterministic) regardless of top_k/top_p. pub fn new(seed: u64, temperature: f64, top_k: usize, top_p: f64) -> Self { let sampling = if temperature <= 0.0 { Sampling::ArgMax } else if top_p > 0.0 && top_p < 1.0 { Sampling::TopKThenTopP { k: top_k, p: top_p, temperature, } } else { Sampling::TopK { k: top_k, temperature, } }; Self { inner: LogitsProcessor::from_sampling(seed, sampling), } } pub fn default_for_csm(seed: u64) -> Self { Self::new(seed, DEFAULT_TEMPERATURE, DEFAULT_TOPK, DEFAULT_TOPP) } pub fn sample(&mut self, logits: &Tensor) -> Result { Ok(self.inner.sample(logits)?) } pub fn inner_mut(&mut self) -> &mut LogitsProcessor { &mut self.inner } }