// classifiers.mbt — classification helpers (v0.42.0).
//
// Standalone classification utilities that don't depend on a specific
// model architecture. Useful for any supervised learning pipeline:
// - `argmax` for prediction
// - `softmax` for probability conversion
// - `classification_accuracy` for evaluation
// - `top_k` for top-k retrieval
//
// `dqn.mbt::q_argmax` already implements the same argmax logic for RL;
// we provide a generalised version here so the classifier code path
// stays out of RL modules.
///|
/// Index of the maximum element in `scores`. First occurrence wins on
/// ties. Returns 0 for an empty array.
pub fn argmax(scores : Array[Float]) -> Int {
let n = scores.length()
if n == 0 {
return 0
}
let mut best = 0
let mut best_v = scores[0]
for i in 1.. best_v {
best_v = scores[i]
best = i
}
}
best
}
///|
/// Softmax: `out[i] = exp(scores[i]) / Σ exp(scores[j])`. Numerically
/// stable — subtracts the max score before exponentiating to avoid
/// overflow on large inputs.
pub fn softmax(scores : Array[Float]) -> Array[Float] {
let n = scores.length()
let out : Array[Float] = Array::make(n, 0.0F)
if n == 0 {
return out
}
// Find max for numerical stability.
let mut m = scores[0]
for i in 1.. m {
m = scores[i]
}
}
// Compute exp(scores[i] - m) and accumulate sum.
let mut sum : Float = 0.0F
for i in 0.. 0.0F {
let inv = 1.0F / sum
for i in 0.. Array[Int] {
let n = scores.length()
let take = if k < n { k } else { n }
// Build (score, index) pairs and sort by score descending.
let pairs : Array[(Float, Int)] = []
for i in 0.. 0 {
if pairs[j - 1].0 < key.0 {
pairs[j] = pairs[j - 1]
} else {
break
}
j = j - 1
}
pairs[j] = key
i = i + 1
}
let out : Array[Int] = []
for i in 0.. Float {
let n = predicted.length()
if n == 0 || n != labels.length() {
return 0.0F
}
let mut correct = 0
for i in 0.. Int {
argmax(scores)
}
///|
/// Predict top-k class indices. Convenience wrapper.
pub fn top_k_prediction(scores : Array[Float], k : Int) -> Array[Int] {
top_k(scores, k)
}