#![allow(ambiguous_glob_reexports)]
pub mod fast_tanh;
pub mod fused;
pub mod hard_swish;
pub mod hard_tanh;
pub mod kernel_macro;
pub mod leaky_hard_tanh;
pub mod prelu;
pub mod relu;
pub mod sigmoid;
pub mod silu;
pub mod softsign;
pub mod tanh;
#[cfg(test)]
mod activations_test;
pub use fast_tanh::*;
pub use fused::*;
pub use hard_swish::*;
pub use hard_tanh::*;
pub use leaky_hard_tanh::*;
pub use prelu::*;
pub use relu::*;
use serde::{Deserialize, Serialize};
pub use sigmoid::*;
pub use silu::*;
pub use softsign::*;
pub use tanh::*;
#[derive(Clone, Copy, Debug, PartialEq, Eq, Default, Serialize, Deserialize)]
#[repr(usize)]
pub enum ActivationPrecision {
#[default]
Standard = 1,
Fast = 0,
}
impl ActivationPrecision {
#[inline]
pub fn from_f32(val: f32) -> Self {
match val.round() as i32 {
0 => Self::Fast,
_ => Self::Standard,
}
}
#[inline]
pub fn to_f32(self) -> f32 {
self as u32 as f32
}
}
static ACTIVATION_MODE: core::sync::atomic::AtomicUsize =
core::sync::atomic::AtomicUsize::new(ActivationPrecision::Standard as usize);
thread_local! {
static ACTIVE_MODEL_PRECISION: std::cell::Cell<Option<ActivationPrecision>> = const { std::cell::Cell::new(None) };
}
pub struct ActivationPrecisionGuard {
_private: (),
}
impl Drop for ActivationPrecisionGuard {
#[inline(always)]
fn drop(&mut self) {
ACTIVE_MODEL_PRECISION.with(|p| p.set(None));
}
}
#[inline]
pub fn set_activation_precision(mode: ActivationPrecision) {
ACTIVATION_MODE.store(mode as usize, core::sync::atomic::Ordering::Relaxed);
}
#[inline]
pub fn set_thread_local_activation_precision(
mode: Option<ActivationPrecision>,
) -> ActivationPrecisionGuard {
ACTIVE_MODEL_PRECISION.with(|p| p.set(mode));
ActivationPrecisionGuard { _private: () }
}
#[inline]
pub fn thread_local_activation_precision() -> Option<ActivationPrecision> {
ACTIVE_MODEL_PRECISION.with(|p| p.get())
}
#[inline]
pub fn activation_precision() -> ActivationPrecision {
if let Some(precision) = ACTIVE_MODEL_PRECISION.with(|p| p.get()) {
precision
} else {
match ACTIVATION_MODE.load(core::sync::atomic::Ordering::Relaxed) {
0 => ActivationPrecision::Fast,
_ => ActivationPrecision::Standard,
}
}
}
#[inline(always)]
pub fn tanh_slice(data: &mut [f32]) {
if activation_precision() == ActivationPrecision::Standard {
crate::math::common::dispatch_simd!(tanh_slice_hf(data));
} else {
crate::math::common::dispatch_simd!(tanh_slice(data));
}
}
#[inline(always)]
pub fn sigmoid_slice(data: &mut [f32]) {
if activation_precision() == ActivationPrecision::Standard {
crate::math::common::dispatch_simd!(sigmoid_slice_hf(data));
} else {
crate::math::common::dispatch_simd!(sigmoid_slice(data));
}
}
#[inline(always)]
pub fn relu_slice(data: &mut [f32]) {
crate::math::common::dispatch_simd!(relu_slice(data));
}
#[inline(always)]
pub fn prelu_slice(data: &mut [f32], slopes: &[f32]) {
crate::math::common::dispatch_simd!(prelu_slice(data, slopes));
}
#[inline(always)]
pub fn softsign_slice(data: &mut [f32]) {
crate::math::common::dispatch_simd!(softsign_slice(data));
}
#[inline(always)]
pub fn silu_slice(data: &mut [f32]) {
crate::math::common::dispatch_simd!(silu_slice(data));
}
#[inline(always)]
pub fn hard_tanh_slice(data: &mut [f32]) {
crate::math::common::dispatch_simd!(hard_tanh_slice(data));
}
#[inline(always)]
pub fn hard_swish_slice(data: &mut [f32]) {
crate::math::common::dispatch_simd!(hard_swish_slice(data));
}
#[inline(always)]
pub fn fast_tanh_slice(data: &mut [f32]) {
crate::math::common::dispatch_simd!(fast_tanh_slice(data));
}
#[inline(always)]
pub fn leaky_hard_tanh_slice(
data: &mut [f32],
min_val: f32,
max_val: f32,
min_slope: f32,
max_slope: f32,
) {
crate::math::common::dispatch_simd!(leaky_hard_tanh_slice(
data, min_val, max_val, min_slope, max_slope
));
}