use super::{MAX_FAST_EXP_INPUT, scalar, sigmoid_scalar};
use crate::objective::GradPair;
#[allow(
clippy::wildcard_imports,
reason = "intrinsic modules are used wholesale"
)]
use std::arch::x86_64::*;
const WIDTH: usize = 8;
const _: () = assert!(crate::objective::GRADIENT_BLOCK_ROWS.is_multiple_of(WIDTH));
#[inline]
#[target_feature(enable = "avx2,fma")]
unsafe fn exp_f32(value: __m256) -> __m256 {
unsafe {
let multiply: unsafe fn(__m256, __m256) -> __m256 = _mm256_mul_ps;
let exponent =
_mm256_cvtps_epi32(multiply(value, _mm256_set1_ps(std::f32::consts::LOG2_E)));
let exponent_f32 = _mm256_cvtepi32_ps(exponent);
let reduced = _mm256_fnmadd_ps(exponent_f32, _mm256_set1_ps(0.693_359_4), value);
let reduced = _mm256_fmadd_ps(exponent_f32, _mm256_set1_ps(2.121_944_4e-4), reduced);
let c = |x: f32| _mm256_set1_ps(x);
let polynomial = {
let squared = multiply(reduced, reduced);
let fourth = multiply(squared, squared);
let pair_0 = _mm256_add_ps(c(1.0), reduced);
let pair_1 = _mm256_fmadd_ps(c(1.0 / 6.0), reduced, c(0.5));
let pair_2 = _mm256_fmadd_ps(c(1.0 / 120.0), reduced, c(1.0 / 24.0));
let pair_3 = _mm256_fmadd_ps(c(1.0 / 5_040.0), reduced, c(1.0 / 720.0));
let low = _mm256_fmadd_ps(pair_1, squared, pair_0);
let high = _mm256_fmadd_ps(pair_3, squared, pair_2);
_mm256_fmadd_ps(high, fourth, low)
};
let exponent_bits =
_mm256_slli_epi32::<23>(_mm256_add_epi32(exponent, _mm256_set1_epi32(127)));
multiply(polynomial, _mm256_castsi256_ps(exponent_bits))
}
}
#[inline]
#[target_feature(enable = "avx2,fma")]
unsafe fn abs_f32(value: __m256) -> __m256 {
unsafe {
let and_not: unsafe fn(__m256, __m256) -> __m256 = _mm256_andnot_ps;
and_not(_mm256_set1_ps(-0.0), value)
}
}
#[inline]
#[target_feature(enable = "avx2,fma")]
unsafe fn sigmoid_f32(value: __m256) -> __m256 {
unsafe {
let one = _mm256_set1_ps(1.0);
let exp = exp_f32(_mm256_sub_ps(_mm256_setzero_ps(), abs_f32(value)));
let denominator = _mm256_add_ps(one, exp);
let positive = _mm256_div_ps(one, denominator);
let negative = _mm256_div_ps(exp, denominator);
let non_negative = _mm256_cmp_ps::<_CMP_GE_OQ>(value, _mm256_setzero_ps());
_mm256_blendv_ps(negative, positive, non_negative)
}
}
#[inline]
#[target_feature(enable = "avx2,fma")]
unsafe fn regular_input(value: __m256) -> bool {
unsafe {
let in_range =
_mm256_cmp_ps::<_CMP_LE_OQ>(abs_f32(value), _mm256_set1_ps(MAX_FAST_EXP_INPUT));
_mm256_movemask_ps(in_range) == 0xFF
}
}
#[inline]
#[target_feature(enable = "avx2,fma")]
unsafe fn store_pairs(dest: *mut GradPair, grad: __m256, hess: __m256) {
unsafe {
let low = _mm256_unpacklo_ps(grad, hess); let high = _mm256_unpackhi_ps(grad, hess); let dest = dest.cast::<f32>();
_mm256_storeu_ps(dest, _mm256_permute2f128_ps::<0x20>(low, high));
_mm256_storeu_ps(dest.add(WIDTH), _mm256_permute2f128_ps::<0x31>(low, high));
}
}
macro_rules! unary_inplace_kernel {
($name:ident, $kernel:expr, $scalar:expr) => {
#[target_feature(enable = "avx2,fma")]
pub(super) unsafe fn $name(values: &mut [f32]) {
unsafe {
let mut index = 0;
while index + WIDTH <= values.len() {
let input = _mm256_loadu_ps(values.as_ptr().add(index));
if regular_input(input) {
_mm256_storeu_ps(values.as_mut_ptr().add(index), ($kernel)(input));
} else {
for value in &mut values[index..index + WIDTH] {
*value = ($scalar)(*value);
}
}
index += WIDTH;
}
for value in &mut values[index..] {
*value = ($scalar)(*value);
}
}
}
};
}
unary_inplace_kernel!(exp_inplace, exp_f32, f32::exp);
unary_inplace_kernel!(sigmoid_inplace, sigmoid_f32, sigmoid_scalar);
#[target_feature(enable = "avx2,fma")]
pub(super) unsafe fn logistic_gradient(
preds: &[f32],
labels: &[f32],
weights: Option<&[f32]>,
scale_pos_weight: f32,
min_hess: f32,
out: &mut [GradPair],
) {
unsafe {
let one = _mm256_set1_ps(1.0);
let scale = _mm256_set1_ps(scale_pos_weight);
let min_hess_vector = _mm256_set1_ps(min_hess);
let scalar_range = |out: &mut [GradPair], range| {
scalar::logistic_gradient(
preds,
labels,
weights,
scale_pos_weight,
min_hess,
out,
range,
);
};
let mut index = 0;
while index + WIDTH <= preds.len() {
let pred = _mm256_loadu_ps(preds.as_ptr().add(index));
if !regular_input(pred) {
scalar_range(out, index..index + WIDTH);
index += WIDTH;
continue;
}
let label = _mm256_loadu_ps(labels.as_ptr().add(index));
let probability = sigmoid_f32(pred);
let mut weight = match weights {
Some(values) => _mm256_loadu_ps(values.as_ptr().add(index)),
None => one,
};
let positive = _mm256_cmp_ps::<_CMP_EQ_OQ>(label, one);
weight = _mm256_mul_ps(weight, _mm256_blendv_ps(one, scale, positive));
let grad = _mm256_mul_ps(_mm256_sub_ps(probability, label), weight);
let hess = _mm256_mul_ps(
_mm256_max_ps(
_mm256_mul_ps(probability, _mm256_sub_ps(one, probability)),
min_hess_vector,
),
weight,
);
store_pairs(out.as_mut_ptr().add(index), grad, hess);
index += WIDTH;
}
scalar_range(out, index..preds.len());
}
}
#[inline]
#[target_feature(enable = "avx2,fma")]
unsafe fn row_reduce<const K: usize>(
value: __m256,
op: unsafe fn(__m256, __m256) -> __m256,
) -> __m256 {
unsafe {
let reduced = op(value, _mm256_permute_ps::<0b10_11_00_01>(value));
if K == 2 {
reduced
} else {
op(reduced, _mm256_permute_ps::<0b01_00_11_10>(reduced))
}
}
}
#[inline]
#[target_feature(enable = "avx2,fma")]
unsafe fn short_softmax_batch<const K: usize, const GRADIENT: bool>(
preds: *const f32,
) -> Option<__m256> {
unsafe {
let values = _mm256_loadu_ps(preds);
let mut maximum = row_reduce::<K>(values, _mm256_max_ps);
let minimum = row_reduce::<K>(values, _mm256_min_ps);
if GRADIENT {
maximum = _mm256_max_ps(_mm256_set1_ps(f32::MIN_POSITIVE), maximum);
}
let regular = _mm256_cmp_ps::<_CMP_LE_OQ>(
_mm256_sub_ps(maximum, minimum),
_mm256_set1_ps(MAX_FAST_EXP_INPUT),
);
if _mm256_movemask_ps(regular) != 0xFF {
return None;
}
let exp = exp_f32(_mm256_sub_ps(values, maximum));
let sum = row_reduce::<K>(exp, _mm256_add_ps);
Some(_mm256_mul_ps(exp, _mm256_div_ps(_mm256_set1_ps(1.0), sum)))
}
}
#[target_feature(enable = "avx2,fma")]
pub(super) unsafe fn short_softmax_rows<const K: usize>(values: &mut [f32]) {
unsafe {
let (batches, remainder) = values.as_chunks_mut::<WIDTH>();
for batch in batches {
match short_softmax_batch::<K, false>(batch.as_ptr()) {
Some(probabilities) => _mm256_storeu_ps(batch.as_mut_ptr(), probabilities),
None => {
for row in batch.chunks_mut(K) {
super::softmax_scalar(row);
}
}
}
}
for row in remainder.chunks_mut(K) {
super::softmax_scalar(row);
}
}
}
#[inline]
#[target_feature(enable = "avx2,fma")]
#[allow(
clippy::cast_ptr_alignment,
reason = "_mm_load_sd has no alignment requirement"
)]
unsafe fn broadcast_rows<const K: usize>(values: *const f32) -> __m256 {
unsafe {
let (loaded, index) = if K == 2 {
(
_mm_loadu_ps(values),
_mm256_setr_epi32(0, 0, 1, 1, 2, 2, 3, 3),
)
} else {
(
_mm_castpd_ps(_mm_load_sd(values.cast::<f64>())),
_mm256_setr_epi32(0, 0, 0, 0, 1, 1, 1, 1),
)
};
_mm256_permutevar8x32_ps(_mm256_castps128_ps256(loaded), index)
}
}
#[inline]
#[target_feature(enable = "avx2,fma")]
unsafe fn class_indicator<const K: usize>(label: __m256) -> __m256 {
unsafe {
let class = if K == 2 {
_mm256_setr_ps(0.0, 1.0, 0.0, 1.0, 0.0, 1.0, 0.0, 1.0)
} else {
_mm256_setr_ps(0.0, 1.0, 2.0, 3.0, 0.0, 1.0, 2.0, 3.0)
};
let one = _mm256_set1_ps(1.0);
let next = _mm256_add_ps(class, one);
let in_range = _mm256_and_ps(
_mm256_cmp_ps::<_CMP_GE_OQ>(label, class),
_mm256_cmp_ps::<_CMP_LT_OQ>(label, next),
);
let is_zero = _mm256_cmp_ps::<_CMP_NGE_UQ>(label, one);
let class_is_zero = _mm256_cmp_ps::<_CMP_EQ_OQ>(class, _mm256_setzero_ps());
let and: unsafe fn(__m256, __m256) -> __m256 = _mm256_and_ps;
and(_mm256_blendv_ps(in_range, is_zero, class_is_zero), one)
}
}
#[target_feature(enable = "avx2,fma")]
pub(super) unsafe fn short_softmax_gradient<const K: usize>(
preds: &[f32],
labels: &[f32],
weights: Option<&[f32]>,
min_hess: f32,
out: &mut [GradPair],
) {
unsafe {
const { assert!(K == 2 || K == 4) };
let rows_per_batch = WIDTH / K;
let one = _mm256_set1_ps(1.0);
let minimum = _mm256_set1_ps(min_hess);
let mut row = 0;
while row + rows_per_batch <= labels.len() {
let base = row * K;
match short_softmax_batch::<K, true>(preds.as_ptr().add(base)) {
Some(probability) => {
let label = broadcast_rows::<K>(labels.as_ptr().add(row));
let weight = match weights {
Some(values) => broadcast_rows::<K>(values.as_ptr().add(row)),
None => one,
};
let indicator = class_indicator::<K>(label);
let grad = _mm256_mul_ps(_mm256_sub_ps(probability, indicator), weight);
let hess = _mm256_max_ps(
_mm256_mul_ps(
_mm256_mul_ps(
_mm256_mul_ps(probability, _mm256_set1_ps(2.0)),
_mm256_sub_ps(one, probability),
),
weight,
),
minimum,
);
store_pairs(out.as_mut_ptr().add(base), grad, hess);
}
None => {
super::softmax_gradient_rows_scalar(
preds,
labels,
weights,
min_hess,
out,
row..row + rows_per_batch,
K,
);
}
}
row += rows_per_batch;
}
super::softmax_gradient_rows_scalar(
preds,
labels,
weights,
min_hess,
out,
row..labels.len(),
K,
);
}
}
pub(super) unsafe fn count_le_16(cuts: &[f32], value: f32) -> usize {
debug_assert_eq!(cuts.len(), 16);
unsafe {
let value = _mm_set1_ps(value);
let ptr = cuts.as_ptr();
let mut count = 0usize;
for quarter in 0..4 {
let mask = _mm_cmple_ps(_mm_loadu_ps(ptr.add(quarter * 4)), value);
count += _mm_movemask_ps(mask).count_ones() as usize;
}
count
}
}