ab_riscv_interpreter/v/zvexx/arith/
zvexx_arith_helpers.rs1use crate::v::vector_registers::{VectorRegisterFile, VectorRegistersExt};
4use crate::v::zvexx::load::zvexx_load_helpers::{mask_bit, snapshot_mask};
5use crate::v::zvexx::zvexx_helpers::INSTRUCTION_SIZE;
6use crate::{ExecutionError, ProgramCounter};
7use ab_riscv_primitives::prelude::*;
8use core::fmt;
9use core::hint::cold_path;
10use core::num::NonZeroU8;
11
12#[inline(always)]
14#[doc(hidden)]
15#[cfg_attr(feature = "no-panic", no_panic_const::no_panic)]
16pub fn check_vreg_group_alignment<Reg, Memory, PC, CustomError>(
17 program_counter: &PC,
18 vreg: VReg,
19 group_regs: NonZeroU8,
20) -> Result<(), ExecutionError<Reg::Type, CustomError>>
21where
22 Reg: Register,
23 PC: ProgramCounter<Reg::Type, Memory, CustomError>,
24{
25 let group_regs = group_regs.get();
26 let vreg_idx = vreg.to_bits();
27 if !vreg_idx.is_multiple_of(group_regs) || vreg_idx + group_regs > 32 {
28 cold_path();
29 return Err(ExecutionError::IllegalInstruction {
30 address: program_counter.old_pc(INSTRUCTION_SIZE),
31 });
32 }
33 Ok(())
34}
35
36#[inline(always)]
42#[doc(hidden)]
43#[cfg_attr(feature = "no-panic", no_panic_const::no_panic)]
44pub fn check_mask_dest_no_overlap<Reg, Memory, PC, CustomError>(
45 program_counter: &PC,
46 vd: VReg,
47 src_base: VReg,
48 group_regs: NonZeroU8,
49) -> Result<(), ExecutionError<Reg::Type, CustomError>>
50where
51 Reg: Register,
52 PC: ProgramCounter<Reg::Type, Memory, CustomError>,
53{
54 let group_regs = group_regs.get();
55 if group_regs > 1 {
56 let vd_idx = vd.to_bits();
57 let src = src_base.to_bits();
58 if vd_idx >= src && vd_idx < src + group_regs {
59 cold_path();
60 return Err(ExecutionError::IllegalInstruction {
61 address: program_counter.old_pc(INSTRUCTION_SIZE),
62 });
63 }
64 }
65 Ok(())
66}
67
68#[inline(always)]
79#[cfg_attr(feature = "no-panic", no_panic_const::no_panic)]
80pub(crate) unsafe fn read_element_u64<const VLEN: Vlen>(
81 vregs: &VectorRegisterFile<VLEN>,
82 base_reg: VReg,
83 elem_i: u16,
84 sew: Vsew,
85) -> u64 {
86 let sew_bytes = u32::from(sew.bytes_width());
87 let elems_per_reg = VLEN.bytes() / sew_bytes;
88 let reg_off = u32::from(elem_i) / elems_per_reg;
89 let byte_off = (u32::from(elem_i) % elems_per_reg) * sew_bytes;
90 let reg = vregs
92 .get(unsafe { VReg::from_bits(base_reg.to_bits() + reg_off as u8).unwrap_unchecked() });
93 let src = unsafe { reg.get_unchecked(byte_off as usize..(byte_off + sew_bytes) as usize) };
96 let mut buf = [0u8; 8];
97 unsafe { buf.get_unchecked_mut(..sew_bytes as usize) }.copy_from_slice(src);
99 u64::from_le_bytes(buf)
100}
101
102#[inline(always)]
108#[cfg_attr(feature = "no-panic", no_panic_const::no_panic)]
109pub(crate) unsafe fn write_element_u64<const VLEN: Vlen>(
110 vregs: &mut VectorRegisterFile<VLEN>,
111 base_reg: VReg,
112 elem_i: u16,
113 sew: Vsew,
114 value: u64,
115) {
116 let sew_bytes = u32::from(sew.bytes_width());
117 let elems_per_reg = VLEN.bytes() / sew_bytes;
118 let reg_off = u32::from(elem_i) / elems_per_reg;
119 let byte_off = (u32::from(elem_i) % elems_per_reg) * sew_bytes;
120 let buf = value.to_le_bytes();
121 let reg = vregs
123 .get_mut(unsafe { VReg::from_bits(base_reg.to_bits() + reg_off as u8).unwrap_unchecked() });
124 let dst = unsafe { reg.get_unchecked_mut(byte_off as usize..(byte_off + sew_bytes) as usize) };
127 dst.copy_from_slice(unsafe { buf.get_unchecked(..sew_bytes as usize) });
129}
130
131#[inline(always)]
141#[cfg_attr(feature = "no-panic", no_panic_const::no_panic)]
142pub(in super::super) unsafe fn write_mask_bit<const VLEN: Vlen>(
143 vregs: &mut VectorRegisterFile<VLEN>,
144 vd: VReg,
145 elem_i: u16,
146 result: bool,
147) {
148 let byte_idx = usize::from(elem_i / u8::BITS as u16);
149 let bit_idx = elem_i % u8::BITS as u16;
150 let byte = unsafe { vregs.get_mut(vd).get_unchecked_mut(byte_idx) };
152 if result {
153 *byte |= 1 << bit_idx;
154 } else {
155 *byte &= !(1 << bit_idx);
156 }
157}
158
159#[derive(Debug)]
161#[doc(hidden)]
162pub enum OpSrc {
163 Vreg(VReg),
165 Scalar(u64),
167}
168
169#[inline(always)]
181#[doc(hidden)]
182pub unsafe fn execute_arith_op<Reg, ExtState, CustomError, F>(
184 ext_state: &mut ExtState,
185 vd: VReg,
186 vs2: VReg,
187 src: OpSrc,
188 vm: bool,
189 sew: Vsew,
190 op: F,
191) where
192 Reg: Register,
193 ExtState: VectorRegistersExt<Reg, CustomError>,
194 [(); SUPPORTED_ELEN_VLEN::<{ ExtState::ELEN }, { ExtState::VLEN }>]:,
195 CustomError: fmt::Debug,
196 F: Fn(u64, u64, Vsew) -> u64,
197{
198 let vl = ext_state.vl();
199 let vstart = ext_state.vstart();
200 let mask_buf = unsafe { snapshot_mask(ext_state.read_vregs(), vm, vl) };
202
203 for i in vstart.range_to(vl) {
204 if !mask_bit(&mask_buf, i) {
205 continue;
206 }
207
208 let a = unsafe { read_element_u64(ext_state.read_vregs(), vs2, i, sew) };
211
212 let b = match src {
213 OpSrc::Vreg(vs1_base) => {
214 unsafe { read_element_u64(ext_state.read_vregs(), vs1_base, i, sew) }
216 }
217 OpSrc::Scalar(val) => val,
218 };
219
220 let result = op(a, b, sew);
221
222 unsafe {
225 write_element_u64(ext_state.write_vregs(), vd, i, sew, result);
226 }
227 }
228
229 ext_state.mark_vs_dirty();
230 ext_state.reset_vstart();
231}
232
233#[inline(always)]
247#[doc(hidden)]
248pub unsafe fn execute_compare_op<Reg, ExtState, CustomError, F>(
250 ext_state: &mut ExtState,
251 vd: VReg,
252 vs2: VReg,
253 src: OpSrc,
254 vm: bool,
255 sew: Vsew,
256 op: F,
257) where
258 Reg: Register,
259 ExtState: VectorRegistersExt<Reg, CustomError>,
260 [(); SUPPORTED_ELEN_VLEN::<{ ExtState::ELEN }, { ExtState::VLEN }>]:,
261 CustomError: fmt::Debug,
262 F: Fn(u64, u64, Vsew) -> bool,
263{
264 let vl = ext_state.vl();
265 let vstart = ext_state.vstart();
266 let mask_buf = unsafe { snapshot_mask(ext_state.read_vregs(), vm, vl) };
268
269 for i in vstart.range_to(vl) {
270 if !mask_bit(&mask_buf, i) {
273 continue;
274 }
275
276 let a = unsafe { read_element_u64(ext_state.read_vregs(), vs2, i, sew) };
278
279 let b = match src {
280 OpSrc::Vreg(vs1_base) => {
281 unsafe { read_element_u64(ext_state.read_vregs(), vs1_base, i, sew) }
283 }
284 OpSrc::Scalar(val) => val,
285 };
286
287 let result = op(a, b, sew);
288
289 unsafe {
291 write_mask_bit(ext_state.write_vregs(), vd, i, result);
292 }
293 }
294
295 ext_state.mark_vs_dirty();
296 ext_state.reset_vstart();
297}
298
299#[inline(always)]
301#[doc(hidden)]
302#[cfg_attr(feature = "no-panic", no_panic_const::no_panic)]
303pub fn sign_extend(val: u64, sew: Vsew) -> i64 {
304 let shift = u64::BITS - u32::from(sew.bits_width());
305 (val.cast_signed() << shift) >> shift
306}
307
308#[inline(always)]
313#[doc(hidden)]
314#[cfg_attr(feature = "no-panic", no_panic_const::no_panic)]
315pub fn sew_mask(sew: Vsew) -> u64 {
316 if u32::from(sew.bits_width()) == u64::BITS {
317 u64::MAX
318 } else {
319 (1u64 << sew.bits_width()) - 1
320 }
321}