Skip to content

Commit 60a289a

Browse files
committed
merge #1010's head 88b0d31 into noepoch/whir
The one-binary whole-block A/B compares #1010's epoch tree and the block tree on the same code: Gruen GKR rounds, no lift, the global-after-last latch. # Conflicts: # crypto/multilinear/src/constraint_argument.rs
2 parents 697c2a5 + 88b0d31 commit 60a289a

21 files changed

Lines changed: 4382 additions & 98 deletions

‎crypto/math-cuda/kernels/sumcheck.cu‎

Lines changed: 237 additions & 17 deletions
Large diffs are not rendered by default.

‎crypto/math-cuda/src/argue_fused.rs‎

Lines changed: 12 additions & 9 deletions
Original file line numberDiff line numberDiff line change
@@ -17,7 +17,7 @@ use std::sync::Arc;
1717

1818
use crate::Result;
1919
use crate::device::{DeviceReservation, alloc_or_trim, alloc_zeros_or_trim, backend, htod_or_trim};
20-
use crate::sumcheck::{DeviceFactors, eq_table_ext3, launch_shape};
20+
use crate::sumcheck::{FactorView, eq_table_ext3, launch_shape};
2121

2222
const BLOCK_DIM: u32 = 256;
2323
const MAX_GRID: u32 = 4096;
@@ -50,7 +50,7 @@ pub struct FusedBus<'a> {
5050
/// One table's fused zerocheck on the card.
5151
pub struct FusedZerocheck<'f> {
5252
stream: Arc<CudaStream>,
53-
source: &'f DeviceFactors,
53+
source: FactorView<'f>,
5454
width: usize,
5555
/// The cube the lifted factors span.
5656
rows: usize,
@@ -151,16 +151,16 @@ impl<'f> FusedZerocheck<'f> {
151151
/// `r_tail` and `rho_tail` are `r[2..]` and `ρ[2..]` (three u64 a
152152
/// coordinate): the grid pass's weights.
153153
pub fn new(
154-
source: &'f DeviceFactors,
154+
source: FactorView<'f>,
155155
program: FusedProgram<'_>,
156156
bus: FusedBus<'_>,
157157
r_tail: &[u64],
158158
rho_tail: &[u64],
159159
grid_rows: usize,
160160
gruen_rows: usize,
161161
) -> Result<Option<Self>> {
162-
let rows = source.len();
163-
let width = source.width();
162+
let rows = source.rows;
163+
let width = source.width;
164164
assert!(
165165
rows >= 8 && rows.is_power_of_two(),
166166
"the grid needs four rows a group"
@@ -194,7 +194,7 @@ impl<'f> FusedZerocheck<'f> {
194194
)) else {
195195
return Ok(None);
196196
};
197-
let stream = source.stream().clone();
197+
let stream = source.stream.clone();
198198
let q = rows / 4;
199199
let nonempty = |v: &[u64]| -> Vec<u64> { if v.is_empty() { vec![0] } else { v.to_vec() } };
200200
let nodes = htod_or_trim(&stream, &nonempty(program.nodes))?;
@@ -316,7 +316,7 @@ impl<'f> FusedZerocheck<'f> {
316316
q.div_ceil(BLOCK_DIM as u64).clamp(1, MAX_GRID as u64) as u32,
317317
BLOCK_DIM,
318318
);
319-
let factors = self.source.factor_ptrs();
319+
let factors = self.source.ptrs;
320320
unsafe {
321321
self.stream
322322
.launch_builder(&be.zc_bus_u)
@@ -328,6 +328,7 @@ impl<'f> FusedZerocheck<'f> {
328328
.arg(&self.constant)
329329
.arg(&self.eq_rho)
330330
.arg(&mut self.partials)
331+
.arg(&self.source.stride)
331332
.launch(LaunchConfig {
332333
grid_dim: (grid, 1, 1),
333334
block_dim: (block, 1, 1),
@@ -357,7 +358,7 @@ impl<'f> FusedZerocheck<'f> {
357358
}
358359
self.stream.memcpy_htod(&[u64::MAX], &mut self.violation)?;
359360
let (grid, block) = self.shape(q, points.len(), self.num_slots as u64)?;
360-
let factors = self.source.factor_ptrs();
361+
let factors = self.source.ptrs;
361362
unsafe {
362363
self.stream
363364
.launch_builder(&be.zc_grid01)
@@ -373,6 +374,7 @@ impl<'f> FusedZerocheck<'f> {
373374
.arg(&mut self.partials)
374375
.arg(&mut self.violation)
375376
.arg(&u32::from(keep_corners))
377+
.arg(&self.source.stride)
376378
.launch(LaunchConfig {
377379
grid_dim: (grid, points.len() as u32, 1),
378380
block_dim: (block, 1, 1),
@@ -396,7 +398,7 @@ impl<'f> FusedZerocheck<'f> {
396398
let width = self.width as u64;
397399
let total = width * q;
398400
let grid = total.div_ceil(BLOCK_DIM as u64).clamp(1, MAX_GRID as u64) as u32;
399-
let factors = self.source.factor_ptrs();
401+
let factors = self.source.ptrs;
400402
unsafe {
401403
self.stream
402404
.launch_builder(&be.zc_fold2)
@@ -405,6 +407,7 @@ impl<'f> FusedZerocheck<'f> {
405407
.arg(&width)
406408
.arg(&self.scratch)
407409
.arg(&mut self.folded)
410+
.arg(&self.source.stride)
408411
.launch(LaunchConfig {
409412
grid_dim: (grid, 1, 1),
410413
block_dim: (BLOCK_DIM, 1, 1),

‎crypto/math-cuda/src/device.rs‎

Lines changed: 10 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -366,6 +366,12 @@ pub struct Backend {
366366
pub zc_bus_column: CudaFunction,
367367
pub zc_halve: CudaFunction,
368368
pub zc_round_gruen: CudaFunction,
369+
// D-ARGUE S1-3's GKR layer rounds (`crate::gkr::GruenLayer`).
370+
pub gkr_eq_levels_ext3: CudaFunction,
371+
pub gkr_round_gruen: CudaFunction,
372+
pub gkr_gruen_finish: CudaFunction,
373+
// D-BATCH M1-2's input layer from the base columns (`crate::gkr::input_from_columns`).
374+
pub gkr_input_from_columns: CudaFunction,
369375
pub sumcheck_fold_ext3: CudaFunction,
370376
pub mle_fold_base_ext3: CudaFunction,
371377
pub eq_expand_level_ext3: CudaFunction,
@@ -1192,6 +1198,10 @@ impl Backend {
11921198
zc_bus_column: sumcheck.load_function("zc_bus_column")?,
11931199
zc_halve: sumcheck.load_function("zc_halve")?,
11941200
zc_round_gruen: sumcheck.load_function("zc_round_gruen")?,
1201+
gkr_eq_levels_ext3: sumcheck.load_function("gkr_eq_levels_ext3")?,
1202+
gkr_round_gruen: sumcheck.load_function("gkr_round_gruen")?,
1203+
gkr_gruen_finish: sumcheck.load_function("gkr_gruen_finish")?,
1204+
gkr_input_from_columns: sumcheck.load_function("gkr_input_from_columns")?,
11951205
sumcheck_fold_ext3: sumcheck.load_function("sumcheck_fold_ext3")?,
11961206
mle_fold_base_ext3: sumcheck.load_function("mle_fold_base_ext3")?,
11971207
eq_expand_level_ext3: sumcheck.load_function("eq_expand_level_ext3")?,

0 commit comments

Comments
 (0)