#![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::{
simd_sigmoid_poly_avx2, simd_sigmoid_poly_avx512,
};
use crate::math::activations::sigmoid::{
simd_sigmoid_avx2, simd_sigmoid_avx512, simd_sigmoid_dual_avx2,
};
use crate::math::activations::tanh::high_fidelity::{simd_tanh_poly_avx2, simd_tanh_poly_avx512};
use crate::math::activations::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 {
unsafe {
fused_lstm_gates_dyn_tail(
gates,
cell_state,
cell_error,
hidden_state,
hidden_size,
is_hf,
j,
);
}
}
}
#[target_feature(enable = "avx2,fma")]
unsafe 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,
) {
let gp = gates.as_mut_ptr();
let cp = cell_state.as_ptr();
let ep = cell_error.as_ptr();
while j + 8 <= hidden_size {
let gi = _mm256_loadu_ps(gp.add(j));
let gf = _mm256_loadu_ps(gp.add(j + hidden_size));
let gg = _mm256_loadu_ps(gp.add(j + 2 * hidden_size));
let go = _mm256_loadu_ps(gp.add(j + 3 * hidden_size));
let cs = _mm256_loadu_ps(cp.add(j));
let cs_err = _mm256_loadu_ps(ep.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 {
let tail_len = hidden_size - j;
let mut temp_gf = [0.0f32; 8];
let mut temp_gi = [0.0f32; 8];
let mut temp_gg = [0.0f32; 8];
let mut temp_go = [0.0f32; 8];
let mut temp_cs = [0.0f32; 8];
let mut temp_ce = [0.0f32; 8];
for k in 0..tail_len {
temp_gf[k] = *gp.add(j + k + hidden_size);
temp_gi[k] = *gp.add(j + k);
temp_gg[k] = *gp.add(j + k + 2 * hidden_size);
temp_go[k] = *gp.add(j + k + 3 * hidden_size);
temp_cs[k] = *cp.add(j + k);
temp_ce[k] = *ep.add(j + k);
}
let gf_v = _mm256_loadu_ps(temp_gf.as_ptr());
let gi_v = _mm256_loadu_ps(temp_gi.as_ptr());
let gg_v = _mm256_loadu_ps(temp_gg.as_ptr());
let go_v = _mm256_loadu_ps(temp_go.as_ptr());
let cs_v = _mm256_loadu_ps(temp_cs.as_ptr());
let ce_v = _mm256_loadu_ps(temp_ce.as_ptr());
let (new_cs_v, new_ce_v, hidden_v) = if is_hf {
fused_lstm_gates_avx2_hf(gf_v, gi_v, gg_v, go_v, cs_v, ce_v)
} else {
fused_lstm_gates_avx2_std(gf_v, gi_v, gg_v, go_v, cs_v, ce_v)
};
let mut out_cs = [0.0f32; 8];
let mut out_ce = [0.0f32; 8];
let mut out_hs = [0.0f32; 8];
_mm256_storeu_ps(out_cs.as_mut_ptr(), new_cs_v);
_mm256_storeu_ps(out_ce.as_mut_ptr(), new_ce_v);
_mm256_storeu_ps(out_hs.as_mut_ptr(), hidden_v);
let cp_mut = cell_state.as_mut_ptr();
let ep_mut = cell_error.as_mut_ptr();
let hp_mut = hidden_state.as_mut_ptr();
for k in 0..tail_len {
*cp_mut.add(j + k) = out_cs[k];
*ep_mut.add(j + k) = out_ce[k];
*hp_mut.add(j + k) = out_hs[k];
}
}
}
#[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 {
unsafe {
fused_lstm_gates_dyn_tail(
gates,
cell_state,
cell_error,
hidden_state,
hidden_size,
is_hf,
j,
);
}
}
}