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