///| Sampling controls applied after a model produces logits. Keeping this
///| logic separate from proposal verification makes draft and target policies
///|
/// explicit, reproducible, and independently testable.
pub enum SamplingError {
InvalidTopK(Int)
InvalidTopP(Double)
InvalidMinP(Double)
InvalidDistribution
ProbabilityFailure
InvalidTemperature
} derive(Eq, Debug)
///| A conservative decoding policy. `top_k = 0` means no rank cutoff, while
///|
/// `top_p = 1.0` and `min_p = 0.0` mean their respective filters are off.
pub struct SamplingConfig {
temperature : Double
top_k : Int
top_p : Double
min_p : Double
}
///|
pub fn SamplingConfig::new(
temperature? : Double = 1.0,
top_k? : Int = 0,
top_p? : Double = 1.0,
min_p? : Double = 0.0,
) -> Result[SamplingConfig, SamplingError] {
if !finite_number(temperature) || temperature <= 0.0 {
return Err(InvalidTemperature)
}
if top_k < 0 {
return Err(InvalidTopK(top_k))
}
if !finite_number(top_p) || top_p <= 0.0 || top_p > 1.0 {
return Err(InvalidTopP(top_p))
}
if !finite_number(min_p) || min_p < 0.0 || min_p > 1.0 {
return Err(InvalidMinP(min_p))
}
Ok({ temperature, top_k, top_p, min_p })
}
///|
/// Default is pure temperature sampling with no vocabulary truncation.
pub fn SamplingConfig::default() -> SamplingConfig {
{ temperature: 1.0, top_k: 0, top_p: 1.0, min_p: 0.0 }
}
///|
pub fn SamplingConfig::temperature(self : SamplingConfig) -> Double {
self.temperature
}
///|
pub fn SamplingConfig::top_k(self : SamplingConfig) -> Int {
self.top_k
}
///| Return vocabulary indices in descending probability order. Ties retain
///|
/// their lower token id, which is useful when test fixtures need stability.
pub fn ranked_indices(
distribution : Array[Double],
) -> Result[Array[Int], SamplingError] {
if distribution.length() == 0 {
return Err(InvalidDistribution)
}
let indices : Array[Int] = []
for index in 0.. distribution[candidate] {
indices.insert(position, index)
inserted = true
break
}
}
if !inserted {
indices.push(index)
}
}
Ok(indices)
}
///| Normalize non-negative weights after one or more filters have masked the
///| vocabulary. It rejects an all-zero mask rather than silently sampling an
///|
/// impossible token.
pub fn normalize_weights(
weights : Array[Double],
) -> Result[Array[Double], SamplingError] {
if weights.length() == 0 {
return Err(InvalidDistribution)
}
let mut total = 0.0
for weight in weights {
if !finite_number(weight) || weight < 0.0 {
return Err(InvalidDistribution)
}
total = total + weight
}
if !finite_number(total) || total <= 0.0 {
return Err(InvalidDistribution)
}
let normalized : Array[Double] = []
for weight in weights {
normalized.push(weight / total)
}
Ok(normalized)
}
///|
/// Keep the `limit` most likely tokens. `limit = 0` keeps every token.
pub fn filter_top_k(
distribution : Array[Double],
limit : Int,
) -> Result[Array[Double], SamplingError] {
if limit < 0 {
return Err(InvalidTopK(limit))
}
let ranking = match ranked_indices(distribution) {
Ok(value) => value
Err(error) => return Err(error)
}
if limit == 0 || limit >= ranking.length() {
return normalize_weights(distribution)
}
let filtered : Array[Double] = []
for _ in distribution {
filtered.push(0.0)
}
for position in 0.. Result[Array[Double], SamplingError] {
if threshold <= 0.0 || threshold > 1.0 {
return Err(InvalidTopP(threshold))
}
let normalized = match normalize_weights(distribution) {
Ok(value) => value
Err(error) => return Err(error)
}
let ranking = match ranked_indices(normalized) {
Ok(value) => value
Err(error) => return Err(error)
}
let filtered : Array[Double] = []
for _ in normalized {
filtered.push(0.0)
}
let mut cumulative = 0.0
for index in ranking {
filtered[index] = normalized[index]
cumulative = cumulative + normalized[index]
if cumulative >= threshold {
break
}
}
normalize_weights(filtered)
}
///| Min-p filtering uses the most likely token as an adaptive reference:
///|
/// every retained probability is at least `threshold * max_probability`.
pub fn filter_min_p(
distribution : Array[Double],
threshold : Double,
) -> Result[Array[Double], SamplingError] {
if threshold < 0.0 || threshold > 1.0 {
return Err(InvalidMinP(threshold))
}
let normalized = match normalize_weights(distribution) {
Ok(value) => value
Err(error) => return Err(error)
}
let mut maximum = normalized[0]
for value in normalized {
if value > maximum {
maximum = value
}
}
let filtered : Array[Double] = []
for value in normalized {
filtered.push(if value >= maximum * threshold { value } else { 0.0 })
}
normalize_weights(filtered)
}
///| Transform logits into a filtered categorical distribution. Filter order is
///| intentionally fixed: temperature, top-k, top-p, then min-p. This makes a
///|
/// configuration portable across a draft worker and a target worker.
pub fn distribution_for_sampling(
logits : Array[Double],
config : SamplingConfig,
) -> Result[Array[Double], SamplingError] {
let initial = match softmax(logits, temperature=config.temperature) {
Ok(value) => value
Err(_) => return Err(ProbabilityFailure)
}
let after_k = match filter_top_k(initial, config.top_k) {
Ok(value) => value
Err(error) => return Err(error)
}
let after_p = match filter_top_p(after_k, config.top_p) {
Ok(value) => value
Err(error) => return Err(error)
}
filter_min_p(after_p, config.min_p)
}
///| Sample a token from filtered logits with a caller-supplied random number.
///| The caller owns randomness so command-line runs and regression tests can
///|
/// be replayed bit-for-bit.
pub fn sample_logits(
logits : Array[Double],
config : SamplingConfig,
unit_interval : Double,
) -> Result[Int, SamplingError] {
let distribution = match distribution_for_sampling(logits, config) {
Ok(value) => value
Err(error) => return Err(error)
}
match sample_categorical(distribution, unit_interval) {
Ok(token) => Ok(token)
Err(_) => Err(ProbabilityFailure)
}
}