///|
/// Numerically stable probability primitives for decoding algorithms.
pub enum ProbabilityError {
EmptyLogits
InvalidTemperature
InvalidProbability
InvalidSample
} derive(Eq, Debug)
///|
/// Also rejects infinities: infinity minus itself is NaN.
fn finite_number(value : Double) -> Bool {
value - value == 0.0
}
///|
fn valid_uniform(value : Double) -> Bool {
finite_number(value) && value >= 0.0 && value < 1.0
}
///|
pub fn log_sum_exp(values : Array[Double]) -> Result[Double, ProbabilityError] {
if values.length() == 0 {
return Err(EmptyLogits)
}
let mut maximum = values[0]
for value in values {
if !finite_number(value) {
return Err(InvalidProbability)
}
if value > maximum {
maximum = value
}
}
let mut total = 0.0
for value in values {
total = total + @math.exp(value - maximum)
}
Ok(maximum + @math.ln(total))
}
///|
pub fn softmax(
logits : Array[Double],
temperature? : Double = 1.0,
) -> Result[Array[Double], ProbabilityError] {
if logits.length() == 0 {
return Err(EmptyLogits)
}
if !finite_number(temperature) || temperature <= 0.0 {
return Err(InvalidTemperature)
}
let scaled : Array[Double] = []
for value in logits {
if !finite_number(value) || !finite_number(value / temperature) {
return Err(InvalidProbability)
}
scaled.push(value / temperature)
}
let mut maximum = scaled[0]
for value in scaled {
maximum = maximum.max(value)
}
let output : Array[Double] = []
let mut total = 0.0
for value in scaled {
let weight = @math.exp(value - maximum)
output.push(weight)
total = total + weight
}
for i in 0.. Double {
if index < 0 || index >= distribution.length() {
0.0
} else {
distribution[index]
}
}
///|
pub fn accept_probability(
target : Double,
draft : Double,
) -> Result[Double, ProbabilityError] {
if !finite_number(target) ||
!finite_number(draft) ||
target < 0.0 ||
target > 1.0 ||
draft <= 0.0 ||
draft > 1.0 {
return Err(InvalidProbability)
}
Ok((target / draft).min(1.0))
}
///|
pub fn residual_distribution(
target : Array[Double],
draft : Array[Double],
) -> Result[Array[Double], ProbabilityError] {
if !distribution_is_valid(target) ||
!distribution_is_valid(draft) ||
target.length() != draft.length() {
return Err(InvalidProbability)
}
let raw : Array[Double] = []
let mut total = 0.0
for index in 0.. Result[Int, ProbabilityError] {
if !valid_uniform(unit_interval) {
return Err(InvalidSample)
}
if !distribution_is_valid(distribution) {
return Err(InvalidProbability)
}
let mut cumulative = 0.0
let mut last_positive = 0
for index in 0.. 0.0 {
last_positive = index
}
cumulative = cumulative + value
if unit_interval < cumulative {
return Ok(index)
}
}
// Rounding near one must never select a zero-probability tail token.
Ok(last_positive)
}