use crate::semiring::Semiring;
use num_traits::Zero;
pub trait SimdOps<W: Semiring> {
fn simd_plus(left: &[W], right: &[W], result: &mut [W]);
fn simd_times(left: &[W], right: &[W], result: &mut [W]);
fn simd_min(weights: &[W]) -> W;
fn simd_max(weights: &[W]) -> W;
}
pub mod vectorized_arcs {
use crate::arc::Arc;
use crate::semiring::Semiring;
pub fn parallel_transform<W, F>(arcs: Vec<Arc<W>>, transform: F) -> Vec<Arc<W>>
where
W: Semiring + Send + Sync,
F: Fn(Arc<W>) -> Arc<W> + Send + Sync,
{
const CHUNK_SIZE: usize = 64;
#[cfg(feature = "rayon")]
{
use rayon::prelude::*;
if arcs.len() > CHUNK_SIZE * 4 {
arcs.par_chunks(CHUNK_SIZE)
.flat_map(|chunk| chunk.iter().cloned().map(&transform).collect::<Vec<_>>())
.collect()
} else {
arcs.into_iter().map(transform).collect()
}
}
#[cfg(not(feature = "rayon"))]
{
arcs.chunks(CHUNK_SIZE)
.flat_map(|chunk| chunk.iter().cloned().map(&transform).collect::<Vec<_>>())
.collect()
}
}
pub fn batch_weight_operation<W, F>(arcs: &mut [Arc<W>], weight_op: F)
where
W: Semiring + Copy,
F: Fn(&mut [W]),
{
let mut weights: Vec<W> = arcs.iter().map(|arc| arc.weight).collect();
weight_op(&mut weights);
for (arc, &weight) in arcs.iter_mut().zip(weights.iter()) {
arc.weight = weight;
}
}
pub fn simd_sort_by_weight<W: Semiring + PartialOrd>(arcs: &mut [Arc<W>]) {
if arcs.len() <= 16 {
for i in 1..arcs.len() {
let mut j = i;
while j > 0 && arcs[j - 1].weight > arcs[j].weight {
arcs.swap(j - 1, j);
j -= 1;
}
}
return;
}
const CHUNK_SIZE: usize = 16;
for chunk in arcs.chunks_mut(CHUNK_SIZE) {
for i in 1..chunk.len() {
let mut j = i;
while j > 0 && chunk[j - 1].weight > chunk[j].weight {
chunk.swap(j - 1, j);
j -= 1;
}
}
}
let mut chunk_size = CHUNK_SIZE;
while chunk_size < arcs.len() {
let mut i = 0;
while i < arcs.len() {
let mid = (i + chunk_size).min(arcs.len());
let end = (i + chunk_size * 2).min(arcs.len());
if mid < end {
merge_sorted_chunks(&mut arcs[i..end], mid - i);
}
i += chunk_size * 2;
}
chunk_size *= 2;
}
}
fn merge_sorted_chunks<W: Semiring + PartialOrd>(slice: &mut [Arc<W>], mid: usize) {
let left: Vec<_> = slice[..mid].to_vec();
let mut i = 0; let mut j = mid; let mut k = 0;
while i < left.len() && j < slice.len() {
if left[i].weight <= slice[j].weight {
slice[k] = left[i].clone();
i += 1;
} else {
slice[k] = slice[j].clone();
j += 1;
}
k += 1;
}
while i < left.len() {
slice[k] = left[i].clone();
i += 1;
k += 1;
}
}
}
pub mod prefetch {
pub fn prefetch_cache_line<T>(data: &T) {
#[cfg(target_arch = "x86_64")]
unsafe {
std::arch::x86_64::_mm_prefetch(
data as *const T as *const i8,
std::arch::x86_64::_MM_HINT_T0,
);
}
#[cfg(not(target_arch = "x86_64"))]
{
let _ = data;
}
}
pub fn prefetch_sequential<T>(data: &[T], start_idx: usize, count: usize) {
let end_idx = (start_idx + count).min(data.len());
let step_size = 64 / std::mem::size_of::<T>();
for i in (start_idx..end_idx).step_by(step_size.max(1)) {
prefetch_cache_line(&data[i]);
}
}
pub fn prefetch_strided<T>(data: &[T], start_idx: usize, stride: usize, count: usize) {
for i in 0..count {
let idx = start_idx + i * stride;
if idx < data.len() {
prefetch_cache_line(&data[idx]);
}
}
}
}
pub use prefetch::prefetch_cache_line;
#[cfg(target_arch = "x86_64")]
mod simd_features {
use std::sync::atomic::{AtomicU8, Ordering};
const UNKNOWN: u8 = 0;
const SSE_ONLY: u8 = 1;
const AVX2: u8 = 2;
const AVX512: u8 = 3;
static DETECTED_LEVEL: AtomicU8 = AtomicU8::new(UNKNOWN);
pub fn get_simd_level() -> u8 {
let cached = DETECTED_LEVEL.load(Ordering::Relaxed);
if cached != UNKNOWN {
return cached;
}
let level = detect_simd_level();
DETECTED_LEVEL.store(level, Ordering::Relaxed);
level
}
fn detect_simd_level() -> u8 {
if is_x86_feature_detected!("avx512f") && is_x86_feature_detected!("avx512vl") {
AVX512
} else if is_x86_feature_detected!("avx2") && is_x86_feature_detected!("fma") {
AVX2
} else {
SSE_ONLY
}
}
#[allow(dead_code)]
pub const fn sse_only() -> u8 {
SSE_ONLY
}
pub const fn avx2() -> u8 {
AVX2
}
pub const fn avx512() -> u8 {
AVX512
}
}
#[cfg(target_arch = "x86_64")]
mod tropical_simd {
use super::simd_features;
use super::*;
use crate::semiring::TropicalWeight;
impl SimdOps<TropicalWeight> for TropicalWeight {
fn simd_plus(
left: &[TropicalWeight],
right: &[TropicalWeight],
result: &mut [TropicalWeight],
) {
assert_eq!(left.len(), right.len());
assert_eq!(left.len(), result.len());
let level = simd_features::get_simd_level();
if level >= simd_features::avx512() {
unsafe { simd_plus_avx512(left, right, result) }
} else if level >= simd_features::avx2() {
unsafe { simd_plus_avx2(left, right, result) }
} else {
unsafe { simd_plus_sse(left, right, result) }
}
}
fn simd_times(
left: &[TropicalWeight],
right: &[TropicalWeight],
result: &mut [TropicalWeight],
) {
assert_eq!(left.len(), right.len());
assert_eq!(left.len(), result.len());
let level = simd_features::get_simd_level();
if level >= simd_features::avx512() {
unsafe { simd_times_avx512(left, right, result) }
} else if level >= simd_features::avx2() {
unsafe { simd_times_avx2(left, right, result) }
} else {
unsafe { simd_times_sse(left, right, result) }
}
}
fn simd_min(weights: &[TropicalWeight]) -> TropicalWeight {
if weights.is_empty() {
return TropicalWeight::zero();
}
let level = simd_features::get_simd_level();
if level >= simd_features::avx512() {
unsafe { simd_min_avx512(weights) }
} else if level >= simd_features::avx2() {
unsafe { simd_min_avx2(weights) }
} else {
unsafe { simd_min_sse(weights) }
}
}
fn simd_max(weights: &[TropicalWeight]) -> TropicalWeight {
if weights.is_empty() {
return TropicalWeight::zero();
}
let level = simd_features::get_simd_level();
if level >= simd_features::avx512() {
unsafe { simd_max_avx512(weights) }
} else if level >= simd_features::avx2() {
unsafe { simd_max_avx2(weights) }
} else {
unsafe { simd_max_sse(weights) }
}
}
}
#[target_feature(enable = "sse")]
unsafe fn simd_plus_sse(
left: &[TropicalWeight],
right: &[TropicalWeight],
result: &mut [TropicalWeight],
) {
use std::arch::x86_64::{_mm_loadu_ps, _mm_min_ps, _mm_storeu_ps};
let len = left.len();
let simd_len = len & !3;
for i in (0..simd_len).step_by(4) {
let left_array = [
*left[i].value(),
*left[i + 1].value(),
*left[i + 2].value(),
*left[i + 3].value(),
];
let right_array = [
*right[i].value(),
*right[i + 1].value(),
*right[i + 2].value(),
*right[i + 3].value(),
];
let left_vals = _mm_loadu_ps(left_array.as_ptr());
let right_vals = _mm_loadu_ps(right_array.as_ptr());
let min_vals = _mm_min_ps(left_vals, right_vals);
let mut result_array = [0.0f32; 4];
_mm_storeu_ps(result_array.as_mut_ptr(), min_vals);
for j in 0..4 {
result[i + j] = TropicalWeight::new(result_array[j]);
}
}
for i in simd_len..len {
result[i] = left[i].plus(&right[i]);
}
}
#[target_feature(enable = "sse")]
unsafe fn simd_times_sse(
left: &[TropicalWeight],
right: &[TropicalWeight],
result: &mut [TropicalWeight],
) {
use std::arch::x86_64::{_mm_add_ps, _mm_loadu_ps, _mm_storeu_ps};
let len = left.len();
let simd_len = len & !3;
for i in (0..simd_len).step_by(4) {
let left_array = [
*left[i].value(),
*left[i + 1].value(),
*left[i + 2].value(),
*left[i + 3].value(),
];
let right_array = [
*right[i].value(),
*right[i + 1].value(),
*right[i + 2].value(),
*right[i + 3].value(),
];
let left_vals = _mm_loadu_ps(left_array.as_ptr());
let right_vals = _mm_loadu_ps(right_array.as_ptr());
let sum_vals = _mm_add_ps(left_vals, right_vals);
let mut result_array = [0.0f32; 4];
_mm_storeu_ps(result_array.as_mut_ptr(), sum_vals);
for j in 0..4 {
result[i + j] = TropicalWeight::new(result_array[j]);
}
}
for i in simd_len..len {
result[i] = left[i].times(&right[i]);
}
}
#[target_feature(enable = "sse")]
unsafe fn simd_min_sse(weights: &[TropicalWeight]) -> TropicalWeight {
use std::arch::x86_64::{_mm_loadu_ps, _mm_min_ps, _mm_set1_ps, _mm_storeu_ps};
let len = weights.len();
let simd_len = len & !3;
let mut min_vec = _mm_set1_ps(f32::INFINITY);
for i in (0..simd_len).step_by(4) {
let vals_array = [
*weights[i].value(),
*weights[i + 1].value(),
*weights[i + 2].value(),
*weights[i + 3].value(),
];
let vals = _mm_loadu_ps(vals_array.as_ptr());
min_vec = _mm_min_ps(min_vec, vals);
}
let mut result_array = [0.0f32; 4];
_mm_storeu_ps(result_array.as_mut_ptr(), min_vec);
let mut min_val = result_array[0];
for &val in &result_array[1..] {
min_val = min_val.min(val);
}
for weight in weights.iter().skip(simd_len) {
min_val = min_val.min(*weight.value());
}
TropicalWeight::new(min_val)
}
#[target_feature(enable = "sse")]
unsafe fn simd_max_sse(weights: &[TropicalWeight]) -> TropicalWeight {
use std::arch::x86_64::{_mm_loadu_ps, _mm_max_ps, _mm_set1_ps, _mm_storeu_ps};
let len = weights.len();
let simd_len = len & !3;
let mut max_vec = _mm_set1_ps(f32::NEG_INFINITY);
for i in (0..simd_len).step_by(4) {
let vals_array = [
*weights[i].value(),
*weights[i + 1].value(),
*weights[i + 2].value(),
*weights[i + 3].value(),
];
let vals = _mm_loadu_ps(vals_array.as_ptr());
max_vec = _mm_max_ps(max_vec, vals);
}
let mut result_array = [0.0f32; 4];
_mm_storeu_ps(result_array.as_mut_ptr(), max_vec);
let mut max_val = result_array[0];
for &val in &result_array[1..] {
max_val = max_val.max(val);
}
for weight in weights.iter().skip(simd_len) {
max_val = max_val.max(*weight.value());
}
TropicalWeight::new(max_val)
}
#[target_feature(enable = "avx2")]
unsafe fn simd_plus_avx2(
left: &[TropicalWeight],
right: &[TropicalWeight],
result: &mut [TropicalWeight],
) {
use std::arch::x86_64::{_mm256_loadu_ps, _mm256_min_ps, _mm256_storeu_ps};
let len = left.len();
let simd_len = len & !7;
for i in (0..simd_len).step_by(8) {
let left_array = [
*left[i].value(),
*left[i + 1].value(),
*left[i + 2].value(),
*left[i + 3].value(),
*left[i + 4].value(),
*left[i + 5].value(),
*left[i + 6].value(),
*left[i + 7].value(),
];
let right_array = [
*right[i].value(),
*right[i + 1].value(),
*right[i + 2].value(),
*right[i + 3].value(),
*right[i + 4].value(),
*right[i + 5].value(),
*right[i + 6].value(),
*right[i + 7].value(),
];
let left_vals = _mm256_loadu_ps(left_array.as_ptr());
let right_vals = _mm256_loadu_ps(right_array.as_ptr());
let min_vals = _mm256_min_ps(left_vals, right_vals);
let mut result_array = [0.0f32; 8];
_mm256_storeu_ps(result_array.as_mut_ptr(), min_vals);
for j in 0..8 {
result[i + j] = TropicalWeight::new(result_array[j]);
}
}
for i in simd_len..len {
result[i] = left[i].plus(&right[i]);
}
}
#[target_feature(enable = "avx2")]
unsafe fn simd_times_avx2(
left: &[TropicalWeight],
right: &[TropicalWeight],
result: &mut [TropicalWeight],
) {
use std::arch::x86_64::{_mm256_add_ps, _mm256_loadu_ps, _mm256_storeu_ps};
let len = left.len();
let simd_len = len & !7;
for i in (0..simd_len).step_by(8) {
let left_array = [
*left[i].value(),
*left[i + 1].value(),
*left[i + 2].value(),
*left[i + 3].value(),
*left[i + 4].value(),
*left[i + 5].value(),
*left[i + 6].value(),
*left[i + 7].value(),
];
let right_array = [
*right[i].value(),
*right[i + 1].value(),
*right[i + 2].value(),
*right[i + 3].value(),
*right[i + 4].value(),
*right[i + 5].value(),
*right[i + 6].value(),
*right[i + 7].value(),
];
let left_vals = _mm256_loadu_ps(left_array.as_ptr());
let right_vals = _mm256_loadu_ps(right_array.as_ptr());
let sum_vals = _mm256_add_ps(left_vals, right_vals);
let mut result_array = [0.0f32; 8];
_mm256_storeu_ps(result_array.as_mut_ptr(), sum_vals);
for j in 0..8 {
result[i + j] = TropicalWeight::new(result_array[j]);
}
}
for i in simd_len..len {
result[i] = left[i].times(&right[i]);
}
}
#[target_feature(enable = "avx2")]
unsafe fn simd_min_avx2(weights: &[TropicalWeight]) -> TropicalWeight {
use std::arch::x86_64::{_mm256_loadu_ps, _mm256_min_ps, _mm256_set1_ps, _mm256_storeu_ps};
let len = weights.len();
let simd_len = len & !7;
let mut min_vec = _mm256_set1_ps(f32::INFINITY);
for i in (0..simd_len).step_by(8) {
let vals_array = [
*weights[i].value(),
*weights[i + 1].value(),
*weights[i + 2].value(),
*weights[i + 3].value(),
*weights[i + 4].value(),
*weights[i + 5].value(),
*weights[i + 6].value(),
*weights[i + 7].value(),
];
let vals = _mm256_loadu_ps(vals_array.as_ptr());
min_vec = _mm256_min_ps(min_vec, vals);
}
let mut result_array = [0.0f32; 8];
_mm256_storeu_ps(result_array.as_mut_ptr(), min_vec);
let mut min_val = result_array[0];
for &val in &result_array[1..] {
min_val = min_val.min(val);
}
for weight in weights.iter().skip(simd_len) {
min_val = min_val.min(*weight.value());
}
TropicalWeight::new(min_val)
}
#[target_feature(enable = "avx2")]
unsafe fn simd_max_avx2(weights: &[TropicalWeight]) -> TropicalWeight {
use std::arch::x86_64::{_mm256_loadu_ps, _mm256_max_ps, _mm256_set1_ps, _mm256_storeu_ps};
let len = weights.len();
let simd_len = len & !7;
let mut max_vec = _mm256_set1_ps(f32::NEG_INFINITY);
for i in (0..simd_len).step_by(8) {
let vals_array = [
*weights[i].value(),
*weights[i + 1].value(),
*weights[i + 2].value(),
*weights[i + 3].value(),
*weights[i + 4].value(),
*weights[i + 5].value(),
*weights[i + 6].value(),
*weights[i + 7].value(),
];
let vals = _mm256_loadu_ps(vals_array.as_ptr());
max_vec = _mm256_max_ps(max_vec, vals);
}
let mut result_array = [0.0f32; 8];
_mm256_storeu_ps(result_array.as_mut_ptr(), max_vec);
let mut max_val = result_array[0];
for &val in &result_array[1..] {
max_val = max_val.max(val);
}
for weight in weights.iter().skip(simd_len) {
max_val = max_val.max(*weight.value());
}
TropicalWeight::new(max_val)
}
#[cfg(target_feature = "avx512f")]
#[target_feature(enable = "avx512f")]
unsafe fn simd_plus_avx512(
left: &[TropicalWeight],
right: &[TropicalWeight],
result: &mut [TropicalWeight],
) {
use std::arch::x86_64::{_mm512_loadu_ps, _mm512_min_ps, _mm512_storeu_ps};
let len = left.len();
let simd_len = len & !15;
for i in (0..simd_len).step_by(16) {
let mut left_array = [0.0f32; 16];
let mut right_array = [0.0f32; 16];
for j in 0..16 {
left_array[j] = *left[i + j].value();
right_array[j] = *right[i + j].value();
}
let left_vals = _mm512_loadu_ps(left_array.as_ptr());
let right_vals = _mm512_loadu_ps(right_array.as_ptr());
let min_vals = _mm512_min_ps(left_vals, right_vals);
let mut result_array = [0.0f32; 16];
_mm512_storeu_ps(result_array.as_mut_ptr(), min_vals);
for j in 0..16 {
result[i + j] = TropicalWeight::new(result_array[j]);
}
}
for i in simd_len..len {
result[i] = left[i].plus(&right[i]);
}
}
#[cfg(not(target_feature = "avx512f"))]
unsafe fn simd_plus_avx512(
left: &[TropicalWeight],
right: &[TropicalWeight],
result: &mut [TropicalWeight],
) {
simd_plus_avx2(left, right, result)
}
#[cfg(target_feature = "avx512f")]
#[target_feature(enable = "avx512f")]
unsafe fn simd_times_avx512(
left: &[TropicalWeight],
right: &[TropicalWeight],
result: &mut [TropicalWeight],
) {
use std::arch::x86_64::{_mm512_add_ps, _mm512_loadu_ps, _mm512_storeu_ps};
let len = left.len();
let simd_len = len & !15;
for i in (0..simd_len).step_by(16) {
let mut left_array = [0.0f32; 16];
let mut right_array = [0.0f32; 16];
for j in 0..16 {
left_array[j] = *left[i + j].value();
right_array[j] = *right[i + j].value();
}
let left_vals = _mm512_loadu_ps(left_array.as_ptr());
let right_vals = _mm512_loadu_ps(right_array.as_ptr());
let sum_vals = _mm512_add_ps(left_vals, right_vals);
let mut result_array = [0.0f32; 16];
_mm512_storeu_ps(result_array.as_mut_ptr(), sum_vals);
for j in 0..16 {
result[i + j] = TropicalWeight::new(result_array[j]);
}
}
for i in simd_len..len {
result[i] = left[i].times(&right[i]);
}
}
#[cfg(not(target_feature = "avx512f"))]
unsafe fn simd_times_avx512(
left: &[TropicalWeight],
right: &[TropicalWeight],
result: &mut [TropicalWeight],
) {
simd_times_avx2(left, right, result)
}
#[cfg(target_feature = "avx512f")]
#[target_feature(enable = "avx512f")]
unsafe fn simd_min_avx512(weights: &[TropicalWeight]) -> TropicalWeight {
use std::arch::x86_64::{_mm512_loadu_ps, _mm512_min_ps, _mm512_set1_ps, _mm512_storeu_ps};
let len = weights.len();
let simd_len = len & !15;
let mut min_vec = _mm512_set1_ps(f32::INFINITY);
for i in (0..simd_len).step_by(16) {
let mut vals_array = [0.0f32; 16];
for j in 0..16 {
vals_array[j] = *weights[i + j].value();
}
let vals = _mm512_loadu_ps(vals_array.as_ptr());
min_vec = _mm512_min_ps(min_vec, vals);
}
let mut result_array = [0.0f32; 16];
_mm512_storeu_ps(result_array.as_mut_ptr(), min_vec);
let mut min_val = result_array[0];
for &val in &result_array[1..] {
min_val = min_val.min(val);
}
for weight in weights.iter().skip(simd_len) {
min_val = min_val.min(*weight.value());
}
TropicalWeight::new(min_val)
}
#[cfg(not(target_feature = "avx512f"))]
unsafe fn simd_min_avx512(weights: &[TropicalWeight]) -> TropicalWeight {
simd_min_avx2(weights)
}
#[cfg(target_feature = "avx512f")]
#[target_feature(enable = "avx512f")]
unsafe fn simd_max_avx512(weights: &[TropicalWeight]) -> TropicalWeight {
use std::arch::x86_64::{_mm512_loadu_ps, _mm512_max_ps, _mm512_set1_ps, _mm512_storeu_ps};
let len = weights.len();
let simd_len = len & !15;
let mut max_vec = _mm512_set1_ps(f32::NEG_INFINITY);
for i in (0..simd_len).step_by(16) {
let mut vals_array = [0.0f32; 16];
for j in 0..16 {
vals_array[j] = *weights[i + j].value();
}
let vals = _mm512_loadu_ps(vals_array.as_ptr());
max_vec = _mm512_max_ps(max_vec, vals);
}
let mut result_array = [0.0f32; 16];
_mm512_storeu_ps(result_array.as_mut_ptr(), max_vec);
let mut max_val = result_array[0];
for &val in &result_array[1..] {
max_val = max_val.max(val);
}
for weight in weights.iter().skip(simd_len) {
max_val = max_val.max(*weight.value());
}
TropicalWeight::new(max_val)
}
#[cfg(not(target_feature = "avx512f"))]
unsafe fn simd_max_avx512(weights: &[TropicalWeight]) -> TropicalWeight {
simd_max_avx2(weights)
}
}
#[cfg(target_arch = "x86_64")]
mod log_simd {
use super::simd_features;
use super::*;
use crate::semiring::LogWeight;
impl SimdOps<LogWeight> for LogWeight {
fn simd_plus(left: &[LogWeight], right: &[LogWeight], result: &mut [LogWeight]) {
assert_eq!(left.len(), right.len());
assert_eq!(left.len(), result.len());
let level = simd_features::get_simd_level();
if level >= simd_features::avx2() {
unsafe { simd_log_sum_exp_avx2(left, right, result) }
} else {
for i in 0..left.len() {
result[i] = left[i].plus(&right[i]);
}
}
}
fn simd_times(left: &[LogWeight], right: &[LogWeight], result: &mut [LogWeight]) {
assert_eq!(left.len(), right.len());
assert_eq!(left.len(), result.len());
let level = simd_features::get_simd_level();
if level >= simd_features::avx2() {
unsafe { simd_log_times_avx2(left, right, result) }
} else {
unsafe { simd_log_times_sse(left, right, result) }
}
}
fn simd_min(weights: &[LogWeight]) -> LogWeight {
if weights.is_empty() {
return LogWeight::zero();
}
let mut min_val = f64::INFINITY;
for weight in weights {
let val = *weight.value();
if val < min_val {
min_val = val;
}
}
LogWeight::new(min_val)
}
fn simd_max(weights: &[LogWeight]) -> LogWeight {
if weights.is_empty() {
return LogWeight::zero();
}
let mut max_val = f64::NEG_INFINITY;
for weight in weights {
let val = *weight.value();
if val > max_val && !val.is_infinite() {
max_val = val;
}
}
if max_val == f64::NEG_INFINITY {
LogWeight::zero()
} else {
LogWeight::new(max_val)
}
}
}
#[target_feature(enable = "sse2")]
unsafe fn simd_log_times_sse(
left: &[LogWeight],
right: &[LogWeight],
result: &mut [LogWeight],
) {
use std::arch::x86_64::{_mm_add_pd, _mm_loadu_pd, _mm_storeu_pd};
let len = left.len();
let simd_len = len & !1;
for i in (0..simd_len).step_by(2) {
let left_array = [*left[i].value(), *left[i + 1].value()];
let right_array = [*right[i].value(), *right[i + 1].value()];
let left_vals = _mm_loadu_pd(left_array.as_ptr());
let right_vals = _mm_loadu_pd(right_array.as_ptr());
let sum_vals = _mm_add_pd(left_vals, right_vals);
let mut result_array = [0.0f64; 2];
_mm_storeu_pd(result_array.as_mut_ptr(), sum_vals);
result[i] = LogWeight::new(result_array[0]);
result[i + 1] = LogWeight::new(result_array[1]);
}
for i in simd_len..len {
result[i] = left[i].times(&right[i]);
}
}
#[target_feature(enable = "avx2")]
unsafe fn simd_log_times_avx2(
left: &[LogWeight],
right: &[LogWeight],
result: &mut [LogWeight],
) {
use std::arch::x86_64::{_mm256_add_pd, _mm256_loadu_pd, _mm256_storeu_pd};
let len = left.len();
let simd_len = len & !3;
for i in (0..simd_len).step_by(4) {
let left_array = [
*left[i].value(),
*left[i + 1].value(),
*left[i + 2].value(),
*left[i + 3].value(),
];
let right_array = [
*right[i].value(),
*right[i + 1].value(),
*right[i + 2].value(),
*right[i + 3].value(),
];
let left_vals = _mm256_loadu_pd(left_array.as_ptr());
let right_vals = _mm256_loadu_pd(right_array.as_ptr());
let sum_vals = _mm256_add_pd(left_vals, right_vals);
let mut result_array = [0.0f64; 4];
_mm256_storeu_pd(result_array.as_mut_ptr(), sum_vals);
for j in 0..4 {
result[i + j] = LogWeight::new(result_array[j]);
}
}
for i in simd_len..len {
result[i] = left[i].times(&right[i]);
}
}
#[target_feature(enable = "avx2")]
unsafe fn simd_log_sum_exp_avx2(
left: &[LogWeight],
right: &[LogWeight],
result: &mut [LogWeight],
) {
use std::arch::x86_64::{
_mm256_add_pd, _mm256_and_pd, _mm256_castsi256_pd, _mm256_loadu_pd, _mm256_max_pd,
_mm256_set1_epi64x, _mm256_storeu_pd, _mm256_sub_pd,
};
let len = left.len();
let simd_len = len & !3;
let abs_mask = _mm256_castsi256_pd(_mm256_set1_epi64x(0x7FFFFFFFFFFFFFFF_u64 as i64));
for i in (0..simd_len).step_by(4) {
let mut has_inf = false;
for j in 0..4 {
if left[i + j].value().is_infinite() || right[i + j].value().is_infinite() {
has_inf = true;
break;
}
}
if has_inf {
for j in 0..4 {
result[i + j] = left[i + j].plus(&right[i + j]);
}
continue;
}
let left_array = [
-*left[i].value(),
-*left[i + 1].value(),
-*left[i + 2].value(),
-*left[i + 3].value(),
];
let right_array = [
-*right[i].value(),
-*right[i + 1].value(),
-*right[i + 2].value(),
-*right[i + 3].value(),
];
let a = _mm256_loadu_pd(left_array.as_ptr());
let b = _mm256_loadu_pd(right_array.as_ptr());
let max_ab = _mm256_max_pd(a, b);
let diff = _mm256_sub_pd(a, b);
let abs_diff = _mm256_and_pd(abs_mask, diff);
let mut max_array = [0.0f64; 4];
let mut abs_diff_array = [0.0f64; 4];
_mm256_storeu_pd(max_array.as_mut_ptr(), max_ab);
_mm256_storeu_pd(abs_diff_array.as_mut_ptr(), abs_diff);
let mut log1pexp = [0.0f64; 4];
for j in 0..4 {
let x = abs_diff_array[j];
log1pexp[j] = (1.0 + (-x).exp()).ln();
}
let log1pexp_vec = _mm256_loadu_pd(log1pexp.as_ptr());
let sum = _mm256_add_pd(max_ab, log1pexp_vec);
let mut result_array = [0.0f64; 4];
_mm256_storeu_pd(result_array.as_mut_ptr(), sum);
for j in 0..4 {
result[i + j] = LogWeight::new(-result_array[j]);
}
}
for i in simd_len..len {
result[i] = left[i].plus(&right[i]);
}
}
}
#[cfg(not(target_arch = "x86_64"))]
mod fallback_simd {
use super::*;
use crate::semiring::{LogWeight, TropicalWeight};
impl SimdOps<TropicalWeight> for TropicalWeight {
fn simd_plus(
left: &[TropicalWeight],
right: &[TropicalWeight],
result: &mut [TropicalWeight],
) {
for ((l, r), res) in left.iter().zip(right.iter()).zip(result.iter_mut()) {
*res = l.plus(r);
}
}
fn simd_times(
left: &[TropicalWeight],
right: &[TropicalWeight],
result: &mut [TropicalWeight],
) {
for ((l, r), res) in left.iter().zip(right.iter()).zip(result.iter_mut()) {
*res = l.times(r);
}
}
fn simd_min(weights: &[TropicalWeight]) -> TropicalWeight {
weights
.iter()
.fold(TropicalWeight::zero(), |acc, w| acc.plus(w))
}
fn simd_max(weights: &[TropicalWeight]) -> TropicalWeight {
weights
.iter()
.fold(TropicalWeight::new(f32::NEG_INFINITY), |acc, w| {
if w.value() > acc.value() {
*w
} else {
acc
}
})
}
}
impl SimdOps<LogWeight> for LogWeight {
fn simd_plus(left: &[LogWeight], right: &[LogWeight], result: &mut [LogWeight]) {
for ((l, r), res) in left.iter().zip(right.iter()).zip(result.iter_mut()) {
*res = l.plus(r);
}
}
fn simd_times(left: &[LogWeight], right: &[LogWeight], result: &mut [LogWeight]) {
for ((l, r), res) in left.iter().zip(right.iter()).zip(result.iter_mut()) {
*res = l.times(r);
}
}
fn simd_min(weights: &[LogWeight]) -> LogWeight {
weights.iter().fold(LogWeight::zero(), |acc, w| {
if w.value() < acc.value() {
*w
} else {
acc
}
})
}
fn simd_max(weights: &[LogWeight]) -> LogWeight {
weights
.iter()
.fold(LogWeight::new(f64::NEG_INFINITY), |acc, w| {
if w.value() > acc.value() && !w.value().is_infinite() {
*w
} else {
acc
}
})
}
}
}
#[cfg(test)]
mod tests {
use super::prefetch;
use super::vectorized_arcs;
use crate::prelude::*;
use crate::semiring::LogWeight;
use num_traits::Zero;
#[test]
fn test_simd_plus() {
let left = vec![TropicalWeight::new(1.0), TropicalWeight::new(2.0)];
let right = vec![TropicalWeight::new(3.0), TropicalWeight::new(1.5)];
let mut result = vec![TropicalWeight::zero(); 2];
TropicalWeight::simd_plus(&left, &right, &mut result);
assert_eq!(result[0], TropicalWeight::new(1.0)); assert_eq!(result[1], TropicalWeight::new(1.5)); }
#[test]
fn test_simd_times() {
let left = vec![TropicalWeight::new(1.0), TropicalWeight::new(2.0)];
let right = vec![TropicalWeight::new(3.0), TropicalWeight::new(1.5)];
let mut result = vec![TropicalWeight::zero(); 2];
TropicalWeight::simd_times(&left, &right, &mut result);
assert_eq!(result[0], TropicalWeight::new(4.0)); assert_eq!(result[1], TropicalWeight::new(3.5)); }
#[test]
fn test_simd_min() {
let weights = vec![
TropicalWeight::new(3.0),
TropicalWeight::new(1.0),
TropicalWeight::new(2.0),
TropicalWeight::new(0.5),
];
let min_weight = TropicalWeight::simd_min(&weights);
assert_eq!(min_weight, TropicalWeight::new(0.5));
}
#[test]
fn test_simd_plus_large() {
let left: Vec<TropicalWeight> = (0..32).map(|i| TropicalWeight::new(i as f32)).collect();
let right: Vec<TropicalWeight> = (0..32)
.map(|i| TropicalWeight::new((31 - i) as f32))
.collect();
let mut result = vec![TropicalWeight::zero(); 32];
TropicalWeight::simd_plus(&left, &right, &mut result);
for (i, res) in result.iter().enumerate().take(32) {
let expected = (i as f32).min((31 - i) as f32);
assert_eq!(*res, TropicalWeight::new(expected));
}
}
#[test]
fn test_simd_times_large() {
let left: Vec<TropicalWeight> = (0..32).map(|i| TropicalWeight::new(i as f32)).collect();
let right: Vec<TropicalWeight> = (0..32).map(|i| TropicalWeight::new(i as f32)).collect();
let mut result = vec![TropicalWeight::zero(); 32];
TropicalWeight::simd_times(&left, &right, &mut result);
for (i, res) in result.iter().enumerate().take(32) {
let expected = (i as f32) + (i as f32);
assert_eq!(*res, TropicalWeight::new(expected));
}
}
#[test]
fn test_simd_min_large() {
let weights: Vec<TropicalWeight> =
(0..100).map(|i| TropicalWeight::new(i as f32)).collect();
let min_weight = TropicalWeight::simd_min(&weights);
assert_eq!(min_weight, TropicalWeight::new(0.0));
let weights: Vec<TropicalWeight> = (0..100)
.map(|i| TropicalWeight::new(100.0 - i as f32))
.collect();
let min_weight = TropicalWeight::simd_min(&weights);
assert_eq!(min_weight, TropicalWeight::new(1.0));
}
#[test]
fn test_simd_max_large() {
let weights: Vec<TropicalWeight> =
(0..100).map(|i| TropicalWeight::new(i as f32)).collect();
let max_weight = TropicalWeight::simd_max(&weights);
assert_eq!(max_weight, TropicalWeight::new(99.0));
}
#[test]
fn test_log_simd_times() {
let left = vec![LogWeight::new(1.0), LogWeight::new(2.0)];
let right = vec![LogWeight::new(3.0), LogWeight::new(1.5)];
let mut result = vec![LogWeight::zero(); 2];
LogWeight::simd_times(&left, &right, &mut result);
assert_eq!(result[0], LogWeight::new(4.0)); assert_eq!(result[1], LogWeight::new(3.5)); }
#[test]
fn test_log_simd_times_large() {
let left: Vec<LogWeight> = (0..32).map(|i| LogWeight::new(i as f64)).collect();
let right: Vec<LogWeight> = (0..32).map(|i| LogWeight::new(i as f64 * 0.5)).collect();
let mut result = vec![LogWeight::zero(); 32];
LogWeight::simd_times(&left, &right, &mut result);
for (i, res) in result.iter().enumerate().take(32) {
let expected = (i as f64) + (i as f64 * 0.5);
assert!(res.approx_eq(&LogWeight::new(expected), 1e-10));
}
}
#[test]
fn test_log_simd_plus() {
let left = vec![LogWeight::new(1.0), LogWeight::new(2.0)];
let right = vec![LogWeight::new(2.0), LogWeight::new(2.0)];
let mut result = vec![LogWeight::zero(); 2];
LogWeight::simd_plus(&left, &right, &mut result);
let expected0 = left[0].plus(&right[0]);
let expected1 = left[1].plus(&right[1]);
assert!(result[0].approx_eq(&expected0, 1e-10));
assert!(result[1].approx_eq(&expected1, 1e-10));
}
#[test]
fn test_log_simd_plus_large() {
let left: Vec<LogWeight> = (0..32).map(|i| LogWeight::new(i as f64)).collect();
let right: Vec<LogWeight> = (0..32).map(|i| LogWeight::new(i as f64 + 0.5)).collect();
let mut result = vec![LogWeight::zero(); 32];
LogWeight::simd_plus(&left, &right, &mut result);
for i in 0..32 {
let expected = left[i].plus(&right[i]);
assert!(
result[i].approx_eq(&expected, 1e-6),
"Mismatch at {}: got {:?}, expected {:?}",
i,
result[i],
expected
);
}
}
#[test]
fn test_log_simd_with_infinity() {
let left = vec![LogWeight::zero(), LogWeight::new(1.0)];
let right = vec![LogWeight::new(2.0), LogWeight::zero()];
let mut result = vec![LogWeight::new(999.0); 2];
LogWeight::simd_plus(&left, &right, &mut result);
assert_eq!(result[0], LogWeight::new(2.0));
assert_eq!(result[1], LogWeight::new(1.0));
}
#[test]
fn test_vectorized_transform() {
let arcs = vec![
Arc::new(1, 1, TropicalWeight::new(1.0), 0),
Arc::new(2, 2, TropicalWeight::new(2.0), 1),
];
let transformed = vectorized_arcs::parallel_transform(arcs, |arc| {
Arc::new(
arc.ilabel,
arc.olabel,
arc.weight.times(&TropicalWeight::new(0.5)),
arc.nextstate,
)
});
assert_eq!(transformed[0].weight, TropicalWeight::new(1.5));
assert_eq!(transformed[1].weight, TropicalWeight::new(2.5));
}
#[test]
fn test_prefetch() {
let data = vec![1, 2, 3, 4, 5];
prefetch_cache_line(&data[0]);
prefetch::prefetch_sequential(&data, 0, 5);
prefetch::prefetch_strided(&data, 0, 2, 3);
}
}