|
1 | 1 | use crate::data::{InfoBuffer, JaggedDenseInfo}; |
| 2 | +use slop_air::BaseAir; |
2 | 3 | use slop_algebra::{AbstractField, ExtensionField, Field}; |
3 | 4 | use slop_alloc::{Backend, Buffer, CpuBackend, HasBackend, Slice}; |
4 | 5 | use slop_commit::Rounds; |
5 | 6 | use slop_multilinear::{Evaluations, MleEval, Point}; |
6 | | -use slop_tensor::Tensor; |
| 7 | +use slop_tensor::{Tensor, TensorView}; |
7 | 8 | use sp1_gpu_cudart::sys::runtime::KernelPtr; |
8 | 9 | use sp1_gpu_cudart::sys::v2_kernels::{ |
9 | 10 | fix_last_variable_jagged_ext, fix_last_variable_jagged_felt, fix_last_variable_jagged_info, |
10 | 11 | initialize_jagged_info, |
11 | 12 | }; |
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 | +}; |
13 | 16 | 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}; |
15 | 20 | use std::iter::once; |
16 | 21 |
|
17 | 22 | pub trait JaggedFixLastVariableKernel<K: Field> { |
@@ -123,20 +128,99 @@ where |
123 | 128 |
|
124 | 129 | #[inline(always)] |
125 | 130 | 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 | + } |
131 | 189 | } |
132 | 190 |
|
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 |
140 | 224 | } |
141 | 225 |
|
142 | 226 | pub fn evaluate_jagged_columns( |
|
0 commit comments