NeuralAmpModeler-rs 3.0.0

An opinionated, high-performance Neural Amp Modeler (NAM) client and core implementation in Rust for Linux/PipeWire and CLAP plugins.
// SPDX-License-Identifier: Apache-2.0
// Copyright (c) 2026 Fábio Henrique de Lima Silva (fhl.bsb@gmail.com) All rights reserved.
#![allow(
    unsafe_op_in_unsafe_fn,
    clippy::missing_safety_doc,
    clippy::too_many_arguments
)]

//! Fused kernels for LSTM gates (AVX2 and AVX-512).
//!
//! Extracted from `activations/fused.rs` and `simd/avx2.rs`/`simd/avx512.rs`.

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::*;

/// Fused kernel for LSTM gates (AVX2) — Standard (exact-grade) accuracy path.
/// Uses polynomial exp-based tanh/sigmoid approximations with Kahan
/// compensated summation for the cell state.
///
/// # Safety
/// Requires AVX2 and FMA support.
#[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)
}

/// Fused kernel for LSTM gates (AVX2) — Fast (production) accuracy path.
/// Uses Padé tanh + minimax sigmoid with interleaved dual-sigmoid and
/// Kahan compensated summation for the cell state.
///
/// # Safety
/// Requires AVX2 and FMA support.
#[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)
}

/// Fused kernel for LSTM gates (AVX2).
/// Thin dispatch wrapper — prefer `_hf` / `_std` variants in hot paths.
///
/// # Safety
/// Requires AVX2 and FMA support.
#[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) }
    }
}

/// Fused kernel for LSTM gates (AVX-512) — Standard (exact-grade) accuracy path.
/// Uses polynomial exp-based tanh/sigmoid approximations with Kahan
/// compensated summation for the cell state.
///
/// # Safety
/// Requires AVX-512F and AVX-512VL support.
#[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)
}

/// Fused kernel for LSTM gates (AVX-512) — Fast (production) accuracy path.
/// Uses Padé tanh + minimax sigmoid with Kahan compensated summation
/// for the cell state.
///
/// # Safety
/// Requires AVX-512F and AVX-512VL support.
#[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)
}

/// Fused kernel for LSTM gates (AVX-512).
/// Thin dispatch wrapper — prefer `_hf` / `_std` variants in hot paths.
///
/// # Safety
/// Requires AVX-512F and AVX-512VL support.
#[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) }
    }
}

/// Fused kernel for dynamic LSTM gate processing via AVX2.
#[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;
    }
}

/// Fused kernel to update the memory (state) of an LSTM network.
/// This function decides what the network should "forget" from the past and what to "learn" from the present,
/// updating the values all at once for 16 memory cells.
#[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 {
        // Load the 4 decisions (forget, learn, etc.) for 16 cells.
        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));

        // Perform the memory computation in a fused manner.
        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,
        );
    }
}