ab_riscv_interpreter/v/zvexx/reduction/
zvexx_reduction_helpers.rs1use 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#[inline(always)]
19#[expect(clippy::too_many_arguments, reason = "Internal API")]
20#[doc(hidden)]
21pub 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 if vl == Vl::ZERO {
42 cold_path();
43 ext_state.reset_vstart();
44 return;
45 }
46 let init = unsafe { read_element_u64(ext_state.read_vregs(), vs1, 0, sew) };
48 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 let elem = unsafe { read_element_u64(ext_state.read_vregs(), vs2, i, sew) };
57 acc = op(acc, elem, sew);
58 }
59 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#[inline(always)]
76#[expect(clippy::too_many_arguments, reason = "Internal API")]
77#[doc(hidden)]
78pub 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 unsafe { core::hint::unreachable_unchecked() }
104 };
105 if vl == Vl::ZERO {
106 cold_path();
107 ext_state.reset_vstart();
108 return;
109 }
110 let init = unsafe { read_element_u64(ext_state.read_vregs(), vs1, 0, wide_sew) };
112 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 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 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}