///|
/// Stateless top-k sampling is useful as a controlled comparison for Mirostat.
pub fn sample_top_k(
  logits : Array[Double],
  k : Int,
  uniform : Double,
) -> Result[Int, SamplingError] {
  let p = match probabilities(logits) {
    Ok(value) => value
    Err(error) => return Err(error)
  }
  if k < 1 || k > p.length() {
    return Err(InvalidParameter("k must be between one and vocabulary size"))
  }
  let order = match top_indices(p, k) {
    Ok(value) => value
    Err(error) => return Err(error)
  }
  draw_candidates(p, order, uniform)
}

///|
/// Return the smallest top-probability prefix with mass at least threshold.
pub fn top_p_prefix_size(
  probabilities : Array[Double],
  order : Array[Int],
  threshold : Double,
) -> Result[Int, SamplingError] {
  if !finite(threshold) ||
    threshold <= 0.0 ||
    threshold > 1.0 ||
    probabilities.length() == 0 ||
    probabilities.length() != order.length() {
    return Err(InvalidParameter("top-p threshold must be in (0, 1]"))
  }
  match check_order(probabilities, order) {
    Ok(_) => ()
    Err(error) => return Err(error)
  }
  let mut mass = 0.0
  for i in 0..= threshold {
      return Ok(i + 1)
    }
  }
  // Floating-point softmax sums may be just below one.
  Ok(order.length())
}

///|
/// Stateless nucleus sampling on the original model distribution.
pub fn sample_top_p(
  logits : Array[Double],
  threshold : Double,
  uniform : Double,
) -> Result[Int, SamplingError] {
  let p = match probabilities(logits) {
    Ok(value) => value
    Err(error) => return Err(error)
  }
  let order = match rank(p) {
    Ok(value) => value
    Err(error) => return Err(error)
  }
  let keep = match top_p_prefix_size(p, order, threshold) {
    Ok(value) => value
    Err(error) => return Err(error)
  }
  draw_prefix(p, order, keep, uniform)
}

///|
/// Scale logits by positive temperature before another sampling operation.
pub fn with_temperature(
  logits : Array[Double],
  temperature : Double,
) -> Result[Array[Double], SamplingError] {
  match check_logits(logits) {
    Ok(_) => ()
    Err(error) => return Err(error)
  }
  if !finite(temperature) || temperature <= 0.0 {
    return Err(InvalidParameter("temperature must be finite and positive"))
  }
  // Subtracting the common maximum leaves softmax unchanged and avoids
  // overflow when a very small temperature scales large model logits.
  let mut maximum = logits[0]
  for value in logits {
    if value > maximum {
      maximum = value
    }
  }
  let scaled : Array[Double] = []
  for value in logits {
    let result = (value - maximum) / temperature
    if !finite(result) && result != -1.0 / 0.0 {
      return Err(NumericalFailure)
    }
    scaled.push(result)
  }
  Ok(scaled)
}

///|
/// Independent temperature sampling without truncation.
pub fn sample_temperature(
  logits : Array[Double],
  temperature : Double,
  uniform : Double,
) -> Result[Int, SamplingError] {
  let scaled = match with_temperature(logits, temperature) {
    Ok(value) => value
    Err(error) => return Err(error)
  }
  sample_top_k(scaled, scaled.length(), uniform)
}