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
112 changes: 106 additions & 6 deletions crates/lean_vm/src/gkr.rs
Original file line number Diff line number Diff line change
Expand Up @@ -10,7 +10,6 @@ use crate::PAR_THRESHOLD;
use crate::transcript::{Challenger, ProverState, Receiver, Transmitter, VerifierState};
use primitives::field::{F192, F192Unreduced, mul_unreduced4, mul2, mul4};
use primitives::multilinear::{eq_table, interp, poly_eval, shrink_eq_low};
#[cfg(target_arch = "x86_64")]
use primitives::stream::Stream;
use zk_alloc::ArenaVec;

Expand Down Expand Up @@ -226,6 +225,68 @@ impl QuaternaryLayerState {
self.logical_rows /= 2;
}

fn fold_and_message(&mut self, challenge: F192, equality: &[F192]) -> [F192; 4] {
let stored_rows = self.values.len() / 4;
let rows = stored_rows.div_ceil(2);
self.next.truncate(4 * rows);
let values = &self.values;
let dst = parallel::SendPtr(self.next.as_mut_ptr());
const PAIRS: usize = 16;
let pairs = rows.div_ceil(2);
let task = |index: usize| {
let first = index * PAIRS;
let end = (first + PAIRS).min(pairs);
let mut stage = [F192::ZERO; 8 * PAIRS];
let end_row = (2 * end).min(rows);
for row in 2 * first..end_row {
let lo = 8 * row;
let left = &values[lo..lo + 4];
let right = values.get(lo + 4..lo + 8).unwrap_or(&[F192::ONE; 4]);
let product = mul4(std::array::from_fn(|i| left[i] + right[i]), [challenge; 4]);
let offset = 4 * (row - 2 * first);
for i in 0..4 {
stage[offset + i] = left[i] + product[i];
}
}
let mut message = [F192Unreduced::ZERO; 4];
for pair in first..end {
let lo = 8 * (pair - first);
let left = &stage[lo..lo + 4];
let right = if 2 * pair + 1 < rows {
&stage[lo + 4..lo + 8]
} else {
&[F192::ONE; 4]
};
let lines = std::array::from_fn(|i| [left[i], left[i] + right[i]]);
let terms = quartic_summand(lines, equality[pair]);
for i in 0..4 {
message[i] ^= terms[i];
}
}
// The next round reads the destination; this round reads only the local stage.
let stream = Stream::new();
let len = 4 * (end_row - 2 * first);
// SAFETY: tasks own disjoint initialized prefixes of the output, covering every row.
unsafe { stream.copy(dst.slice(8 * first, len), &stage[..len]) };
message
};
let xor = |mut a: [F192Unreduced; 4], b: [F192Unreduced; 4]| {
for i in 0..4 {
a[i] ^= b[i];
}
a
};
let tasks = pairs.div_ceil(PAIRS);
let message = if rows >= PAR_THRESHOLD {
parallel::map_reduce(tasks, || [F192Unreduced::ZERO; 4], task, xor)
} else {
(0..tasks).map(task).fold([F192Unreduced::ZERO; 4], xor)
};
std::mem::swap(&mut self.values, &mut self.next);
self.logical_rows /= 2;
message.map(F192Unreduced::reduce)
}

