Skip to content

Commit f66cd47

Browse files
committed
working, but requiring code cleanup
1 parent 45f7af8 commit f66cd47

3 files changed

Lines changed: 102 additions & 18 deletions

File tree

sp1-gpu/crates/logup_gkr/src/lib.rs

Lines changed: 2 additions & 2 deletions
Original file line numberDiff line numberDiff line change
@@ -552,8 +552,8 @@ mod tests {
552552
#[serial]
553553
fn test_logup_gkr_e2e() {
554554
let rt = tokio::runtime::Runtime::new().unwrap();
555-
let (machine, record, program) =
556-
rt.block_on(tracegen_setup::setup(&test_artifacts::FIBONACCI_ELF, SP1Stdin::new()));
555+
let (machine, record, program) = rt
556+
.block_on(tracegen_setup::setup(&test_artifacts::SSZ_WITHDRAWALS_ELF, SP1Stdin::new()));
557557

558558
run_sync_in_place(|scope| {
559559
// *********** Generate traces using the host tracegen. ***********

sp1-gpu/crates/shard_prover/src/prover.rs

Lines changed: 1 addition & 1 deletion
Original file line numberDiff line numberDiff line change
@@ -810,7 +810,7 @@ mod tests {
810810
let buffer = PinnedBuffer::<Felt>::with_capacity(capacity);
811811
let queue = Arc::new(WorkerQueue::new(vec![buffer]));
812812
let buffer = queue.pop().await.unwrap();
813-
let (_public_values, jagged_trace_data, _shard_chips, _permit) = full_tracegen(
813+
let (_public_values, jagged_trace_data, shard_chips, _permit) = full_tracegen(
814814
&machine,
815815
program.clone(),
816816
Arc::new(record),

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

Lines changed: 99 additions & 15 deletions
Original file line numberDiff line numberDiff line change
@@ -1,17 +1,22 @@
11
use crate::data::{InfoBuffer, JaggedDenseInfo};
2+
use slop_air::BaseAir;
23
use slop_algebra::{AbstractField, ExtensionField, Field};
34
use slop_alloc::{Backend, Buffer, CpuBackend, HasBackend, Slice};
45
use slop_commit::Rounds;
56
use slop_multilinear::{Evaluations, MleEval, Point};
6-
use slop_tensor::Tensor;
7+
use slop_tensor::{Tensor, TensorView};
78
use sp1_gpu_cudart::sys::runtime::KernelPtr;
89
use sp1_gpu_cudart::sys::v2_kernels::{
910
fix_last_variable_jagged_ext, fix_last_variable_jagged_felt, fix_last_variable_jagged_info,
1011
initialize_jagged_info,
1112
};
12-
use sp1_gpu_cudart::{args, DeviceBuffer, DevicePoint, DeviceTensor, TaskScope};
13+
use sp1_gpu_cudart::{
14+
args, dot_along_dim_view, DeviceBuffer, DevicePoint, DeviceTensor, TaskScope,
15+
};
1316
use sp1_gpu_utils::{Ext, Felt, JaggedMle, JaggedTraceMle, TraceDenseData, TraceOffset};
14-
use std::collections::BTreeMap;
17+
use sp1_hypercube::air::MachineAir;
18+
use sp1_hypercube::Chip;
19+
use std::collections::{BTreeMap, BTreeSet};
1520
use std::iter::once;
1621

1722
pub trait JaggedFixLastVariableKernel<K: Field> {
@@ -123,20 +128,99 @@ where
123128

124129
#[inline(always)]
125130
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);
131+
let trace_data = traces.dense();
132+
let backend = traces.backend();
133+
let device_point = DevicePoint::from_host(point, backend).unwrap();
134+
let partial_lagrange = device_point.partial_lagrange();
135+
let total_cols = trace_data
136+
.preprocessed_table_index
137+
.values()
138+
.map(|index| index.num_polys)
139+
.chain(trace_data.main_table_index.values().map(|index| index.num_polys))
140+
.sum::<usize>();
141+
let mut result_buffer =
142+
DeviceBuffer::with_capacity_in(total_cols + 2, backend.clone()).into_inner();
143+
let mut preprocessed_buffer =
144+
DeviceBuffer::with_capacity_in(trace_data.preprocessed_cols, backend.clone()).into_inner();
145+
let mut main_buffer = DeviceBuffer::with_capacity_in(total_cols, backend.clone()).into_inner();
146+
147+
// println!("Preprocessed table index: {:?}", trace_data.preprocessed_table_index);
148+
// println!("Main table index: {:?}", trace_data.main_table_index);
149+
150+
let trace_ptr = trace_data.dense.as_ptr();
151+
for index in trace_data.preprocessed_table_index.values() {
152+
let chip_preprocessed_ptr = unsafe { trace_ptr.add(index.dense_offset.start) };
153+
let chip_preprocessed_view = unsafe {
154+
TensorView::from_raw_parts(
155+
chip_preprocessed_ptr,
156+
[index.num_polys, index.poly_size].try_into().unwrap(),
157+
backend.clone(),
158+
)
159+
};
160+
if index.dense_offset.start != index.dense_offset.end {
161+
let result =
162+
dot_along_dim_view(chip_preprocessed_view, partial_lagrange.guts().as_view(), 1);
163+
// println!(
164+
// "Result dimensions: {:?}, chip num_polys: {}",
165+
// result.shape(),
166+
// index.num_polys
167+
// );
168+
preprocessed_buffer.extend_from_device_slice(result.as_buffer()).unwrap();
169+
}
170+
}
171+
for index in trace_data.main_table_index.values() {
172+
let chip_main_ptr = unsafe { trace_ptr.add(index.dense_offset.start) };
173+
let chip_main_view = unsafe {
174+
TensorView::from_raw_parts(
175+
chip_main_ptr,
176+
[index.num_polys, index.poly_size].try_into().unwrap(),
177+
backend.clone(),
178+
)
179+
};
180+
if index.dense_offset.start != index.dense_offset.end {
181+
let result = dot_along_dim_view(chip_main_view, partial_lagrange.guts().as_view(), 1);
182+
// println!(
183+
// "Result dimensions: {:?}, chip num_polys: {}",
184+
// result.shape(),
185+
// index.num_polys
186+
// );
187+
main_buffer.extend_from_device_slice(result.as_buffer()).unwrap();
188+
}
131189
}
132190

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<_>>()
191+
result_buffer.extend_from_device_slice(&preprocessed_buffer).unwrap();
192+
let padding = Tensor::zeros_in([1], backend.clone());
193+
result_buffer.extend_from_device_slice(padding.as_buffer()).unwrap();
194+
result_buffer.extend_from_device_slice(&main_buffer).unwrap();
195+
let host_result = DeviceBuffer::from_raw(result_buffer).to_host().unwrap();
196+
197+
// println!("Host result first 100 values: {:?}", host_result[..10].to_vec());
198+
199+
// let mut next_input_jagged_trace_mle =
200+
// evaluate_jagged_fix_last_variable(traces, *point.last().unwrap());
201+
// for alpha in point.iter().rev().skip(1) {
202+
// next_input_jagged_trace_mle =
203+
// evaluate_jagged_fix_last_variable(&next_input_jagged_trace_mle, *alpha);
204+
// }
205+
206+
// let host_dense = DeviceBuffer::from_raw(next_input_jagged_trace_mle.dense_data.dense.clone())
207+
// .to_host()
208+
// .unwrap()
209+
// .to_vec();
210+
211+
// // Only every four elements is not padding.
212+
// let expected_host_dense = host_dense.into_iter().step_by(4).collect::<Vec<_>>();
213+
// println!("Expected host dense first 100 values: {:?}", expected_host_dense[..10].to_vec());
214+
215+
// for (i, (elem, expected_elem)) in host_result.iter().zip(expected_host_dense.iter()).enumerate()
216+
// {
217+
// assert_eq!(
218+
// *elem, *expected_elem,
219+
// "element mismatch: elem={} expected={}, index={}",
220+
// *elem, *expected_elem, i
221+
// );
222+
// }
223+
host_result
140224
}
141225

142226
pub fn evaluate_jagged_columns(

0 commit comments

Comments
 (0)