///|
/// Mirostat feedback configuration. Tau and mu are measured in bits.
pub struct Config {
tau : Double
eta : Double
initial_mu : Double
m : Int
} derive(Eq, Debug)
///|
pub fn Config::new(
tau : Double,
eta? : Double = 0.1,
initial_mu? : Double = 2.0 * tau,
m? : Int = 100,
) -> Result[Config, SamplingError] {
if !finite(tau) || tau <= 0.0 {
return Err(InvalidParameter("tau must be finite and positive"))
}
if !finite(eta) || eta <= 0.0 {
return Err(InvalidParameter("eta must be finite and positive"))
}
if !finite(initial_mu) || initial_mu < 0.0 {
return Err(InvalidParameter("initial_mu must be finite and non-negative"))
}
if m < 2 {
return Err(InvalidParameter("m must be at least two"))
}
Ok({ tau, eta, initial_mu, m, })
}
///|
pub fn Config::tau(self : Config) -> Double {
self.tau
}
///|
pub fn Config::eta(self : Config) -> Double {
self.eta
}
///|
pub fn Config::initial_mu(self : Config) -> Double {
self.initial_mu
}
///|
pub fn Config::m(self : Config) -> Int {
self.m
}
///|
pub enum Version {
V1
V2
} derive(Eq, Debug)
///|
pub fn Version::v1() -> Version {
V1
}
///|
pub fn Version::v2() -> Version {
V2
}
///|
pub struct Step {
token : Int
original_probability : Double
observed_surprise : Double
kept_tokens : Int
mu_before : Double
mu_after : Double
} derive(Eq, Debug)
///|
pub fn Step::token(self : Step) -> Int {
self.token
}
///|
pub fn Step::original_probability(self : Step) -> Double {
self.original_probability
}
///|
pub fn Step::observed_surprise(self : Step) -> Double {
self.observed_surprise
}
///|
pub fn Step::kept_tokens(self : Step) -> Int {
self.kept_tokens
}
///|
pub fn Step::mu_before(self : Step) -> Double {
self.mu_before
}
///|
pub fn Step::mu_after(self : Step) -> Double {
self.mu_after
}
///|
/// A session owns only feedback state; the model owns logits and the caller owns RNG.
pub struct Sampler {
config : Config
version : Version
mut mu : Double
mut steps : Int
} derive(Debug)
///|
pub fn Sampler::new(config : Config, version? : Version = V1) -> Sampler {
{ config, version, mu: config.initial_mu, steps: 0, }
}
///|
pub fn Sampler::mu(self : Sampler) -> Double {
self.mu
}
///|
pub fn Sampler::steps(self : Sampler) -> Int {
self.steps
}
///|
pub fn Sampler::version(self : Sampler) -> Version {
self.version
}
///|
pub fn Sampler::reset(self : Sampler) -> Unit {
self.mu = self.config.initial_mu
self.steps = 0
}
///|
fn feedback(mu : Double, observed : Double, config : Config) -> Double {
mu - config.eta * (observed - config.tau)
}