openvm_rv32im_circuit/mulh/
core.rs

1use std::{
2    array,
3    borrow::{Borrow, BorrowMut},
4};
5
6use openvm_circuit::{
7    arch::*,
8    system::memory::{online::TracingMemory, MemoryAuxColsFactory},
9};
10use openvm_circuit_primitives::{
11    bitwise_op_lookup::{BitwiseOperationLookupBus, SharedBitwiseOperationLookupChip},
12    range_tuple::{RangeTupleCheckerBus, SharedRangeTupleCheckerChip},
13    AlignedBytesBorrow, ColumnsAir, StructReflection, StructReflectionHelper,
14};
15use openvm_circuit_primitives_derive::AlignedBorrow;
16use openvm_instructions::{instruction::Instruction, program::DEFAULT_PC_STEP, LocalOpcode};
17use openvm_rv32im_transpiler::MulHOpcode;
18use openvm_stark_backend::{
19    interaction::InteractionBuilder,
20    p3_air::{AirBuilder, BaseAir},
21    p3_field::{Field, PrimeCharacteristicRing, PrimeField32},
22    BaseAirWithPublicValues,
23};
24use strum::IntoEnumIterator;
25
26#[repr(C)]
27#[derive(AlignedBorrow, StructReflection)]
28pub struct MulHCoreCols<T, const NUM_LIMBS: usize, const LIMB_BITS: usize> {
29    pub a: [T; NUM_LIMBS],
30    pub b: [T; NUM_LIMBS],
31    pub c: [T; NUM_LIMBS],
32
33    pub a_mul: [T; NUM_LIMBS],
34    pub b_ext: T,
35    pub c_ext: T,
36
37    pub opcode_mulh_flag: T,
38    pub opcode_mulhsu_flag: T,
39    pub opcode_mulhu_flag: T,
40}
41
42#[derive(Copy, Clone, Debug, derive_new::new, ColumnsAir)]
43#[columns_via(MulHCoreCols<u8, NUM_LIMBS, LIMB_BITS>)]
44pub struct MulHCoreAir<const NUM_LIMBS: usize, const LIMB_BITS: usize> {
45    pub bitwise_lookup_bus: BitwiseOperationLookupBus,
46    pub range_tuple_bus: RangeTupleCheckerBus<2>,
47}
48
49impl<F: Field, const NUM_LIMBS: usize, const LIMB_BITS: usize> BaseAir<F>
50    for MulHCoreAir<NUM_LIMBS, LIMB_BITS>
51{
52    fn width(&self) -> usize {
53        MulHCoreCols::<F, NUM_LIMBS, LIMB_BITS>::width()
54    }
55}
56impl<F: Field, const NUM_LIMBS: usize, const LIMB_BITS: usize> BaseAirWithPublicValues<F>
57    for MulHCoreAir<NUM_LIMBS, LIMB_BITS>
58{
59}
60
61impl<AB, I, const NUM_LIMBS: usize, const LIMB_BITS: usize> VmCoreAir<AB, I>
62    for MulHCoreAir<NUM_LIMBS, LIMB_BITS>
63where
64    AB: InteractionBuilder,
65    I: VmAdapterInterface<AB::Expr>,
66    I::Reads: From<[[AB::Expr; NUM_LIMBS]; 2]>,
67    I::Writes: From<[[AB::Expr; NUM_LIMBS]; 1]>,
68    I::ProcessedInstruction: From<MinimalInstruction<AB::Expr>>,
69{
70    fn eval(
71        &self,
72        builder: &mut AB,
73        local_core: &[AB::Var],
74        _from_pc: AB::Var,
75    ) -> AdapterAirContext<AB::Expr, I> {
76        let cols: &MulHCoreCols<_, NUM_LIMBS, LIMB_BITS> = local_core.borrow();
77        let flags = [
78            cols.opcode_mulh_flag,
79            cols.opcode_mulhsu_flag,
80            cols.opcode_mulhu_flag,
81        ];
82
83        let is_valid = flags.iter().fold(AB::Expr::ZERO, |acc, &flag| {
84            builder.assert_bool(flag);
85            acc + flag.into()
86        });
87        builder.assert_bool(is_valid.clone());
88
89        let b = &cols.b;
90        let c = &cols.c;
91        let carry_divide = AB::F::from_u32(1 << LIMB_BITS).inverse();
92
93        // Note b * c = a << LIMB_BITS + a_mul, in order to constrain that a is correct we
94        // need to compute the carries generated by a_mul.
95        let a_mul = &cols.a_mul;
96        let mut carry_mul: [AB::Expr; NUM_LIMBS] = array::from_fn(|_| AB::Expr::ZERO);
97
98        for i in 0..NUM_LIMBS {
99            let expected_limb = if i == 0 {
100                AB::Expr::ZERO
101            } else {
102                carry_mul[i - 1].clone()
103            } + (0..=i).fold(AB::Expr::ZERO, |ac, k| ac + (b[k] * c[i - k]));
104            carry_mul[i] = AB::Expr::from(carry_divide) * (expected_limb - a_mul[i]);
105        }
106
107        for (a_mul, carry_mul) in a_mul.iter().zip(carry_mul.iter()) {
108            self.range_tuple_bus
109                .send(vec![(*a_mul).into(), carry_mul.clone()])
110                .eval(builder, is_valid.clone());
111        }
112
113        // We can now constrain that a is correct using carry_mul[NUM_LIMBS - 1]
114        let a = &cols.a;
115        let mut carry: [AB::Expr; NUM_LIMBS] = array::from_fn(|_| AB::Expr::ZERO);
116
117        for j in 0..NUM_LIMBS {
118            let expected_limb = if j == 0 {
119                carry_mul[NUM_LIMBS - 1].clone()
120            } else {
121                carry[j - 1].clone()
122            } + ((j + 1)..NUM_LIMBS)
123                .fold(AB::Expr::ZERO, |acc, k| acc + (b[k] * c[NUM_LIMBS + j - k]))
124                + (0..(j + 1)).fold(AB::Expr::ZERO, |acc, k| {
125                    acc + (b[k] * cols.c_ext) + (c[k] * cols.b_ext)
126                });
127            carry[j] = AB::Expr::from(carry_divide) * (expected_limb - a[j]);
128        }
129
130        for (a, carry) in a.iter().zip(carry.iter()) {
131            self.range_tuple_bus
132                .send(vec![(*a).into(), carry.clone()])
133                .eval(builder, is_valid.clone());
134        }
135
136        // Check that b_ext and c_ext are correct using bitwise lookup. We check
137        // both b and c when the opcode is MULH, and only b when MULHSU.
138        let sign_mask = AB::F::from_u32(1 << (LIMB_BITS - 1));
139        let ext_inv = AB::F::from_u32((1 << LIMB_BITS) - 1).inverse();
140        let b_sign = cols.b_ext * ext_inv;
141        let c_sign = cols.c_ext * ext_inv;
142
143        builder.assert_bool(b_sign.clone());
144        builder.assert_bool(c_sign.clone());
145        builder
146            .when(cols.opcode_mulhu_flag)
147            .assert_zero(b_sign.clone());
148        builder
149            .when(cols.opcode_mulhu_flag + cols.opcode_mulhsu_flag)
150            .assert_zero(c_sign.clone());
151
152        self.bitwise_lookup_bus
153            .send_range(
154                AB::Expr::from_u32(2) * (b[NUM_LIMBS - 1] - b_sign * sign_mask),
155                (cols.opcode_mulh_flag + AB::Expr::ONE) * (c[NUM_LIMBS - 1] - c_sign * sign_mask),
156            )
157            .eval(builder, cols.opcode_mulh_flag + cols.opcode_mulhsu_flag);
158
159        let expected_opcode = VmCoreAir::<AB, I>::expr_to_global_expr(
160            self,
161            flags.iter().zip(MulHOpcode::iter()).fold(
162                AB::Expr::ZERO,
163                |acc, (flag, local_opcode)| {
164                    acc + (*flag).into() * AB::Expr::from_u8(local_opcode as u8)
165                },
166            ),
167        );
168
169        AdapterAirContext {
170            to_pc: None,
171            reads: [cols.b.map(Into::into), cols.c.map(Into::into)].into(),
172            writes: [cols.a.map(Into::into)].into(),
173            instruction: MinimalInstruction {
174                is_valid,
175                opcode: expected_opcode,
176            }
177            .into(),
178        }
179    }
180
181    fn start_offset(&self) -> usize {
182        MulHOpcode::CLASS_OFFSET
183    }
184}
185
186#[repr(C)]
187#[derive(AlignedBytesBorrow, Debug)]
188pub struct MulHCoreRecord<const NUM_LIMBS: usize, const LIMB_BITS: usize> {
189    pub b: [u8; NUM_LIMBS],
190    pub c: [u8; NUM_LIMBS],
191    pub local_opcode: u8,
192}
193
194#[derive(Clone, Copy, derive_new::new)]
195pub struct MulHExecutor<A, const NUM_LIMBS: usize, const LIMB_BITS: usize> {
196    adapter: A,
197    pub offset: usize,
198}
199
200#[derive(Clone)]
201pub struct MulHFiller<A, const NUM_LIMBS: usize, const LIMB_BITS: usize> {
202    adapter: A,
203    pub bitwise_lookup_chip: SharedBitwiseOperationLookupChip<LIMB_BITS>,
204    pub range_tuple_chip: SharedRangeTupleCheckerChip<2>,
205}
206
207impl<A, const NUM_LIMBS: usize, const LIMB_BITS: usize> MulHFiller<A, NUM_LIMBS, LIMB_BITS> {
208    pub fn new(
209        adapter: A,
210        bitwise_lookup_chip: SharedBitwiseOperationLookupChip<LIMB_BITS>,
211        range_tuple_chip: SharedRangeTupleCheckerChip<2>,
212    ) -> Self {
213        // The RangeTupleChecker is used to range check (a[i], carry[i]) pairs where 0 <= i
214        // < 2 * NUM_LIMBS. a[i] must have LIMB_BITS bits and carry[i] is the sum of i + 1
215        // bytes (with LIMB_BITS bits). BitwiseOperationLookup is used to sign check bytes.
216        debug_assert!(
217            range_tuple_chip.sizes()[0] == 1 << LIMB_BITS,
218            "First element of RangeTupleChecker must have size {}",
219            1 << LIMB_BITS
220        );
221        debug_assert!(
222            range_tuple_chip.sizes()[1] >= (1 << LIMB_BITS) * 2 * NUM_LIMBS as u32,
223            "Second element of RangeTupleChecker must have size of at least {}",
224            (1 << LIMB_BITS) * 2 * NUM_LIMBS as u32
225        );
226
227        Self {
228            adapter,
229            bitwise_lookup_chip,
230            range_tuple_chip,
231        }
232    }
233}
234
235impl<F, A, RA, const NUM_LIMBS: usize, const LIMB_BITS: usize> PreflightExecutor<F, RA>
236    for MulHExecutor<A, NUM_LIMBS, LIMB_BITS>
237where
238    F: PrimeField32,
239    A: 'static
240        + AdapterTraceExecutor<
241            F,
242            ReadData: Into<[[u8; NUM_LIMBS]; 2]>,
243            WriteData: From<[[u8; NUM_LIMBS]; 1]>,
244        >,
245    for<'buf> RA: RecordArena<
246        'buf,
247        EmptyAdapterCoreLayout<F, A>,
248        (
249            A::RecordMut<'buf>,
250            &'buf mut MulHCoreRecord<NUM_LIMBS, LIMB_BITS>,
251        ),
252    >,
253{
254    fn get_opcode_name(&self, opcode: usize) -> String {
255        format!(
256            "{:?}",
257            MulHOpcode::from_usize(opcode - MulHOpcode::CLASS_OFFSET)
258        )
259    }
260
261    fn execute(
262        &self,
263        state: VmStateMut<F, TracingMemory, RA>,
264        instruction: &Instruction<F>,
265    ) -> Result<(), ExecutionError> {
266        let Instruction { opcode, .. } = instruction;
267
268        let (mut adapter_record, core_record) = state.ctx.alloc(EmptyAdapterCoreLayout::new());
269
270        A::start(*state.pc, state.memory, &mut adapter_record);
271
272        core_record.local_opcode = opcode.local_opcode_idx(MulHOpcode::CLASS_OFFSET) as u8;
273        let mulh_opcode = MulHOpcode::from_usize(core_record.local_opcode as usize);
274
275        [core_record.b, core_record.c] = self
276            .adapter
277            .read(state.memory, instruction, &mut adapter_record)
278            .into();
279
280        let (a, _, _, _, _) = run_mulh::<NUM_LIMBS, LIMB_BITS>(
281            mulh_opcode,
282            &core_record.b.map(u32::from),
283            &core_record.c.map(u32::from),
284        );
285
286        let a = a.map(|x| x as u8);
287        self.adapter
288            .write(state.memory, instruction, [a].into(), &mut adapter_record);
289
290        *state.pc = state.pc.wrapping_add(DEFAULT_PC_STEP);
291
292        Ok(())
293    }
294}
295
296impl<F, A, const NUM_LIMBS: usize, const LIMB_BITS: usize> TraceFiller<F>
297    for MulHFiller<A, NUM_LIMBS, LIMB_BITS>
298where
299    F: PrimeField32,
300    A: 'static + AdapterTraceFiller<F>,
301{
302    fn fill_trace_row(&self, mem_helper: &MemoryAuxColsFactory<F>, row_slice: &mut [F]) {
303        // SAFETY: row_slice is guaranteed by the caller to have at least A::WIDTH +
304        // MulHCoreCols::width() elements
305        let (adapter_row, mut core_row) = unsafe { row_slice.split_at_mut_unchecked(A::WIDTH) };
306        self.adapter.fill_trace_row(mem_helper, adapter_row);
307        // SAFETY: core_row contains a valid MulHCoreRecord written by the executor
308        // during trace generation
309        let record: &MulHCoreRecord<NUM_LIMBS, LIMB_BITS> =
310            unsafe { get_record_from_slice(&mut core_row, ()) };
311        let core_row: &mut MulHCoreCols<F, NUM_LIMBS, LIMB_BITS> = core_row.borrow_mut();
312
313        let opcode = MulHOpcode::from_usize(record.local_opcode as usize);
314        let (a, a_mul, carry, b_ext, c_ext) = run_mulh::<NUM_LIMBS, LIMB_BITS>(
315            opcode,
316            &record.b.map(u32::from),
317            &record.c.map(u32::from),
318        );
319
320        for i in 0..NUM_LIMBS {
321            self.range_tuple_chip.add_count(&[a_mul[i], carry[i]]);
322            self.range_tuple_chip
323                .add_count(&[a[i], carry[NUM_LIMBS + i]]);
324        }
325
326        if opcode != MulHOpcode::MULHU {
327            let b_sign_mask = if b_ext == 0 { 0 } else { 1 << (LIMB_BITS - 1) };
328            let c_sign_mask = if c_ext == 0 { 0 } else { 1 << (LIMB_BITS - 1) };
329            self.bitwise_lookup_chip.request_range(
330                (record.b[NUM_LIMBS - 1] as u32 - b_sign_mask) << 1,
331                (record.c[NUM_LIMBS - 1] as u32 - c_sign_mask)
332                    << ((opcode == MulHOpcode::MULH) as u32),
333            );
334        }
335
336        // Write in reverse order
337        core_row.opcode_mulhu_flag = F::from_bool(opcode == MulHOpcode::MULHU);
338        core_row.opcode_mulhsu_flag = F::from_bool(opcode == MulHOpcode::MULHSU);
339        core_row.opcode_mulh_flag = F::from_bool(opcode == MulHOpcode::MULH);
340        core_row.c_ext = F::from_u32(c_ext);
341        core_row.b_ext = F::from_u32(b_ext);
342        core_row.a_mul = a_mul.map(F::from_u32);
343        core_row.c = record.c.map(F::from_u8);
344        core_row.b = record.b.map(F::from_u8);
345        core_row.a = a.map(F::from_u32);
346    }
347}
348
349// returns mulh[[s]u], mul, carry, x_ext, y_ext
350#[inline(always)]
351pub(super) fn run_mulh<const NUM_LIMBS: usize, const LIMB_BITS: usize>(
352    opcode: MulHOpcode,
353    x: &[u32; NUM_LIMBS],
354    y: &[u32; NUM_LIMBS],
355) -> ([u32; NUM_LIMBS], [u32; NUM_LIMBS], Vec<u32>, u32, u32) {
356    let mut mul = [0; NUM_LIMBS];
357    let mut carry = vec![0; 2 * NUM_LIMBS];
358    for i in 0..NUM_LIMBS {
359        if i > 0 {
360            mul[i] = carry[i - 1];
361        }
362        for j in 0..=i {
363            mul[i] += x[j] * y[i - j];
364        }
365        carry[i] = mul[i] >> LIMB_BITS;
366        mul[i] %= 1 << LIMB_BITS;
367    }
368
369    let x_ext = (x[NUM_LIMBS - 1] >> (LIMB_BITS - 1))
370        * if opcode == MulHOpcode::MULHU {
371            0
372        } else {
373            (1 << LIMB_BITS) - 1
374        };
375    let y_ext = (y[NUM_LIMBS - 1] >> (LIMB_BITS - 1))
376        * if opcode == MulHOpcode::MULH {
377            (1 << LIMB_BITS) - 1
378        } else {
379            0
380        };
381
382    let mut mulh = [0; NUM_LIMBS];
383    let mut x_prefix = 0;
384    let mut y_prefix = 0;
385
386    for i in 0..NUM_LIMBS {
387        x_prefix += x[i];
388        y_prefix += y[i];
389        mulh[i] = carry[NUM_LIMBS + i - 1] + x_prefix * y_ext + y_prefix * x_ext;
390        for j in (i + 1)..NUM_LIMBS {
391            mulh[i] += x[j] * y[NUM_LIMBS + i - j];
392        }
393        carry[NUM_LIMBS + i] = mulh[i] >> LIMB_BITS;
394        mulh[i] %= 1 << LIMB_BITS;
395    }
396
397    (mulh, mul, carry, x_ext, y_ext)
398}