diff --git a/crates/yscv-kernels/README.md b/crates/yscv-kernels/README.md index 387128e..a8c9ffd 100644 --- a/crates/yscv-kernels/README.md +++ b/crates/yscv-kernels/README.md @@ -61,7 +61,9 @@ packing loop. Cascade dispatch: 48→32→16→8→4→scalar columns. k-loop unrolled by 2 with doubled accumulator sets to break FMA latency dependency chains (FMA latency = 4 cycles, 2 FMA ports on Zen 4 — spacing accumulator -reuse to 4+ cycles eliminates pipeline stalls). +reuse to 4+ cycles eliminates pipeline stalls). Four-row transposed-A tiles +prefetch strided B panels four K iterations ahead for the measured +`M >= 128, K <= 32` AVX/FMA range and in-order Cortex-A53/A55 kernels. ### Depthwise Conv SIMD diff --git a/crates/yscv-kernels/benches/kernels_cpu_ops.rs b/crates/yscv-kernels/benches/kernels_cpu_ops.rs index a06254e..4291e7c 100644 --- a/crates/yscv-kernels/benches/kernels_cpu_ops.rs +++ b/crates/yscv-kernels/benches/kernels_cpu_ops.rs @@ -2,12 +2,12 @@ use std::num::NonZeroUsize; use criterion::{Criterion, black_box, criterion_group, criterion_main}; use yscv_kernels::{ - Backend, BatchNorm2dParams, LayerNormLastDimParams, ParallelElementwiseConfig, + Backend, BatchNorm2dParams, BinaryKind, LayerNormLastDimParams, ParallelElementwiseConfig, ParallelMatmulConfig, SeparableConv2dParams, ThreadedCpuBackend, ThreadedCpuBackendConfig, add, - avg_pool2d_nhwc, batch_norm2d_nhwc, conv2d_nhwc, conv2d_nhwc_indirect_padded, - conv2d_nhwc_padded, depthwise_conv2d_nhwc, layer_norm_last_dim, log_softmax_last_dim, - logsumexp_last_dim, matmul_2d, matmul_2d_sequential, max_pool2d_nhwc, relu, - separable_conv2d_nhwc, sigmoid, softmax_last_dim, + avg_pool2d_nhwc, batch_norm2d_nhwc, binary_same_shape_dispatch, conv2d_nhwc, + conv2d_nhwc_indirect_padded, conv2d_nhwc_padded, depthwise_conv2d_nhwc, layer_norm_last_dim, + log_softmax_last_dim, logsumexp_last_dim, matmul_2d, matmul_2d_sequential, max_pool2d_nhwc, + relu, separable_conv2d_nhwc, sigmoid, softmax_last_dim, }; use yscv_tensor::Tensor; @@ -170,6 +170,27 @@ fn bench_elementwise_modes(c: &mut Criterion) { black_box(out); }); }); + let mut raw_out = vec![0.0; lhs.data().len()]; + group.bench_function("add_same_shape_raw_slice", |b| { + b.iter(|| { + binary_same_shape_dispatch( + black_box(lhs.data()), + black_box(rhs.data()), + black_box(&mut raw_out), + BinaryKind::Add, + ); + }); + }); + group.bench_function("mul_same_shape_raw_slice", |b| { + b.iter(|| { + binary_same_shape_dispatch( + black_box(lhs.data()), + black_box(rhs.data()), + black_box(&mut raw_out), + BinaryKind::Mul, + ); + }); + }); group.bench_function("add_same_shape_threaded_2", |b| { b.iter(|| { let out = threaded_backend diff --git a/crates/yscv-kernels/src/ops/matmul/trans_a.rs b/crates/yscv-kernels/src/ops/matmul/trans_a.rs index 66caafe..9a32f34 100644 --- a/crates/yscv-kernels/src/ops/matmul/trans_a.rs +++ b/crates/yscv-kernels/src/ops/matmul/trans_a.rs @@ -4,6 +4,25 @@ use super::*; +#[cfg(any(target_arch = "x86", target_arch = "x86_64", target_arch = "aarch64"))] +const MATMUL_PREFETCH_AHEAD: usize = 4; + +#[cfg(any(target_arch = "x86", target_arch = "x86_64", target_arch = "aarch64"))] +#[inline(always)] +#[allow(unsafe_code, unsafe_op_in_unsafe_fn)] +unsafe fn prefetch_l1_keep(p: *const f32) { + #[cfg(target_arch = "x86")] + std::arch::x86::_mm_prefetch::<{ std::arch::x86::_MM_HINT_T0 }>(p as *const i8); + #[cfg(target_arch = "x86_64")] + std::arch::x86_64::_mm_prefetch::<{ std::arch::x86_64::_MM_HINT_T0 }>(p as *const i8); + #[cfg(target_arch = "aarch64")] + core::arch::asm!( + "prfm pldl1keep, [{p}]", + p = in(reg) p, + options(nostack, preserves_flags, readonly), + ); +} + pub(super) fn non_trans_4row_disabled() -> bool { static CACHED: std::sync::OnceLock = std::sync::OnceLock::new(); *CACHED.get_or_init(|| std::env::var_os("YSCV_NON_TRANS_4ROW_OFF").is_some()) @@ -929,6 +948,8 @@ unsafe fn trans_a_4row_avx2( use std::arch::x86::*; #[cfg(target_arch = "x86_64")] use std::arch::x86_64::*; + let prefetch = m >= 128 && n > 16 && k <= 32; + let prefetch_end = k.saturating_sub(MATMUL_PREFETCH_AHEAD); unsafe { let r0 = out_4rows.as_mut_ptr(); let r1 = r0.add(n); @@ -955,6 +976,9 @@ unsafe fn trans_a_4row_avx2( let a1 = _mm256_set1_ps(*a_row.add(1)); let a2 = _mm256_set1_ps(*a_row.add(2)); let a3 = _mm256_set1_ps(*a_row.add(3)); + if prefetch && ki < prefetch_end { + prefetch_l1_keep(b.as_ptr().add((ki + MATMUL_PREFETCH_AHEAD) * n + col)); + } let bptr = b.as_ptr().add(ki * n + col); let b0 = _mm256_loadu_ps(bptr); let b1 = _mm256_loadu_ps(bptr.add(8)); @@ -1015,6 +1039,8 @@ unsafe fn trans_a_4row_neon( n: usize, out_4rows: &mut [f32], ) { + let prefetch = n > 16 && crate::host_cpu().uarch.is_in_order(); + let prefetch_end = k.saturating_sub(MATMUL_PREFETCH_AHEAD); unsafe { let r0 = out_4rows.as_mut_ptr(); let r1 = r0.add(n); @@ -1048,6 +1074,9 @@ unsafe fn trans_a_4row_neon( let a1 = vdupq_n_f32(*a_row.add(1)); let a2 = vdupq_n_f32(*a_row.add(2)); let a3 = vdupq_n_f32(*a_row.add(3)); + if prefetch && ki < prefetch_end { + prefetch_l1_keep(b.as_ptr().add((ki + MATMUL_PREFETCH_AHEAD) * n + col)); + } let bptr = b.as_ptr().add(ki * n + col); let b0 = vld1q_f32(bptr); let b1 = vld1q_f32(bptr.add(4)); diff --git a/crates/yscv-kernels/src/ops/simd/binary.rs b/crates/yscv-kernels/src/ops/simd/binary.rs index d9de455..dc4ab0d 100644 --- a/crates/yscv-kernels/src/ops/simd/binary.rs +++ b/crates/yscv-kernels/src/ops/simd/binary.rs @@ -507,23 +507,10 @@ unsafe fn binary_same_shape_avx(lhs: &[f32], rhs: &[f32], out: &mut [f32], kind: let out_ptr = out.as_mut_ptr(); let mut index = 0usize; - // 4x unrolled: process 32 floats per iteration with software prefetch. - // Matches vDSP throughput by keeping the OoO pipeline fully saturated. + // 4x unrolled: process 32 floats per iteration. match kind { BinaryKind::Add => { while index + 32 <= len { - #[cfg(target_arch = "x86")] - { - use std::arch::x86::_mm_prefetch; - _mm_prefetch::<3>(left_ptr.add(index + 32) as *const i8); - _mm_prefetch::<3>(right_ptr.add(index + 32) as *const i8); - } - #[cfg(target_arch = "x86_64")] - { - use std::arch::x86_64::_mm_prefetch; - _mm_prefetch::<3>(left_ptr.add(index + 32) as *const i8); - _mm_prefetch::<3>(right_ptr.add(index + 32) as *const i8); - } let a0 = _mm256_loadu_ps(left_ptr.add(index)); let b0 = _mm256_loadu_ps(right_ptr.add(index)); let a1 = _mm256_loadu_ps(left_ptr.add(index + 8)); @@ -541,18 +528,6 @@ unsafe fn binary_same_shape_avx(lhs: &[f32], rhs: &[f32], out: &mut [f32], kind: } BinaryKind::Sub => { while index + 32 <= len { - #[cfg(target_arch = "x86")] - { - use std::arch::x86::_mm_prefetch; - _mm_prefetch::<3>(left_ptr.add(index + 32) as *const i8); - _mm_prefetch::<3>(right_ptr.add(index + 32) as *const i8); - } - #[cfg(target_arch = "x86_64")] - { - use std::arch::x86_64::_mm_prefetch; - _mm_prefetch::<3>(left_ptr.add(index + 32) as *const i8); - _mm_prefetch::<3>(right_ptr.add(index + 32) as *const i8); - } let a0 = _mm256_loadu_ps(left_ptr.add(index)); let b0 = _mm256_loadu_ps(right_ptr.add(index)); let a1 = _mm256_loadu_ps(left_ptr.add(index + 8)); @@ -570,18 +545,6 @@ unsafe fn binary_same_shape_avx(lhs: &[f32], rhs: &[f32], out: &mut [f32], kind: } BinaryKind::Mul => { while index + 32 <= len { - #[cfg(target_arch = "x86")] - { - use std::arch::x86::_mm_prefetch; - _mm_prefetch::<3>(left_ptr.add(index + 32) as *const i8); - _mm_prefetch::<3>(right_ptr.add(index + 32) as *const i8); - } - #[cfg(target_arch = "x86_64")] - { - use std::arch::x86_64::_mm_prefetch; - _mm_prefetch::<3>(left_ptr.add(index + 32) as *const i8); - _mm_prefetch::<3>(right_ptr.add(index + 32) as *const i8); - } let a0 = _mm256_loadu_ps(left_ptr.add(index)); let b0 = _mm256_loadu_ps(right_ptr.add(index)); let a1 = _mm256_loadu_ps(left_ptr.add(index + 8)); diff --git a/crates/yscv-kernels/src/ops/simd/fma.rs b/crates/yscv-kernels/src/ops/simd/fma.rs index 01fa3a0..17e2ff7 100644 --- a/crates/yscv-kernels/src/ops/simd/fma.rs +++ b/crates/yscv-kernels/src/ops/simd/fma.rs @@ -18,6 +18,17 @@ use std::arch::x86_64::{ _mm256_storeu_ps, }; +#[cfg(target_arch = "aarch64")] +#[inline(always)] +#[allow(unsafe_code, unsafe_op_in_unsafe_fn)] +unsafe fn prefetch_l1_keep(p: *const f32) { + core::arch::asm!( + "prfm pldl1keep, [{p}]", + p = in(reg) p, + options(nostack, preserves_flags, readonly), + ); +} + // =========================================================================== // FMA dispatch (conv2d inner loop helper) // =========================================================================== @@ -698,12 +709,25 @@ unsafe fn matmul_row_set_avx_fma( n: usize, ) { #[cfg(target_arch = "x86")] - use std::arch::x86::_mm256_fmadd_ps; + use std::arch::x86::{_MM_HINT_T0, _mm_prefetch, _mm256_fmadd_ps}; #[cfg(target_arch = "x86_64")] - use std::arch::x86_64::_mm256_fmadd_ps; + use std::arch::x86_64::{_MM_HINT_T0, _mm_prefetch, _mm256_fmadd_ps}; let mut col = 0usize; while col + 48 <= n { + let k2 = k & !1; + let mut p = 0usize; + let mut bb = right; + let mut bc = right; + if k != 0 { + bb = right.add(col); + _mm_prefetch::<_MM_HINT_T0>(bb as *const i8); + if k2 != 0 { + bc = right.add(n + col); + _mm_prefetch::<_MM_HINT_T0>(bc as *const i8); + } + } + let mut a0 = _mm256_setzero_ps(); let mut a1 = _mm256_setzero_ps(); let mut a2 = _mm256_setzero_ps(); @@ -716,11 +740,19 @@ unsafe fn matmul_row_set_avx_fma( let mut b3 = _mm256_setzero_ps(); let mut b4 = _mm256_setzero_ps(); let mut b5 = _mm256_setzero_ps(); - let k2 = k & !1; - let mut p = 0usize; while p < k2 { + let mut next_bb = bb; + let mut next_bc = bc; + if p + 2 < k { + next_bb = bb.add(2 * n); + _mm_prefetch::<_MM_HINT_T0>(next_bb as *const i8); + if p + 3 < k { + next_bc = bc.add(2 * n); + _mm_prefetch::<_MM_HINT_T0>(next_bc as *const i8); + } + } + let va = _mm256_set1_ps(*left_row.add(p)); - let bb = right.add(p * n + col); a0 = _mm256_fmadd_ps(va, _mm256_loadu_ps(bb), a0); a1 = _mm256_fmadd_ps(va, _mm256_loadu_ps(bb.add(8)), a1); a2 = _mm256_fmadd_ps(va, _mm256_loadu_ps(bb.add(16)), a2); @@ -728,18 +760,18 @@ unsafe fn matmul_row_set_avx_fma( a4 = _mm256_fmadd_ps(va, _mm256_loadu_ps(bb.add(32)), a4); a5 = _mm256_fmadd_ps(va, _mm256_loadu_ps(bb.add(40)), a5); let vb = _mm256_set1_ps(*left_row.add(p + 1)); - let bc = right.add((p + 1) * n + col); b0 = _mm256_fmadd_ps(vb, _mm256_loadu_ps(bc), b0); b1 = _mm256_fmadd_ps(vb, _mm256_loadu_ps(bc.add(8)), b1); b2 = _mm256_fmadd_ps(vb, _mm256_loadu_ps(bc.add(16)), b2); b3 = _mm256_fmadd_ps(vb, _mm256_loadu_ps(bc.add(24)), b3); b4 = _mm256_fmadd_ps(vb, _mm256_loadu_ps(bc.add(32)), b4); b5 = _mm256_fmadd_ps(vb, _mm256_loadu_ps(bc.add(40)), b5); + bb = next_bb; + bc = next_bc; p += 2; } if p < k { let va = _mm256_set1_ps(*left_row.add(p)); - let bb = right.add(p * n + col); a0 = _mm256_fmadd_ps(va, _mm256_loadu_ps(bb), a0); a1 = _mm256_fmadd_ps(va, _mm256_loadu_ps(bb.add(8)), a1); a2 = _mm256_fmadd_ps(va, _mm256_loadu_ps(bb.add(16)), a2); @@ -756,6 +788,19 @@ unsafe fn matmul_row_set_avx_fma( col += 48; } while col + 32 <= n { + let k2 = k & !1; + let mut p = 0usize; + let mut bb = right; + let mut bc = right; + if k != 0 { + bb = right.add(col); + _mm_prefetch::<_MM_HINT_T0>(bb as *const i8); + if k2 != 0 { + bc = right.add(n + col); + _mm_prefetch::<_MM_HINT_T0>(bc as *const i8); + } + } + let mut a0 = _mm256_setzero_ps(); let mut a1 = _mm256_setzero_ps(); let mut a2 = _mm256_setzero_ps(); @@ -764,26 +809,34 @@ unsafe fn matmul_row_set_avx_fma( let mut b1 = _mm256_setzero_ps(); let mut b2 = _mm256_setzero_ps(); let mut b3 = _mm256_setzero_ps(); - let k2 = k & !1; - let mut p = 0usize; while p < k2 { + let mut next_bb = bb; + let mut next_bc = bc; + if p + 2 < k { + next_bb = bb.add(2 * n); + _mm_prefetch::<_MM_HINT_T0>(next_bb as *const i8); + if p + 3 < k { + next_bc = bc.add(2 * n); + _mm_prefetch::<_MM_HINT_T0>(next_bc as *const i8); + } + } + let va = _mm256_set1_ps(*left_row.add(p)); - let bb = right.add(p * n + col); a0 = _mm256_fmadd_ps(va, _mm256_loadu_ps(bb), a0); a1 = _mm256_fmadd_ps(va, _mm256_loadu_ps(bb.add(8)), a1); a2 = _mm256_fmadd_ps(va, _mm256_loadu_ps(bb.add(16)), a2); a3 = _mm256_fmadd_ps(va, _mm256_loadu_ps(bb.add(24)), a3); let vb = _mm256_set1_ps(*left_row.add(p + 1)); - let bc = right.add((p + 1) * n + col); b0 = _mm256_fmadd_ps(vb, _mm256_loadu_ps(bc), b0); b1 = _mm256_fmadd_ps(vb, _mm256_loadu_ps(bc.add(8)), b1); b2 = _mm256_fmadd_ps(vb, _mm256_loadu_ps(bc.add(16)), b2); b3 = _mm256_fmadd_ps(vb, _mm256_loadu_ps(bc.add(24)), b3); + bb = next_bb; + bc = next_bc; p += 2; } if p < k { let va = _mm256_set1_ps(*left_row.add(p)); - let bb = right.add(p * n + col); a0 = _mm256_fmadd_ps(va, _mm256_loadu_ps(bb), a0); a1 = _mm256_fmadd_ps(va, _mm256_loadu_ps(bb.add(8)), a1); a2 = _mm256_fmadd_ps(va, _mm256_loadu_ps(bb.add(16)), a2); @@ -796,26 +849,47 @@ unsafe fn matmul_row_set_avx_fma( col += 32; } while col + 16 <= n { + let k2 = k & !1; + let mut p = 0usize; + let mut bb = right; + let mut bc = right; + if k != 0 { + bb = right.add(col); + _mm_prefetch::<_MM_HINT_T0>(bb as *const i8); + if k2 != 0 { + bc = right.add(n + col); + _mm_prefetch::<_MM_HINT_T0>(bc as *const i8); + } + } + let mut a0 = _mm256_setzero_ps(); let mut a1 = _mm256_setzero_ps(); let mut b0 = _mm256_setzero_ps(); let mut b1 = _mm256_setzero_ps(); - let k2 = k & !1; - let mut p = 0usize; while p < k2 { + let mut next_bb = bb; + let mut next_bc = bc; + if p + 2 < k { + next_bb = bb.add(2 * n); + _mm_prefetch::<_MM_HINT_T0>(next_bb as *const i8); + if p + 3 < k { + next_bc = bc.add(2 * n); + _mm_prefetch::<_MM_HINT_T0>(next_bc as *const i8); + } + } + let va = _mm256_set1_ps(*left_row.add(p)); - let bb = right.add(p * n + col); a0 = _mm256_fmadd_ps(va, _mm256_loadu_ps(bb), a0); a1 = _mm256_fmadd_ps(va, _mm256_loadu_ps(bb.add(8)), a1); let vb = _mm256_set1_ps(*left_row.add(p + 1)); - let bc = right.add((p + 1) * n + col); b0 = _mm256_fmadd_ps(vb, _mm256_loadu_ps(bc), b0); b1 = _mm256_fmadd_ps(vb, _mm256_loadu_ps(bc.add(8)), b1); + bb = next_bb; + bc = next_bc; p += 2; } if p < k { let va = _mm256_set1_ps(*left_row.add(p)); - let bb = right.add(p * n + col); a0 = _mm256_fmadd_ps(va, _mm256_loadu_ps(bb), a0); a1 = _mm256_fmadd_ps(va, _mm256_loadu_ps(bb.add(8)), a1); } @@ -970,13 +1044,25 @@ unsafe fn matmul_row_set_neon( k: usize, n: usize, ) { + const PREFETCH_AHEAD: usize = 8; + let mut col = 0usize; while col + 4 <= n { let mut acc = vdupq_n_f32(0.0); - for p in 0..k { + let prefetch_end = k.saturating_sub(PREFETCH_AHEAD); + let mut p = 0usize; + while p < prefetch_end { + prefetch_l1_keep(right.add((p + PREFETCH_AHEAD) * n + col)); + let a = vdupq_n_f32(*left_row.add(p)); + let b = vld1q_f32(right.add(p * n + col)); + acc = vfmaq_f32(acc, a, b); + p += 1; + } + while p < k { let a = vdupq_n_f32(*left_row.add(p)); let b = vld1q_f32(right.add(p * n + col)); acc = vfmaq_f32(acc, a, b); + p += 1; } vst1q_f32(out_row.add(col), acc); col += 4;