use antecedent_core::KernelPolicy;
use crate::view::{BitMaskView, F64VectorView};
#[derive(Clone, Copy, Debug, Eq, PartialEq, Hash)]
pub enum KernelImpl {
Scalar,
PortableOptimized,
ArchSimd,
}
#[must_use]
pub fn arch_simd_available() -> bool {
false
}
#[must_use]
pub fn select_impl(policy: &KernelPolicy) -> KernelImpl {
if policy.force_scalar {
return KernelImpl::Scalar;
}
let arch_ok = policy.allow_arch_simd && arch_simd_available();
let portable_ok = policy.allow_portable_optimized;
if !arch_ok && !portable_ok {
return KernelImpl::Scalar;
}
if arch_ok {
return KernelImpl::ArchSimd;
}
if portable_ok { KernelImpl::PortableOptimized } else { KernelImpl::Scalar }
}
fn portable_or_scalar_reductions(policy: KernelPolicy) -> KernelImpl {
match select_impl(&policy) {
KernelImpl::ArchSimd => KernelImpl::PortableOptimized,
other => other,
}
}
#[must_use]
pub fn masked_sum(
policy: &KernelPolicy,
x: F64VectorView<'_>,
mask: Option<BitMaskView<'_>>,
) -> f64 {
match portable_or_scalar_reductions(*policy) {
KernelImpl::Scalar => crate::scalar::masked_sum(x, mask),
KernelImpl::PortableOptimized | KernelImpl::ArchSimd => {
crate::portable::masked_sum(x, mask)
}
}
}
#[must_use]
pub fn masked_mean(
policy: &KernelPolicy,
x: F64VectorView<'_>,
mask: Option<BitMaskView<'_>>,
) -> Option<f64> {
match portable_or_scalar_reductions(*policy) {
KernelImpl::Scalar => crate::scalar::masked_mean(x, mask),
KernelImpl::PortableOptimized | KernelImpl::ArchSimd => {
crate::portable::masked_mean(x, mask)
}
}
}
#[must_use]
pub fn masked_variance(
policy: &KernelPolicy,
x: F64VectorView<'_>,
mask: Option<BitMaskView<'_>>,
) -> Option<f64> {
match portable_or_scalar_reductions(*policy) {
KernelImpl::Scalar => crate::scalar::masked_variance(x, mask),
KernelImpl::PortableOptimized | KernelImpl::ArchSimd => {
crate::portable::masked_variance(x, mask)
}
}
}
pub fn gather(policy: &KernelPolicy, src: F64VectorView<'_>, indices: &[usize], out: &mut [f64]) {
match portable_or_scalar_reductions(*policy) {
KernelImpl::Scalar => crate::scalar::gather(src, indices, out),
KernelImpl::PortableOptimized | KernelImpl::ArchSimd => {
crate::portable::gather(src, indices, out);
}
}
}
pub fn copy_vec(policy: &KernelPolicy, src: F64VectorView<'_>, dst: &mut [f64]) {
match portable_or_scalar_reductions(*policy) {
KernelImpl::Scalar => crate::scalar::copy_vec(src, dst),
KernelImpl::PortableOptimized | KernelImpl::ArchSimd => crate::portable::copy_vec(src, dst),
}
}
#[must_use]
pub fn partial_correlation(
policy: &KernelPolicy,
x: &[f64],
y: &[f64],
z_cols: &[&[f64]],
workspace: &mut crate::parcorr::ParCorrWorkspace,
) -> Option<f64> {
match portable_or_scalar_reductions(*policy) {
KernelImpl::Scalar => crate::parcorr::partial_correlation_scalar(x, y, z_cols, workspace),
KernelImpl::PortableOptimized | KernelImpl::ArchSimd => {
crate::parcorr::partial_correlation_portable(x, y, z_cols, workspace)
}
}
}
#[must_use]
pub fn masked_covariance(
policy: &KernelPolicy,
x: F64VectorView<'_>,
y: F64VectorView<'_>,
mask: Option<BitMaskView<'_>>,
) -> Option<f64> {
match portable_or_scalar_reductions(*policy) {
KernelImpl::Scalar => crate::scalar::masked_covariance(x, y, mask),
KernelImpl::PortableOptimized | KernelImpl::ArchSimd => {
crate::portable::masked_covariance(x, y, mask)
}
}
}
#[must_use]
pub fn standardize_inplace(policy: &KernelPolicy, x: &mut [f64], eps: f64) -> (f64, f64) {
match portable_or_scalar_reductions(*policy) {
KernelImpl::Scalar => crate::scalar::standardize_inplace(x, eps),
KernelImpl::PortableOptimized | KernelImpl::ArchSimd => {
crate::portable::standardize_inplace(x, eps)
}
}
}
pub fn pairwise_l1_fill(policy: &KernelPolicy, x: &[f64], out: &mut [f64]) {
match portable_or_scalar_reductions(*policy) {
KernelImpl::Scalar => crate::scalar::pairwise_l1_fill(x, out),
KernelImpl::PortableOptimized | KernelImpl::ArchSimd => {
crate::portable::pairwise_l1_fill(x, out);
}
}
}
pub fn accumulate_contingency(
policy: &KernelPolicy,
x_codes: &[u32],
y_codes: &[u32],
out: &mut [f64],
n_y_levels: usize,
) {
match portable_or_scalar_reductions(*policy) {
KernelImpl::Scalar => {
crate::scalar::accumulate_contingency(x_codes, y_codes, out, n_y_levels);
}
KernelImpl::PortableOptimized | KernelImpl::ArchSimd => {
crate::portable::accumulate_contingency(x_codes, y_codes, out, n_y_levels);
}
}
}
pub fn accumulate_contingency_rows(
policy: &KernelPolicy,
x_codes: &[u32],
y_codes: &[u32],
rows: &[usize],
out: &mut [f64],
n_y_levels: usize,
) {
match portable_or_scalar_reductions(*policy) {
KernelImpl::Scalar => {
crate::scalar::accumulate_contingency_rows(x_codes, y_codes, rows, out, n_y_levels);
}
KernelImpl::PortableOptimized | KernelImpl::ArchSimd => {
crate::portable::accumulate_contingency_rows(x_codes, y_codes, rows, out, n_y_levels);
}
}
}
#[must_use]
pub fn weighted_sum(policy: &KernelPolicy, x: &[f64], weights: &[f64]) -> f64 {
match portable_or_scalar_reductions(*policy) {
KernelImpl::Scalar => crate::scalar::weighted_sum(x, weights),
KernelImpl::PortableOptimized | KernelImpl::ArchSimd => {
crate::portable::weighted_sum(x, weights)
}
}
}
#[must_use]
pub fn weighted_mean(policy: &KernelPolicy, x: &[f64], weights: &[f64]) -> Option<f64> {
match portable_or_scalar_reductions(*policy) {
KernelImpl::Scalar => crate::scalar::weighted_mean(x, weights),
KernelImpl::PortableOptimized | KernelImpl::ArchSimd => {
crate::portable::weighted_mean(x, weights)
}
}
}
#[must_use]
pub fn weighted_dot(policy: &KernelPolicy, x: &[f64], y: &[f64], weights: &[f64]) -> f64 {
match portable_or_scalar_reductions(*policy) {
KernelImpl::Scalar => crate::scalar::weighted_dot(x, y, weights),
KernelImpl::PortableOptimized | KernelImpl::ArchSimd => {
crate::portable::weighted_dot(x, y, weights)
}
}
}