ab_riscv_interpreter/v/zvexx/carry/zvexx_carry_helpers.rs
1//! Opaque helpers for ZveXx extension
2
3use crate::v::vector_registers::{VectorRegisterFile, VectorRegistersExt};
4pub use crate::v::zvexx::arith::zvexx_arith_helpers::{
5 OpSrc, check_mask_dest_no_overlap, check_vreg_group_alignment,
6};
7use crate::v::zvexx::arith::zvexx_arith_helpers::{
8 read_element_u64, sew_mask, write_element_u64, write_mask_bit,
9};
10use crate::v::zvexx::load::zvexx_load_helpers::mask_bit;
11use ab_riscv_primitives::prelude::*;
12use core::fmt;
13
14// TODO: Safety comment here doesn't make sense
15/// Read a single mask bit from vector register `v0` at element index `i`.
16///
17/// Used to retrieve the per-element carry-in or borrow-in for vadc/vsbc.
18///
19/// # Safety
20/// `i / 8 < VLEN.bytes()` must hold, guaranteed when `i < vl <= VLEN`.
21#[inline(always)]
22#[cfg_attr(feature = "no-panic", no_panic_const::no_panic)]
23pub(in super::super) unsafe fn carry_bit<const VLEN: Vlen>(
24 vregs: &VectorRegisterFile<VLEN>,
25 i: u16,
26) -> u64 {
27 let v0 = vregs.get(VReg::V0);
28 u64::from(mask_bit(v0, i))
29}
30
31/// Execute an element-wise add-with-carry over `vstart..vl`, writing SEW-wide data results into
32/// `vd`.
33///
34/// Carry-in for each element is read from `v0[i]` when `WITH_CARRY` is true. All elements in
35/// `vstart..vl` are processed unconditionally (no execution mask).
36///
37/// # Safety
38/// - `vd.to_bits() % group_regs == 0` and `vd.to_bits() + group_regs <= 32`
39/// - `vs2.to_bits() % group_regs == 0` and `vs2.to_bits() + group_regs <= 32`
40/// - `src` register satisfies the same alignment (verified by caller)
41/// - `vd.to_bits() != 0` (vd must not overlap v0, which holds the carry-in)
42/// - `vl <= group_regs * VLEN.bytes() / sew_bytes`
43#[inline(always)]
44#[doc(hidden)]
45// TODO: #[cfg_attr(feature = "no-panic", no_panic_const::no_panic)]
46pub unsafe fn execute_carry_add<const WITH_CARRY: bool, Reg, ExtState, CustomError>(
47 ext_state: &mut ExtState,
48 vd: VReg,
49 vs2: VReg,
50 src: OpSrc,
51 sew: Vsew,
52) where
53 Reg: Register,
54 ExtState: VectorRegistersExt<Reg, CustomError>,
55 [(); SUPPORTED_ELEN_VLEN::<{ ExtState::ELEN }, { ExtState::VLEN }>]:,
56 CustomError: fmt::Debug,
57{
58 let vl = ext_state.vl();
59 let vstart = ext_state.vstart();
60 for i in vstart.range_to(vl) {
61 // SAFETY: `vs2 % group_regs == 0` and `vs2 + group_regs <= 32` (caller precondition);
62 // `i < vl <= group_regs * elems_per_reg`, so
63 // `vs2 + i / elems_per_reg < vs2 + group_regs <= 32`
64 let a = unsafe { read_element_u64(ext_state.read_vregs(), vs2, i, sew) };
65 let b = match src {
66 OpSrc::Vreg(vs1_base) => {
67 // SAFETY: caller verified that the vs1 register group satisfies the same alignment
68 // constraint as vs2; the index argument is identical, so the same bound holds:
69 // `vs1_base + i / elems_per_reg < 32`
70 unsafe { read_element_u64(ext_state.read_vregs(), vs1_base, i, sew) }
71 }
72 OpSrc::Scalar(val) => val,
73 };
74 let c = if WITH_CARRY {
75 // SAFETY: `i < vl <= VLEN`, so `i / 8 < VLEN.bytes()`
76 unsafe { carry_bit(ext_state.read_vregs(), i) }
77 } else {
78 0
79 };
80
81 // Wrap naturally: write_element_u64 writes only the low sew_bytes
82 let result = a.wrapping_add(b).wrapping_add(c);
83 // SAFETY: `vd % group_regs == 0` and `vd + group_regs <= 32` (caller precondition);
84 // `i < vl <= group_regs * elems_per_reg`, so
85 // `vd + i / elems_per_reg < vd + group_regs <= 32`
86 unsafe {
87 write_element_u64(ext_state.write_vregs(), vd, i, sew, result);
88 }
89 }
90
91 ext_state.mark_vs_dirty();
92 ext_state.reset_vstart();
93}
94
95/// Execute an element-wise subtract-with-borrow over `vstart..vl`, writing SEW-wide data results
96/// into `vd`.
97///
98/// Borrow-in for each element is read from `v0[i]` (always true for vsbc). All elements in
99/// `vstart..vl` are processed unconditionally.
100///
101/// # Safety
102/// Same as [`execute_carry_add()`].
103#[inline(always)]
104#[doc(hidden)]
105// TODO: #[cfg_attr(feature = "no-panic", no_panic_const::no_panic)]
106pub unsafe fn execute_carry_sub<Reg, ExtState, CustomError>(
107 ext_state: &mut ExtState,
108 vd: VReg,
109 vs2: VReg,
110 src: OpSrc,
111 sew: Vsew,
112) where
113 Reg: Register,
114 ExtState: VectorRegistersExt<Reg, CustomError>,
115 [(); SUPPORTED_ELEN_VLEN::<{ ExtState::ELEN }, { ExtState::VLEN }>]:,
116 CustomError: fmt::Debug,
117{
118 let vl = ext_state.vl();
119 let vstart = ext_state.vstart();
120 for i in vstart.range_to(vl) {
121 // SAFETY: `vs2 % group_regs == 0` and `vs2 + group_regs <= 32` (caller precondition);
122 // `i < vl <= group_regs * elems_per_reg`, so
123 // `vs2 + i / elems_per_reg < vs2 + group_regs <= 32`
124 let a = unsafe { read_element_u64(ext_state.read_vregs(), vs2, i, sew) };
125 let b = match src {
126 OpSrc::Vreg(vs1_base) => {
127 // SAFETY: caller verified that the vs1 register group satisfies the same alignment
128 // constraint as vs2; the index argument is identical, so the same bound holds:
129 // `vs1_base + i / elems_per_reg < 32`
130 unsafe { read_element_u64(ext_state.read_vregs(), vs1_base, i, sew) }
131 }
132 OpSrc::Scalar(val) => val,
133 };
134 // SAFETY: `i < vl <= VLEN`, so `i / 8 < VLEN.bytes()`
135 let borrow = unsafe { carry_bit(ext_state.read_vregs(), i) };
136
137 let result = a.wrapping_sub(b).wrapping_sub(borrow);
138 // SAFETY: `vd % group_regs == 0` and `vd + group_regs <= 32` (caller precondition);
139 // `i < vl <= group_regs * elems_per_reg`, so
140 // `vd + i / elems_per_reg < vd + group_regs <= 32`
141 unsafe {
142 write_element_u64(ext_state.write_vregs(), vd, i, sew, result);
143 }
144 }
145
146 ext_state.mark_vs_dirty();
147 ext_state.reset_vstart();
148}
149
150/// Execute an element-wise add-with-carry over `vstart..vl`, writing the carry-out as a single mask
151/// bit per element into `vd`.
152///
153/// When `WITH_CARRY` is true, carry-in for element `i` is read from `v0[i]`. When false, carry-in
154/// is treated as zero.
155///
156/// All elements are processed unconditionally (no execution mask).
157///
158/// Tail mask bits (indices `>= vl`) are left undisturbed per spec ยง5.3.
159///
160/// # Safety
161/// - `vs2.to_bits() % group_regs == 0` and `vs2.to_bits() + group_regs <= 32`
162/// - `src` register satisfies the same alignment
163/// - `vl <= group_regs * VLEN.bytes() / sew_bytes` and `vl <= VLEN`
164/// - vd overlap constraints checked by caller
165#[inline(always)]
166#[doc(hidden)]
167// TODO: #[cfg_attr(feature = "no-panic", no_panic_const::no_panic)]
168pub unsafe fn execute_carry_add_mask<const WITH_CARRY: bool, Reg, ExtState, CustomError>(
169 ext_state: &mut ExtState,
170 vd: VReg,
171 vs2: VReg,
172 src: OpSrc,
173 sew: Vsew,
174) where
175 Reg: Register,
176 ExtState: VectorRegistersExt<Reg, CustomError>,
177 [(); SUPPORTED_ELEN_VLEN::<{ ExtState::ELEN }, { ExtState::VLEN }>]:,
178 CustomError: fmt::Debug,
179{
180 let vl = ext_state.vl();
181 let vstart = ext_state.vstart();
182 let mask = sew_mask(sew);
183
184 for i in vstart.range_to(vl) {
185 // SAFETY: `vs2 % group_regs == 0` and `vs2 + group_regs <= 32` (caller precondition);
186 // `i < vl <= group_regs * elems_per_reg`, so
187 // `vs2 + i / elems_per_reg < vs2 + group_regs <= 32`
188 let a = unsafe { read_element_u64(ext_state.read_vregs(), vs2, i, sew) };
189 let b = match src {
190 OpSrc::Vreg(vs1_base) => {
191 // SAFETY: caller verified that the vs1 register group satisfies the same alignment
192 // constraint as vs2; the index argument is identical, so the same bound holds:
193 // `vs1_base + i / elems_per_reg < 32`
194 unsafe { read_element_u64(ext_state.read_vregs(), vs1_base, i, sew) }
195 }
196 OpSrc::Scalar(val) => val,
197 };
198 let c = if WITH_CARRY {
199 // SAFETY: `i < vl <= VLEN`, so `i / 8 < VLEN.bytes()`
200 unsafe { carry_bit(ext_state.read_vregs(), i) }
201 } else {
202 0
203 };
204
205 // Use u128 to capture the carry-out bit beyond SEW
206 let sum = u128::from(a & mask) + u128::from(b & mask) + u128::from(c);
207 let carry_out = (sum >> sew.bits_width()) & 1 != 0;
208
209 // SAFETY: `i < vl <= VLEN`, so `i / 8 < VLEN.bytes()`
210 unsafe {
211 write_mask_bit(ext_state.write_vregs(), vd, i, carry_out);
212 }
213 }
214
215 ext_state.mark_vs_dirty();
216 ext_state.reset_vstart();
217}
218
219/// Execute an element-wise subtract-with-borrow over `vstart..vl`, writing the borrow-out as a
220/// single mask bit per element into `vd`.
221///
222/// When `WITH_BORROW` is true, borrow-in for element `i` is read from `v0[i]`. When false,
223/// borrow-in is treated as zero.
224///
225/// Borrow-out is 1 when the subtraction underflows unsigned:
226/// `borrow_out = (b + borrow_in) > a` (compared as SEW-wide unsigned values).
227///
228/// # Safety
229/// Same as [`execute_carry_add_mask()`].
230#[inline(always)]
231#[doc(hidden)]
232// TODO: #[cfg_attr(feature = "no-panic", no_panic_const::no_panic)]
233pub unsafe fn execute_carry_sub_mask<const WITH_BORROW: bool, Reg, ExtState, CustomError>(
234 ext_state: &mut ExtState,
235 vd: VReg,
236 vs2: VReg,
237 src: OpSrc,
238 sew: Vsew,
239) where
240 Reg: Register,
241 ExtState: VectorRegistersExt<Reg, CustomError>,
242 [(); SUPPORTED_ELEN_VLEN::<{ ExtState::ELEN }, { ExtState::VLEN }>]:,
243 CustomError: fmt::Debug,
244{
245 let vl = ext_state.vl();
246 let vstart = ext_state.vstart();
247 let mask = sew_mask(sew);
248
249 for i in vstart.range_to(vl) {
250 // SAFETY: `vs2 % group_regs == 0` and `vs2 + group_regs <= 32` (caller precondition);
251 // `i < vl <= group_regs * elems_per_reg`, so
252 // `vs2 + i / elems_per_reg < vs2 + group_regs <= 32`
253 let a = unsafe { read_element_u64(ext_state.read_vregs(), vs2, i, sew) };
254 let b = match src {
255 OpSrc::Vreg(vs1_base) => {
256 // SAFETY: caller verified that the vs1 register group satisfies the same alignment
257 // constraint as vs2; the index argument is identical, so the same bound holds:
258 // `vs1_base + i / elems_per_reg < 32`
259 unsafe { read_element_u64(ext_state.read_vregs(), vs1_base, i, sew) }
260 }
261 OpSrc::Scalar(val) => val,
262 };
263 let borrow_in = if WITH_BORROW {
264 // SAFETY: `i < vl <= VLEN`, so `i / 8 < VLEN.bytes()`
265 unsafe { carry_bit(ext_state.read_vregs(), i) }
266 } else {
267 0
268 };
269
270 let a_m = u128::from(a & mask);
271 let rhs = u128::from(b & mask) + u128::from(borrow_in);
272 let borrow_out = a_m < rhs;
273
274 // SAFETY: `i < vl <= VLEN`, so `i / 8 < VLEN.bytes()`
275 unsafe {
276 write_mask_bit(ext_state.write_vregs(), vd, i, borrow_out);
277 }
278 }
279
280 ext_state.mark_vs_dirty();
281 ext_state.reset_vstart();
282}