use super::*;
use crate::dispatch::ComplexScalar;
#[cfg(feature = "epilogue")]
use crate::kernel::epilogue::BiasDim;
#[cfg(feature = "complex")]
#[allow(clippy::too_many_arguments)]
pub fn gemm_cplx<T: ComplexScalar>(
alpha: T,
a: MatRef<'_, T>,
conj_a: bool,
b: MatRef<'_, T>,
conj_b: bool,
beta: T,
c: MatMut<'_, T>,
par: Parallelism,
) {
workspace::with_thread_pool(|ws| gemm_cplx_with(ws, alpha, a, conj_a, b, conj_b, beta, c, par));
}
#[cfg(feature = "complex")]
#[allow(clippy::too_many_arguments)]
pub fn gemm_cplx_with<T: ComplexScalar>(
ws: &mut Workspace,
alpha: T,
a: MatRef<'_, T>,
conj_a: bool,
b: MatRef<'_, T>,
conj_b: bool,
beta: T,
c: MatMut<'_, T>,
par: Parallelism,
) {
validate_gemm_views(&a, &b, &c);
unsafe {
dispatch::execute_complex(
conj_a,
conj_b,
Task {
m: a.rows,
k: a.cols,
n: b.cols,
alpha,
a: a.data.as_ptr(),
rsa: a.rs,
csa: a.cs,
b: b.data.as_ptr(),
rsb: b.rs,
csb: b.cs,
beta,
c: c.data.as_mut_ptr(),
rsc: c.rs,
csc: c.cs,
},
par,
ws,
);
}
}
#[cfg(all(feature = "complex", feature = "epilogue"))]
#[allow(clippy::too_many_arguments)]
pub fn gemm_cplx_fused<T: ComplexScalar>(
alpha: T,
a: MatRef<'_, T>,
conj_a: bool,
b: MatRef<'_, T>,
conj_b: bool,
beta: T,
c: MatMut<'_, T>,
bias: Option<Bias<'_, T>>,
par: Parallelism,
) {
workspace::with_thread_pool(|ws| {
gemm_cplx_fused_with(ws, alpha, a, conj_a, b, conj_b, beta, c, bias, par)
});
}
#[cfg(all(feature = "complex", feature = "epilogue"))]
#[allow(clippy::too_many_arguments)]
pub fn gemm_cplx_fused_with<T: ComplexScalar>(
ws: &mut Workspace,
alpha: T,
a: MatRef<'_, T>,
conj_a: bool,
b: MatRef<'_, T>,
conj_b: bool,
beta: T,
c: MatMut<'_, T>,
bias: Option<Bias<'_, T>>,
par: Parallelism,
) {
let Some(bias) = bias else {
gemm_cplx_with(ws, alpha, a, conj_a, b, conj_b, beta, c, par);
return;
};
validate_gemm_views(&a, &b, &c);
validate_bias(&Some(bias), a.rows, b.cols, &c);
let epi = to_fused_epi(Some(bias), None);
unsafe {
dispatch::execute_complex_fused(
conj_a,
conj_b,
Task {
m: a.rows,
k: a.cols,
n: b.cols,
alpha,
a: a.data.as_ptr(),
rsa: a.rs,
csa: a.cs,
b: b.data.as_ptr(),
rsb: b.rs,
csb: b.cs,
beta,
c: c.data.as_mut_ptr(),
rsc: c.rs,
csc: c.cs,
},
epi,
par,
ws,
);
}
}
#[cfg(all(feature = "complex", feature = "epilogue"))]
#[allow(clippy::too_many_arguments)]
pub unsafe fn gemm_cplx_fused_unchecked<T: ComplexScalar>(
m: usize,
k: usize,
n: usize,
alpha: T,
a: *const T,
rsa: isize,
csa: isize,
conj_a: bool,
b: *const T,
rsb: isize,
csb: isize,
conj_b: bool,
beta: T,
c: *mut T,
rsc: isize,
csc: isize,
bias: *const T,
bias_dim: BiasDim,
has_bias: bool,
par: Parallelism,
) {
unsafe {
workspace::with_thread_pool(|ws| {
gemm_cplx_fused_unchecked_with(
ws, m, k, n, alpha, a, rsa, csa, conj_a, b, rsb, csb, conj_b, beta, c, rsc, csc,
bias, bias_dim, has_bias, par,
);
});
}
}
#[cfg(all(feature = "complex", feature = "epilogue"))]
#[allow(clippy::too_many_arguments)]
pub unsafe fn gemm_cplx_fused_unchecked_with<T: ComplexScalar>(
ws: &mut Workspace,
m: usize,
k: usize,
n: usize,
alpha: T,
a: *const T,
rsa: isize,
csa: isize,
conj_a: bool,
b: *const T,
rsb: isize,
csb: isize,
conj_b: bool,
beta: T,
c: *mut T,
rsc: isize,
csc: isize,
bias: *const T,
bias_dim: BiasDim,
has_bias: bool,
par: Parallelism,
) {
let epi = to_fused_epi_raw(bias, bias_dim, has_bias, None);
unsafe {
dispatch::execute_complex_fused(
conj_a,
conj_b,
Task {
m,
k,
n,
alpha,
a,
rsa,
csa,
b,
rsb,
csb,
beta,
c,
rsc,
csc,
},
epi,
par,
ws,
);
}
}
#[cfg(feature = "complex")]
#[allow(clippy::too_many_arguments)]
pub unsafe fn gemm_cplx_unchecked<T: ComplexScalar>(
m: usize,
k: usize,
n: usize,
alpha: T,
a: *const T,
rsa: isize,
csa: isize,
conj_a: bool,
b: *const T,
rsb: isize,
csb: isize,
conj_b: bool,
beta: T,
c: *mut T,
rsc: isize,
csc: isize,
par: Parallelism,
) {
unsafe {
workspace::with_thread_pool(|ws| {
gemm_cplx_unchecked_with(
ws, m, k, n, alpha, a, rsa, csa, conj_a, b, rsb, csb, conj_b, beta, c, rsc, csc,
par,
);
});
}
}
#[cfg(feature = "complex")]
#[allow(clippy::too_many_arguments)]
pub unsafe fn gemm_cplx_unchecked_with<T: ComplexScalar>(
ws: &mut Workspace,
m: usize,
k: usize,
n: usize,
alpha: T,
a: *const T,
rsa: isize,
csa: isize,
conj_a: bool,
b: *const T,
rsb: isize,
csb: isize,
conj_b: bool,
beta: T,
c: *mut T,
rsc: isize,
csc: isize,
par: Parallelism,
) {
unsafe {
dispatch::execute_complex(
conj_a,
conj_b,
Task {
m,
k,
n,
alpha,
a,
rsa,
csa,
b,
rsb,
csb,
beta,
c,
rsc,
csc,
},
par,
ws,
);
}
}