Skip to content

Commit 83f4d62

Browse files
erabinovclaude
andauthored
perf(gpu): more efficient column evaluation claims at the end of GKR (#2754)
Co-authored-by: Claude Opus 4.7 (1M context) <noreply@anthropic.com>
1 parent 08d2b2f commit 83f4d62

2 files changed

Lines changed: 48 additions & 71 deletions

File tree

sp1-gpu/crates/logup_gkr/benches/gkr.rs

Lines changed: 11 additions & 55 deletions
Original file line numberDiff line numberDiff line change
@@ -18,19 +18,18 @@ use std::sync::Arc;
1818
use criterion::{black_box, criterion_group, criterion_main, BatchSize, BenchmarkId, Criterion};
1919
use rand::{rngs::StdRng, Rng, SeedableRng};
2020
use slop_challenger::{FieldChallenger, IopCtx};
21-
use slop_multilinear::{MultilinearPcsChallenger, Point};
21+
use slop_multilinear::Point;
2222
use sp1_core_machine::riscv::RiscvAir;
23-
use sp1_gpu_cudart::{DevicePoint, TaskScope};
23+
use sp1_gpu_cudart::TaskScope;
2424
use sp1_gpu_jagged_tracegen::test_utils::bench_utils::{
2525
with_trace_source, FullKind, RealTraceData,
2626
};
2727
use sp1_gpu_jagged_tracegen::test_utils::tracegen_setup::CORE_MAX_LOG_ROW_COUNT;
28-
use sp1_gpu_logup_gkr::{
29-
generate_gkr_circuit, prove_gkr_circuit, CudaLogUpGkrOptions, Interactions,
30-
};
28+
use sp1_gpu_logup_gkr::{generate_gkr_circuit, prove_logup_gkr, CudaLogUpGkrOptions, Interactions};
3129
use sp1_gpu_utils::{Ext, Felt, TestGC};
3230
use sp1_hypercube::air::MachineAir;
3331
use sp1_hypercube::Chip;
32+
use sp1_primitives::SP1GlobalContext;
3433

3534
/// `prove_gkr_circuit` flag: when true, the prover recomputes the first layer on demand from
3635
/// the raw trace each time it walks back to the leaves; when false, it caches the materialized
@@ -125,66 +124,23 @@ fn run_prove<R: Rng>(
125124
) {
126125
let RealTraceData { machine: _, cluster, public_values: _, device_mle } = data;
127126
let interactions = build_interactions(&cluster, scope);
128-
let beta_dim = beta_seed_dim(&cluster);
129-
130-
// initial_number_of_variables = num_interaction_variables + 1, where
131-
// num_interaction_variables is `log2(num_total_interactions).next_power_of_two()`.
132-
let num_interactions: usize =
133-
cluster.iter().map(|chip| chip.sends().len() + chip.receives().len()).sum();
134-
let initial_number_of_variables = num_interactions.next_power_of_two().trailing_zeros() + 1;
135127

136128
let mut group = c.benchmark_group("prove");
137129
group.sample_size(10);
138130
group.bench_with_input(id, &(), |b, _| {
139131
b.iter_batched(
140132
|| {
141-
// Per-iteration setup. Mirrors `prove_logup_gkr` lines 198-244 but skips the
142-
// observe-output-claims step (which only affects challenger state, not the
143-
// prove call's runtime). We still build the circuit and compute the initial
144-
// numerator/denominator evaluations because `prove_gkr_circuit` consumes both.
145-
let mut challenger = TestGC::default_challenger();
146-
let alpha: Ext = challenger.sample_ext_element();
147-
let beta_seed: Point<Ext> =
148-
(0..beta_dim).map(|_| challenger.sample_ext_element::<Ext>()).collect();
149-
150-
let (output, circuit) = generate_gkr_circuit(
133+
let challenger = TestGC::default_challenger();
134+
scope.synchronize_blocking().unwrap();
135+
(interactions.clone(), challenger)
136+
},
137+
|(interactions, mut challenger)| {
138+
let result = prove_logup_gkr::<SP1GlobalContext, _>(
151139
&cluster,
152-
interactions.clone(),
140+
interactions,
153141
&device_mle,
154-
alpha,
155-
beta_seed,
156142
GKR_OPTIONS,
157-
scope.clone(),
158-
);
159-
160-
let first_eval_point = challenger.sample_point::<Ext>(initial_number_of_variables);
161-
let first_point =
162-
DevicePoint::from_host(&first_eval_point, output.numerator.backend())
163-
.unwrap()
164-
.into_inner();
165-
let first_point_eq = DevicePoint::new(first_point).partial_lagrange();
166-
let first_numerator_eval =
167-
output.numerator.eval_at_eq(&first_point_eq).to_host_vec().unwrap()[0];
168-
let first_denominator_eval =
169-
output.denominator.eval_at_eq(&first_point_eq).to_host_vec().unwrap()[0];
170-
171-
scope.synchronize_blocking().unwrap();
172-
(
173-
first_numerator_eval,
174-
first_denominator_eval,
175-
first_eval_point,
176-
circuit,
177-
challenger,
178-
)
179-
},
180-
|(num_eval, den_eval, eval_point, circuit, mut challenger)| {
181-
let result = prove_gkr_circuit(
182-
num_eval,
183-
den_eval,
184-
eval_point,
185-
circuit,
186143
&mut challenger,
187-
RECOMPUTE_FIRST_LAYER,
188144
);
189145
scope.synchronize_blocking().unwrap();
190146
black_box(result)

sp1-gpu/crates/zerocheck/src/primitives.rs

Lines changed: 37 additions & 16 deletions
Original file line numberDiff line numberDiff line change
@@ -2,14 +2,18 @@ use crate::data::{InfoBuffer, JaggedDenseInfo};
22
use slop_algebra::{AbstractField, ExtensionField, Field};
33
use slop_alloc::{Backend, Buffer, CpuBackend, HasBackend, Slice};
44
use slop_commit::Rounds;
5+
56
use slop_multilinear::{MleEval, Point};
6-
use slop_tensor::Tensor;
7+
use slop_tensor::{Tensor, TensorView};
8+
79
use sp1_gpu_cudart::sys::runtime::KernelPtr;
810
use sp1_gpu_cudart::sys::v2_kernels::{
911
fix_last_variable_jagged_ext, fix_last_variable_jagged_felt, fix_last_variable_jagged_info,
1012
initialize_jagged_info,
1113
};
12-
use sp1_gpu_cudart::{args, DeviceBuffer, DevicePoint, DeviceTensor, TaskScope};
14+
use sp1_gpu_cudart::{
15+
args, dot_along_dim_view, DeviceBuffer, DevicePoint, DeviceTensor, TaskScope,
16+
};
1317
use sp1_gpu_utils::{Ext, Felt, JaggedMle, JaggedTraceMle, TraceDenseData, TraceOffset};
1418
use std::collections::BTreeMap;
1519
use std::iter::once;
@@ -123,20 +127,39 @@ where
123127

124128
#[inline(always)]
125129
pub fn evaluate_traces(traces: &JaggedTraceMle<Felt, TaskScope>, point: &Point<Ext>) -> Vec<Ext> {
126-
let mut next_input_jagged_trace_mle =
127-
evaluate_jagged_fix_last_variable(traces, *point.last().unwrap());
128-
for alpha in point.iter().rev().skip(1) {
129-
next_input_jagged_trace_mle =
130-
evaluate_jagged_fix_last_variable(&next_input_jagged_trace_mle, *alpha);
130+
let trace_data = traces.dense();
131+
let backend = traces.backend();
132+
let device_point = DevicePoint::from_host(point, backend).unwrap();
133+
let partial_lagrange = device_point.partial_lagrange();
134+
let total_cols = trace_data
135+
.preprocessed_table_index
136+
.values()
137+
.chain(trace_data.main_table_index.values())
138+
.map(|index| index.num_polys)
139+
.sum::<usize>();
140+
let mut result_buffer =
141+
DeviceBuffer::with_capacity_in(total_cols, backend.clone()).into_inner();
142+
143+
let trace_ptr = trace_data.dense.as_ptr();
144+
let chip_indices =
145+
trace_data.preprocessed_table_index.values().chain(trace_data.main_table_index.values());
146+
for index in chip_indices {
147+
if index.dense_offset.start == index.dense_offset.end {
148+
continue;
149+
}
150+
let chip_ptr = unsafe { trace_ptr.add(index.dense_offset.start) };
151+
let chip_view = unsafe {
152+
TensorView::from_raw_parts(
153+
chip_ptr,
154+
[index.num_polys, index.poly_size].try_into().unwrap(),
155+
backend.clone(),
156+
)
157+
};
158+
let result = dot_along_dim_view(chip_view, partial_lagrange.guts().as_view(), 1);
159+
result_buffer.extend_from_device_slice(result.as_buffer()).unwrap();
131160
}
132161

133-
let host_dense = DeviceBuffer::from_raw(next_input_jagged_trace_mle.dense_data.dense.clone())
134-
.to_host()
135-
.unwrap()
136-
.to_vec();
137-
138-
// Only every four elements is not padding.
139-
host_dense.into_iter().step_by(4).collect::<Vec<_>>()
162+
DeviceBuffer::from_raw(result_buffer).to_host().unwrap()
140163
}
141164

142165
pub fn evaluate_jagged_columns(
@@ -353,8 +376,6 @@ pub fn round_batch_evaluations(
353376
let preprocessed_host_evaluations =
354377
preprocessed_host_evaluations.into_iter().collect::<Vec<_>>();
355378

356-
// Skip the padding column, if it exists.
357-
evals_so_far = jagged_trace_mle.dense().preprocessed_cols;
358379
let mut main_host_evaluations = Vec::new();
359380
for offset in jagged_trace_mle.dense().main_table_index.values() {
360381
if offset.poly_size == 0 {

0 commit comments

Comments
 (0)