#[inline]
pub fn row_max_and_any_nan(row: &[f32]) -> (f32, bool) {
let mut max_val = f32::NEG_INFINITY;
let mut any_nan = false;
for &v in row {
if v.is_nan() {
any_nan = true;
} else {
max_val = max_val.max(v);
}
}
(max_val, any_nan)
}
#[inline]
pub fn row_fails_closed_pre_exp(max_val: f32, any_nan: bool) -> bool {
any_nan || !max_val.is_finite()
}
#[inline]
pub fn is_masked_neg_inf(v: f32) -> bool {
v == f32::NEG_INFINITY
}
#[inline]
pub fn finalize_row(row: &mut [f32], sum: f32) {
if sum.is_finite() && sum > 0.0 {
let inv = 1.0 / sum;
for v in row.iter_mut() {
*v *= inv;
}
} else {
row.fill(0.0);
}
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn row_max_and_any_nan_ignores_nan_in_max_but_flags_it() {
let row = [1.0f32, f32::NAN, 3.0, 2.0];
let (max_val, any_nan) = row_max_and_any_nan(&row);
assert_eq!(max_val, 3.0);
assert!(any_nan);
}
#[test]
fn row_max_and_any_nan_all_neg_inf() {
let row = [f32::NEG_INFINITY; 4];
let (max_val, any_nan) = row_max_and_any_nan(&row);
assert_eq!(max_val, f32::NEG_INFINITY);
assert!(!any_nan);
}
#[test]
fn row_fails_closed_pre_exp_nan() {
assert!(row_fails_closed_pre_exp(3.0, true));
}
#[test]
fn row_fails_closed_pre_exp_pos_inf_max() {
assert!(row_fails_closed_pre_exp(f32::INFINITY, false));
}
#[test]
fn row_fails_closed_pre_exp_neg_inf_max() {
assert!(row_fails_closed_pre_exp(f32::NEG_INFINITY, false));
}
#[test]
fn row_fails_closed_pre_exp_finite_ok() {
assert!(!row_fails_closed_pre_exp(5.0, false));
}
#[test]
fn is_masked_neg_inf_exact() {
assert!(is_masked_neg_inf(f32::NEG_INFINITY));
assert!(!is_masked_neg_inf(-1.0e30));
assert!(!is_masked_neg_inf(f32::NAN));
}
#[test]
fn finalize_row_normalizes_well_formed_sum() {
let mut row = [1.0f32, 1.0, 2.0];
finalize_row(&mut row, 4.0);
assert_eq!(row, [0.25, 0.25, 0.5]);
}
#[test]
fn finalize_row_zeros_on_nan_sum() {
let mut row = [1.0f32, f32::NAN, 2.0];
finalize_row(&mut row, f32::NAN);
assert_eq!(row, [0.0, 0.0, 0.0]);
}
#[test]
fn finalize_row_zeros_on_zero_sum() {
let mut row = [0.0f32, 0.0, 0.0];
finalize_row(&mut row, 0.0);
assert_eq!(row, [0.0, 0.0, 0.0]);
}
#[test]
fn finalize_row_zeros_on_negative_sum() {
let mut row = [1.0f32, 2.0];
finalize_row(&mut row, -1.0);
assert_eq!(row, [0.0, 0.0]);
}
#[test]
fn finalize_row_zeros_on_infinite_sum() {
let mut row = [1.0f32, 2.0];
finalize_row(&mut row, f32::INFINITY);
assert_eq!(row, [0.0, 0.0]);
}
}