Skip to content
Closed
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
106 changes: 79 additions & 27 deletions crates/sumcheck/src/frontload.rs
Original file line number Diff line number Diff line change
Expand Up @@ -607,13 +607,19 @@ impl<'a, E: ExtensionField> WorkingState<'a, E> {
if !self.worker_matches_frontload_tail(term) {
return vec![E::ZERO; self.poly.aux_info.max_degree + 1];
}
match degree {
1 => return sumcheck_macro::frontload_mixed_sumcheck_code_gen!(1, self, term, round),
2 => return sumcheck_macro::frontload_mixed_sumcheck_code_gen!(2, self, term, round),
3 => return sumcheck_macro::frontload_mixed_sumcheck_code_gen!(3, self, term, round),
4 => return sumcheck_macro::frontload_mixed_sumcheck_code_gen!(4, self, term, round),
5 => return sumcheck_macro::frontload_mixed_sumcheck_code_gen!(5, self, term, round),
_ => {}
let has_worker_bit_round = term
.product
.iter()
.any(|&idx| self.mle_round_is_worker_bit(idx, round));
if !has_worker_bit_round {
match degree {
1 => return sumcheck_macro::frontload_mixed_sumcheck_code_gen!(1, self, term, round),
2 => return sumcheck_macro::frontload_mixed_sumcheck_code_gen!(2, self, term, round),
3 => return sumcheck_macro::frontload_mixed_sumcheck_code_gen!(3, self, term, round),
4 => return sumcheck_macro::frontload_mixed_sumcheck_code_gen!(4, self, term, round),
5 => return sumcheck_macro::frontload_mixed_sumcheck_code_gen!(5, self, term, round),
_ => {}
}
}

let live_vars = term
Expand Down Expand Up @@ -722,12 +728,22 @@ impl<'a, E: ExtensionField> WorkingState<'a, E> {
}

fn mle_round_endpoints(&self, mle_idx: usize, round: usize, lane: usize) -> (E, E) {
let original_num_vars = self.poly.flattened_ml_extensions[mle_idx].num_vars();
if round >= original_num_vars {
return (E::ZERO, read_eval_or_zero(&self.mles[mle_idx], 0));
let local_num_vars = self.poly.flattened_ml_extensions[mle_idx].num_vars();
let global_num_vars = self.global_mle_num_vars[mle_idx];
if round >= local_num_vars {
let value = read_eval_or_zero(&self.mles[mle_idx], 0);
if round < global_num_vars {
let worker_bit = self.worker_round_bit(mle_idx, round);
return if worker_bit == 0 {
(value, E::ZERO)
} else {
(E::ZERO, value)
};
}
return (E::ZERO, value);
}

let remaining_vars = original_num_vars - round;
let remaining_vars = local_num_vars - round;
let suffix_mask = if remaining_vars == 1 {
0
} else {
Expand All @@ -741,20 +757,44 @@ impl<'a, E: ExtensionField> WorkingState<'a, E> {
}

fn fixed_frontload_factor(&self, mle_idx: usize, round: usize) -> E {
let tail_start = self.frontload_tail_start(mle_idx);
if round <= tail_start {
let local_num_vars = self.poly.flattened_ml_extensions[mle_idx].num_vars();
let global_num_vars = self.global_mle_num_vars[mle_idx];
if round <= local_num_vars {
return E::ONE;
}
self.challenges[tail_start..round]
.iter()
.fold(E::ONE, |acc, challenge| acc * challenge.elements)
(local_num_vars..round).fold(E::ONE, |acc, fixed_round| {
let challenge = self.challenges[fixed_round].elements;
if fixed_round < global_num_vars {
let worker_bit = self.worker_round_bit(mle_idx, fixed_round);
let eq = if worker_bit == 0 {
E::ONE - challenge
} else {
challenge
};
acc * eq
} else {
acc * challenge
}
})
}

fn frontload_tail_start(&self, mle_idx: usize) -> usize {
if self.worker.is_some() {
self.poly.flattened_ml_extensions[mle_idx].num_vars()
self.global_mle_num_vars[mle_idx]
}

fn mle_round_is_worker_bit(&self, mle_idx: usize, round: usize) -> bool {
let local_num_vars = self.poly.flattened_ml_extensions[mle_idx].num_vars();
let global_num_vars = self.global_mle_num_vars[mle_idx];
round >= local_num_vars && round < global_num_vars
}

fn worker_round_bit(&self, mle_idx: usize, round: usize) -> usize {
let local_num_vars = self.poly.flattened_ml_extensions[mle_idx].num_vars();
debug_assert!(round >= local_num_vars);
if let Some((worker_id, _log_num_workers)) = self.worker {
(worker_id >> (round - local_num_vars)) & 1
} else {
self.global_mle_num_vars[mle_idx]
1
}
}

Expand Down Expand Up @@ -797,16 +837,28 @@ fn build_phase2_poly<'a, E: ExtensionField>(
poly.aux_info.max_degree = first_worker.poly.aux_info.max_degree;

for (mle_idx, meta) in poly_meta.iter().enumerate() {
let global_num_vars = first_worker.global_mle_num_vars[mle_idx];
let mle = match meta {
FrontloadPolyMeta::Normal => {
let values = workers
.iter()
.map(|worker| {
read_eval(&worker.mles[mle_idx], 0)
* worker.fixed_frontload_factor(mle_idx, local_num_vars)
})
.collect_vec();
MultilinearExtension::from_evaluations_ext_vec(log_num_workers, values)
if global_num_vars <= local_num_vars {
let value = workers
.iter()
.map(|worker| {
read_eval(&worker.mles[mle_idx], 0)
* worker.fixed_frontload_factor(mle_idx, local_num_vars)
})
.sum();
MultilinearExtension::from_evaluations_ext_vec(0, vec![value])
} else {
let values = workers
.iter()
.map(|worker| {
read_eval(&worker.mles[mle_idx], 0)
* worker.fixed_frontload_factor(mle_idx, local_num_vars)
})
.collect_vec();
MultilinearExtension::from_evaluations_ext_vec(log_num_workers, values)
}
}
FrontloadPolyMeta::Phase1Only => {
let value = read_eval(&first_worker.mles[mle_idx], 0)
Expand Down
96 changes: 53 additions & 43 deletions crates/sumcheck/src/test.rs
Original file line number Diff line number Diff line change
Expand Up @@ -76,6 +76,8 @@ fn test_frontload_2phase_sum_keeps_small_mle_compact() {
let large = multilinear_extensions::mle::MultilinearExtension::<GoldilocksExt2>::random(
num_vars, &mut rng,
);
let medium =
multilinear_extensions::mle::MultilinearExtension::<GoldilocksExt2>::random(5, &mut rng);
let small =
multilinear_extensions::mle::MultilinearExtension::<GoldilocksExt2>::random(2, &mut rng);
let poly = VirtualPolynomials::new_from_monimials(
Expand All @@ -86,6 +88,10 @@ fn test_frontload_2phase_sum_keeps_small_mle_compact() {
scalar: Either::Right(GoldilocksExt2::ONE),
product: vec![Either::Left(&large)],
},
Term {
scalar: Either::Right(GoldilocksExt2::ONE),
product: vec![Either::Left(&medium)],
},
Term {
scalar: Either::Right(GoldilocksExt2::ONE),
product: vec![Either::Left(&small)],
Expand All @@ -99,6 +105,7 @@ fn test_frontload_2phase_sum_keeps_small_mle_compact() {

let mut direct_poly = VirtualPolynomial::new(num_vars);
let large_idx = direct_poly.register_mle(Arc::new(large));
let medium_idx = direct_poly.register_mle(Arc::new(medium));
let small_idx = direct_poly.register_mle(Arc::new(small));
direct_poly.aux_info.max_degree = 1;
direct_poly
Expand All @@ -109,6 +116,10 @@ fn test_frontload_2phase_sum_keeps_small_mle_compact() {
scalar: Either::Right(GoldilocksExt2::ONE),
product: vec![large_idx],
},
Term {
scalar: Either::Right(GoldilocksExt2::ONE),
product: vec![medium_idx],
},
Term {
scalar: Either::Right(GoldilocksExt2::ONE),
product: vec![small_idx],
Expand Down Expand Up @@ -211,6 +222,27 @@ fn test_random_monimials_use_frontload_sum() {
&mut rng,
);
let max_num_variables = *nv.iter().max().unwrap();

// Build a single-worker VirtualPolynomial for natural frontload evaluation check.
// Must be built before the mutable borrow in new_from_monimials below.
let mut direct_poly = VirtualPolynomial::new(max_num_variables);
direct_poly.aux_info.max_degree = degree;
for term in &monimials {
let indices: Vec<usize> = term
.product
.iter()
.map(|mle| direct_poly.register_mle(Arc::new(mle.clone())))
.collect_vec();
direct_poly
.products
.push(multilinear_extensions::virtual_poly::MonomialTerms {
terms: vec![Term {
scalar: Either::Right(term.scalar),
product: indices,
}],
});
}

let poly = VirtualPolynomials::<GoldilocksExt2>::new_from_monimials(
4,
max_num_variables,
Expand Down Expand Up @@ -243,8 +275,10 @@ fn test_random_monimials_use_frontload_sum() {
.map(|challenge| challenge.elements)
.collect_vec();
assert_eq!(
worker_aware_frontload_evaluate(4, max_num_variables, &monimials, &point),
subclaim.expected_evaluation
frontload::evaluate(&direct_poly, &point),
subclaim.expected_evaluation,
"frontload 2phase final evaluation mismatch: natural frontload evaluation \
must agree with the verifier's expected evaluation"
);
}

Expand Down Expand Up @@ -312,54 +346,30 @@ fn test_frontload_2phase_mle_category_combinations() {
.iter()
.map(|challenge| challenge.elements)
.collect_vec();
let mut direct_poly = VirtualPolynomial::new(max_num_variables);
direct_poly.aux_info.max_degree = degree;
for Term { scalar, product } in &monomials {
let indices = product
.iter()
.map(|mle| direct_poly.register_mle(Arc::new(mle.clone())))
.collect_vec();
direct_poly
.products
.push(multilinear_extensions::virtual_poly::MonomialTerms {
terms: vec![Term {
scalar: Either::Right(*scalar),
product: indices,
}],
});
}
assert_eq!(
worker_aware_frontload_evaluate(num_threads, max_num_variables, &monomials, &point),
frontload::evaluate(&direct_poly, &point),
subclaim.expected_evaluation,
"frontload 2phase failed for {selected_names}"
);
}
}

fn worker_aware_frontload_evaluate<E: ExtensionField>(
num_threads: usize,
max_num_variables: usize,
monomials: &[Term<E, MultilinearExtension<'_, E>>],
point: &[E],
) -> E {
let log_num_workers = p3::util::log2_strict_usize(num_threads);
let local_num_vars = max_num_variables - log_num_workers;
monomials
.iter()
.map(|Term { scalar, product }| {
product
.iter()
.map(|mle| {
let mle_num_vars = mle.num_vars();
if mle_num_vars > log_num_workers {
let local_real_vars = mle_num_vars - log_num_workers;
let mle_point = point[..local_real_vars]
.iter()
.chain(&point[local_num_vars..max_num_variables])
.copied()
.collect_vec();
let local_tail = point[local_real_vars..local_num_vars]
.iter()
.copied()
.product::<E>();
mle.evaluate(&mle_point) * local_tail
} else {
let local_eval = mle.evaluate(&point[..mle_num_vars]);
point[mle_num_vars..max_num_variables]
.iter()
.fold(local_eval, |acc, point| acc * *point)
}
})
.product::<E>()
* *scalar
})
.sum()
}

// test polynomial mixed with different num_var
#[test]
fn test_sumcheck_with_different_degree() {
Expand Down
Loading