///|
/// The current adaptive candidate distribution for an external categorical
/// sampler. Arrays align by index: `tokens[i]` has `conditional[i]` probability
/// after truncation and `original[i]` probability before truncation.
pub struct CandidateSet {
tokens : Array[Int]
conditional : Array[Double]
original : Array[Double]
mu : Double
} derive(Debug)
///|
pub fn CandidateSet::tokens(self : CandidateSet) -> Array[Int] {
self.tokens.copy()
}
///|
pub fn CandidateSet::conditional(self : CandidateSet) -> Array[Double] {
self.conditional.copy()
}
///|
pub fn CandidateSet::original(self : CandidateSet) -> Array[Double] {
self.original.copy()
}
///|
pub fn CandidateSet::mu(self : CandidateSet) -> Double {
self.mu
}
///|
/// Return the candidate-list position for a vocabulary token. This position
/// indexes the arrays returned by `conditional()` and `original()`.
pub fn CandidateSet::position(self : CandidateSet, token : Int) -> Int? {
for index in 0.. Double? {
match self.position(token) {
Some(index) => Some(self.conditional[index])
None => None
}
}
///|
/// Probability under the original model distribution, or None if this token
/// is not in the current candidate set.
pub fn CandidateSet::original_probability(
self : CandidateSet,
token : Int,
) -> Double? {
match self.position(token) {
Some(index) => Some(self.original[index])
None => None
}
}
///|
pub fn Sampler::candidates(
self : Sampler,
logits : Array[Double],
) -> Result[CandidateSet, SamplingError] {
let p = match probabilities(logits) {
Ok(value) => value
Err(error) => return Err(error)
}
self.candidates_distribution(p)
}
///|
pub fn Sampler::candidates_weights(
self : Sampler,
weights : Array[Double],
) -> Result[CandidateSet, SamplingError] {
let p = match normalize_weights(weights) {
Ok(value) => value
Err(error) => return Err(error)
}
self.candidates_distribution(p)
}
///|
fn Sampler::candidates_distribution(
self : Sampler,
p : Array[Double],
) -> Result[CandidateSet, SamplingError] {
let tokens = match self.select_candidates(p) {
Ok(value) => value
Err(error) => return Err(error)
}
let original : Array[Double] = []
let mut mass = 0.0
for token in tokens {
let value = p[token]
original.push(value)
mass = mass + value
}
if !finite(mass) || mass <= 0.0 {
return Err(NumericalFailure)
}
let conditional = original.map(fn(value) { value / mass })
Ok({ tokens, conditional, original, mu: self.mu, })
}