fn children(&self) -> [F192; 4] {
debug_assert_eq!(self.values.len(), 4);
debug_assert_eq!(self.logical_rows, 1);
Expand Down Expand Up @@ -314,8 +375,12 @@ pub fn prove_product_triple(leaves: [ArenaVec<F192>; 3], ps: &mut ProverState, s
Vec::new()
};
let mut round_point = Vec::with_capacity(round_count);
for _ in 0..round_count {
let messages = [0, 1, 2].map(|tree| trees[tree].round_message(&equality));
let mut messages = if round_count > 0 {
trees.each_ref().map(|tree| tree.round_message(&equality))
} else {
[[F192::ZERO; 4]; 3]
};
for round in 0..round_count {
let mut coeffs = [0, 1, 2, 3].map(|coefficient| {
messages[0][coefficient] + lambda * (messages[1][coefficient] + lambda * messages[2][coefficient])
});
Expand All @@ -326,10 +391,14 @@ pub fn prove_product_triple(leaves: [ArenaVec<F192>; 3], ps: &mut ProverState, s
ps.add_scalars(&coeffs);
let challenge = ps.sample();
round_point.push(challenge);
for tree in &mut trees {
tree.fold(challenge);
}
shrink_eq_low(&mut equality);
if round + 1 < round_count {
messages = trees.each_mut().map(|tree| tree.fold_and_message(challenge, &equality));
} else {
for tree in &mut trees {
tree.fold(challenge);
}
}
}

for tree in &trees {
Expand Down Expand Up @@ -475,6 +544,37 @@ mod tests {
}
}

#[test]
fn fused_fold_matches_separate_fold_and_message() {
for width in [4usize, 16, 1 << 14] {
for len in [1, 4, 5, 7, 8, 9, 31, 32, 33, 4 * width - 5, 4 * width - 1, 4 * width] {
if len > 4 * width {
continue;
}
let values: ArenaVec<F192> = (0..len)
.map(|i| F192::new((17 * i + 1) as u64, (i * i + 3) as u64, (5 * i + 7) as u64))
.collect();
let mut reference = QuaternaryLayerState::new(ArenaVec::from_slice(&values), width);
let mut fused = QuaternaryLayerState::new(values, width);
let point: Vec<F192> = (0..width.ilog2() - 1)
.map(|i| F192::new(31 + u64::from(i), 7, 11))
.collect();
let mut equality = eq_table(&point);
while reference.logical_rows > 2 {
let challenge = F192::new(reference.logical_rows as u64, 13, 19);
reference.fold(challenge);
shrink_eq_low(&mut equality);
let message = fused.fold_and_message(challenge, &equality);
assert_eq!(message, reference.round_message(&equality), "width={width}, len={len}");
assert_eq!(&*fused.values, &*reference.values, "width={width}, len={len}");
}
reference.fold(F192::Y);
fused.fold(F192::Y);
assert_eq!(fused.children(), reference.children());
}
}
}

#[test]
fn radix_four_roundtrip_at_even_and_odd_depths() {
for mu in 0..=10 {
Expand Down
32 changes: 22 additions & 10 deletions crates/pcs/src/ring_switch.rs
Original file line number Diff line number Diff line change
Expand Up @@ -461,19 +461,31 @@ pub(crate) fn combine_deferred_into(outputs: &[DeferredRingSwitchOutput], out: &
);

parallel::chunks_mut(out, block_len, |hi, out_block| {
for (claim_idx, claim) in outputs.iter().enumerate() {
let e_hi = claim.eq_hi[hi];
if claim_idx == 0 {
for (slot, &e_lo) in out_block.iter_mut().zip(&claim.eq_lo) {
*slot = fold_one_slot_ext(e_lo * e_hi, &claim.table);
}
} else {
for (slot, &e_lo) in out_block.iter_mut().zip(&claim.eq_lo) {
*slot += fold_one_slot_ext(e_lo * e_hi, &claim.table);
combine_deferred_chunk(outputs, hi * block_len, out_block);
});
}

pub(crate) fn combine_deferred_chunk(outputs: &[DeferredRingSwitchOutput], start: usize, out: &mut [F192]) {
for (claim_idx, claim) in outputs.iter().enumerate() {
let block_len = claim.eq_lo.len();
assert!(start + out.len() <= block_len * claim.eq_hi.len());
let mut done = 0;
while done < out.len() {
let index = start + done;
let lo = index % block_len;
let len = (block_len - lo).min(out.len() - done);
let e_hi = claim.eq_hi[index / block_len];
for (slot, &e_lo) in out[done..done + len].iter_mut().zip(&claim.eq_lo[lo..lo + len]) {
let value = fold_one_slot_ext(e_lo * e_hi, &claim.table);
if claim_idx == 0 {
*slot = value;
} else {
*slot += value;
}
}
done += len;
}
});
}
}

/// Split point for the factored eq build: low half sized ~n/2 (min 4, the
Expand Down
94 changes: 94 additions & 0 deletions crates/pcs/src/stack_open.rs
Original file line number Diff line number Diff line change
Expand Up @@ -66,6 +66,8 @@ use super::ring_switch;
use super::whir::{ProverConfig, VerifierConfig};
use super::whir::{ProverData, recursive_prover_with_basis, recursive_verifier_with_basis_succinct};

mod basis;

// ---------------------------------------------------------------------------
// Claim types
// ---------------------------------------------------------------------------
Expand Down Expand Up @@ -413,6 +415,27 @@ pub fn open_batch_mixed_whir_stacked(
let mut target = rs_outputs
.iter()
.fold(F192::ZERO, |acc, out| acc + out.batched_sumcheck_claim);
if stack.len() >= 1 << 16 {
target += point_claims
.iter()
.zip(lambdas_pd)
.fold(F192::ZERO, |sum, (claim, &lambda)| sum + lambda * claim.value());
let lane_block = 1usize << (log_n - config.initial_k);
let (b_stack, message) = basis::build(stack, lane_block, point_claims, lambdas_pd, ring, &rs_outputs);
mark("basis + initial message", &mut t);
super::whir::recursive_prover_with_prepared_basis(
config,
log_n,
stack,
b_stack,
target,
&prover_data.codeword,
&prover_data.merkle_tree,
Some(message),
ps,
);
return;
}
// Parallel first-touch wins for the tower stack: its many scattered point
// claims otherwise fault pages one claim at a time. A scatter that lands on
// slots an earlier one already touched has to accumulate, so those slots
Expand Down Expand Up @@ -563,6 +586,77 @@ mod tests {

const DOMAIN: &[u8] = b"stack-open-test";

#[test]
fn fused_basis_matches_dense_weights() {
let mut rng = Rng::new(0xBA515);
for (lane_vars, lanes) in [(6usize, 1usize), (6, 3), (10, 15), (10, 37)] {
let lane_block = 1 << lane_vars;
let stack: Vec<F64> = (0..lanes * lane_block).map(|_| F64(rng.next_u64())).collect();
let qflock_vars = lane_vars + usize::from(lanes > 1);
let qflock_len = 1 << qflock_vars;
let offset = if stack.len() >= 2 * qflock_len { qflock_len } else { 0 };
let ring = RingSwitchOpen {
offset,
qflock_vars,
claims: (0..2)
.map(|_| RingSwitchClaim {
suffix_point: rng.ext_vec(qflock_vars),
s_hat_v: None,
})
.collect(),
};
let coordinates = rng.ext_vec(192);
let rs_outputs: Vec<_> = ring
.claims
.iter()
.map(|claim| {
let state =
ring_switch::prove_prepare(&stack[offset..offset + qflock_len], &claim.suffix_point, None);
ring_switch::prove_finish_deferred(state, &coordinates, rng.ext())
})
.collect();
let mut claims: Vec<_> = [
(offset, qflock_vars),
((lanes - 1) * lane_block, lane_vars),
(8, 3),
(0, 0),
]
.into_iter()
.map(|(offset, vars)| StackClaim::Point {
offset,
low_point: rng.ext_vec(vars),
value: rng.ext(),
})
.collect();
for stride_log in [0, 1, 3, qflock_vars - 1, qflock_vars] {
claims.push(StackClaim::Strided {
offset,
slot: (1 << stride_log) - 1,
stride_log,
point: rng.ext_vec(qflock_vars - stride_log),
value: rng.ext(),
});
}
let lambdas = rng.ext_vec(claims.len());
let mut expected = vec![F192::ZERO; stack.len()];
ring_switch::combine_deferred_into(&rs_outputs, &mut expected[offset..offset + qflock_len]);
let mut target = F192::ZERO;
fold_stacked_point_claims(
&mut expected,
&mut target,
&claims,
&lambdas,
&vec![false; claims.len()],
);
let (actual, message) = basis::build(&stack, lane_block, &claims, &lambdas, &ring, &rs_outputs);
assert_eq!(&*actual, expected, "lane_vars={lane_vars}, lanes={lanes}");
let (_, expected_message) = super::super::whir::build_initial_basis(&stack, lane_block, |start, dst| {
dst.copy_from_slice(&expected[start..start + dst.len()]);
});
assert_eq!(message, expected_message);
}
}

struct Instance {
vc: VerifierConfig,
log_n: usize,
Expand Down
Loading
Loading