1use std::{cmp::max, iter};
5
6use binius_compute::Allocator;
7use binius_core::word::Word;
8use binius_field::{BinaryField, Field, PackedField, WideMul};
9use binius_ip::sumcheck::{RoundCoeffs, SumcheckOutput};
10use binius_ip_prover::{
11 channel::IPProverChannel,
12 sumcheck::{
13 ProveSingleOutput, bivariate_product_prover, prove_single, round_evals::RoundEvals,
14 },
15};
16use binius_math::{
17 FieldVec,
18 multilinear::eq::{eq_ind_partial_eval, eq_ind_zero},
19};
20use binius_utils::{
21 checked_arithmetics::log2_ceil_usize,
22 rayon::{
23 prelude::*,
24 task_size::{IndexedParallelIteratorExt, WorkPerItem},
25 },
26};
27use binius_verifier::protocols::shift::evaluate_words_mle;
28use tracing::instrument;
29
30use super::{
31 SegmentWords, claims::PreparedOperandClaims, key_collection::KeyCollection,
32 phase_1::Phase1Output,
33};
34use crate::fold_word::BitAxisFolder;
35
36#[allow(clippy::too_many_arguments)]
60#[instrument(skip_all, name = "prove_phase_2")]
61pub fn prove_phase_2<F, P, Channel, A>(
62 key_collection: &KeyCollection,
63 words: SegmentWords<'_>,
64 prepared: &PreparedOperandClaims<F>,
65 phase_1_output: Phase1Output<F>,
66 shift_ind_eval: F,
67 epsilon: F,
68 channel: &mut Channel,
69 alloc: &A,
70) -> ShiftOutput<F>
71where
72 F: BinaryField,
73 P: PackedField<Scalar = F>,
74 Channel: IPProverChannel<F>,
75 A: Allocator,
76{
77 let Phase1Output {
78 r_j,
79 inner,
80 outer,
81 psi: _,
82 gamma: _,
83 g_eval: _,
84 } = phase_1_output;
85
86 let r_j_tensor = eq_ind_partial_eval::<F>(&r_j);
87
88 let folder = BitAxisFolder::new(r_j_tensor.as_ref());
93 let public_folded = folder.fold::<P, _>(alloc, words.public);
94 let hidden_folded = folder.fold::<P, _>(alloc, words.hidden);
95
96 let (public_monster, hidden_monster) =
97 key_collection.build_monster_segments(alloc, prepared, shift_ind_eval, &inner, &outer);
98
99 let log_segment_words = max(public_folded.log_len(), hidden_folded.log_len());
104 let hidden_folded = hidden_folded.zero_extend_in(alloc, log_segment_words);
105 let hidden_monster = hidden_monster.zero_extend_in(alloc, log_segment_words);
106
107 run_sumcheck(
108 &public_folded,
109 hidden_folded,
110 &public_monster,
111 hidden_monster,
112 shift_ind_eval,
113 words.public,
114 r_j,
115 epsilon,
116 channel,
117 alloc,
118 )
119}
120
121struct SegmentPair<'a, P: PackedField, A: Allocator> {
127 public: &'a FieldVec<P, A>,
129 hidden: FieldVec<P, A>,
131}
132
133impl<'a, P: PackedField, A: Allocator> SegmentPair<'a, P, A> {
134 const fn new(public: &'a FieldVec<P, A>, hidden: FieldVec<P, A>) -> Self {
136 Self { public, hidden }
137 }
138
139 fn fold<F>(self, alpha: F) -> FieldVec<P, A>
145 where
146 F: Field,
147 P: PackedField<Scalar = F>,
148 {
149 let Self { public, mut hidden } = self;
150
151 let alpha_broadcast = P::broadcast(alpha);
153 hidden
154 .as_mut()
155 .par_iter_mut()
156 .with_min_task(WorkPerItem::FieldMuls)
157 .for_each(|hidden_i| *hidden_i *= alpha_broadcast);
158
159 let one_minus_alpha = P::broadcast(F::ONE - alpha);
163 let n_public_packed = public.as_ref().len();
164 for (value, &public_i) in
165 iter::zip(&mut hidden.as_mut()[..n_public_packed], public.as_ref())
166 {
167 *value += public_i * one_minus_alpha;
168 }
169
170 hidden
171 }
172}
173
174fn first_round_coeffs<F, P: PackedField<Scalar = F>, A: Allocator>(
188 witness: &SegmentPair<'_, P, A>,
189 monster: &SegmentPair<'_, P, A>,
190 gamma: F,
191) -> RoundCoeffs<F>
192where
193 F: BinaryField,
194{
195 let wide_dense = (witness.hidden.as_ref(), monster.hidden.as_ref())
197 .into_par_iter()
198 .with_min_task(WorkPerItem::FieldMuls)
199 .map(|(&hidden_i, &monster_i)| P::wide_mul(hidden_i, monster_i))
200 .reduce(<P as WideMul>::Output::default, |lhs, rhs| lhs + rhs);
201
202 let n_public_packed = witness.public.as_ref().len();
204 let (wide_low_hidden, wide_low_cross) = iter::zip(
205 iter::zip(witness.public.as_ref(), &witness.hidden.as_ref()[..n_public_packed]),
206 iter::zip(monster.public.as_ref(), &monster.hidden.as_ref()[..n_public_packed]),
207 )
208 .map(|((&public_i, &hidden_i), (&public_monster_i, &hidden_monster_i))| {
209 (
210 P::wide_mul(hidden_i, hidden_monster_i),
211 P::wide_mul(public_i + hidden_i, public_monster_i + hidden_monster_i),
212 )
213 })
214 .fold(
215 (<P as WideMul>::Output::default(), <P as WideMul>::Output::default()),
216 |(acc_hidden, acc_cross), (hidden_term, cross_term)| {
217 (acc_hidden + hidden_term, acc_cross + cross_term)
218 },
219 );
220
221 let sum_lanes = |wide: <P as WideMul>::Output| P::reduce(wide).iter().sum::<F>();
222 let y_1 = sum_lanes(wide_dense);
223 let y_inf = y_1 + sum_lanes(wide_low_hidden) + sum_lanes(wide_low_cross);
224
225 RoundEvals([y_1, y_inf]).interpolate(gamma)
226}
227
228#[allow(clippy::too_many_arguments)]
253#[instrument(skip_all, name = "run_sumcheck")]
254pub fn run_sumcheck<F, P: PackedField<Scalar = F>, Channel: IPProverChannel<F>, A: Allocator>(
255 public_folded: &FieldVec<P, A>,
256 hidden_folded: FieldVec<P, A>,
257 public_monster: &FieldVec<P, A>,
258 hidden_monster: FieldVec<P, A>,
259 shift_ind_eval: F,
260 public_words: &[Word],
261 r_j: Vec<F>,
262 gamma: F,
263 channel: &mut Channel,
264 alloc: &A,
265) -> ShiftOutput<F>
266where
267 F: BinaryField,
268{
269 let log_hidden = hidden_folded.log_len();
272 assert_eq!(hidden_monster.log_len(), log_hidden);
273 assert_eq!(public_monster.log_len(), public_folded.log_len());
274 assert!(public_folded.log_len() <= log_hidden);
275
276 let witness = SegmentPair::<'_, P, A>::new(public_folded, hidden_folded);
277 let monster = SegmentPair::<'_, P, A>::new(public_monster, hidden_monster);
278
279 let round_coeffs = first_round_coeffs(&witness, &monster, gamma);
281 channel.send_many(round_coeffs.clone().truncate().coeffs());
282 let alpha = channel.sample();
283 let round_sum = round_coeffs.evaluate(&alpha);
284
285 let folded_witness = witness.fold(alpha);
288 let folded_monster = monster.fold(alpha);
289 let prover = bivariate_product_prover(alloc, [folded_witness, folded_monster], round_sum);
290
291 let ProveSingleOutput {
292 multilinear_evals,
293 challenges,
294 } = prove_single(prover, channel);
295
296 let mut r_y = iter::once(alpha).chain(challenges).collect::<Vec<_>>();
297 r_y.reverse();
299
300 let [trace_eval, monster_eval] = multilinear_evals
301 .try_into()
302 .expect("prover has 2 multilinear polynomials");
303
304 let wiring_eval = monster_eval * shift_ind_eval.invert_or_zero();
310
311 let log_half = r_y.len() - 1;
317 let r_segment = r_y[log_half];
318 let log_public_words = log2_ceil_usize(public_words.len());
322 let public_eval = evaluate_words_mle::<F, F>(public_words, &r_j, &r_y[..log_public_words]);
323 let padded_public_eval = eq_ind_zero(&r_y[log_public_words..log_half]) * public_eval;
324 let witness_eval =
325 (trace_eval - (F::ONE - r_segment) * padded_public_eval) * r_segment.invert_or_zero();
326 channel.send_one(witness_eval);
327
328 ShiftOutput {
329 sumcheck: SumcheckOutput {
330 challenges: [r_j, r_y].concat(),
331 eval: witness_eval,
332 },
333 wiring_eval,
334 }
335}
336
337#[derive(Debug)]
342pub struct ShiftOutput<F> {
343 pub sumcheck: SumcheckOutput<F>,
345 pub wiring_eval: F,
347}