mod scalar;
const LOG_LOSS_EPSILON: f64 = 1e-15;
const BINARY_LOG_LOSS_EPSILON: f64 = 1e-16f32 as f64;
const MIN_POSITIVE_PREDICTION: f64 = 1e-8;
use crate::objective::GradPair;
#[cfg(any(target_arch = "aarch64", target_arch = "x86_64"))]
use std::sync::LazyLock;
#[cfg(target_arch = "aarch64")]
mod aarch64;
#[cfg(target_arch = "x86_64")]
mod x86_64;
#[cfg(any(target_arch = "aarch64", target_arch = "x86_64"))]
const MIN_SIMD_LEN: usize = 16;
#[cfg(any(target_arch = "aarch64", target_arch = "x86_64"))]
const _: () = assert!(MIN_SIMD_LEN <= crate::objective::GRADIENT_BLOCK_ROWS);
const _: () = assert!(std::mem::size_of::<GradPair>() == 2 * std::mem::size_of::<f32>());
#[cfg(any(target_arch = "aarch64", target_arch = "x86_64"))]
const MAX_FAST_EXP_INPUT: f32 = 80.0;
#[cfg(target_arch = "aarch64")]
static NEON_AVAILABLE: LazyLock<bool> =
LazyLock::new(|| std::arch::is_aarch64_feature_detected!("neon"));
#[cfg(target_arch = "x86_64")]
static AVX2_FMA_AVAILABLE: LazyLock<bool> = LazyLock::new(|| {
std::arch::is_x86_feature_detected!("avx2") && std::arch::is_x86_feature_detected!("fma")
});
macro_rules! dispatch_unary_inplace {
($values:expr, $kernel:ident) => {
#[cfg(target_arch = "aarch64")]
if $values.len() >= MIN_SIMD_LEN && neon_available() {
unsafe { aarch64::$kernel($values) };
return;
}
#[cfg(target_arch = "x86_64")]
if $values.len() >= MIN_SIMD_LEN && avx2_fma_available() {
unsafe { x86_64::$kernel($values) };
return;
}
};
}
macro_rules! dispatch_gradient {
($gate:expr, $neon_call:expr) => {
#[cfg(target_arch = "aarch64")]
if $gate && neon_available() {
return unsafe { $neon_call };
}
};
($gate:expr, $neon_call:expr, $avx_call:expr) => {
dispatch_gradient!($gate, $neon_call);
#[cfg(target_arch = "x86_64")]
if $gate && avx2_fma_available() {
unsafe { $avx_call };
return;
}
};
}
macro_rules! dispatch_softmax {
($num_class:expr, $eligible:expr, wide: $wide:expr, short: |$k:ident| $short:expr) => {
#[cfg(target_arch = "aarch64")]
if $eligible && neon_available() {
if $num_class >= 8 {
unsafe { $wide };
return;
}
if (2..=4).contains(&$num_class) {
unsafe {
match $num_class {
2 => {
use aarch64 as arch;
const $k: usize = 2;
$short
}
3 => {
use aarch64 as arch;
const $k: usize = 3;
$short
}
_ => {
use aarch64 as arch;
const $k: usize = 4;
$short
}
}
}
return;
}
}
#[cfg(target_arch = "x86_64")]
if ($num_class == 2 || $num_class == 4) && $eligible && avx2_fma_available() {
unsafe {
if $num_class == 2 {
use x86_64 as arch;
const $k: usize = 2;
$short
} else {
use x86_64 as arch;
const $k: usize = 4;
$short
}
}
return;
}
};
}
#[cfg(target_arch = "aarch64")]
#[inline]
fn neon_available() -> bool {
cfg!(target_feature = "neon") || *NEON_AVAILABLE
}
#[cfg(target_arch = "x86_64")]
#[inline]
fn avx2_fma_available() -> bool {
*AVX2_FMA_AVAILABLE
}
#[inline(always)]
pub(crate) fn prefetch_read<T>(value: &T) {
#[cfg(target_arch = "aarch64")]
unsafe {
std::arch::asm!(
"prfm pldl1keep, [{ptr}]",
ptr = in(reg) std::ptr::from_ref::<T>(value),
options(nostack, readonly, preserves_flags)
);
}
#[cfg(target_arch = "x86_64")]
unsafe {
std::arch::x86_64::_mm_prefetch::<{ std::arch::x86_64::_MM_HINT_T0 }>(
std::ptr::from_ref::<T>(value).cast::<i8>(),
);
}
#[cfg(not(any(target_arch = "aarch64", target_arch = "x86_64")))]
let _ = value;
}
#[inline(always)]
pub(crate) fn step_if_greater(base: usize, a: u32, b: u32) -> usize {
#[cfg(target_arch = "aarch64")]
{
let out: usize;
unsafe {
std::arch::asm!(
"cmp {a:w}, {b:w}",
"cinc {out}, {base}, hi",
a = in(reg) a,
b = in(reg) b,
base = in(reg) base,
out = lateout(reg) out,
options(pure, nomem, nostack)
);
}
out
}
#[cfg(not(target_arch = "aarch64"))]
{
base + usize::from(a > b)
}
}
#[inline]
pub(crate) fn count_le(cuts: &[f32], value: f32) -> usize {
#[cfg(target_arch = "aarch64")]
if cuts.len() == 16 && neon_available() {
return unsafe { aarch64::count_le_16(cuts, value) };
}
#[cfg(target_arch = "x86_64")]
if cuts.len() == 16 {
return unsafe { x86_64::count_le_16(cuts, value) };
}
cuts.iter().filter(|&&cut| cut <= value).count()
}
#[inline(always)]
pub(crate) fn shap_edge_terms(alpha: f32, h: &[f32; 8], u: &[f32; 8]) -> [f32; 8] {
#[cfg(target_arch = "aarch64")]
if neon_available() {
return unsafe { aarch64::edge_terms_8(alpha, h, u) };
}
std::array::from_fn(|i| alpha * h[i] / shap_denominator(alpha, u[i]))
}
#[inline(always)]
pub(crate) fn shap_scaled_basis(alpha: f32, c: &[f32; 8], u: &[f32; 8]) -> [f32; 8] {
#[cfg(target_arch = "aarch64")]
if neon_available() {
return unsafe { aarch64::scaled_basis_8(alpha, c, u) };
}
std::array::from_fn(|i| c[i] * shap_denominator(alpha, u[i]))
}
#[inline(always)]
pub(crate) fn shap_divided_basis(alpha: f32, c: &[f32; 8], u: &[f32; 8]) -> Option<[f32; 8]> {
#[cfg(target_arch = "aarch64")]
if neon_available() {
return unsafe { aarch64::divided_basis_8(alpha, c, u) };
}
let old: [f32; 8] = std::array::from_fn(|i| shap_denominator(alpha, u[i]));
c.iter()
.zip(&old)
.all(|(c, o)| c.is_finite() && o.is_finite())
.then(|| std::array::from_fn(|i| c[i] / old[i]))
}
#[inline(always)]
fn shap_denominator(alpha: f32, u: f32) -> f32 {
if cfg!(target_arch = "aarch64") {
alpha.mul_add(u, 1.0)
} else {
alpha * u + 1.0
}
}
#[inline]
#[cfg(any(target_arch = "aarch64", target_arch = "x86_64"))]
fn gradient_gate(preds: &[f32], labels: &[f32], weights: Option<&[f32]>, out: &[GradPair]) -> bool {
metric_gate(preds, labels, weights.map(RowWeights::from)) && out.len() >= preds.len()
}
#[inline]
#[cfg(any(target_arch = "aarch64", target_arch = "x86_64"))]
fn metric_gate(preds: &[f32], labels: &[f32], weights: Option<RowWeights<'_>>) -> bool {
let len = preds.len();
len >= MIN_SIMD_LEN
&& labels.len() >= len
&& weights.is_none_or(|weights| weights.cells().is_some_and(|cells| cells >= len))
}
#[derive(Debug, Clone, Copy)]
pub(crate) struct RowWeights<'a> {
values: &'a [f32],
stride: usize,
}
impl<'a> RowWeights<'a> {
pub(crate) fn new(values: &'a [f32], stride: usize) -> Self {
debug_assert!(stride > 0, "row weight stride must be positive");
RowWeights { values, stride }
}
#[inline]
pub(crate) fn get(self, cell: usize) -> f32 {
if self.stride == 1 {
self.values[cell]
} else {
self.values[cell / self.stride]
}
}
#[inline]
pub(crate) fn cells(self) -> Option<usize> {
self.values.len().checked_mul(self.stride)
}
}
impl<'a> From<&'a [f32]> for RowWeights<'a> {
fn from(values: &'a [f32]) -> Self {
Self::new(values, 1)
}
}
#[inline]
fn class_rows_cover(
preds: &[f32],
labels: &[f32],
weights: Option<&[f32]>,
num_class: usize,
) -> bool {
labels
.len()
.checked_mul(num_class)
.is_some_and(|len| preds.len() >= len)
&& weights.is_none_or(|values| values.len() >= labels.len())
}
#[inline]
pub(crate) fn sigmoid_scalar(x: f32) -> f32 {
1.0 / ((-x).min(88.7).exp() + 1.0)
}
#[inline]
pub(crate) fn exp_inplace(values: &mut [f32]) {
dispatch_unary_inplace!(values, exp_inplace);
for value in values.iter_mut() {
*value = value.exp();
}
}
#[inline]
pub(crate) fn sigmoid_inplace(values: &mut [f32]) {
dispatch_unary_inplace!(values, sigmoid_inplace);
for value in values.iter_mut() {
*value = sigmoid_scalar(*value);
}
}
pub(crate) fn logistic_gradient(
preds: &[f32],
labels: &[f32],
weights: Option<&[f32]>,
scale_pos_weight: f32,
min_hess: f32,
out: &mut [GradPair],
) {
dispatch_gradient!(
gradient_gate(preds, labels, weights, out),
aarch64::logistic_gradient(preds, labels, weights, scale_pos_weight, min_hess, out),
x86_64::logistic_gradient(preds, labels, weights, scale_pos_weight, min_hess, out)
);
scalar::logistic_gradient(
preds,
labels,
weights,
scale_pos_weight,
min_hess,
out,
0..preds.len(),
);
}
pub(crate) fn poisson_gradient(
preds: &[f32],
labels: &[f32],
weights: Option<&[f32]>,
max_delta_step: f32,
out: &mut [GradPair],
) {
dispatch_gradient!(
gradient_gate(preds, labels, weights, out),
aarch64::poisson_gradient(preds, labels, weights, max_delta_step, out)
);
scalar::poisson_gradient(preds, labels, weights, max_delta_step, out, 0..preds.len());
}
pub(crate) fn gamma_gradient(
preds: &[f32],
labels: &[f32],
weights: Option<&[f32]>,
scale_pos_weight: f32,
out: &mut [GradPair],
) {
dispatch_gradient!(
gradient_gate(preds, labels, weights, out),
aarch64::gamma_gradient(preds, labels, weights, scale_pos_weight, out)
);
scalar::gamma_gradient(
preds,
labels,
weights,
scale_pos_weight,
out,
0..preds.len(),
);
}
pub(crate) fn tweedie_gradient(
preds: &[f32],
labels: &[f32],
weights: Option<&[f32]>,
rho: f32,
out: &mut [GradPair],
) {
dispatch_gradient!(
gradient_gate(preds, labels, weights, out),
aarch64::tweedie_gradient(preds, labels, weights, rho, out)
);
scalar::tweedie_gradient(preds, labels, weights, rho, out, 0..preds.len());
}
pub(crate) fn softmax_rows_inplace(values: &mut [f32], num_class: usize) {
dispatch_softmax!(
num_class,
values.len() >= MIN_SIMD_LEN,
wide: aarch64::softmax_rows_inplace(values, num_class),
short: |K| arch::short_softmax_rows::<K>(values)
);
for row in values.chunks_mut(num_class) {
softmax_scalar(row);
}
}
pub(crate) fn softmax_gradient(
preds: &[f32],
labels: &[f32],
weights: Option<&[f32]>,
num_class: usize,
min_hess: f32,
out: &mut [GradPair],
) {
let complete = labels
.len()
.checked_mul(num_class)
.is_some_and(|len| len == preds.len() && out.len() >= len)
&& weights.is_none_or(|values| values.len() >= labels.len());
dispatch_softmax!(
num_class,
preds.len() >= MIN_SIMD_LEN && complete,
wide: aarch64::softmax_gradient(preds, labels, weights, num_class, min_hess, out),
short: |K| arch::short_softmax_gradient::<K>(preds, labels, weights, min_hess, out)
);
debug_assert!(complete);
softmax_gradient_rows_scalar(
preds,
labels,
weights,
min_hess,
out,
0..labels.len(),
num_class,
);
}
pub(super) fn softmax_gradient_rows_scalar(
preds: &[f32],
labels: &[f32],
weights: Option<&[f32]>,
min_hess: f32,
out: &mut [GradPair],
rows: std::ops::Range<usize>,
k: usize,
) {
for current in rows {
let base = current * k;
softmax_gradient_row_scalar(
&preds[base..base + k],
labels[current] as usize,
weights.map_or(1.0, |values| values[current]),
min_hess,
&mut out[base..base + k],
);
}
}
pub(super) fn softmax_scalar(values: &mut [f32]) {
let Some(&first) = values.first() else {
return;
};
let wmax = values[1..].iter().fold(first, |m, &v| v.max(m));
let mut wsum = 0f64;
for value in values.iter_mut() {
*value = (*value - wmax).exp();
wsum += f64::from(*value);
}
let wsum = wsum as f32;
for value in values.iter_mut() {
*value /= wsum;
}
}
pub(super) fn softmax_gradient_row_scalar(
preds: &[f32],
label: usize,
weight: f32,
min_hess: f32,
out: &mut [GradPair],
) {
let mut wmax = f32::MIN_POSITIVE;
for &value in preds {
wmax = value.max(wmax);
}
let mut wsum = 0f64;
for &prediction in preds {
wsum += f64::from((prediction - wmax).exp());
}
let wsum = wsum as f32;
for (class, (output, &prediction)) in out.iter_mut().zip(preds).enumerate() {
let p = (prediction - wmax).exp() / wsum;
let h = (2.0 * p * (1.0 - p) * weight).max(min_hess);
let g = if class == label { p - 1.0 } else { p };
*output = GradPair::new(g * weight, h);
}
}
pub(crate) fn squared_error_sum(
preds: &[f32],
labels: &[f32],
weights: Option<RowWeights<'_>>,
) -> (f64, f64) {
distance_sum::<true>(preds, labels, weights)
}
pub(crate) fn absolute_error_sum(
preds: &[f32],
labels: &[f32],
weights: Option<RowWeights<'_>>,
) -> (f64, f64) {
distance_sum::<false>(preds, labels, weights)
}
fn distance_sum<const SQUARED: bool>(
preds: &[f32],
labels: &[f32],
weights: Option<RowWeights<'_>>,
) -> (f64, f64) {
dispatch_gradient!(
metric_gate(preds, labels, weights),
aarch64::distance_sum::<SQUARED>(preds, labels, weights)
);
scalar::distance_sum::<SQUARED>(preds, labels, weights, 0..preds.len())
}
pub(crate) fn classification_error_sum(
preds: &[f32],
labels: &[f32],
weights: Option<RowWeights<'_>>,
) -> (f64, f64) {
dispatch_gradient!(
metric_gate(preds, labels, weights),
aarch64::classification_error_sum(preds, labels, weights)
);
scalar::classification_error_sum(preds, labels, weights, 0..preds.len())
}
pub(crate) fn log_loss_sum(
preds: &[f32],
labels: &[f32],
weights: Option<RowWeights<'_>>,
) -> (f64, f64) {
dispatch_gradient!(
metric_gate(preds, labels, weights),
aarch64::log_loss_sum(preds, labels, weights)
);
scalar::log_loss(preds, labels, weights, 0..preds.len())
}
pub(crate) fn positive_nloglik_sum<const GAMMA: bool>(
preds: &[f32],
labels: &[f32],
weights: Option<RowWeights<'_>>,
) -> (f64, f64) {
dispatch_gradient!(
metric_gate(preds, labels, weights),
aarch64::positive_nloglik_sum::<GAMMA>(preds, labels, weights)
);
scalar::positive_nloglik::<GAMMA>(preds, labels, weights, 0..preds.len())
}
pub(crate) fn tweedie_nloglik_sum(
preds: &[f32],
labels: &[f32],
weights: Option<RowWeights<'_>>,
rho: f64,
) -> (f64, f64) {
dispatch_gradient!(
rho.is_finite() && rho > 1.0 && rho < 2.0 && metric_gate(preds, labels, weights),
aarch64::tweedie_nloglik_sum(preds, labels, weights, rho)
);
scalar::tweedie_nloglik(preds, labels, weights, rho, 0..preds.len())
}
pub(crate) fn multiclass_log_loss_sum(
preds: &[f32],
labels: &[f32],
weights: Option<&[f32]>,
num_class: usize,
) -> (f64, f64) {
let complete = class_rows_cover(preds, labels, weights, num_class);
dispatch_gradient!(
labels.len() >= MIN_SIMD_LEN && complete,
aarch64::multiclass_log_loss_sum(preds, labels, weights, num_class)
);
debug_assert!(complete);
scalar::multiclass_log_loss(preds, labels, weights, num_class, 0..labels.len())
}
pub(crate) fn multiclass_error_sum(
preds: &[f32],
labels: &[f32],
weights: Option<&[f32]>,
num_class: usize,
) -> (f64, f64) {
let complete = class_rows_cover(preds, labels, weights, num_class);
dispatch_gradient!(
num_class >= 8
&& u32::try_from(num_class).is_ok()
&& labels.len() >= MIN_SIMD_LEN
&& complete,
aarch64::multiclass_error_sum(preds, labels, weights, num_class)
);
debug_assert!(complete);
multiclass_error_sum_rows(preds, labels, weights, num_class, argmax_scalar)
}
pub(super) fn multiclass_error_sum_rows(
preds: &[f32],
labels: &[f32],
weights: Option<&[f32]>,
num_class: usize,
argmax: impl Fn(&[f32]) -> usize,
) -> (f64, f64) {
let mut wrong = 0.0;
let mut weight_sum = 0.0;
for (row_index, &label) in labels.iter().enumerate() {
let weight = weights.map_or(1.0, |values| f64::from(values[row_index]));
let row = &preds[row_index * num_class..(row_index + 1) * num_class];
let best = argmax(row);
if best != label as usize {
wrong += weight;
}
weight_sum += weight;
}
(wrong, weight_sum)
}
pub(crate) fn argmax_scalar(values: &[f32]) -> usize {
let mut best = 0;
for index in 1..values.len() {
if values[index] > values[best] {
best = index;
}
}
best
}
#[cfg(test)]
mod tests;