Skip to main content

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}