tract_linalg/generic/
sigmoid.rs1#![allow(clippy::excessive_precision)]
2use tract_data::internal::*;
3
4pub const LOW: f32 = -14.5;
10
11pub fn ssigmoid(x: f32) -> f32 {
25 const HIGH: f32 = -LOW;
26
27 const ALPHA_13: f32 = -4.433153405e-18;
28 const ALPHA_11: f32 = 1.169974371e-14;
29 const ALPHA_9: f32 = -1.875289645e-11;
30 const ALPHA_7: f32 = 4.257889523e-8;
31 const ALPHA_5: f32 = 0.00004811817576;
32 const ALPHA_3: f32 = 0.008163842030;
33 const ALPHA_1: f32 = 0.2499999971;
34 const BETA_6: f32 = 3.922935744e-6;
35 const BETA_4: f32 = 0.001524872358;
36 const BETA_2: f32 = 0.1159886749;
37 const BETA_0: f32 = 1.0;
38
39 let x = x.clamp(LOW, HIGH);
40
41 let x2 = x * x;
42
43 let p = ALPHA_13;
44 let p = x2 * p + ALPHA_11;
45 let p = x2 * p + ALPHA_9;
46 let p = x2 * p + ALPHA_7;
47 let p = x2 * p + ALPHA_5;
48 let p = x2 * p + ALPHA_3;
49 let p = x2 * p + ALPHA_1;
50 let p = p * x;
51
52 let q = BETA_6;
53 let q = x2 * q + BETA_4;
54 let q = x2 * q + BETA_2;
55 let q = x2 * q + BETA_0;
56
57 p / q + 0.5
58}
59
60pub fn hsigmoid(x: f16) -> f16 {
65 const LOW: f16 = f16::from_f32_const(-6.92);
72 const HIGH: f16 = f16::from_f32_const(6.92);
73
74 const ALPHA_5: f16 = f16::from_f32_const(-0.0000124702);
75 const ALPHA_3: f16 = f16::from_f32_const(0.00400222);
76 const ALPHA_1: f16 = f16::from_f32_const(0.249895);
77
78 const BETA_2: f16 = f16::from_f32_const(0.098734);
79 const BETA_0: f16 = f16::from_f32_const(1.0);
80
81 let x = x.clamp(LOW, HIGH);
82
83 let x2 = x * x;
84
85 let p = ALPHA_5;
86 let p = x2 * p + ALPHA_3;
87 let p = x2 * p + ALPHA_1;
88 let p = p * x;
89
90 let q = BETA_2;
91 let q = x2 * q + BETA_0;
92
93 p / q + f16::from_f32_const(0.5)
94}
95
96routine_ew_rust!(generic;
97 f32,
98 generic_sigmoid_f32_4n,
99 4,
100 4,
101 fn run(x: &mut [f32], _: ()) {
102 debug_assert!(x.len() % Self::nr() == 0);
103 debug_assert!(x.as_ptr() as usize % Self::alignment_bytes() == 0);
104 x.iter_mut().for_each(|px| *px = ssigmoid(*px))
105 },
106 func(Sigmoid)
107);
108
109routine_ew_rust!(generic;
110 f16,
111 generic_sigmoid_f16_8n,
112 8,
113 8,
114 fn run(x: &mut [f16], _: ()) {
115 debug_assert!(x.len() % Self::nr() == 0);
116 debug_assert!(x.as_ptr() as usize % Self::alignment_bytes() == 0);
117 x.iter_mut().for_each(|px| *px = hsigmoid(*px))
118 },
119 func(Sigmoid)
120);