diff --git a/prover/src/tables/bitwise.rs b/prover/src/tables/bitwise.rs index c4871765f..45bddb636 100644 --- a/prover/src/tables/bitwise.rs +++ b/prover/src/tables/bitwise.rs @@ -342,7 +342,7 @@ pub fn preprocessed_commitment(options: &ProofOptions) -> Commitment { /// to zero and will be updated when other tables send lookups. pub fn generate_bitwise_trace() -> TraceTable { let mut trace = TraceTable::new_main( - vec![FE::zero(); NUM_ROWS * cols::NUM_COLUMNS], + crate::tables::types::zeroed_fe_vec(NUM_ROWS * cols::NUM_COLUMNS), cols::NUM_COLUMNS, 1, ); @@ -383,8 +383,9 @@ pub fn generate_bitwise_trace() -> TraceTable cols::MU_MSB8, - BitwiseOperationType::Msb16 => cols::MU_MSB16, - BitwiseOperationType::Zero => cols::MU_ZERO, - BitwiseOperationType::AreBytes => cols::MU_ARE_BYTES, - BitwiseOperationType::IsHalf => cols::MU_IS_HALF, - BitwiseOperationType::IsB20 => cols::MU_IS_B20, - BitwiseOperationType::Hwsl => cols::MU_HWSL, - BitwiseOperationType::ByteAluAnd => cols::MU_BYTE_ALU_AND, - BitwiseOperationType::ByteAluOr => cols::MU_BYTE_ALU_OR, - BitwiseOperationType::ByteAluXor => cols::MU_BYTE_ALU_XOR, - }; + let mu_col = mu_column(op.lookup_type); // Increment multiplicity let current = trace.main_table.get_row(row)[mu_col]; @@ -431,6 +421,180 @@ pub fn update_multiplicities( } } +/// Number of distinct BITWISE lookup types (one multiplicity column each). +/// Derived from [`BitwiseOperationType::ALL`], which the compile-time guard +/// below keeps in lockstep with [`lookup_type_index`]. +pub(crate) const NUM_LOOKUP_TYPES: usize = BitwiseOperationType::ALL.len(); + +/// Dense index in `[0, NUM_LOOKUP_TYPES)` for a lookup type. Ordering is an +/// internal detail of the histogram; [`BitwiseOperationType::ALL`] is its +/// inverse, enforced at compile time. +#[inline] +pub(crate) const fn lookup_type_index(t: BitwiseOperationType) -> usize { + match t { + BitwiseOperationType::Msb8 => 0, + BitwiseOperationType::Msb16 => 1, + BitwiseOperationType::Zero => 2, + BitwiseOperationType::AreBytes => 3, + BitwiseOperationType::IsHalf => 4, + BitwiseOperationType::IsB20 => 5, + BitwiseOperationType::Hwsl => 6, + BitwiseOperationType::ByteAluAnd => 7, + BitwiseOperationType::ByteAluOr => 8, + BitwiseOperationType::ByteAluXor => 9, + } +} + +/// The MU_* multiplicity column for each lookup type, in [`lookup_type_index`] +/// order. This is the single source of truth for the type→column mapping: both +/// the per-op path ([`mu_column`]) and the histogram fill ([`type_mu_column`]) +/// index into this one array. The compile-time block below checks the entries +/// are pairwise distinct, so a duplicate column is a build error rather than a +/// silent overwrite in [`BitwiseHistogram::fill_multiplicities`]. +const MU_COLUMNS: [usize; NUM_LOOKUP_TYPES] = [ + cols::MU_MSB8, // Msb8 + cols::MU_MSB16, // Msb16 + cols::MU_ZERO, // Zero + cols::MU_ARE_BYTES, // AreBytes + cols::MU_IS_HALF, // IsHalf + cols::MU_IS_B20, // IsB20 + cols::MU_HWSL, // Hwsl + cols::MU_BYTE_ALU_AND, // ByteAluAnd + cols::MU_BYTE_ALU_OR, // ByteAluOr + cols::MU_BYTE_ALU_XOR, // ByteAluXor +]; + +/// Multiplicity column for a lookup type. Used by the per-op path +/// ([`update_multiplicities`]), which is still live production code: continuation +/// epochs add their L2G lookups through it on top of the histogram-filled trace. +#[inline] +pub(crate) const fn mu_column(t: BitwiseOperationType) -> usize { + MU_COLUMNS[lookup_type_index(t)] +} + +/// Multiplicity column for the histogram lane at dense index `type_idx` +/// (inverse of [`lookup_type_index`]). Used by [`BitwiseHistogram::fill_multiplicities`]. +/// +/// Reads directly from [`MU_COLUMNS`], the single type→column source of truth. +#[inline] +const fn type_mu_column(type_idx: usize) -> usize { + MU_COLUMNS[type_idx] +} + +// Compile-time guards on the type↔column bookkeeping. +// +// 1. `ALL` must list every lookup type exactly once, in `lookup_type_index` +// order (i.e. it is the exact inverse of that mapping). Adding a variant +// forces the `lookup_type_index` match to be extended (exhaustiveness), and +// this assert then forces `ALL` — and with it `NUM_LOOKUP_TYPES` — to follow. +// 2. The type→column map is now derived from the single `MU_COLUMNS` array, and +// its entries are checked pairwise distinct (injective). A wrong or duplicated +// MU column would silently unbalance the BITWISE bus, so both are compile +// errors, not test failures. +const _: () = { + let mut i = 0; + while i < NUM_LOOKUP_TYPES { + assert!(lookup_type_index(BitwiseOperationType::ALL[i]) == i); + let mut j = i + 1; + while j < NUM_LOOKUP_TYPES { + assert!( + MU_COLUMNS[i] != MU_COLUMNS[j], + "MU_COLUMNS entries must map distinct lookup types to distinct columns" + ); + j += 1; + } + i += 1; + } +}; + +/// "Histogram-on-the-fly" accumulator for BITWISE lookup multiplicities. +/// +/// Replaces materializing the giant `Vec` (whose only consumer +/// is the multiplicity count) with a dense counter array. Each lookup increments +/// `counters[type_idx * NUM_ROWS + row_index(x, y, z)]`. +/// +/// The histogram is a commutative monoid: increments and [`merge`](Self::merge) +/// are order-independent, so per-thread histograms can be tree-reduced and the +/// resulting multiplicities are byte-identical to the serial per-op count that +/// [`update_multiplicities`] produces (both just sum the same lookups per cell). +/// +/// Memory: `NUM_ROWS * NUM_LOOKUP_TYPES * 8` bytes = 2^20 * 10 * 8 = 80 MiB. +pub(crate) struct BitwiseHistogram { + counters: Box<[u64]>, +} + +impl BitwiseHistogram { + /// Allocate a zeroed histogram (80 MiB). + // No `Default` impl on purpose: `new()` allocates 80 MiB, so a stray + // `..Default::default()` / `#[derive(Default)]` must not silently do that. + #[allow(clippy::new_without_default)] + pub(crate) fn new() -> Self { + Self { + counters: vec![0u64; NUM_ROWS * NUM_LOOKUP_TYPES].into_boxed_slice(), + } + } + + /// Increment the counter for one lookup. + #[inline] + pub(crate) fn bump(&mut self, op: BitwiseOperation) { + self.bump_n(op, 1); + } + + /// Add `n` occurrences of one lookup in a single step (e.g. CPU padding rows, + /// which all send identical all-zero lookups). + #[inline] + pub(crate) fn bump_n(&mut self, op: BitwiseOperation, n: u64) { + let idx = lookup_type_index(op.lookup_type) * NUM_ROWS + row_index(op.x, op.y, op.z); + // (x, y) are u8, and row_index debug-asserts z < 16, so in debug builds a + // corrupt op fails loudly here. In release an out-of-domain z would NOT + // panic: the flat index can land in another type's lane and silently + // mis-count both cells — the proof then fails verification instead of the + // prover crashing. What actually upholds the invariant is that every + // `BitwiseOperation` constructor masks or debug-asserts z < 16. + self.counters[idx] += n; + } + + /// Fold a slice of lookups into the histogram. + #[inline] + pub(crate) fn add_ops(&mut self, ops: &[BitwiseOperation]) { + for &op in ops { + self.bump(op); + } + } + + /// Merge another histogram into this one (commutative, order-independent). + pub(crate) fn merge(&mut self, other: &BitwiseHistogram) { + for (a, b) in self.counters.iter_mut().zip(other.counters.iter()) { + *a += *b; + } + } + + /// Write the accumulated multiplicities into the BITWISE trace's MU columns. + /// + /// OVERWRITES each nonzero cell with its count (it does not add to what is + /// there), so it assumes the MU columns are still zero — true for a fresh + /// [`generate_bitwise_trace`] output, where it produces exactly the same MU + /// columns as calling [`update_multiplicities`] with the full op vector. + /// Callers that layer additional lookups on top (continuation epochs add + /// their L2G lookups via `update_multiplicities`, which increments) must do + /// so strictly AFTER this fill, never before. + pub(crate) fn fill_multiplicities( + &self, + trace: &mut TraceTable, + ) { + for type_idx in 0..NUM_LOOKUP_TYPES { + let mu_col = type_mu_column(type_idx); + let base = type_idx * NUM_ROWS; + for row in 0..NUM_ROWS { + let count = self.counters[base + row]; + if count != 0 { + trace.main_table.set_fe(row, mu_col, FE::from(count)); + } + } + } + } +} + /// Types of lookups the BITWISE table provides. #[derive(Debug, Clone, Copy, PartialEq, Eq, Hash)] pub enum BitwiseOperationType { @@ -446,6 +610,24 @@ pub enum BitwiseOperationType { ByteAluXor, } +impl BitwiseOperationType { + /// Every lookup type exactly once, in [`lookup_type_index`] order (the + /// compile-time guard next to [`type_mu_column`] enforces this). The array + /// length is the single origin of [`NUM_LOOKUP_TYPES`]. + pub(crate) const ALL: [Self; 10] = [ + Self::Msb8, + Self::Msb16, + Self::Zero, + Self::AreBytes, + Self::IsHalf, + Self::IsB20, + Self::Hwsl, + Self::ByteAluAnd, + Self::ByteAluOr, + Self::ByteAluXor, + ]; +} + /// A lookup request to the BITWISE precomputed table. /// /// The BITWISE table has 2^20 rows indexed by `(x, y, z)`. diff --git a/prover/src/tables/branch.rs b/prover/src/tables/branch.rs index d4baf10c5..b5bfe83b6 100644 --- a/prover/src/tables/branch.rs +++ b/prover/src/tables/branch.rs @@ -32,7 +32,7 @@ use stark::trace::TraceTable; use std::collections::HashMap; -use super::types::{BusId, FE, GoldilocksExtension, GoldilocksField, SHIFT_16, VmTable, alu_op}; +use super::types::{BusId, GoldilocksExtension, GoldilocksField, SHIFT_16, VmTable, alu_op}; // ========================================================================= // Column indices for BRANCH table @@ -166,7 +166,7 @@ pub fn generate_branch_trace( let unique_ops: Vec<_> = op_map.into_iter().collect(); let num_rows = unique_ops.len().next_power_of_two().max(4); let mut trace = TraceTable::new_main( - vec![FE::zero(); num_rows * cols::NUM_COLUMNS], + crate::tables::types::zeroed_fe_vec(num_rows * cols::NUM_COLUMNS), cols::NUM_COLUMNS, 1, ); diff --git a/prover/src/tables/bytewise.rs b/prover/src/tables/bytewise.rs index 82d7c8772..2808365c6 100644 --- a/prover/src/tables/bytewise.rs +++ b/prover/src/tables/bytewise.rs @@ -19,7 +19,7 @@ use stark::lookup::{BusInteraction, BusValue, Multiplicity, Packing}; use stark::trace::TraceTable; -use super::types::{BusId, FE, GoldilocksExtension, GoldilocksField, VmTable, alu_op}; +use super::types::{BusId, GoldilocksExtension, GoldilocksField, VmTable, alu_op}; // ========================================================================= // Column indices for BYTEWISE table @@ -107,7 +107,7 @@ pub fn generate_bytewise_trace( let unique_ops: Vec<_> = op_map.into_iter().collect(); let num_rows = unique_ops.len().next_power_of_two().max(4); let mut trace = TraceTable::new_main( - vec![FE::zero(); num_rows * cols::NUM_COLUMNS], + crate::tables::types::zeroed_fe_vec(num_rows * cols::NUM_COLUMNS), cols::NUM_COLUMNS, 1, ); diff --git a/prover/src/tables/commit.rs b/prover/src/tables/commit.rs index cd1ca264b..4660c7fb0 100644 --- a/prover/src/tables/commit.rs +++ b/prover/src/tables/commit.rs @@ -163,7 +163,7 @@ pub fn generate_commit_trace( let n = ops.len(); let num_rows = n.next_power_of_two().max(4); let mut trace = TraceTable::new_main( - vec![FE::zero(); num_rows * cols::NUM_COLUMNS], + crate::tables::types::zeroed_fe_vec(num_rows * cols::NUM_COLUMNS), cols::NUM_COLUMNS, 1, ); diff --git a/prover/src/tables/cpu.rs b/prover/src/tables/cpu.rs index 42d197942..781bb02b0 100644 --- a/prover/src/tables/cpu.rs +++ b/prover/src/tables/cpu.rs @@ -24,7 +24,7 @@ //! JALR bit (the memory-width bits are 0), so `mem_flags ∈ {0,1} = JALR` and the //! `mem_flags` column is used directly as `JALR` wherever it is gated by `BRANCH`. -use super::types::{BusId, DecodeEntry, FE, GoldilocksExtension, GoldilocksField, VmTable, alu_op}; +use super::types::{BusId, DecodeEntry, GoldilocksExtension, GoldilocksField, VmTable, alu_op}; use crate::Error; use executor::vm::{ instruction::{decoding::Instruction, execution::SyscallNumbers}, @@ -440,7 +440,7 @@ pub fn generate_cpu_trace( let n = operations.len(); let num_rows = n.next_power_of_two().max(4); let mut trace = TraceTable::new_main( - vec![FE::zero(); num_rows * cols::NUM_COLUMNS], + crate::tables::types::zeroed_fe_vec(num_rows * cols::NUM_COLUMNS), cols::NUM_COLUMNS, 1, ); diff --git a/prover/src/tables/cpu32.rs b/prover/src/tables/cpu32.rs index e0931ff3a..8b1bf86d6 100644 --- a/prover/src/tables/cpu32.rs +++ b/prover/src/tables/cpu32.rs @@ -197,7 +197,7 @@ pub fn generate_cpu32_trace( ) -> TraceTable { let num_rows = operations.len().next_power_of_two().max(4); let mut trace = TraceTable::new_main( - vec![FE::zero(); num_rows * cols::NUM_COLUMNS], + crate::tables::types::zeroed_fe_vec(num_rows * cols::NUM_COLUMNS), cols::NUM_COLUMNS, 1, ); diff --git a/prover/src/tables/decode.rs b/prover/src/tables/decode.rs index 509f86991..bfd1ddb90 100644 --- a/prover/src/tables/decode.rs +++ b/prover/src/tables/decode.rs @@ -128,7 +128,7 @@ pub fn generate_decode_trace( let num_entries = entries.len() + 1; let num_rows = num_entries.next_power_of_two().max(2); let mut trace = TraceTable::new_main( - vec![FE::zero(); num_rows * cols::NUM_COLUMNS], + crate::tables::types::zeroed_fe_vec(num_rows * cols::NUM_COLUMNS), cols::NUM_COLUMNS, 1, ); @@ -393,7 +393,7 @@ fn build_decode_table( let num_entries = entries.len() + 1; let num_rows = num_entries.next_power_of_two().max(2); let mut trace = TraceTable::new_main( - vec![FE::zero(); num_rows * cols::NUM_COLUMNS], + crate::tables::types::zeroed_fe_vec(num_rows * cols::NUM_COLUMNS), cols::NUM_COLUMNS, 1, ); diff --git a/prover/src/tables/dvrm.rs b/prover/src/tables/dvrm.rs index 032963c82..9f979742b 100644 --- a/prover/src/tables/dvrm.rs +++ b/prover/src/tables/dvrm.rs @@ -35,7 +35,7 @@ use stark::trace::TraceTable; use std::collections::HashMap; use super::types::{ - BusId, FE, GoldilocksExtension, GoldilocksField, NEG_INV_2_16, NEG_INV_2_32, NEG_INV_2_48, + BusId, GoldilocksExtension, GoldilocksField, NEG_INV_2_16, NEG_INV_2_32, NEG_INV_2_48, NEG_INV_2_64, SHIFT_16, VmTable, alu_op, }; @@ -298,7 +298,7 @@ pub fn generate_dvrm_trace( let unique_ops: Vec<_> = op_map.into_iter().collect(); let num_rows = unique_ops.len().next_power_of_two().max(4); let mut trace = TraceTable::new_main( - vec![FE::zero(); num_rows * cols::NUM_COLUMNS], + crate::tables::types::zeroed_fe_vec(num_rows * cols::NUM_COLUMNS), cols::NUM_COLUMNS, 1, ); diff --git a/prover/src/tables/ec_scalar.rs b/prover/src/tables/ec_scalar.rs index 5589a14e5..08f797f47 100644 --- a/prover/src/tables/ec_scalar.rs +++ b/prover/src/tables/ec_scalar.rs @@ -83,7 +83,7 @@ pub fn generate_ec_scalar_trace( let n = ops.len(); let num_rows = n.next_power_of_two().max(4); let mut trace = TraceTable::new_main( - vec![FE::zero(); num_rows * cols::NUM_COLUMNS], + crate::tables::types::zeroed_fe_vec(num_rows * cols::NUM_COLUMNS), cols::NUM_COLUMNS, 1, ); diff --git a/prover/src/tables/ecdas.rs b/prover/src/tables/ecdas.rs index 26bbe44e4..2ffde164d 100644 --- a/prover/src/tables/ecdas.rs +++ b/prover/src/tables/ecdas.rs @@ -94,7 +94,7 @@ pub fn generate_ecdas_trace( let n = ops.len(); let num_rows = n.next_power_of_two().max(4); let mut trace = TraceTable::new_main( - vec![FE::zero(); num_rows * cols::NUM_COLUMNS], + crate::tables::types::zeroed_fe_vec(num_rows * cols::NUM_COLUMNS), cols::NUM_COLUMNS, 1, ); diff --git a/prover/src/tables/ecsm.rs b/prover/src/tables/ecsm.rs index 0dba13910..846c31812 100644 --- a/prover/src/tables/ecsm.rs +++ b/prover/src/tables/ecsm.rs @@ -148,7 +148,7 @@ pub fn generate_ecsm_trace( let n = ops.len(); let num_rows = n.next_power_of_two().max(4); let mut trace = TraceTable::new_main( - vec![FE::zero(); num_rows * cols::NUM_COLUMNS], + crate::tables::types::zeroed_fe_vec(num_rows * cols::NUM_COLUMNS), cols::NUM_COLUMNS, 1, ); diff --git a/prover/src/tables/eq.rs b/prover/src/tables/eq.rs index 117d8426b..0f20ca695 100644 --- a/prover/src/tables/eq.rs +++ b/prover/src/tables/eq.rs @@ -26,7 +26,7 @@ use stark::trace::TraceTable; use stark::constraints::builder::{ConstraintBuilder, ConstraintSet}; -use super::types::{BusId, FE, GoldilocksExtension, GoldilocksField, VmTable, alu_op}; +use super::types::{BusId, GoldilocksExtension, GoldilocksField, VmTable, alu_op}; use crate::constraints::templates::{AddOperand, emit_add_pair, emit_is_bit}; // ========================================================================= @@ -128,7 +128,7 @@ pub fn generate_eq_trace( let unique_ops: Vec<_> = op_map.into_iter().collect(); let num_rows = unique_ops.len().next_power_of_two().max(4); let mut trace = TraceTable::new_main( - vec![FE::zero(); num_rows * cols::NUM_COLUMNS], + crate::tables::types::zeroed_fe_vec(num_rows * cols::NUM_COLUMNS), cols::NUM_COLUMNS, 1, ); diff --git a/prover/src/tables/global_memory.rs b/prover/src/tables/global_memory.rs index de6d95d7d..d7bc2ebe9 100644 --- a/prover/src/tables/global_memory.rs +++ b/prover/src/tables/global_memory.rs @@ -123,7 +123,7 @@ pub fn generate_global_trace( ); let num_rows = page_size; // One row per byte in the page - let mut data = vec![FE::zero(); num_rows * cols::NUM_COLUMNS]; + let mut data = crate::tables::types::zeroed_fe_vec(num_rows * cols::NUM_COLUMNS); for offset in 0..page_size { let byte_addr = page_base + (offset as u64); diff --git a/prover/src/tables/halt.rs b/prover/src/tables/halt.rs index 44bbf26cb..675e10170 100644 --- a/prover/src/tables/halt.rs +++ b/prover/src/tables/halt.rs @@ -30,7 +30,7 @@ use stark::lookup::{BusInteraction, BusValue, LinearTerm, Multiplicity, Packing}; use stark::trace::TraceTable; -use super::types::{BusId, FE, GoldilocksExtension, GoldilocksField, VmTable}; +use super::types::{BusId, GoldilocksExtension, GoldilocksField, VmTable}; // ========================================================================= // Column indices for HALT table @@ -72,7 +72,11 @@ pub fn generate_halt_trace( timestamp <= u32::MAX as u64, "HALT timestamp {timestamp} exceeds u32 range" ); - let mut trace = TraceTable::new_main(vec![FE::zero(); cols::NUM_COLUMNS], cols::NUM_COLUMNS, 1); + let mut trace = TraceTable::new_main( + crate::tables::types::zeroed_fe_vec(cols::NUM_COLUMNS), + cols::NUM_COLUMNS, + 1, + ); let table = &mut trace.main_table; table.set_dword_wl(0, cols::TIMESTAMP_0, timestamp); diff --git a/prover/src/tables/keccak.rs b/prover/src/tables/keccak.rs index 832869012..9626b7e3b 100644 --- a/prover/src/tables/keccak.rs +++ b/prover/src/tables/keccak.rs @@ -98,7 +98,7 @@ pub fn generate_keccak_trace( let n = ops.len(); let num_rows = n.next_power_of_two().max(4); let mut trace = TraceTable::new_main( - vec![FE::zero(); num_rows * cols::NUM_COLUMNS], + crate::tables::types::zeroed_fe_vec(num_rows * cols::NUM_COLUMNS), cols::NUM_COLUMNS, 1, ); diff --git a/prover/src/tables/keccak_rc.rs b/prover/src/tables/keccak_rc.rs index f9f0d1cc4..142b5bdde 100644 --- a/prover/src/tables/keccak_rc.rs +++ b/prover/src/tables/keccak_rc.rs @@ -190,7 +190,7 @@ pub fn preprocessed_commitment(options: &ProofOptions) -> Commitment { /// updated via `update_multiplicities` after all round-chip lookups are known. pub fn generate_keccak_rc_trace() -> TraceTable { let mut trace = TraceTable::new_main( - vec![FE::zero(); NUM_ROWS * cols::NUM_COLUMNS], + crate::tables::types::zeroed_fe_vec(NUM_ROWS * cols::NUM_COLUMNS), cols::NUM_COLUMNS, 1, ); diff --git a/prover/src/tables/keccak_rnd.rs b/prover/src/tables/keccak_rnd.rs index 30a50e0b2..afc5dee3a 100644 --- a/prover/src/tables/keccak_rnd.rs +++ b/prover/src/tables/keccak_rnd.rs @@ -244,7 +244,7 @@ pub fn generate_keccak_rnd_trace( ) -> TraceTable { let n_rows = (ops.len() * 24).next_power_of_two().max(4); let mut trace = TraceTable::new_main( - vec![FE::zero(); n_rows * cols::NUM_COLUMNS], + crate::tables::types::zeroed_fe_vec(n_rows * cols::NUM_COLUMNS), cols::NUM_COLUMNS, 1, ); diff --git a/prover/src/tables/load.rs b/prover/src/tables/load.rs index c2bf389dc..1da3cf564 100644 --- a/prover/src/tables/load.rs +++ b/prover/src/tables/load.rs @@ -180,7 +180,7 @@ pub fn generate_load_trace( ) -> TraceTable { let num_rows = operations.len().next_power_of_two().max(4); let mut trace = TraceTable::new_main( - vec![FE::zero(); num_rows * cols::NUM_COLUMNS], + crate::tables::types::zeroed_fe_vec(num_rows * cols::NUM_COLUMNS), cols::NUM_COLUMNS, 1, ); diff --git a/prover/src/tables/local_to_global.rs b/prover/src/tables/local_to_global.rs index 668dc1353..3de72f62a 100644 --- a/prover/src/tables/local_to_global.rs +++ b/prover/src/tables/local_to_global.rs @@ -267,7 +267,7 @@ pub fn generate_local_to_global_trace( boundaries: &[CellBoundary], ) -> TraceTable { let num_rows = boundaries.len().next_power_of_two().max(1); - let mut data = vec![FE::zero(); num_rows * cols::NUM_COLUMNS]; + let mut data = crate::tables::types::zeroed_fe_vec(num_rows * cols::NUM_COLUMNS); for (row, b) in boundaries.iter().enumerate() { let base = row * cols::NUM_COLUMNS; diff --git a/prover/src/tables/lt.rs b/prover/src/tables/lt.rs index a68191f37..86e88a9f7 100644 --- a/prover/src/tables/lt.rs +++ b/prover/src/tables/lt.rs @@ -31,7 +31,7 @@ use stark::trace::TraceTable; use std::collections::HashMap; -use super::types::{BusId, FE, GoldilocksExtension, GoldilocksField, SHIFT_16, VmTable, alu_op}; +use super::types::{BusId, GoldilocksExtension, GoldilocksField, SHIFT_16, VmTable, alu_op}; // ========================================================================= // Column indices for LT table @@ -168,7 +168,7 @@ pub fn generate_lt_trace( let unique_ops: Vec<_> = op_map.into_iter().collect(); let num_rows = unique_ops.len().next_power_of_two().max(4); let mut trace = TraceTable::new_main( - vec![FE::zero(); num_rows * cols::NUM_COLUMNS], + crate::tables::types::zeroed_fe_vec(num_rows * cols::NUM_COLUMNS), cols::NUM_COLUMNS, 1, ); diff --git a/prover/src/tables/memw.rs b/prover/src/tables/memw.rs index 4f775f535..338b85467 100644 --- a/prover/src/tables/memw.rs +++ b/prover/src/tables/memw.rs @@ -35,7 +35,7 @@ use stark::trace::TraceTable; use stark::constraints::builder::{ConstraintBuilder, ConstraintSet}; -use super::types::{BusId, FE, GoldilocksExtension, GoldilocksField, VmTable, alu_op}; +use super::types::{BusId, GoldilocksExtension, GoldilocksField, VmTable, alu_op}; use crate::constraints::templates::emit_is_bit; /// Maximum number of rows per MEMW table chunk. @@ -107,17 +107,22 @@ pub struct MemwOperation { pub is_register: bool, /// Base address (64-bit) pub base_address: u64, - /// Values to write (8 bytes) - pub value: [u64; 8], + /// Values to write. Each element is one memory byte (0-255) or, for register + /// accesses, a 32-bit half of the register word — both fit in u32, so this is + /// `[u32; 8]` rather than `[u64; 8]` to halve the struct's footprint (the walk + /// materializes tens of millions of these; it is memory-bandwidth-bound). + pub value: [u32; 8], /// Timestamp of this access pub timestamp: u64, /// Access width: 1, 2, 4, or 8 bytes pub width: u8, /// Whether this is a read (true) or write (false) pub is_read: bool, - /// Previous values at the addresses (filled by memory model) - pub old: [u64; 8], - /// Previous timestamps at the addresses (filled by memory model) + /// Previous values at the addresses (filled by memory model). Same element + /// domain as `value` (byte or 32-bit register half) → `[u32; 8]`. + pub old: [u32; 8], + /// Previous timestamps at the addresses (filled by memory model). Timestamps + /// can reach u64::MAX (HALT), so these stay `[u64; 8]`. pub old_timestamp: [u64; 8], } @@ -126,7 +131,7 @@ impl MemwOperation { pub fn new( is_register: bool, base_address: u64, - value: [u64; 8], + value: [u32; 8], timestamp: u64, width: u8, is_read: bool, @@ -144,7 +149,7 @@ impl MemwOperation { } /// Set the old values (from memory model). - pub fn with_old(mut self, old: [u64; 8], old_timestamp: [u64; 8]) -> Self { + pub fn with_old(mut self, old: [u32; 8], old_timestamp: [u64; 8]) -> Self { self.old = old; self.old_timestamp = old_timestamp; self @@ -175,7 +180,7 @@ pub fn generate_memw_trace( ) -> TraceTable { let num_rows = operations.len().next_power_of_two().max(4); let mut trace = TraceTable::new_main( - vec![FE::zero(); num_rows * cols::NUM_COLUMNS], + crate::tables::types::zeroed_fe_vec(num_rows * cols::NUM_COLUMNS), cols::NUM_COLUMNS, 1, ); @@ -191,7 +196,7 @@ pub fn generate_memw_trace( // value[8] for i in 0..8 { - table.set_u64(row_idx, cols::VALUE[i], op.value[i]); + table.set_u64(row_idx, cols::VALUE[i], op.value[i] as u64); } // timestamp as DWordWL (2 words) @@ -205,7 +210,7 @@ pub fn generate_memw_trace( // Output: old[8] for i in 0..8 { - table.set_u64(row_idx, cols::OLD[i], op.old[i]); + table.set_u64(row_idx, cols::OLD[i], op.old[i] as u64); } // Auxiliary: carry[7] diff --git a/prover/src/tables/memw_aligned.rs b/prover/src/tables/memw_aligned.rs index b9517ec91..0853bf5ff 100644 --- a/prover/src/tables/memw_aligned.rs +++ b/prover/src/tables/memw_aligned.rs @@ -42,7 +42,7 @@ use stark::trace::TraceTable; use stark::constraints::builder::{ConstraintBuilder, ConstraintSet}; use super::memw::MemwOperation; -use super::types::{BusId, FE, GoldilocksExtension, GoldilocksField, VmTable, alu_op}; +use super::types::{BusId, GoldilocksExtension, GoldilocksField, VmTable, alu_op}; use crate::constraints::templates::emit_is_bit; /// Maximum number of rows per MEMW_A table chunk. @@ -95,7 +95,7 @@ pub fn generate_memw_aligned_trace( ) -> TraceTable { let num_rows = operations.len().next_power_of_two().max(4); let mut trace = TraceTable::new_main( - vec![FE::zero(); num_rows * cols::NUM_COLUMNS], + crate::tables::types::zeroed_fe_vec(num_rows * cols::NUM_COLUMNS), cols::NUM_COLUMNS, 1, ); @@ -107,7 +107,7 @@ pub fn generate_memw_aligned_trace( table.set_dword_whh(row_idx, cols::BASE_ADDRESS[0], op.base_address); for i in 0..8 { - table.set_u64(row_idx, cols::VALUE[i], op.value[i]); + table.set_u64(row_idx, cols::VALUE[i], op.value[i] as u64); } table.set_dword_wl(row_idx, cols::TIMESTAMP_0, op.timestamp); @@ -118,7 +118,7 @@ pub fn generate_memw_aligned_trace( table.set_bool(row_idx, cols::WRITE8, w8); for i in 0..8 { - table.set_u64(row_idx, cols::OLD[i], op.old[i]); + table.set_u64(row_idx, cols::OLD[i], op.old[i] as u64); } // Single old_timestamp (from old_timestamp[0], verified equal for all bytes) diff --git a/prover/src/tables/memw_register.rs b/prover/src/tables/memw_register.rs index c02380c5f..590a55100 100644 --- a/prover/src/tables/memw_register.rs +++ b/prover/src/tables/memw_register.rs @@ -17,7 +17,7 @@ //! //! ## Column layout (10 columns) //! -//! - `ADDRESS`: Byte (register index 0-31) +//! - `ADDRESS`: Byte (register index 0-255: x0-x31, plus x254/x255) //! - `TIMESTAMP_0`: Word (low 32 bits) //! - `TIMESTAMP_1`: Word (high 32 bits) //! - `VAL_0`: Word (low 32 bits of register value) @@ -43,8 +43,9 @@ use stark::trace::TraceTable; use stark::constraints::builder::{ConstraintBuilder, ConstraintSet}; +use super::bitwise::{BitwiseHistogram, BitwiseOperation, BitwiseOperationType}; use super::memw::MemwOperation; -use super::types::{BusId, FE, GoldilocksExtension, GoldilocksField, VmTable}; +use super::types::{BusId, GoldilocksExtension, GoldilocksField, VmTable}; use crate::constraints::templates::emit_is_bit; // ========================================================================= @@ -52,7 +53,7 @@ use crate::constraints::templates::emit_is_bit; // ========================================================================= pub mod cols { - /// Register index (0-31). CPU sends base_address = 2*reg_index. + /// Register index (0-255: x0-x31, plus x254/x255). CPU sends base_address = 2*reg_index. pub const ADDRESS: usize = 0; /// Timestamp low 32 bits @@ -85,28 +86,87 @@ pub mod cols { // Trace generation // ========================================================================= -/// Generates the MEMW_R trace table from register operations. +/// Compact, already-decomposed record for one MEMW_R (register fast-path) access. /// -/// Reuses `MemwOperation` -- the trace generator divides `base_address` by 2 -/// to recover the register index (CPU sends `2 * register_index`). -pub fn generate_memw_register_trace( - operations: &[MemwOperation], -) -> TraceTable { - let num_rows = operations.len().next_power_of_two().max(4); - let mut trace = TraceTable::new_main( - vec![FE::zero(); num_rows * cols::NUM_COLUMNS], - cols::NUM_COLUMNS, - 1, - ); - let table = &mut trace.main_table; +/// This is the "direct-to-column" carrier: it holds exactly the fields the MEMW_R +/// column fill ([`generate_memw_register_trace_from_rows`]) and its IS_HALFWORD +/// bitwise collector ([`collect_bitwise_from_memw_register`]) need, and nothing +/// else. It replaces the full `MemwOperation` (~152 B after the `[u32; 8]` +/// value/old shrink, but still 8-element arrays) for register accesses — the +/// largest table by rows — so the walk never materializes a `MemwOperation` for +/// the register fast path. +/// +/// Field domains mirror `MemwOperation`'s: +/// - `address` = `base_address / 2` (the register index 0..=255; ADDRESS column, +/// and `2*ADDRESS` on the memory/MEMW buses) +/// - `val0/val1` = `value[0]`/`value[1]` (the 32-bit register halves) +/// - `old0/old1` = `old[0]`/`old[1]` +/// - `old_ts_lo` = `old_timestamp[0] & 0xFFFF_FFFF` (the two words share old_timestamp, +/// enforced by `is_register_op`; the upper limb is TIMESTAMP_1 = timestamp>>32) +#[derive(Debug, Clone, Copy)] +pub(crate) struct RegRow { + /// Register index 0..=255 (`base_address / 2`); u16 keeps the struct at + /// 32 bytes — it is the largest persisted array of the walk. + address: u16, + timestamp: u64, + val0: u32, + val1: u32, + old0: u32, + old1: u32, + old_ts_lo: u32, + is_read: bool, +} - for (row_idx, op) in operations.iter().enumerate() { +impl RegRow { + /// Build a `RegRow` from pre-decomposed register-access fields. + /// + /// `reg_addr` is `2 * reg_index` as sent by the CPU; `old_ts` is the (shared) + /// old_timestamp of both register words. This is the ONLY place the MEMW_R + /// row encoding (halved address, masked `old_ts_lo`) is defined — + /// [`Self::from_memw`] delegates here, so the walk fast path and the + /// `MemwOperation` paths cannot drift. + #[inline] + #[allow(clippy::too_many_arguments)] + pub(crate) fn new( + reg_addr: u64, + timestamp: u64, + val0: u32, + val1: u32, + old0: u32, + old1: u32, + old_ts: u64, + is_read: bool, + ) -> Self { debug_assert_eq!( - op.base_address % 2, + reg_addr % 2, 0, - "register base_address must be even (got {})", - op.base_address + "register base_address must be even (got {reg_addr})" ); + debug_assert!( + reg_addr / 2 <= u16::MAX as u64, + "register index exceeds u16 (got base_address {reg_addr})" + ); + RegRow { + address: (reg_addr / 2) as u16, + timestamp, + val0, + val1, + old0, + old1, + old_ts_lo: (old_ts & 0xFFFF_FFFF) as u32, + is_read, + } + } + + /// Build a `RegRow` from a fully-formed register `MemwOperation`. Used on the + /// precompile / commit / keccak / halt paths, which construct a `MemwOperation` + /// first and only convert to the compact row once the op is known to route to + /// MEMW_R. + /// + /// Only valid for ops for which `is_register_op` is true (width==2, atomic + /// old_timestamp). + #[inline] + pub(crate) fn from_memw(op: &MemwOperation) -> Self { // Both register words must have been last accessed at the same timestamp. // MEMW_R stores a single old_timestamp_lo and shares TIMESTAMP_1 as the // upper limb, so if the two words differ, the wrong token would be sent @@ -116,34 +176,101 @@ pub fn generate_memw_register_trace( "register words must share old_timestamp ({} != {})", op.old_timestamp[0], op.old_timestamp[1] ); + Self::new( + op.base_address, + op.timestamp, + op.value[0], + op.value[1], + op.old[0], + op.old[1], + op.old_timestamp[0], + op.is_read, + ) + } +} - // ADDRESS = base_address / 2 (CPU sends 2 * register_index) - table.set_u64(row_idx, cols::ADDRESS, op.base_address / 2); +/// Generates the MEMW_R trace table from register operations. +/// +/// Thin wrapper over [`generate_memw_register_trace_from_rows`] (via +/// [`RegRow::from_memw`]) so there is exactly one MEMW_R column-write sequence. +/// +/// Test-only: production code fills MEMW_R directly from [`RegRow`]s, so the walk +/// never routes through this `MemwOperation`-based entry point. +#[cfg(test)] +pub(crate) fn generate_memw_register_trace( + operations: &[MemwOperation], +) -> TraceTable { + let rows: Vec = operations.iter().map(RegRow::from_memw).collect(); + generate_memw_register_trace_from_rows(&rows) +} - // Timestamp split into lo/hi 32-bit words - table.set_dword_wl(row_idx, cols::TIMESTAMP_0, op.timestamp); +/// The MEMW_R column fill from compact [`RegRow`]s. This is the single source of +/// truth for the MEMW_R trace layout; both the walk's direct fast path and the +/// `MemwOperation`-based `generate_memw_register_trace` test wrapper land here. +pub(crate) fn generate_memw_register_trace_from_rows( + rows: &[RegRow], +) -> TraceTable { + let num_rows = rows.len().next_power_of_two().max(4); + let mut trace = TraceTable::new_main( + crate::tables::types::zeroed_fe_vec(num_rows * cols::NUM_COLUMNS), + cols::NUM_COLUMNS, + 1, + ); + let table = &mut trace.main_table; - // Value: registers are DWordWL = 2 words - table.set_u64(row_idx, cols::VAL_0, op.value[0]); - table.set_u64(row_idx, cols::VAL_1, op.value[1]); + for (row_idx, r) in rows.iter().enumerate() { + // ADDRESS = base_address / 2 (already divided in RegRow). + table.set_u64(row_idx, cols::ADDRESS, r.address as u64); + // Timestamp split into lo/hi 32-bit words. + table.set_dword_wl(row_idx, cols::TIMESTAMP_0, r.timestamp); + // Value: registers are DWordWL = 2 words. + table.set_u64(row_idx, cols::VAL_0, r.val0 as u64); + table.set_u64(row_idx, cols::VAL_1, r.val1 as u64); + // Old value. + table.set_u64(row_idx, cols::OLD_0, r.old0 as u64); + table.set_u64(row_idx, cols::OLD_1, r.old1 as u64); + // Old timestamp low (upper limb shared with TIMESTAMP_1). + table.set_u64(row_idx, cols::OLD_TIMESTAMP_LO, r.old_ts_lo as u64); + // Multiplicity. + table.set_bool(row_idx, cols::MU_READ, r.is_read); + table.set_bool(row_idx, cols::MU_WRITE, !r.is_read); + } - // Old value - table.set_u64(row_idx, cols::OLD_0, op.old[0]); - table.set_u64(row_idx, cols::OLD_1, op.old[1]); + trace +} - // Old timestamp low (upper limb shared with TIMESTAMP_1) - table.set_u64( - row_idx, - cols::OLD_TIMESTAMP_LO, - op.old_timestamp[0] & 0xFFFF_FFFF, - ); +/// The single IS_HALFWORD lookup a MEMW_R access sends: proves the timestamp delta +/// `ts_lo - old_ts_lo` is in [1, 2^16] by decomposing `ts_lo - old_ts_lo - 1` into +/// two bytes. +/// +/// Must stay in lockstep with the IS_HALFWORD send in [`bus_interactions`]: the +/// lookup counted here has to be exactly the lookup each MEMW_R row sends, or the +/// BITWISE bus goes unbalanced. +#[inline] +fn memw_register_is_half_lookup(ts_lo: u32, old_ts_lo: u32) -> BitwiseOperation { + debug_assert!( + ts_lo > old_ts_lo, + "ts_lo must exceed old_ts_lo (enforced by the MEMW_R routing predicate)" + ); + let diff_minus_1 = (ts_lo - old_ts_lo - 1) as u16; + BitwiseOperation::halfword( + BitwiseOperationType::IsHalf, + (diff_minus_1 & 0xFF) as u8, + (diff_minus_1 >> 8) as u8, + ) +} - // Multiplicity - table.set_bool(row_idx, cols::MU_READ, op.is_read); - table.set_bool(row_idx, cols::MU_WRITE, !op.is_read); +/// IS_HALFWORD bitwise lookups for MEMW_R, bumped straight into the histogram +/// via the shared [`memw_register_is_half_lookup`] helper (the same lookup the +/// MEMW_R trace fill uses), one per row. No intermediate op vector: register +/// rows number in the tens of millions and the histogram is the only consumer. +pub(crate) fn collect_bitwise_from_memw_register(rows: &[RegRow], hist: &mut BitwiseHistogram) { + for r in rows { + hist.bump(memw_register_is_half_lookup( + (r.timestamp & 0xFFFF_FFFF) as u32, + r.old_ts_lo, + )); } - - trace } // ========================================================================= diff --git a/prover/src/tables/mul.rs b/prover/src/tables/mul.rs index 2f0fa1d0e..181fba514 100644 --- a/prover/src/tables/mul.rs +++ b/prover/src/tables/mul.rs @@ -36,7 +36,7 @@ use stark::trace::TraceTable; use std::collections::HashMap; use super::types::{ - BusId, FE, GoldilocksExtension, GoldilocksField, INV_2_32, INV_2_64, INV_2_96, INV_2_128, + BusId, GoldilocksExtension, GoldilocksField, INV_2_32, INV_2_64, INV_2_96, INV_2_128, NEG_INV_2_16, NEG_INV_2_32, NEG_INV_2_48, NEG_INV_2_64, NEG_INV_2_80, NEG_INV_2_96, NEG_INV_2_112, NEG_INV_2_128, SHIFT_16, VmTable, alu_op, }; @@ -306,7 +306,7 @@ pub fn generate_mul_trace( let unique_ops: Vec<_> = op_map.into_iter().collect(); let num_rows = unique_ops.len().next_power_of_two().max(4); let mut trace = TraceTable::new_main( - vec![FE::zero(); num_rows * cols::NUM_COLUMNS], + crate::tables::types::zeroed_fe_vec(num_rows * cols::NUM_COLUMNS), cols::NUM_COLUMNS, 1, ); diff --git a/prover/src/tables/page.rs b/prover/src/tables/page.rs index 5cba30435..18ce6b52b 100644 --- a/prover/src/tables/page.rs +++ b/prover/src/tables/page.rs @@ -249,7 +249,7 @@ pub fn generate_page_trace( let num_rows = page_size; // One row per byte in the page let mut trace = TraceTable::new_main( - vec![FE::zero(); num_rows * cols::NUM_COLUMNS], + crate::tables::types::zeroed_fe_vec(num_rows * cols::NUM_COLUMNS), cols::NUM_COLUMNS, 1, ); @@ -361,8 +361,8 @@ pub fn compute_precomputed_commitment(config: &PageConfig, options: &ProofOption // bytes loaded from the binary. Either way the column is fully determined // before execution, so the verifier can check it against a preprocessed // commitment instead of including it in the main trace. - let mut offset_col = vec![FE::zero(); num_rows]; - let mut init_col = vec![FE::zero(); num_rows]; + let mut offset_col = crate::tables::types::zeroed_fe_vec(num_rows); + let mut init_col = crate::tables::types::zeroed_fe_vec(num_rows); for i in 0..page_size { offset_col[i] = FE::from(i as u64); diff --git a/prover/src/tables/register.rs b/prover/src/tables/register.rs index 46c675b65..4da1f5efe 100644 --- a/prover/src/tables/register.rs +++ b/prover/src/tables/register.rs @@ -220,7 +220,7 @@ pub fn generate_register_trace( ) -> TraceTable { let num_rows = NUM_REGISTER_ADDRESSES.next_power_of_two(); let mut trace = TraceTable::new_main( - vec![FE::zero(); num_rows * cols::NUM_COLUMNS], + crate::tables::types::zeroed_fe_vec(num_rows * cols::NUM_COLUMNS), cols::NUM_COLUMNS, 1, ); @@ -281,8 +281,8 @@ pub fn compute_precomputed_commitment(options: &ProofOptions, init: &[u32]) -> C let num_rows = NUM_REGISTER_ADDRESSES.next_power_of_two(); let addr_list = register_word_address_list(); - let mut offset_col = vec![FE::zero(); num_rows]; - let mut init_col = vec![FE::zero(); num_rows]; + let mut offset_col = crate::tables::types::zeroed_fe_vec(num_rows); + let mut init_col = crate::tables::types::zeroed_fe_vec(num_rows); for i in 0..NUM_REGISTER_ADDRESSES { offset_col[i] = FE::from(addr_list[i]); @@ -308,9 +308,9 @@ pub fn compute_precomputed_commitment_with_fini( let num_rows = NUM_REGISTER_ADDRESSES.next_power_of_two(); let addr_list = register_word_address_list(); - let mut offset_col = vec![FE::zero(); num_rows]; - let mut init_col = vec![FE::zero(); num_rows]; - let mut fini_col = vec![FE::zero(); num_rows]; + let mut offset_col = crate::tables::types::zeroed_fe_vec(num_rows); + let mut init_col = crate::tables::types::zeroed_fe_vec(num_rows); + let mut fini_col = crate::tables::types::zeroed_fe_vec(num_rows); for i in 0..NUM_REGISTER_ADDRESSES { offset_col[i] = FE::from(addr_list[i]); diff --git a/prover/src/tables/shift.rs b/prover/src/tables/shift.rs index 77a8ae32a..5ac5a393f 100644 --- a/prover/src/tables/shift.rs +++ b/prover/src/tables/shift.rs @@ -21,7 +21,7 @@ use stark::lookup::{BusInteraction, BusValue, LinearTerm, Multiplicity, Packing}; use stark::trace::TraceTable; -use super::types::{BusId, FE, GoldilocksExtension, GoldilocksField, SHIFT_16, VmTable, alu_op}; +use super::types::{BusId, GoldilocksExtension, GoldilocksField, SHIFT_16, VmTable, alu_op}; // ========================================================================= // Column indices @@ -357,7 +357,7 @@ pub fn generate_shift_trace( // Spec declares μ: Bit. let num_rows = operations.len().next_power_of_two().max(4); let mut trace = TraceTable::new_main( - vec![FE::zero(); num_rows * cols::NUM_COLUMNS], + crate::tables::types::zeroed_fe_vec(num_rows * cols::NUM_COLUMNS), cols::NUM_COLUMNS, 1, ); diff --git a/prover/src/tables/store.rs b/prover/src/tables/store.rs index c1dfc937a..ac30832e1 100644 --- a/prover/src/tables/store.rs +++ b/prover/src/tables/store.rs @@ -97,7 +97,7 @@ pub fn generate_store_trace( ) -> TraceTable { let num_rows = operations.len().next_power_of_two().max(4); let mut trace = TraceTable::new_main( - vec![FE::zero(); num_rows * cols::NUM_COLUMNS], + crate::tables::types::zeroed_fe_vec(num_rows * cols::NUM_COLUMNS), cols::NUM_COLUMNS, 1, ); diff --git a/prover/src/tables/trace_builder.rs b/prover/src/tables/trace_builder.rs index ecd0b87ab..36fd107c3 100644 --- a/prover/src/tables/trace_builder.rs +++ b/prover/src/tables/trace_builder.rs @@ -60,7 +60,7 @@ use super::local_to_global; use super::lt::{self, LtOperation}; use super::memw::{self, MemwOperation}; use super::memw_aligned; -use super::memw_register; +use super::memw_register::{self, RegRow}; use super::mul::{self, MulOperation}; use super::page::{self, FinalByteState, FinalStateMap, PageConfig}; use super::register::{self, FinalRegisterStateMap, FinalRegisterWordState}; @@ -133,12 +133,12 @@ impl MemoryState { } /// Read multiple bytes. Returns arrays of values and timestamps. - fn read_bytes(&self, base_address: u64, count: usize) -> ([u64; 8], [u64; 8]) { - let mut values = [0u64; 8]; + fn read_bytes(&self, base_address: u64, count: usize) -> ([u32; 8], [u64; 8]) { + let mut values = [0u32; 8]; let mut timestamps = [0u64; 8]; for i in 0..count { let (val, ts) = self.read_byte(base_address.wrapping_add(i as u64)); - values[i] = val as u64; + values[i] = val as u32; timestamps[i] = ts; } (values, timestamps) @@ -313,8 +313,17 @@ fn cpu_op_to_bytes_and_signed(op: &CpuOperation) -> (usize, bool) { /// Pack a 64-bit register value into the MEMW value format. /// /// For register operations, values are packed as [lo32, hi32, 0, 0, 0, 0, 0, 0]. -fn pack_register_value(value: u64) -> [u64; 8] { - [value & 0xFFFF_FFFF, value >> 32, 0, 0, 0, 0, 0, 0] +fn pack_register_value(value: u64) -> [u32; 8] { + [ + (value & 0xFFFF_FFFF) as u32, + (value >> 32) as u32, + 0, + 0, + 0, + 0, + 0, + 0, + ] } // ============================================================================= @@ -352,25 +361,186 @@ fn collect_cpu_ops( // Phase 2: CPU ops → MEMW, LOAD, LT, Bitwise // ============================================================================= +/// Destination table for a `MemwOperation`. +/// +/// The order of the checks matters and must never change: register ops would +/// also pass `is_aligned_op`, so MEMW_R is decided first, then MEMW_A, and the +/// rest goes to the general MEMW table. +#[derive(Debug, Clone, Copy, PartialEq, Eq)] +enum MemwRoute { + Register, + Aligned, + General, +} + +/// The single classification used everywhere a `MemwOperation` is routed to a +/// table — the walk's [`MemwBuckets`] and the sizing pass (`count_table_lengths`) +/// share it, so their routing cannot drift. +#[inline] +fn classify_memw(op: &MemwOperation) -> MemwRoute { + if is_register_op(op) { + MemwRoute::Register + } else if is_aligned_op(op) { + MemwRoute::Aligned + } else { + MemwRoute::General + } +} + +/// Routes each `MemwOperation` into its destination table bucket at CREATION time +/// (register fast-path / aligned / general), so the walk fills the three buckets directly +/// and no separate routing pass is needed downstream. Classification order is register +/// first, then aligned (see [`classify_memw`]), and push order within each bucket is the +/// walk's insertion order — the buckets are fully deterministic, which the per-cell +/// multiplicity counts rely on. +/// +/// ## Direct-to-column register fill +/// +/// For the register fast path we do NOT materialize a `Vec`. Ops that route +/// to MEMW_R are stored as compact [`RegRow`]s (`register_rows`) and later filled directly +/// into the MEMW_R columns. The `aligned` / `general` buckets hold `MemwOperation`s — an op +/// that FAILS `is_register_op` is routed there (aligned if `is_aligned_op`, else general). +#[derive(Default)] +struct MemwBuckets { + /// Compact register rows (filled directly into the MEMW_R columns). + register_rows: Vec, + aligned: Vec, + general: Vec, +} + +impl MemwBuckets { + fn with_register_capacity(n: usize) -> Self { + Self { + register_rows: Vec::with_capacity(n), + aligned: Vec::new(), + general: Vec::new(), + } + } + + #[inline] + fn push(&mut self, op: MemwOperation) { + match classify_memw(&op) { + MemwRoute::Register => self.register_rows.push(RegRow::from_memw(&op)), + MemwRoute::Aligned => self.aligned.push(op), + MemwRoute::General => self.general.push(op), + } + } + fn extend_ops(&mut self, ops: impl IntoIterator) { + for op in ops { + self.push(op); + } + } +} + +/// Sink for `MemwOperation`s so `collect_register_ops_from_cpu` can feed either a plain +/// `Vec` (the `count_table_lengths` trace-sizing pass) or the classifying +/// [`MemwBuckets`] (the walk). +trait MemwSink { + fn push_op(&mut self, op: MemwOperation); + + /// Fast path for a 2-word register access (M1/M3/M5 and precompile register I/O). + /// + /// The caller passes the compact, pre-decomposed fields. The sink decides routing + /// (via the same predicate as `is_register_op`): if the timestamp delta admits the + /// op into MEMW_R it fills a compact [`RegRow`] DIRECTLY — no `MemwOperation` is + /// built. Only on the (rare) fallback (delta out of IS_HALF range, or upper-limb + /// mismatch) does it build the `MemwOperation` (via [`build_reg_fallback`]) and + /// route it to the aligned/general bucket exactly as before. + /// + /// `reg_addr` is `2 * reg_index`; `val`/`old` are the two 32-bit halves of the new + /// and previous register words; `old_ts` is the (shared) old_timestamp of both words. + #[inline] + fn push_reg_access( + &mut self, + reg_addr: u64, + val: [u32; 2], + old: [u32; 2], + timestamp: u64, + old_ts: u64, + is_read: bool, + ) { + // Default impl (plain Vec): register accesses are still ordinary MemwOperations. + self.push_op(build_reg_fallback( + reg_addr, val, old, timestamp, old_ts, is_read, + )); + } +} +impl MemwSink for Vec { + #[inline] + fn push_op(&mut self, op: MemwOperation) { + self.push(op); + } +} +impl MemwSink for MemwBuckets { + #[inline] + fn push_op(&mut self, op: MemwOperation) { + self.push(op); + } + + #[inline] + fn push_reg_access( + &mut self, + reg_addr: u64, + val: [u32; 2], + old: [u32; 2], + timestamp: u64, + old_ts: u64, + is_read: bool, + ) { + // Mirror `is_register_op` for a width-2 register access whose two words share + // `old_ts` (always true here by construction). If it passes, fill a RegRow + // directly; otherwise fall back to the general/aligned MemwOperation path. + if reg_ts_delta_in_range(timestamp, old_ts) { + self.register_rows.push(RegRow::new( + reg_addr, timestamp, val[0], val[1], old[0], old[1], old_ts, is_read, + )); + } else { + let op = build_reg_fallback(reg_addr, val, old, timestamp, old_ts, is_read); + debug_assert!(!is_register_op(&op), "reg fallback must not be MEMW_R"); + self.push(op); + } + } +} + +/// Materialize the aligned/general `MemwOperation` for a register access that does +/// NOT fit the MEMW_R fast path. Register values pack as `[lo, hi, 0, …]` (see +/// [`pack_register_value`]) and both words share `old_ts`, so this rebuilds exactly +/// the op the fast-path callers would otherwise have routed to the buckets. +fn build_reg_fallback( + reg_addr: u64, + val: [u32; 2], + old: [u32; 2], + timestamp: u64, + old_ts: u64, + is_read: bool, +) -> MemwOperation { + let value = [val[0], val[1], 0, 0, 0, 0, 0, 0]; + let old_value = [old[0], old[1], 0, 0, 0, 0, 0, 0]; + let old_timestamps = [old_ts, old_ts, 0, 0, 0, 0, 0, 0]; + MemwOperation::new(true, reg_addr, value, timestamp, 2, is_read) + .with_old(old_value, old_timestamps) +} + /// Collects all derived operations from CPU operations in a single pass. /// /// This includes: -/// - MEMW ops (register reads/writes M1/M3/M5, memory loads/stores M6/M7) +/// - MEMW ops (register reads/writes M1/M3/M5, memory loads/stores M6/M7), +/// already routed into their MEMW_R / MEMW_A / MEMW buckets (see [`MemwBuckets`]) /// - LOAD ops (memory loads with sign/zero extension) /// - LT ops (from SLT/BLT instructions) /// - Bitwise lookups (from CPU operations) /// /// MEMW and LOAD collection requires sequential processing with state tracking. /// -/// Returns: (memw_ops, load_ops, lt_ops, shift_ops, bitwise_ops, commit_ops, keccak_ops, -/// cpu32_ops, ecsm_ops, ec_scalar_ops, ecdas_ops) +/// Returns: (memw_buckets, load_ops, lt_ops, shift_ops, bitwise_ops, commit_ops, +/// keccak_ops, cpu32_ops, ecsm_ops, ec_scalar_ops, ecdas_ops) #[allow(clippy::type_complexity)] fn collect_ops_from_cpu( cpu_ops: &[CpuOperation], memory_state: &mut MemoryState, register_state: &mut RegisterState, ) -> ( - Vec, + MemwBuckets, Vec, Vec, Vec, @@ -382,7 +552,7 @@ fn collect_ops_from_cpu( Vec, Vec, ) { - let mut memw_ops = Vec::with_capacity(cpu_ops.len() * 3); + let mut memw = MemwBuckets::with_register_capacity(cpu_ops.len() * 3); let mut load_ops = Vec::with_capacity(cpu_ops.len() / 8 + 1); let mut lt_ops = Vec::with_capacity(cpu_ops.len() / 10 + 1); let mut shift_ops = Vec::with_capacity(cpu_ops.len() / 10 + 1); @@ -414,16 +584,16 @@ fn collect_ops_from_cpu( // Collect memory operations for Load/Store instructions if op.decode.fields.is_load() { let (memw_op, load_op, lookups) = collect_load_op_from_cpu(op, memory_state); - memw_ops.push(memw_op); + memw.push(memw_op); load_ops.push(load_op); bitwise_ops.extend(lookups); } else if op.decode.fields.is_store() { let memw_op = collect_store_op_from_cpu(op, memory_state); - memw_ops.push(memw_op); + memw.push(memw_op); } // Collect register operations (M1, M3, M5) - collect_register_ops_from_cpu(op, register_state, &mut memw_ops); + collect_register_ops_from_cpu(op, register_state, &mut memw); // Collect COMMIT ECALL memory operations (register reads/writes + byte reads) if op.ecall_commit { @@ -433,7 +603,7 @@ fn collect_ops_from_cpu( current_commit_index as u64, )); let reg_commit_ops = collect_commit_memw_ops(op, register_state, memory_state); - memw_ops.extend(reg_commit_ops); + memw.extend_ops(reg_commit_ops); let count = u32::try_from(op.commit_count).expect("commit_count exceeds u32 range"); current_commit_index = current_commit_index .checked_add(count) @@ -469,7 +639,7 @@ fn collect_ops_from_cpu( // collect_keccak_memw_ops handles memory_state + register_state updates let keccak_memw_ops = collect_keccak_memw_ops(op, &input, &output, memory_state, register_state); - memw_ops.extend(keccak_memw_ops); + memw.extend_ops(keccak_memw_ops); keccak_ops.push(KeccakOperation { timestamp: op.timestamp, state_addr, @@ -482,7 +652,7 @@ fn collect_ops_from_cpu( if op.ecall_ecsm { let (ecsm_memw, ecsm_op, ec_scalar_rows, ecdas_rows) = collect_ecsm_ops(op, memory_state, register_state); - memw_ops.extend(ecsm_memw); + memw.extend_ops(ecsm_memw); ecsm_ops.push(ecsm_op); ec_scalar_ops.extend(ec_scalar_rows); ecdas_ops.extend(ecdas_rows); @@ -518,7 +688,9 @@ fn collect_ops_from_cpu( } } - // Collect CPU range-check bitwise lookups (ARE_BYTES + IS_HALF). + // Collect CPU range-check bitwise lookups (ARE_BYTES + IS_HALF). Kept serial here: + // it's only ~110 ms (a serial `.extend` into one growing Vec), and moving it to a + // rayon `flat_map`-collect over 6.8 M per-op Vecs regressed p4 ~4× (alloc + merge). bitwise_ops.extend(op.collect_bitwise_ops()); } @@ -531,7 +703,7 @@ fn collect_ops_from_cpu( ); ( - memw_ops, + memw, load_ops, lt_ops, shift_ops, @@ -562,9 +734,9 @@ fn collect_load_op_from_cpu( let (_old_values, old_timestamps) = memory_state.read_bytes(base_address, 8); // Extract individual bytes from loaded value - let mut value_bytes = [0u64; 8]; + let mut value_bytes = [0u32; 8]; for (j, byte) in value_bytes.iter_mut().take(byte_count).enumerate() { - *byte = (loaded_value >> (j * 8)) & 0xFF; + *byte = ((loaded_value >> (j * 8)) & 0xFF) as u32; } // Sign/zero extend the upper bytes @@ -595,7 +767,7 @@ fn collect_load_op_from_cpu( op.timestamp, byte_count as u8, signed, - res_bytes, + res_bytes.map(u64::from), ); // Collect MSB8 lookups for sign bit extraction @@ -621,13 +793,15 @@ fn collect_store_op_from_cpu(op: &CpuOperation, memory_state: &mut MemoryState) let (old_values, old_timestamps) = memory_state.read_bytes(base_address, 8); // Pack ALL 8 bytes of store_value into value_bytes. - // Bus 14: MEMW Memory Write receiver reconstructs lo32/hi32 via linear combination - // of all 8 bytes. Must match CPU M7 which sends full rv2 as [lo32, hi32]. + // Bus 14: the MEMW Memory Write receiver reconstructs lo32/hi32 via a linear + // combination of all 8 bytes, so it must match the store value the CPU sends + // as [lo32, hi32] on the MEMORY bus (MEMOP) and that this STORE chip forwards + // as the MEMW write (the CPU no longer emits an inline store MEMW — see below). // Bus 16: only positions 0..byte_count participate (controlled by w2/w4/write8 // multiplicities), so extra bytes don't affect memory consistency. - let mut value_bytes = [0u64; 8]; + let mut value_bytes = [0u32; 8]; for (j, byte) in value_bytes.iter_mut().enumerate() { - *byte = (store_value >> (j * 8)) & 0xFF; + *byte = ((store_value >> (j * 8)) & 0xFF) as u32; } // The STORE chip now owns this MEMW write (the CPU sends MEMORY instead of @@ -700,10 +874,10 @@ fn collect_ecsm_ops( for (base, bytes) in [(addr_xg, &witness.x_g), (addr_k, &witness.k)] { for i in 0..4 { let addr = base.wrapping_add((8 * i) as u64); - let mut value = [0u64; 8]; + let mut value = [0u32; 8]; let mut dword = 0u64; for j in 0..8 { - value[j] = bytes[8 * i + j] as u64; + value[j] = bytes[8 * i + j] as u32; dword |= (bytes[8 * i + j] as u64) << (8 * j); } let (_old, old_ts) = memory_state.read_bytes(addr, 8); @@ -728,7 +902,7 @@ fn collect_ecsm_ops( for offset in 0..32u64 { let addr = addr_k.wrapping_add(offset); let byte = k[offset as usize]; - let value = [byte as u64, 0, 0, 0, 0, 0, 0, 0]; + let value = [byte as u32, 0, 0, 0, 0, 0, 0, 0]; let (_v, old_ts) = memory_state.read_byte(addr); memw_ops.push( MemwOperation::new(false, addr, value, t + 1, 1, true) @@ -740,10 +914,10 @@ fn collect_ecsm_ops( // xR writes at T + 2 (4 doublewords). for i in 0..4 { let addr = addr_xr.wrapping_add((8 * i) as u64); - let mut value = [0u64; 8]; + let mut value = [0u32; 8]; let mut dword = 0u64; for j in 0..8 { - value[j] = witness.x_r[8 * i + j] as u64; + value[j] = witness.x_r[8 * i + j] as u32; dword |= (witness.x_r[8 * i + j] as u64) << (8 * j); } let (old_vals, old_ts) = memory_state.read_bytes(addr, 8); @@ -773,10 +947,10 @@ fn collect_ecsm_ops( /// Collects register read/write operations (M1, M3, M5) from CpuOperation, /// pushing them into `memw_ops`. -fn collect_register_ops_from_cpu( +fn collect_register_ops_from_cpu( op: &CpuOperation, register_state: &mut RegisterState, - memw_ops: &mut Vec, + memw_ops: &mut S, ) { let d = &op.decode.fields; // These register accesses happen for every real instruction. For non-word @@ -795,12 +969,18 @@ fn collect_register_ops_from_cpu( } else { register_state.read(d.rs1) }; - // old_timestamps array is 8 elements but only first 2 are used for registers - let old_timestamps = [old_ts, old_ts, 0, 0, 0, 0, 0, 0]; - - let memw_op = MemwOperation::new(true, reg_addr, reg_value, op.timestamp, 2, true) - .with_old(reg_value, old_timestamps); - memw_ops.push(memw_op); + let ts = op.timestamp; + // Direct fast path: fill a RegRow when routing to MEMW_R; push_reg_access + // rebuilds the identical MemwOperation only on the (rare) general/aligned + // fallback. Reads leave the value unchanged, so old == new here. + memw_ops.push_reg_access( + reg_addr, + [reg_value[0], reg_value[1]], + [reg_value[0], reg_value[1]], + ts, + old_ts, + true, + ); if d.rs1 == 255 { register_state.write_pc(op.rv1, op.timestamp); } else { @@ -813,12 +993,15 @@ fn collect_register_ops_from_cpu( let reg_value = pack_register_value(op.rv2); let reg_addr = 2 * d.rs2 as u64; let (_old_val, old_ts) = register_state.read(d.rs2); - // old_timestamps array is 8 elements but only first 2 are used for registers - let old_timestamps = [old_ts, old_ts, 0, 0, 0, 0, 0, 0]; - - let memw_op = MemwOperation::new(true, reg_addr, reg_value, op.timestamp + 1, 2, true) - .with_old(reg_value, old_timestamps); - memw_ops.push(memw_op); + let ts = op.timestamp + 1; + memw_ops.push_reg_access( + reg_addr, + [reg_value[0], reg_value[1]], + [reg_value[0], reg_value[1]], + ts, + old_ts, + true, + ); register_state.write(d.rs2, op.rv2, op.timestamp + 1); } @@ -828,12 +1011,15 @@ fn collect_register_ops_from_cpu( let reg_addr = 2 * d.rd as u64; let (old_val, old_ts) = register_state.read(d.rd); let old_value = pack_register_value(old_val); - // old_timestamps array is 8 elements but only first 2 are used for registers - let old_timestamps = [old_ts, old_ts, 0, 0, 0, 0, 0, 0]; - - let memw_op = MemwOperation::new(true, reg_addr, reg_value, op.timestamp + 2, 2, false) - .with_old(old_value, old_timestamps); - memw_ops.push(memw_op); + let ts = op.timestamp + 2; + memw_ops.push_reg_access( + reg_addr, + [reg_value[0], reg_value[1]], + [old_value[0], old_value[1]], + ts, + old_ts, + false, + ); register_state.write(d.rd, op.rvd, op.timestamp + 2); } @@ -1068,8 +1254,8 @@ fn collect_commit_memw_ops( let new_index = old_index .checked_add(u32::try_from(count).expect("commit_count exceeds u32 range")) .expect("commit index exceeds u32 range"); - let old_value = [old_index as u64, 0, 0, 0, 0, 0, 0, 0]; - let new_value = [new_index as u64, 0, 0, 0, 0, 0, 0, 0]; + let old_value = [old_index, 0, 0, 0, 0, 0, 0, 0]; + let new_value = [new_index, 0, 0, 0, 0, 0, 0, 0]; let old_timestamps = [old_ts, 0, 0, 0, 0, 0, 0, 0]; let memw_op = MemwOperation::new( true, @@ -1088,7 +1274,7 @@ fn collect_commit_memw_ops( for i in 0..count { let addr = buf_addr.wrapping_add(i); let (byte_val, old_ts) = memory_state.read_byte(addr); - let value = [byte_val as u64, 0, 0, 0, 0, 0, 0, 0]; + let value = [byte_val as u32, 0, 0, 0, 0, 0, 0, 0]; let old_timestamps = [old_ts, 0, 0, 0, 0, 0, 0, 0]; let memw_op = MemwOperation::new(false, addr, value, ts, 1, true).with_old(value, old_timestamps); @@ -1196,10 +1382,10 @@ fn collect_keccak_memw_ops( .checked_add(lane_idx as u64 * 8) .expect("keccak state address range must be validated by the executor"); - let mut old_bytes = [0u64; 8]; + let mut old_bytes = [0u32; 8]; let mut old_timestamps = [0u64; 8]; for b in 0..8 { - old_bytes[b] = (in_lane >> (b * 8)) & 0xFF; + old_bytes[b] = ((in_lane >> (b * 8)) & 0xFF) as u32; let byte_addr = lane_addr .checked_add(b as u64) .expect("keccak state address range must be validated by the executor"); @@ -1207,9 +1393,9 @@ fn collect_keccak_memw_ops( old_timestamps[b] = old_ts; } - let mut value_bytes = [0u64; 8]; + let mut value_bytes = [0u32; 8]; for (b, byte) in value_bytes.iter_mut().enumerate() { - *byte = (out_lane >> (b * 8)) & 0xFF; + *byte = ((out_lane >> (b * 8)) & 0xFF) as u32; } let memw_op = MemwOperation::new(false, lane_addr, value_bytes, ts, 8, true) @@ -1377,36 +1563,25 @@ pub(crate) fn is_register_op(op: &MemwOperation) -> bool { if op.old_timestamp[0] != op.old_timestamp[1] { return false; } - let ts = op.timestamp; - let old_ts = op.old_timestamp[0]; - let ts_lo = ts & 0xFFFF_FFFF; - let old_ts_lo = old_ts & 0xFFFF_FFFF; - let ts_hi = ts >> 32; - let old_ts_hi = old_ts >> 32; - ts_hi == old_ts_hi && ts_lo > old_ts_lo && (ts_lo - old_ts_lo) <= 0x10000 + reg_ts_delta_in_range(op.timestamp, op.old_timestamp[0]) } -/// Collects IS_HALFWORD bitwise lookups for MEMW_R operations. +/// The timestamp-delta admission test for MEMW_R (conditions 3-5 of `is_register_op`), +/// factored out so the direct fast path (`push_reg_access`) and the `MemwOperation` +/// classifier (`is_register_op`) share EXACTLY the same routing logic: +/// - `ts_hi == old_ts_hi` (upper limbs match) +/// - `ts_lo > old_ts_lo` (lower-limb ordering) +/// - `ts_lo - old_ts_lo <= 2^16` (delta fits the IS_HALF range [1, 2^16]) /// -/// For each register op: checks that `timestamp[0] - old_timestamp_lo - 1` fits -/// in a halfword (proving the timestamp delta is in range [1, 2^16]). -fn collect_bitwise_from_memw_register(ops: &[MemwOperation]) -> Vec { - ops.iter() - .map(|op| { - let ts_lo = op.timestamp & 0xFFFF_FFFF; - let old_ts_lo = op.old_timestamp[0] & 0xFFFF_FFFF; - debug_assert!( - ts_lo > old_ts_lo, - "ts_lo must exceed old_ts_lo (enforced by is_register_op)" - ); - let diff_minus_1 = (ts_lo - old_ts_lo - 1) as u16; - BitwiseOperation::halfword( - BitwiseOperationType::IsHalf, - (diff_minus_1 & 0xFF) as u8, - (diff_minus_1 >> 8) as u8, - ) - }) - .collect() +/// The fast path only calls this for width-2 register accesses whose two words share +/// `old_ts` by construction, so conditions 1-2 of `is_register_op` always hold there. +#[inline] +fn reg_ts_delta_in_range(timestamp: u64, old_ts: u64) -> bool { + let ts_lo = timestamp & 0xFFFF_FFFF; + let old_ts_lo = old_ts & 0xFFFF_FFFF; + let ts_hi = timestamp >> 32; + let old_ts_hi = old_ts >> 32; + ts_hi == old_ts_hi && ts_lo > old_ts_lo && (ts_lo - old_ts_lo) <= 0x10000 } // ============================================================================= @@ -1802,34 +1977,21 @@ fn collect_bitwise_from_branch(branch_ops: &[BranchOperation]) -> Vec Vec { - if num_padding_rows == 0 { - return Vec::new(); - } - - let mut ops = Vec::with_capacity(num_padding_rows * 14); - for _ in 0..num_padding_rows { - // The shrunk CPU sends, per row (incl. padding where all values are 0): - // 3× ARE_BYTES (rs1/rs2, rd/instruction_length, alu_flags/mem_flags) and - // 4× IS_HALF (the four `res` halves). - for _ in 0..3 { - ops.push(BitwiseOperation::byte_op( - BitwiseOperationType::AreBytes, - 0, - 0, - )); - } - for _ in 0..4 { - ops.push(BitwiseOperation::halfword( - BitwiseOperationType::IsHalf, - 0, - 0, - )); - } - } - ops +/// Add the BITWISE lookups every CPU padding row sends. A padding row has all +/// values zero, so per row it sends 3× ARE_BYTES[0,0] (rs1/rs2, +/// rd/instruction_length, alu_flags/mem_flags) and 4× IS_HALF[0] (the four `res` +/// halves). Every padding row sends the same lookups, so their whole +/// contribution is a per-cell count: two `bump_n`s, no per-row work. +fn add_padding_byte_checks(hist: &mut bitwise::BitwiseHistogram, num_padding_rows: usize) { + let n = num_padding_rows as u64; + hist.bump_n( + BitwiseOperation::byte_op(BitwiseOperationType::AreBytes, 0, 0), + 3 * n, + ); + hist.bump_n( + BitwiseOperation::halfword(BitwiseOperationType::IsHalf, 0, 0), + 4 * n, + ); } /// Collects ARE_BYTES lookups from PAGE data (init and fini values). @@ -1949,11 +2111,11 @@ fn collect_bitwise_from_page( image: &I, memory_state: &MemoryState, exclude_touched: bool, -) -> Vec { + hist: &mut bitwise::BitwiseHistogram, +) { use std::collections::BTreeSet; let page_size = page::DEFAULT_PAGE_SIZE; - let mut bitwise_ops = Vec::new(); let init_page_data = build_init_page_data(image); @@ -1985,16 +2147,16 @@ fn collect_bitwise_from_page( // Get fini value (from final_state or init if never accessed) let fini = final_state.get(&addr).map_or(init, |state| state.value); - // C1+C2: ARE_BYTES[init, fini] — batched range check for both bytes - bitwise_ops.push(BitwiseOperation::byte_op( + // C1+C2: ARE_BYTES[init, fini] — batched range check for both bytes. + // Bumped straight into the histogram: this loop visits every byte of + // every touched page, and the histogram is the only consumer. + hist.bump(BitwiseOperation::byte_op( BitwiseOperationType::AreBytes, init, fini, )); } } - - bitwise_ops } // ============================================================================= @@ -2577,7 +2739,8 @@ struct CollectedOps { cpu_ops: Vec, memw_ops: Vec, memw_aligned_ops: Vec, - memw_register_ops: Vec, + /// Direct-fill MEMW_R rows (register fast path). + memw_register_rows: Vec, load_ops: Vec, lt_ops: Vec, shift_ops: Vec, @@ -2640,7 +2803,7 @@ fn chunk_and_generate( #[allow(clippy::too_many_arguments)] fn collect_all_ops( cpu_ops: Vec, - mut memw_ops: Vec, + mut memw: MemwBuckets, load_ops: Vec, mut lt_ops: Vec, mut shift_ops: Vec, @@ -2659,16 +2822,20 @@ fn collect_all_ops( // Only the final epoch terminates; intermediate epochs keep their boundary // register state (no zeroizing) so it can seed the next epoch. if is_final { - let halt_memw_ops = collect_halt_ops(register_state); - memw_ops.extend(halt_memw_ops); + // Route halt ops through the same classifier; they append to the end of their + // buckets. + memw.extend_ops(collect_halt_ops(register_state)); } - // Route MEMW_R (register fast-path) first, then MEMW_A (aligned), rest → MEMW. - // Order matters: register ops would also pass is_aligned_op, so check first. - let (memw_register_ops, memw_ops): (Vec<_>, Vec<_>) = - memw_ops.into_iter().partition(is_register_op); - let (memw_aligned_ops, memw_ops): (Vec<_>, Vec<_>) = - memw_ops.into_iter().partition(is_aligned_op); + // The walk (`collect_ops_from_cpu`) already routed every MemwOperation into its bucket at + // creation via `MemwBuckets`, so there is no separate routing pass here: the ops are not + // moved a second time. Order within each bucket is the walk's insertion order, which the + // multiplicity counts depend on being deterministic. + let MemwBuckets { + register_rows: memw_register_rows, + aligned: memw_aligned_ops, + general: memw_ops, + } = memw; // Collect BRANCH operations from CPU ops where branch_cond = true let branch_ops: Vec = cpu_ops @@ -2773,7 +2940,7 @@ fn collect_all_ops( cpu_ops, memw_ops, memw_aligned_ops, - memw_register_ops, + memw_register_rows, load_ops, lt_ops, shift_ops, @@ -2817,11 +2984,11 @@ fn build_traces( cpu_ops, memw_ops, memw_aligned_ops, - memw_register_ops, + memw_register_rows, load_ops, mut lt_ops, shift_ops, - mut bitwise_ops, + bitwise_ops, branch_ops, mul_ops, dvrm_ops, @@ -2847,51 +3014,12 @@ fn build_traces( // ===================================================================== #[cfg(feature = "instruments")] let __sp = stark::instruments::span("p4_bitwise_collect"); - bitwise_ops.extend(collect_bitwise_from_lt(<_ops)); - // MUL/DVRM dedup their per-unique bit-gated lookups PER CHIP INSTANCE, so pass - // the same chunk size used to split them into instances (see chunk_and_generate - // below) so the BITWISE multiplicity matches the per-instance sends. - bitwise_ops.extend(collect_bitwise_from_mul(&mul_ops, max_rows.mul)); - bitwise_ops.extend(collect_bitwise_from_dvrm(&dvrm_ops, max_rows.dvrm)); - bitwise_ops.extend(collect_bitwise_from_branch(&branch_ops)); - bitwise_ops.extend(shift::collect_bitwise_from_shift(&shift_ops)); - // Auxiliary chips: BYTEWISE sends 8× BYTE_ALU/op; EQ sends 4× IS_HALF + ZERO. - for op in &bytewise_ops { - bitwise_ops.extend(op.collect_bitwise_ops()); - } - for op in &eq_ops { - bitwise_ops.extend(op.collect_bitwise_ops()); - } - for op in &store_ops { - bitwise_ops.extend(op.collect_bitwise_ops()); - } - bitwise_ops.extend(collect_bitwise_from_memw_aligned(&memw_aligned_ops)); - // MEMW_R sends IS_HALFWORD[timestamp_0 - old_timestamp_lo - 1] - bitwise_ops.extend(collect_bitwise_from_memw_register(&memw_register_ops)); - // PAGE tables do a batched ARE_BYTES[init, fini] lookup per row (C1+C2). - // Continuation epochs (l2g_memory_bookend) skip PAGE entirely (see the - // generate_page_tables call below), so they skip its AreBytes lookups too. - if let Some(image) = initial_image - && !l2g_memory_bookend - { - bitwise_ops.extend(collect_bitwise_from_page( - image, - memory_state, - l2g_memory_bookend, - )); - } let public_output_bytes: Vec = commit_ops .iter() .filter(|op| !op.end) .map(|op| op.value) .collect(); - // COMMIT table sends AreBytes and IsHalfword lookups - bitwise_ops.extend(collect_bitwise_from_commit(&commit_ops)); - // KECCAK_RND sends XOR/AND/ARE_BYTES/HWSL; KECCAK core sends IS_HALF - bitwise_ops.extend(collect_bitwise_from_keccak(&keccak_ops)); - bitwise_ops.extend(collect_bitwise_from_ecsm(&ecsm_ops)); - bitwise_ops.extend(collect_bitwise_from_ecdas(&ecdas_ops)); // CPU padding rows send ARE_BYTES with all-zero values. // Add corresponding ops so the bitwise table multiplicities balance. @@ -2899,7 +3027,124 @@ fn build_traces( .chunks(max_rows.cpu) .map(|chunk| chunk.len().next_power_of_two().max(4) - chunk.len()) .sum(); - bitwise_ops.extend(collect_byte_check_ops_for_padding(num_padding_rows)); + + // The per-source bitwise collectors are all pure functions of their inputs, and the + // BITWISE multiplicities are order-independent (they ride a permutation-invariant bus), + // so every source can be collected in parallel and the per-worker histograms summed in + // any order. + // + // MUL/DVRM dedup their per-unique bit-gated lookups PER CHIP INSTANCE, so pass the same + // chunk size used to split them into instances so multiplicities match the per-instance + // sends. MEMW_R sends IS_HALFWORD[timestamp_0 - old_timestamp_lo - 1]. PAGE does a + // batched ARE_BYTES[init, fini] per row (skipped in continuation epochs, which the L2G + // table owns). COMMIT sends AreBytes+IsHalfword; KECCAK_RND sends XOR/AND/ARE_BYTES/HWSL. + // We never concatenate the lookups into one giant `Vec` (~140 M ops / + // ~560 MB at 10-tx whose only consumer is the multiplicity count). Each collector bumps + // the `BitwiseHistogram` it is handed: the heavy sources (MEMW_R one-per-row, PAGE + // one-per-byte, padding) count directly with no per-source Vec at all, and the small + // sources fold their transient `collect_*` Vec in and drop it. The histogram is a + // commutative monoid, so per-worker histograms tree-reduce to multiplicities that are + // independent of accumulation order. + type Collector<'a> = Box; + let mul_chunk = max_rows.mul; + let dvrm_chunk = max_rows.dvrm; + // Every source except the two dominant ones (the in-walk lookups and MEMW_R, which are + // split into row-ranges in the parallel path below) stays a single whole-source collector. + let mut collectors: Vec = vec![ + Box::new(|h| h.add_ops(&collect_bitwise_from_lt(<_ops))), + Box::new(|h| h.add_ops(&collect_bitwise_from_mul(&mul_ops, mul_chunk))), + Box::new(|h| h.add_ops(&collect_bitwise_from_dvrm(&dvrm_ops, dvrm_chunk))), + Box::new(|h| h.add_ops(&collect_bitwise_from_branch(&branch_ops))), + Box::new(|h| h.add_ops(&shift::collect_bitwise_from_shift(&shift_ops))), + Box::new(|h| { + for op in &bytewise_ops { + h.add_ops(&op.collect_bitwise_ops()); + } + }), + Box::new(|h| { + for op in &eq_ops { + h.add_ops(&op.collect_bitwise_ops()); + } + }), + Box::new(|h| { + for op in &store_ops { + h.add_ops(&op.collect_bitwise_ops()); + } + }), + Box::new(|h| h.add_ops(&collect_bitwise_from_memw_aligned(&memw_aligned_ops))), + Box::new(|h| h.add_ops(&collect_bitwise_from_commit(&commit_ops))), + Box::new(|h| h.add_ops(&collect_bitwise_from_keccak(&keccak_ops))), + Box::new(|h| h.add_ops(&collect_bitwise_from_ecsm(&ecsm_ops))), + Box::new(|h| h.add_ops(&collect_bitwise_from_ecdas(&ecdas_ops))), + Box::new(|h| add_padding_byte_checks(h, num_padding_rows)), + ]; + if let Some(image) = initial_image + && !l2g_memory_bookend + { + collectors.push(Box::new(move |h| { + collect_bitwise_from_page(image, memory_state, l2g_memory_bookend, h) + })); + } + + let mut base = bitwise::BitwiseHistogram::new(); + + #[cfg(feature = "parallel")] + { + use rayon::prelude::*; + // Cap concurrent 80 MiB histograms at `cap` to bound peak memory. The two dominant + // sources — the in-walk lookups and MEMW_R (each tens of millions of items) — are + // split into ~`cap` row-range slices so they parallelize INTERNALLY instead of each + // pinning one core while the rest idle. Every unit (whole collectors + the heavy + // slices) is round-robined into exactly `cap` buckets, one histogram each, so the + // split heavy work is spread across buckets rather than piled into one. + // add_ops/bump/merge form a commutative monoid, so any partition yields + // byte-identical multiplicities (same as the serial fallback below). + let cap = rayon::current_num_threads().clamp(1, 8); + let mut units: Vec = Vec::with_capacity(collectors.len() + 2 * cap); + let iw_chunk = bitwise_ops.len().div_ceil(cap).max(1); + for slice in bitwise_ops.chunks(iw_chunk) { + units.push(Box::new(move |h| h.add_ops(slice))); + } + let reg_chunk = memw_register_rows.len().div_ceil(cap).max(1); + for slice in memw_register_rows.chunks(reg_chunk) { + units.push(Box::new(move |h| { + memw_register::collect_bitwise_from_memw_register(slice, h) + })); + } + units.extend(collectors); + + let mut buckets: Vec> = (0..cap).map(|_| Vec::new()).collect(); + for (i, unit) in units.into_iter().enumerate() { + buckets[i % cap].push(unit); + } + if let Some(reduced) = buckets + .par_iter() + .map(|bucket| { + let mut h = bitwise::BitwiseHistogram::new(); + for f in bucket { + f(&mut h); + } + h + }) + .reduce_with(|mut a, b| { + a.merge(&b); + a + }) + { + base.merge(&reduced); + } + } + #[cfg(not(feature = "parallel"))] + { + base.add_ops(&bitwise_ops); + memw_register::collect_bitwise_from_memw_register(&memw_register_rows, &mut base); + for f in &collectors { + f(&mut base); + } + } + let bitwise_histogram = base; + // The in-walk lookup Vec has been counted into the histogram; free it now. + drop(bitwise_ops); #[cfg(feature = "instruments")] drop(__sp); @@ -2966,10 +3211,12 @@ fn build_traces( ) }; let gen_memw_registers = || { + // Direct-to-column fill from compact RegRows — the register fast path never + // materializes a `Vec`. chunk_and_generate( - &memw_register_ops, + &memw_register_rows, max_rows.memw_register, - memw_register::generate_memw_register_trace, + memw_register::generate_memw_register_trace_from_rows, #[cfg(feature = "disk-spill")] storage_mode, ) @@ -3069,7 +3316,8 @@ fn build_traces( }; let gen_bitwise = || { let mut bitwise = bitwise::generate_bitwise_trace(); - bitwise::update_multiplicities(&mut bitwise, &bitwise_ops); + // Fill the MU columns (11..=20) from the accumulated histogram. + bitwise_histogram.fill_multiplicities(&mut bitwise); bitwise }; // Each CPU operation looks up the DECODE table once; padding rows look up @@ -3388,19 +3636,21 @@ pub fn count_table_lengths( by_width: &mut [usize; 4], aligned: &mut usize, register: &mut usize| { - if is_register_op(op) { - *register += 1; - } else if is_aligned_op(op) { - *aligned += 1; - } else { - let idx = match op.width { - 1 => 0, - 2 => 1, - 4 => 2, - 8 => 3, - _ => return, - }; - by_width[idx] += 1; + // Same classifier as the walk's MemwBuckets, so the sizing pass counts + // exactly the rows trace generation will produce. + match classify_memw(op) { + MemwRoute::Register => *register += 1, + MemwRoute::Aligned => *aligned += 1, + MemwRoute::General => { + let idx = match op.width { + 1 => 0, + 2 => 1, + 4 => 2, + 8 => 3, + _ => return, + }; + by_width[idx] += 1; + } } }; diff --git a/prover/src/tables/types.rs b/prover/src/tables/types.rs index 71c83284a..85eaca17a 100644 --- a/prover/src/tables/types.rs +++ b/prover/src/tables/types.rs @@ -43,6 +43,52 @@ pub fn dword_wl(x: u64) -> [FE; 2] { [FE::from(x & 0xFFFF_FFFF), FE::from(x >> 32)] } +/// Allocate a zeroed `Vec` via the allocator's calloc path (demand-zeroed OS +/// pages, touched lazily on first write) instead of the element-wise clone loop +/// that `vec![FE::zero(); n]` runs. Trace fill is memory-bandwidth-bound, so +/// skipping the eager zeroing sweep of the (multi-GB) trace is a real win. +/// +/// Sound because `FE` is `#[repr(transparent)]` over its `u64` base value +/// (Goldilocks has no Montgomery form), so `FE`'s canonical zero is the all-zero +/// bit pattern — identical to `0u64`. The `zeroed_fe_vec_matches_fe_zero` test +/// guards this invariant. Technique borrowed from SP1's `zeroed_f_vec`. +#[inline] +pub fn zeroed_fe_vec(len: usize) -> Vec { + const _: () = assert!(core::mem::size_of::() == core::mem::size_of::()); + const _: () = assert!(core::mem::align_of::() == core::mem::align_of::()); + let zeros: Vec = vec![0u64; len]; + // Reinterpret the buffer as `Vec` via its raw parts rather than + // `mem::transmute::, Vec>`. `Vec`'s field layout is unspecified + // and may depend on its element type, so transmuting one `Vec` to another + // relies on that unspecified layout (the std `mem::transmute` docs call this + // out and recommend `from_raw_parts`). Rebuilding from `(ptr, len, cap)` + // reuses the same allocation and carries no `Vec`-layout assumption. + let mut zeros = core::mem::ManuallyDrop::new(zeros); + // SAFETY: `FE` is `#[repr(transparent)]` over `u64` with identical size and + // alignment (asserted above), and `0u64` is exactly `FE::zero()`'s bit + // pattern (Goldilocks has no Montgomery form), so the zeroed `u64` buffer is + // a valid `[FE]` of all `FE::zero()`. `len`/`capacity` are element counts and + // the element sizes are equal, so they carry over unchanged; the eventual + // dealloc uses the same `size * capacity` and alignment as the original + // allocation. `ManuallyDrop` stops the source `Vec` from freeing the buffer + // that the returned `Vec` now owns. + unsafe { Vec::from_raw_parts(zeros.as_mut_ptr() as *mut FE, zeros.len(), zeros.capacity()) } +} + +#[cfg(test)] +mod zeroed_fe_vec_tests { + use super::*; + + /// Guards the `zeroed_fe_vec` invariant: a calloc'd all-zero buffer + /// reinterpreted as `Vec` must equal an element-wise `FE::zero()` fill. + #[test] + fn zeroed_fe_vec_matches_fe_zero() { + for len in [0usize, 1, 7, 64, 1024] { + assert_eq!(zeroed_fe_vec(len), vec![FE::zero(); len], "len={len}"); + } + } +} + /// Decompose a `u64` into its four little-endian 16-bit limbs as field elements: /// `[x[0..16], x[16..32], x[32..48], x[48..64]]` (the `DWordHL` column encoding). #[inline]