Skip to main content

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}