///|
/// Errors are values so callers can decide whether to skip a malformed model row.
pub enum SamplingError {
EmptyLogits
AllMasked
NonFiniteLogit(Int)
InvalidUniform
InvalidProbability(Int)
InvalidParameter(String)
NumericalFailure
} derive(Eq, Debug)
///|
pub fn SamplingError::invalid_parameter(message : String) -> SamplingError {
InvalidParameter(message)
}
///|
pub fn SamplingError::message(self : SamplingError) -> String {
match self {
EmptyLogits => "logits or weights must not be empty"
AllMasked => "at least one finite logit must remain unmasked"
NonFiniteLogit(index) =>
"logit at token " + index.to_string() + " is NaN or positive infinity"
InvalidUniform => "uniform draw must be finite and in [0, 1)"
InvalidProbability(index) =>
"weight at token " + index.to_string() + " is invalid"
InvalidParameter(message) => message
NumericalFailure => "sampling calculation failed numerically"
}
}
///|
fn finite(value : Double) -> Bool {
value == value && value != 1.0 / 0.0 && value != -1.0 / 0.0
}
///|
fn check_logits(logits : Array[Double]) -> Result[Unit, SamplingError] {
if logits.length() == 0 {
return Err(EmptyLogits)
}
let mut usable = 0
for i in 0.. Result[Array[Double], SamplingError] {
match check_logits(logits) {
Err(e) => return Err(e)
Ok(_) => ()
}
let mut maximum = logits[0]
for value in logits {
if value > maximum {
maximum = value
}
}
let weights : Array[Double] = []
let mut total = 0.0
for value in logits {
let weight = @math.exp(value - maximum)
weights.push(weight)
total = total + weight
}
if !finite(total) || total <= 0.0 {
return Err(NumericalFailure)
}
weights.map(fn(weight) { weight / total }) |> Ok
}
///|
/// Normalize externally supplied non-negative weights. This is useful when a
/// model backend already supplies probabilities rather than logits.
pub fn normalize_weights(
weights : Array[Double],
) -> Result[Array[Double], SamplingError] {
if weights.length() == 0 {
return Err(EmptyLogits)
}
let mut maximum = 0.0
for i in 0.. maximum {
maximum = value
}
}
if maximum == 0.0 {
return Err(NumericalFailure)
}
let scaled : Array[Double] = []
let mut total = 0.0
for weight in weights {
let value = weight / maximum
scaled.push(value)
total = total + value
}
if !finite(total) || total <= 0.0 {
return Err(NumericalFailure)
}
Ok(scaled.map(fn(value) { value / total }))
}
///|
/// Sort token ids by descending probability; ties favor the lower id.
pub fn rank(probabilities : Array[Double]) -> Result[Array[Int], SamplingError] {
if probabilities.length() == 0 {
return Err(EmptyLogits)
}
let order : Array[Int] = []
for i in 0.. probabilities[right] {
-1
} else if probabilities[left] < probabilities[right] {
1
} else {
left.compare(right)
}
})
Ok(order)
}
///|
fn check_order(
probabilities : Array[Double],
order : Array[Int],
) -> Result[Unit, SamplingError] {
if probabilities.length() == 0 || probabilities.length() != order.length() {
return Err(InvalidParameter("order must cover the vocabulary"))
}
let seen = Array::make(order.length(), false)
for token in order {
if token < 0 || token >= order.length() || seen[token] {
return Err(InvalidParameter("order must be a token permutation"))
}
seen[token] = true
if !finite(probabilities[token]) || probabilities[token] < 0.0 {
return Err(InvalidProbability(token))
}
}
Ok(())
}
///|
/// Draw from a ranked prefix using one externally supplied uniform number.
/// The caller controls randomness, making a run exactly replayable.
pub fn draw_prefix(
probabilities : Array[Double],
order : Array[Int],
keep : Int,
uniform : Double,
) -> Result[Int, SamplingError] {
if !finite(uniform) || uniform < 0.0 || uniform >= 1.0 {
return Err(InvalidUniform)
}
if keep <= 0 ||
keep > order.length() ||
probabilities.length() != order.length() {
return Err(InvalidParameter("keep must be between one and vocabulary size"))
}
match check_order(probabilities, order) {
Ok(_) => ()
Err(error) => return Err(error)
}
draw_valid_prefix(probabilities, order, keep, uniform)
}
///|
/// Draw from an arbitrary list of distinct candidate token indices. This
/// accepts a partial vocabulary and retains the supplied candidate order.
pub fn draw_candidates(
probabilities : Array[Double],
candidates : Array[Int],
uniform : Double,
) -> Result[Int, SamplingError] {
if !finite(uniform) || uniform < 0.0 || uniform >= 1.0 {
return Err(InvalidUniform)
}
if probabilities.length() == 0 ||
candidates.length() == 0 ||
candidates.length() > probabilities.length() {
return Err(InvalidParameter("candidate list must be nonempty"))
}
let seen = Array::make(probabilities.length(), false)
for token in candidates {
if token < 0 || token >= probabilities.length() || seen[token] {
return Err(InvalidParameter("invalid or repeated candidate token"))
}
seen[token] = true
if !finite(probabilities[token]) || probabilities[token] < 0.0 {
return Err(InvalidProbability(token))
}
}
draw_valid_prefix(probabilities, candidates, candidates.length(), uniform)
}
///|
fn draw_valid_prefix(
probabilities : Array[Double],
order : Array[Int],
keep : Int,
uniform : Double,
) -> Result[Int, SamplingError] {
let mut total = 0.0
for i in 0.. 0.0 {
last_positive = token
}
cumulative = cumulative + weight
if threshold < cumulative {
return Ok(token)
}
}
Ok(last_positive)
}
///|
/// Information content in bits under the untruncated model distribution.
pub fn surprise(probability : Double) -> Result[Double, SamplingError] {
if !finite(probability) || probability <= 0.0 || probability > 1.0 {
return Err(InvalidProbability(0))
}
Ok(-@math.ln(probability) / @math.ln(2.0))
}