Skip to main content

ab_riscv_interpreter/v/zvexx/reduction/
zvexx_reduction_helpers.rs

1//! Opaque helpers for ZveXx extension
2use crate::v::vector_registers::VectorRegistersExt;
3use crate::v::zvexx::arith::zvexx_arith_helpers::{
4    read_element_u64, sign_extend, write_element_u64,
5};
6use crate::v::zvexx::load::zvexx_load_helpers::{mask_bit, snapshot_mask};
7use ab_riscv_primitives::prelude::*;
8use core::fmt;
9use core::hint::cold_path;
10
11/// Execute a single-width integer reduction.
12///
13/// # Safety
14/// - `vs2.to_bits() % group_regs == 0` and `vs2.to_bits() + group_regs <= 32` (verified by caller)
15/// - `vstart == 0` (verified by caller; reductions with non-zero vstart are illegal)
16/// - `vl <= group_regs * VLEN.bytes() / sew_bytes`
17/// - `vl <= VLEN`
18#[inline(always)]
19#[expect(clippy::too_many_arguments, reason = "Internal API")]
20#[doc(hidden)]
21// TODO: #[cfg_attr(feature = "no-panic", no_panic_const::no_panic)]
22pub unsafe fn execute_reduce_op<Reg, ExtState, CustomError, F>(
23    ext_state: &mut ExtState,
24    vd: VReg,
25    vs2: VReg,
26    vs1: VReg,
27    vm: bool,
28    vl: Vl,
29    sew: Vsew,
30    op: F,
31) where
32    Reg: Register,
33    ExtState: VectorRegistersExt<Reg, CustomError>,
34    [(); SUPPORTED_ELEN_VLEN::<{ ExtState::ELEN }, { ExtState::VLEN }>]:,
35    CustomError: fmt::Debug,
36    F: Fn(u64, u64, Vsew) -> u64,
37{
38    // Spec ยง5.4: when vstart >= vl, no element of vd is updated. For reductions this means
39    // vl == 0 (since caller has verified vstart == 0). In that case we must not write vd and
40    // must not mark vs dirty.
41    if vl == Vl::ZERO {
42        cold_path();
43        ext_state.reset_vstart();
44        return;
45    }
46    // SAFETY: element 0 always fits within register vs1
47    let init = unsafe { read_element_u64(ext_state.read_vregs(), vs1, 0, sew) };
48    // SAFETY: `vl <= VLEN`
49    let mask_buf = unsafe { snapshot_mask(ext_state.read_vregs(), vm, vl) };
50    let mut acc = init;
51    for i in Vstart::ZERO.range_to(vl) {
52        if !mask_bit(&mask_buf, i) {
53            continue;
54        }
55        // SAFETY: `vs2 % group_regs == 0` and `i < vl <= group_regs * elems_per_reg`
56        let elem = unsafe { read_element_u64(ext_state.read_vregs(), vs2, i, sew) };
57        acc = op(acc, elem, sew);
58    }
59    // SAFETY: element 0 always fits within register vd
60    unsafe {
61        write_element_u64(ext_state.write_vregs(), vd, 0, sew, acc);
62    }
63    ext_state.mark_vs_dirty();
64    ext_state.reset_vstart();
65}
66
67/// Execute a widening integer sum reduction.
68///
69/// # Safety
70/// - `vs2.to_bits() % group_regs == 0` and `vs2.to_bits() + group_regs <= 32` (verified by caller)
71/// - `sew.double_width().is_some()` (verified by caller)
72/// - `vstart == 0` (verified by caller)
73/// - `vl <= group_regs * VLEN.bytes() / sew_bytes`
74/// - `vl <= VLEN`
75#[inline(always)]
76#[expect(clippy::too_many_arguments, reason = "Internal API")]
77#[doc(hidden)]
78// TODO: #[cfg_attr(feature = "no-panic", no_panic_const::no_panic)]
79pub unsafe fn execute_widening_reduce_op<
80    const SIGN_EXTEND_SRC: bool,
81    Reg,
82    ExtState,
83    CustomError,
84    F,
85>(
86    ext_state: &mut ExtState,
87    vd: VReg,
88    vs2: VReg,
89    vs1: VReg,
90    vm: bool,
91    vl: Vl,
92    sew: Vsew,
93    op: F,
94) where
95    Reg: Register,
96    ExtState: VectorRegistersExt<Reg, CustomError>,
97    [(); SUPPORTED_ELEN_VLEN::<{ ExtState::ELEN }, { ExtState::VLEN }>]:,
98    CustomError: fmt::Debug,
99    F: Fn(u64, u64, Vsew) -> u64,
100{
101    let Some(wide_sew) = sew.double_width() else {
102        // SAFETY: caller verified `2*SEW <= ELEN`; E64 widening is unreachable here
103        unsafe { core::hint::unreachable_unchecked() }
104    };
105    if vl == Vl::ZERO {
106        cold_path();
107        ext_state.reset_vstart();
108        return;
109    }
110    // SAFETY: element 0 always fits within register vs1
111    let init = unsafe { read_element_u64(ext_state.read_vregs(), vs1, 0, wide_sew) };
112    // SAFETY: `vl <= VLEN`
113    let mask_buf = unsafe { snapshot_mask(ext_state.read_vregs(), vm, vl) };
114    let mut acc = init;
115    for i in Vstart::ZERO.range_to(vl) {
116        if !mask_bit(&mask_buf, i) {
117            continue;
118        }
119        // SAFETY: same bounds argument as `execute_reduce_op`
120        let raw = unsafe { read_element_u64(ext_state.read_vregs(), vs2, i, sew) };
121        let elem = if SIGN_EXTEND_SRC {
122            sign_extend(raw, sew).cast_unsigned()
123        } else {
124            raw
125        };
126        acc = op(acc, elem, wide_sew);
127    }
128    // SAFETY: element 0 always fits within register vd
129    unsafe {
130        write_element_u64(ext_state.write_vregs(), vd, 0, wide_sew, acc);
131    }
132    ext_state.mark_vs_dirty();
133    ext_state.reset_vstart();
134}