#![allow(
unsafe_op_in_unsafe_fn,
clippy::missing_safety_doc,
clippy::too_many_arguments
)]
use crate::math::activations::ActivationPrecision;
use crate::math::activations::activation_precision;
use crate::math::activations::sigmoid::high_fidelity::{
scalar_sigmoid_poly, simd_sigmoid_poly_avx2, simd_sigmoid_poly_avx512,
};
use crate::math::activations::sigmoid::{
scalar_minimax_sigmoid, simd_sigmoid_avx2, simd_sigmoid_avx512, simd_sigmoid_dual_avx2,
};
use crate::math::activations::tanh::high_fidelity::{
scalar_tanh_poly, simd_tanh_poly_avx2, simd_tanh_poly_avx512,
};
use crate::math::activations::tanh::{scalar_pade_tanh, simd_tanh_avx2, simd_tanh_avx512};
use core::arch::x86_64::*;
#[inline]
#[target_feature(enable = "avx2,fma")]
pub unsafe fn fused_lstm_gates_avx2_hf(
gf: __m256,
gi: __m256,
gg: __m256,
go: __m256,
cs: __m256,
cs_err: __m256,
) -> (__m256, __m256, __m256) {
let sig_f = unsafe { simd_sigmoid_poly_avx2(gf) };
let sig_i = unsafe { simd_sigmoid_poly_avx2(gi) };
let sig_o = unsafe { simd_sigmoid_poly_avx2(go) };
let tanh_g = unsafe { simd_tanh_poly_avx2(gg) };
let f_cs = _mm256_mul_ps(sig_f, cs);
let f_err = _mm256_mul_ps(sig_f, cs_err);
let i_g = _mm256_mul_ps(sig_i, tanh_g);
let y = _mm256_sub_ps(i_g, f_err);
let new_cs = _mm256_add_ps(f_cs, y);
let new_cs_err = _mm256_sub_ps(_mm256_sub_ps(new_cs, f_cs), y);
let hidden = _mm256_mul_ps(sig_o, unsafe { simd_tanh_poly_avx2(new_cs) });
(new_cs, new_cs_err, hidden)
}
#[inline]
#[target_feature(enable = "avx2,fma")]
pub unsafe fn fused_lstm_gates_avx2_std(
gf: __m256,
gi: __m256,
gg: __m256,
go: __m256,
cs: __m256,
cs_err: __m256,
) -> (__m256, __m256, __m256) {
let (sig_f, sig_i) = unsafe { simd_sigmoid_dual_avx2(gf, gi) };
let sig_o = unsafe { simd_sigmoid_avx2(go) };
let tanh_g = unsafe { simd_tanh_avx2(gg) };
let f_cs = _mm256_mul_ps(sig_f, cs);
let f_err = _mm256_mul_ps(sig_f, cs_err);
let i_g = _mm256_mul_ps(sig_i, tanh_g);
let y = _mm256_sub_ps(i_g, f_err);
let new_cs = _mm256_add_ps(f_cs, y);
let new_cs_err = _mm256_sub_ps(_mm256_sub_ps(new_cs, f_cs), y);
let hidden = _mm256_mul_ps(sig_o, unsafe { simd_tanh_avx2(new_cs) });
(new_cs, new_cs_err, hidden)
}
#[inline]
#[target_feature(enable = "avx2,fma")]
pub unsafe fn fused_lstm_gates_avx2(
gf: __m256,
gi: __m256,
gg: __m256,
go: __m256,
cs: __m256,
cs_err: __m256,
) -> (__m256, __m256, __m256) {
if activation_precision() == ActivationPrecision::Standard {
unsafe { fused_lstm_gates_avx2_hf(gf, gi, gg, go, cs, cs_err) }
} else {
unsafe { fused_lstm_gates_avx2_std(gf, gi, gg, go, cs, cs_err) }
}
}
#[inline]
#[target_feature(enable = "avx512f,avx512vl")]
pub unsafe fn fused_lstm_gates_avx512_hf(
gf: __m512,
gi: __m512,
gg: __m512,
go: __m512,
cs: __m512,
cs_err: __m512,
) -> (__m512, __m512, __m512) {
let sig_f = unsafe { simd_sigmoid_poly_avx512(gf) };
let sig_i = unsafe { simd_sigmoid_poly_avx512(gi) };
let sig_o = unsafe { simd_sigmoid_poly_avx512(go) };
let tanh_g = unsafe { simd_tanh_poly_avx512(gg) };
let f_cs = _mm512_mul_ps(sig_f, cs);
let f_err = _mm512_mul_ps(sig_f, cs_err);
let i_g = _mm512_mul_ps(sig_i, tanh_g);
let y = _mm512_sub_ps(i_g, f_err);
let new_cs = _mm512_add_ps(f_cs, y);
let new_cs_err = _mm512_sub_ps(_mm512_sub_ps(new_cs, f_cs), y);
let hidden = _mm512_mul_ps(sig_o, unsafe { simd_tanh_poly_avx512(new_cs) });
(new_cs, new_cs_err, hidden)
}
#[inline]
#[target_feature(enable = "avx512f,avx512vl")]
pub unsafe fn fused_lstm_gates_avx512_std(
gf: __m512,
gi: __m512,
gg: __m512,
go: __m512,
cs: __m512,
cs_err: __m512,
) -> (__m512, __m512, __m512) {
let sig_f = unsafe { simd_sigmoid_avx512(gf) };
let sig_i = unsafe { simd_sigmoid_avx512(gi) };
let sig_o = unsafe { simd_sigmoid_avx512(go) };
let tanh_g = unsafe { simd_tanh_avx512(gg) };
let f_cs = _mm512_mul_ps(sig_f, cs);
let f_err = _mm512_mul_ps(sig_f, cs_err);
let i_g = _mm512_mul_ps(sig_i, tanh_g);
let y = _mm512_sub_ps(i_g, f_err);
let new_cs = _mm512_add_ps(f_cs, y);
let new_cs_err = _mm512_sub_ps(_mm512_sub_ps(new_cs, f_cs), y);
let hidden = _mm512_mul_ps(sig_o, unsafe { simd_tanh_avx512(new_cs) });
(new_cs, new_cs_err, hidden)
}
#[inline]
#[target_feature(enable = "avx512f,avx512vl")]
pub unsafe fn fused_lstm_gates_avx512(
gf: __m512,
gi: __m512,
gg: __m512,
go: __m512,
cs: __m512,
cs_err: __m512,
) -> (__m512, __m512, __m512) {
if activation_precision() == ActivationPrecision::Standard {
unsafe { fused_lstm_gates_avx512_hf(gf, gi, gg, go, cs, cs_err) }
} else {
unsafe { fused_lstm_gates_avx512_std(gf, gi, gg, go, cs, cs_err) }
}
}
#[inline]
#[target_feature(enable = "avx2,fma")]
pub unsafe fn fused_lstm_gates_dyn_avx2(
gates: &mut [f32],
cell_state: &mut [f32],
cell_error: &mut [f32],
hidden_state: &mut [f32],
hidden_size: usize,
) {
let is_hf = activation_precision() == ActivationPrecision::Standard;
let mut j = 0;
while j + 8 <= hidden_size {
let gi = _mm256_loadu_ps(gates.as_ptr().add(j));
let gf = _mm256_loadu_ps(gates.as_ptr().add(j + hidden_size));
let gg = _mm256_loadu_ps(gates.as_ptr().add(j + 2 * hidden_size));
let go = _mm256_loadu_ps(gates.as_ptr().add(j + 3 * hidden_size));
let cs = _mm256_loadu_ps(cell_state.as_ptr().add(j));
let cs_err = _mm256_loadu_ps(cell_error.as_ptr().add(j));
let (new_cs, new_cs_err, hidden) = if is_hf {
fused_lstm_gates_avx2_hf(gf, gi, gg, go, cs, cs_err)
} else {
fused_lstm_gates_avx2_std(gf, gi, gg, go, cs, cs_err)
};
_mm256_storeu_ps(cell_state.as_mut_ptr().add(j), new_cs);
_mm256_storeu_ps(cell_error.as_mut_ptr().add(j), new_cs_err);
_mm256_storeu_ps(hidden_state.as_mut_ptr().add(j), hidden);
j += 8;
}
if j < hidden_size {
fused_lstm_gates_dyn_tail(
gates,
cell_state,
cell_error,
hidden_state,
hidden_size,
is_hf,
j,
);
}
}
#[cold]
#[inline(never)]
fn fused_lstm_gates_dyn_tail(
gates: &mut [f32],
cell_state: &mut [f32],
cell_error: &mut [f32],
hidden_state: &mut [f32],
hidden_size: usize,
is_hf: bool,
mut j: usize,
) {
while j < hidden_size {
if is_hf {
let sig_i = scalar_sigmoid_poly(gates[j]);
let sig_f = scalar_sigmoid_poly(gates[j + hidden_size]);
let tanh_g = scalar_tanh_poly(gates[j + 2 * hidden_size]);
let sig_o = scalar_sigmoid_poly(gates[j + 3 * hidden_size]);
let f_cs = sig_f * cell_state[j];
let f_err = sig_f * cell_error[j];
let i_g = sig_i * tanh_g;
let y = i_g - f_err;
let new_cs = f_cs + y;
let new_cs_err = (new_cs - f_cs) - y;
cell_state[j] = new_cs;
cell_error[j] = new_cs_err;
hidden_state[j] = sig_o * scalar_tanh_poly(new_cs);
} else {
let sig_i = scalar_minimax_sigmoid(gates[j]);
let sig_f = scalar_minimax_sigmoid(gates[j + hidden_size]);
let tanh_g = scalar_pade_tanh(gates[j + 2 * hidden_size]);
let sig_o = scalar_minimax_sigmoid(gates[j + 3 * hidden_size]);
let f_cs = sig_f * cell_state[j];
let f_err = sig_f * cell_error[j];
let i_g = sig_i * tanh_g;
let y = i_g - f_err;
let new_cs = f_cs + y;
let new_cs_err = (new_cs - f_cs) - y;
cell_state[j] = new_cs;
cell_error[j] = new_cs_err;
hidden_state[j] = sig_o * scalar_pade_tanh(new_cs);
}
j += 1;
}
}
#[inline]
#[target_feature(enable = "avx512f,avx512vl")]
pub unsafe fn fused_lstm_gates_dyn_avx512(
gates: &mut [f32],
cell_state: &mut [f32],
cell_error: &mut [f32],
hidden_state: &mut [f32],
hidden_size: usize,
) {
let is_hf = activation_precision() == ActivationPrecision::Standard;
let mut j = 0;
while j + 16 <= hidden_size {
let gi = _mm512_loadu_ps(gates.as_ptr().add(j));
let gf = _mm512_loadu_ps(gates.as_ptr().add(j + hidden_size));
let gg = _mm512_loadu_ps(gates.as_ptr().add(j + 2 * hidden_size));
let go = _mm512_loadu_ps(gates.as_ptr().add(j + 3 * hidden_size));
let cs = _mm512_loadu_ps(cell_state.as_ptr().add(j));
let cs_err = _mm512_loadu_ps(cell_error.as_ptr().add(j));
let (new_cs, new_cs_err, hidden) = if is_hf {
fused_lstm_gates_avx512_hf(gf, gi, gg, go, cs, cs_err)
} else {
fused_lstm_gates_avx512_std(gf, gi, gg, go, cs, cs_err)
};
_mm512_storeu_ps(cell_state.as_mut_ptr().add(j), new_cs);
_mm512_storeu_ps(cell_error.as_mut_ptr().add(j), new_cs_err);
_mm512_storeu_ps(hidden_state.as_mut_ptr().add(j), hidden);
j += 16;
}
if j < hidden_size {
fused_lstm_gates_dyn_tail(
gates,
cell_state,
cell_error,
hidden_state,
hidden_size,
is_hf,
j,
);
}
}