///|
/// Estimate the Zipf exponent from adjacent probability ratios in the top m.
/// This follows the paper's least-squares estimator and uses log base two.
pub fn zipf_exponent(
  probabilities : Array[Double],
  order : Array[Int],
  m : Int,
) -> Result[Double, SamplingError] {
  if m < 2 || order.length() != probabilities.length() || order.length() < 2 {
    return Err(InvalidParameter("Zipf estimation needs two ranked tokens"))
  }
  match check_order(probabilities, order) {
    Ok(_) => ()
    Err(error) => return Err(error)
  }
  estimate_zipf_prefix(probabilities, order, m)
}

///|
fn estimate_zipf_prefix(
  probabilities : Array[Double],
  order : Array[Int],
  m : Int,
) -> Result[Double, SamplingError] {
  let pairs = if m < order.length() { m - 1 } else { order.length() - 1 }
  let mut numerator = 0.0
  let mut denominator = 0.0
  for i in 0.. Result[Int, SamplingError] {
  if !finite(mu) || !finite(exponent) || exponent <= 0.0 || vocabulary <= 0 {
    return Err(InvalidParameter("invalid Mirostat k inputs"))
  }
  if vocabulary == 1 {
    return Ok(1)
  }
  let epsilon = exponent - 1.0
  let denominator = if epsilon > -0.000001 && epsilon < 0.000001 {
    @math.ln(vocabulary.to_double())
  } else {
    (1.0 - @math.exp(-epsilon * @math.ln(vocabulary.to_double()))) / epsilon
  }
  if !finite(denominator) || denominator <= 0.0 {
    return Err(NumericalFailure)
  }
  let log_k = (mu * @math.ln(2.0) - @math.ln(denominator)) / exponent
  if !finite(log_k) {
    return Err(NumericalFailure)
  }
  if log_k <= 0.0 {
    return Ok(1)
  }
  if log_k >= @math.ln(vocabulary.to_double()) {
    return Ok(vocabulary)
  }
  let raw = @math.exp(log_k)
  let rounded = (raw + 0.5).to_int()
  Ok(
    if rounded < 1 {
      1
    } else if rounded > vocabulary {
      vocabulary
    } else {
      rounded
    },
  )
}

///|
/// V2 uses the current surprise target as a cutoff on original token
/// probabilities. It always keeps at least the most likely token.
pub fn v2_prefix_size(
  probabilities : Array[Double],
  order : Array[Int],
  mu : Double,
) -> Result[Int, SamplingError] {
  if !finite(mu) ||
    order.length() == 0 ||
    order.length() != probabilities.length() {
    return Err(InvalidParameter("invalid V2 prefix inputs"))
  }
  match check_order(probabilities, order) {
    Ok(_) => ()
    Err(error) => return Err(error)
  }
  let mut kept = 0
  for token in order {
    let p = probabilities[token]
    if p <= 0.0 {
      break
    }
    let information = match surprise(p) {
      Ok(value) => value
      Err(e) => return Err(e)
    }
    if information <= mu {
      kept = kept + 1
    } else {
      break
    }
  }
  Ok(if kept == 0 { 1 } else { kept })
}

///|
/// One complete Mirostat step. State changes only after every calculation
/// succeeds, so malformed logits cannot corrupt an ongoing generation.
pub fn Sampler::sample(
  self : Sampler,
  logits : Array[Double],
  uniform : Double,
) -> Result[Step, SamplingError] {
  let p = match probabilities(logits) {
    Ok(value) => value
    Err(e) => return Err(e)
  }
  self.sample_distribution(p, uniform)
}

///|
/// Sample from an already computed model distribution, supplied as weights.
/// The weights are normalized internally and need not sum to one.
pub fn Sampler::sample_weights(
  self : Sampler,
  weights : Array[Double],
  uniform : Double,
) -> Result[Step, SamplingError] {
  let p = match normalize_weights(weights) {
    Ok(value) => value
    Err(error) => return Err(error)
  }
  self.sample_distribution(p, uniform)
}

///|
fn Sampler::sample_distribution(
  self : Sampler,
  p : Array[Double],
  uniform : Double,
) -> Result[Step, SamplingError] {
  let order = match self.select_candidates(p) {
    Ok(value) => value
    Err(e) => return Err(e)
  }
  let token = match draw_candidates(p, order, uniform) {
    Ok(value) => value
    Err(e) => return Err(e)
  }
  self.commit_step(p, token, order.length())
}

///|
fn Sampler::commit_step(
  self : Sampler,
  p : Array[Double],
  token : Int,
  keep : Int,
) -> Result[Step, SamplingError] {
  let observed = match surprise(p[token]) {
    Ok(value) => value
    Err(e) => return Err(e)
  }
  let before = self.mu
  let after = feedback(before, observed, self.config)
  if !finite(after) {
    return Err(NumericalFailure)
  }
  self.mu = after
  self.steps = self.steps + 1
  Ok({
    token,
    original_probability: p[token],
    observed_surprise: observed,
    kept_tokens: keep,
    mu_before: before,
    mu_after: after,
  })
}

///|
fn Sampler::select_candidates(
  self : Sampler,
  p : Array[Double],
) -> Result[Array[Int], SamplingError] {
  match self.version {
    V1 =>
      if p.length() == 1 {
        Ok([0])
      } else {
        let m = if self.config.m < p.length() {
          self.config.m
        } else {
          p.length()
        }
        let top_m = match top_indices(p, m) {
          Ok(value) => value
          Err(error) => return Err(error)
        }
        let s = match estimate_zipf_prefix(p, top_m, m) {
          Ok(value) => value
          Err(e) => return Err(e)
        }
        let k = match estimated_k(self.mu, s, p.length()) {
          Ok(value) => value
          Err(error) => return Err(error)
        }
        top_indices(p, k)
      }
    V2 => {
      let candidates : Array[Int] = []
      for token in 0.. 0.0 {
          let information = match surprise(p[token]) {
            Ok(value) => value
            Err(error) => return Err(error)
          }
          if information <= self.mu {
            candidates.push(token)
          }
        }
      }
      if candidates.length() == 0 {
        return top_indices(p, 1)
      }
      candidates.sort_by(fn(left, right) {
        if better_token(left, right, p) {
          -1
        } else if better_token(right, left, p) {
          1
        } else {
          0
        }
      })
      Ok(candidates)
    }
  }
}