ab_riscv_interpreter/zvbb/zvbb_helpers.rs
1//! Opaque helpers for Zvbb extension
2
3use crate::v::vector_registers::VectorRegistersExt;
4pub use crate::v::zvexx::arith::zvexx_arith_helpers::{OpSrc, check_vreg_group_alignment};
5use crate::v::zvexx::arith::zvexx_arith_helpers::{read_element_u64, write_element_u64};
6use crate::v::zvexx::load::zvexx_load_helpers::mask_bit;
7use ab_riscv_primitives::prelude::*;
8use core::fmt;
9
10/// Execute element-wise full bit-reversal over `vstart..vl`, writing SEW-wide results into `vd`.
11///
12/// For each active element i: all bits within `vs2[i]` are reversed end-to-end
13/// (bit 0 <-> bit SEW-1). This differs from `vbrev8`, which reverses bits within each byte while
14/// preserving byte order; `vbrev` also inverts the byte order as a side effect of reversing the
15/// whole element.
16///
17/// When `vm=false`, masked-off elements are left undisturbed.
18///
19/// # Safety
20/// - `vd.to_bits() % group_regs == 0` and `vd.to_bits() + group_regs <= 32`
21/// - `vs2.to_bits() % group_regs == 0` and `vs2.to_bits() + group_regs <= 32`
22/// - `vl <= group_regs * VLEN.bytes() / sew_bytes`
23#[inline(always)]
24#[doc(hidden)]
25#[cfg_attr(feature = "no-panic", no_panic_const::no_panic)]
26pub unsafe fn execute_vbrev<Reg, ExtState, CustomError>(
27 ext_state: &mut ExtState,
28 vd: VReg,
29 vs2: VReg,
30 sew: Vsew,
31 vm: bool,
32) where
33 Reg: Register,
34 ExtState: VectorRegistersExt<Reg, CustomError>,
35 [(); SUPPORTED_ELEN_VLEN::<{ ExtState::ELEN }, { ExtState::VLEN }>]:,
36 CustomError: fmt::Debug,
37{
38 let vl = ext_state.vl();
39 let vstart = ext_state.vstart();
40 for i in vstart.range_to(vl) {
41 if !vm && !mask_bit(ext_state.read_vregs().get(VReg::V0), i) {
42 continue;
43 }
44 // SAFETY: `vs2 % group_regs == 0` and `vs2 + group_regs <= 32`; `i < vl`
45 let elem = unsafe { read_element_u64(ext_state.read_vregs(), vs2, i, sew) };
46 // `elem` is zero-extended from SEW bits to u64; reverse_bits() on the primitive type
47 // of exactly SEW width naturally handles the upper zero bits from zero-extension
48 let result = match sew {
49 Vsew::E8 => u64::from((elem as u8).reverse_bits()),
50 Vsew::E16 => u64::from((elem as u16).reverse_bits()),
51 Vsew::E32 => u64::from((elem as u32).reverse_bits()),
52 Vsew::E64 => elem.reverse_bits(),
53 };
54 // SAFETY: `vd % group_regs == 0` and `vd + group_regs <= 32`; `i < vl`
55 unsafe {
56 write_element_u64(ext_state.write_vregs(), vd, i, sew, result);
57 }
58 }
59 ext_state.mark_vs_dirty();
60 ext_state.reset_vstart();
61}
62
63/// Execute element-wise count-leading-zeros over `vstart..vl`, writing SEW-wide results into `vd`.
64///
65/// For each active element i: `vd[i] = clz(vs2[i])`, counting within the SEW-wide field. An
66/// all-zero element produces SEW, not 64.
67///
68/// When `vm=false`, masked-off elements are left undisturbed.
69///
70/// # Safety
71/// Same register-group constraints as [`execute_vbrev`].
72#[inline(always)]
73#[doc(hidden)]
74#[cfg_attr(feature = "no-panic", no_panic_const::no_panic)]
75pub unsafe fn execute_vclz<Reg, ExtState, CustomError>(
76 ext_state: &mut ExtState,
77 vd: VReg,
78 vs2: VReg,
79 sew: Vsew,
80 vm: bool,
81) where
82 Reg: Register,
83 ExtState: VectorRegistersExt<Reg, CustomError>,
84 [(); SUPPORTED_ELEN_VLEN::<{ ExtState::ELEN }, { ExtState::VLEN }>]:,
85 CustomError: fmt::Debug,
86{
87 let vl = ext_state.vl();
88 let vstart = ext_state.vstart();
89 let sew_bits = u32::from(sew.bits_width());
90 for i in vstart.range_to(vl) {
91 if !vm && !mask_bit(ext_state.read_vregs().get(VReg::V0), i) {
92 continue;
93 }
94 // SAFETY: `vs2 % group_regs == 0` and `vs2 + group_regs <= 32`; `i < vl`
95 let elem = unsafe { read_element_u64(ext_state.read_vregs(), vs2, i, sew) };
96 // `elem` is zero-extended from SEW bits to u64; `leading_zeros()` on a u64 therefore counts
97 // the extra (64 - SEW) upper zero bits introduced by zero-extension. Subtracting them gives
98 // the count within the SEW-wide field.
99 let clz = elem.leading_zeros() - (64 - sew_bits);
100 // SAFETY: `vd % group_regs == 0` and `vd + group_regs <= 32`; `i < vl`
101 unsafe {
102 write_element_u64(ext_state.write_vregs(), vd, i, sew, u64::from(clz));
103 }
104 }
105 ext_state.mark_vs_dirty();
106 ext_state.reset_vstart();
107}
108
109/// Execute element-wise count-trailing-zeros over `vstart..vl`, writing SEW-wide results into `vd`.
110///
111/// For each active element i: `vd[i] = ctz(vs2[i])`, counting within the SEW-wide field. An
112/// all-zero element produces SEW, not 64.
113///
114/// When `vm=false`, masked-off elements are left undisturbed.
115///
116/// # Safety
117/// Same register-group constraints as [`execute_vbrev`].
118#[inline(always)]
119#[doc(hidden)]
120#[cfg_attr(feature = "no-panic", no_panic_const::no_panic)]
121pub unsafe fn execute_vctz<Reg, ExtState, CustomError>(
122 ext_state: &mut ExtState,
123 vd: VReg,
124 vs2: VReg,
125 sew: Vsew,
126 vm: bool,
127) where
128 Reg: Register,
129 ExtState: VectorRegistersExt<Reg, CustomError>,
130 [(); SUPPORTED_ELEN_VLEN::<{ ExtState::ELEN }, { ExtState::VLEN }>]:,
131 CustomError: fmt::Debug,
132{
133 let vl = ext_state.vl();
134 let vstart = ext_state.vstart();
135 let sew_bits = u32::from(sew.bits_width());
136 for i in vstart.range_to(vl) {
137 if !vm && !mask_bit(ext_state.read_vregs().get(VReg::V0), i) {
138 continue;
139 }
140 // SAFETY: `vs2 % group_regs == 0` and `vs2 + group_regs <= 32`; `i < vl`
141 let elem = unsafe { read_element_u64(ext_state.read_vregs(), vs2, i, sew) };
142 // For non-zero `elem`, `trailing_zeros()` on the zero-extended u64 value is correct: the
143 // upper zero bits do not affect the trailing count. For zero, `trailing_zeros()` returns
144 // 64, but the spec result is SEW; cap at `sew_bits` handles both cases.
145 let ctz = elem.trailing_zeros().min(sew_bits);
146 // SAFETY: `vd % group_regs == 0` and `vd + group_regs <= 32`; `i < vl`
147 unsafe {
148 write_element_u64(ext_state.write_vregs(), vd, i, sew, u64::from(ctz));
149 }
150 }
151 ext_state.mark_vs_dirty();
152 ext_state.reset_vstart();
153}
154
155/// Execute element-wise population count over `vstart..vl`, writing SEW-wide results into `vd`.
156///
157/// For each active element i: `vd[i] = popcount(vs2[i])`, in range `[0, SEW]`.
158///
159/// When `vm=false`, masked-off elements are left undisturbed.
160///
161/// # Safety
162/// Same register-group constraints as [`execute_vbrev`].
163#[inline(always)]
164#[doc(hidden)]
165#[cfg_attr(feature = "no-panic", no_panic_const::no_panic)]
166pub unsafe fn execute_vcpop<Reg, ExtState, CustomError>(
167 ext_state: &mut ExtState,
168 vd: VReg,
169 vs2: VReg,
170 sew: Vsew,
171 vm: bool,
172) where
173 Reg: Register,
174 ExtState: VectorRegistersExt<Reg, CustomError>,
175 [(); SUPPORTED_ELEN_VLEN::<{ ExtState::ELEN }, { ExtState::VLEN }>]:,
176 CustomError: fmt::Debug,
177{
178 let vl = ext_state.vl();
179 let vstart = ext_state.vstart();
180 for i in vstart.range_to(vl) {
181 if !vm && !mask_bit(ext_state.read_vregs().get(VReg::V0), i) {
182 continue;
183 }
184 // SAFETY: `vs2 % group_regs == 0` and `vs2 + group_regs <= 32`; `i < vl`
185 let elem = unsafe { read_element_u64(ext_state.read_vregs(), vs2, i, sew) };
186 // `elem` is zero-extended from SEW bits; upper bits are already zero, so `count_ones()`
187 // directly gives the population count within the SEW-wide field
188 let cpop = elem.count_ones();
189 // SAFETY: `vd % group_regs == 0` and `vd + group_regs <= 32`; `i < vl`
190 unsafe {
191 write_element_u64(ext_state.write_vregs(), vd, i, sew, u64::from(cpop));
192 }
193 }
194 ext_state.mark_vs_dirty();
195 ext_state.reset_vstart();
196}
197
198/// Execute element-wise widening shift-left-logical over `vstart..vl`, writing 2*SEW-wide
199/// results into `vd`.
200///
201/// For each active element i: `vd[i] = zero_extend_to_2SEW(vs2[i]) << (src[i] % (2*SEW))`.
202/// The source operand width is SEW; the destination element width is `double_sew` (2*SEW).
203///
204/// The caller must ensure SEW <= E32 (i.e., `sew.double_width()` is `Some`); passing SEW=E64 is a
205/// programming error that would produce a result wider than u64.
206///
207/// When `vm=false`, masked-off destination elements are left undisturbed.
208///
209/// # Safety
210/// - `vd` register group satisfies alignment for EMUL = 2*LMUL: `vd.to_bits() % dest_group_regs ==
211/// 0` and `vd.to_bits() + dest_group_regs <= 32`
212/// - `vs2` register group satisfies alignment for LMUL
213/// - `src` register (if `Vreg`) satisfies the same alignment as `vs2`
214/// - `vl <= dest_group_regs * VLEN.bytes() / double_sew_bytes`
215#[inline(always)]
216#[doc(hidden)]
217#[cfg_attr(feature = "no-panic", no_panic_const::no_panic)]
218pub unsafe fn execute_vwsll<Reg, ExtState, CustomError>(
219 ext_state: &mut ExtState,
220 vd: VReg,
221 vs2: VReg,
222 src: OpSrc,
223 sew: Vsew,
224 double_sew: Vsew,
225 vm: bool,
226) where
227 Reg: Register,
228 ExtState: VectorRegistersExt<Reg, CustomError>,
229 [(); SUPPORTED_ELEN_VLEN::<{ ExtState::ELEN }, { ExtState::VLEN }>]:,
230 CustomError: fmt::Debug,
231{
232 let vl = ext_state.vl();
233 let vstart = ext_state.vstart();
234 // `double_sew_bits` is always a power of two (16, 32, or 64); `& (bits - 1)` is equivalent to
235 // `% bits` and avoids a division
236 let double_sew_bits = u64::from(double_sew.bits_width());
237 for i in vstart.range_to(vl) {
238 if !vm && !mask_bit(ext_state.read_vregs().get(VReg::V0), i) {
239 continue;
240 }
241 // SAFETY: `vs2 % group_regs == 0` and `vs2 + group_regs <= 32`; `i < vl`
242 let a = unsafe { read_element_u64(ext_state.read_vregs(), vs2, i, sew) };
243 let amount = match src {
244 OpSrc::Vreg(vs1_base) => {
245 // SAFETY: same alignment constraint as vs2; same index bound
246 unsafe { read_element_u64(ext_state.read_vregs(), vs1_base, i, sew) }
247 }
248 OpSrc::Scalar(val) => val,
249 };
250 let shift = (amount & (double_sew_bits - 1)) as u32;
251 // `a` is zero-extended from SEW bits; `shift < double_sew_bits <= 64`, so this never shifts
252 // by >= 64. The caller guarantees SEW <= E32, hence `double_sew_bits <= 64`.
253 let result = a << shift;
254 // SAFETY: `vd % dest_group_regs == 0` and `vd + dest_group_regs <= 32`; `i < vl`;
255 // `write_element_u64` with `double_sew` writes exactly 2*SEW bits of `result`
256 unsafe {
257 write_element_u64(ext_state.write_vregs(), vd, i, double_sew, result);
258 }
259 }
260 ext_state.mark_vs_dirty();
261 ext_state.reset_vstart();
262}