Skip to main content

sp1_gpu_basefold/
encoder.rs

1use std::sync::Arc;
2
3use slop_challenger::IopCtx;
4use slop_dft::DftOrdering;
5use sp1_gpu_cudart::TaskScope;
6use sp1_gpu_merkle_tree::MerkleTree;
7use sp1_gpu_utils::Felt;
8
9use slop_algebra::{AbstractField, Field};
10use slop_tensor::{Tensor, TensorView};
11use sp1_gpu_cudart::{
12    sys::dft::{batch_coset_dft, dft_init_default_stream, dft_init_twiddles},
13    CudaError, DeviceCopy,
14};
15use sp1_primitives::SP1Field;
16
17pub fn encode_batch<'a>(
18    dft: CudaDftKoalaBear,
19    log_blowup: u32,
20    data: TensorView<'a, Felt, TaskScope>,
21    dst: &mut Tensor<Felt, TaskScope>,
22) -> Result<(), CudaError> {
23    dft.coset_dft_into(
24        data,
25        dst,
26        <Felt as AbstractField>::one(),
27        log_blowup as usize,
28        DftOrdering::BitReversed,
29        1,
30    )
31    .unwrap();
32    Ok(())
33}
34
35pub trait CudaDftSys<T: DeviceCopy>: 'static + Send + Sync {
36    /// # Safety
37    ///
38    /// The caller must ensure the validity of pointers, allocation size, and lifetimes.
39    #[allow(clippy::too_many_arguments)]
40    unsafe fn dft_unchecked(
41        &self,
42        d_out: *mut T,
43        d_in: *mut T,
44        lg_domain_size: u32,
45        lg_blowup: u32,
46        shift: T,
47        batch_size: u32,
48        bit_rev_output: bool,
49        backend: &TaskScope,
50    ) -> Result<(), CudaError>;
51}
52
53#[derive(Debug, Clone, Copy, Default, PartialEq, Eq, PartialOrd, Ord, Hash)]
54pub struct CudaDft<F, T>(pub F, std::marker::PhantomData<T>);
55
56#[derive(Clone)]
57pub struct CudaStackedPcsProverData<GC: IopCtx> {
58    /// The usizes are the height of the Merkle tree and the number of elements in a leaf.
59    pub merkle_tree_tcs_data: (MerkleTree<GC::Digest, TaskScope>, GC::Digest, usize, usize),
60    /// The codeword (encoded polynomial). This is `None` when `drop_traces` is true.
61    pub codeword_mle: Option<Arc<Tensor<GC::F, TaskScope>>>,
62}
63
64impl<T: Field, F: CudaDftSys<T>> CudaDft<F, T> {
65    /// Performs a discrete Fourier transform along the last dimension of the input tensor.
66    fn coset_dft_into<'a>(
67        &self,
68        src: TensorView<'a, T, TaskScope>,
69        dst: &mut Tensor<T, TaskScope>,
70        shift: T,
71        log_blowup: usize,
72        ordering: DftOrdering,
73        dim: usize,
74    ) -> Result<(), CudaError> {
75        let backend = src.backend();
76        let d_in = src.as_ptr() as *mut T;
77        let d_out = dst.as_mut_ptr();
78        let src_dimensions = src.sizes();
79        let dst_dimensions = dst.sizes();
80
81        let shift = shift / T::generator();
82
83        assert_eq!(
84            src_dimensions[0], dst_dimensions[0],
85            "dimension mismatch along the first dimension"
86        );
87        assert_eq!(src.sizes().len(), 2);
88        assert_eq!(dst.sizes().len(), 2);
89        assert_eq!(dim, 1);
90
91        let lg_domain_size = src_dimensions[1].ilog2();
92        let lg_blowup = dst_dimensions[1].ilog2() - lg_domain_size;
93        assert_eq!(log_blowup, lg_blowup as usize);
94        let batch_size = src_dimensions[0] as u32;
95        let bit_rev_output = ordering == DftOrdering::BitReversed;
96
97        unsafe {
98            // Set the correct length for the output tensor
99            dst.assume_init();
100            // Call the function.
101            self.0.dft_unchecked(
102                d_out,
103                d_in,
104                lg_domain_size,
105                lg_blowup,
106                shift,
107                batch_size,
108                bit_rev_output,
109                backend,
110            )
111        }
112    }
113}
114
115#[derive(Copy, Clone, Debug)]
116pub struct CudaB31Kernels;
117
118pub type CudaDftKoalaBear = CudaDft<CudaB31Kernels, Felt>;
119
120impl CudaB31Kernels {
121    pub fn initialize_twiddles(max_log_size: u32, backend: &TaskScope) -> Result<(), CudaError> {
122        CudaError::result_from_ffi(unsafe { dft_init_twiddles(max_log_size, backend.handle()) })
123    }
124}
125
126impl Default for CudaB31Kernels {
127    fn default() -> Self {
128        CudaError::result_from_ffi(unsafe { dft_init_default_stream() }).unwrap();
129        Self
130    }
131}
132
133impl CudaDftSys<SP1Field> for CudaB31Kernels {
134    unsafe fn dft_unchecked(
135        &self,
136        d_out: *mut SP1Field,
137        d_in: *mut SP1Field,
138        lg_domain_size: u32,
139        lg_blowup: u32,
140        shift: SP1Field,
141        batch_size: u32,
142        bit_rev_output: bool,
143        scope: &TaskScope,
144    ) -> Result<(), CudaError> {
145        CudaError::result_from_ffi(batch_coset_dft(
146            d_out,
147            d_in,
148            lg_domain_size,
149            lg_blowup,
150            shift,
151            batch_size,
152            bit_rev_output,
153            scope.handle(),
154        ))
155    }
156}
157
158#[cfg(test)]
159mod tests {
160    use itertools::Itertools;
161    use rand::thread_rng;
162    use slop_algebra::AbstractField;
163    use slop_dft::{p3::Radix2DitParallel, Dft};
164
165    use sp1_gpu_cudart::{run_sync_in_place, DeviceTensor};
166
167    use super::*;
168
169    #[test]
170    fn test_batch_coset_dft() {
171        let mut rng = thread_rng();
172
173        let log_degrees = [1, 2, 3, 4, 5, 6, 7, 8, 9, 10, 11, 12, 13, 14, 15];
174        let log_blowup = 1;
175        let shift = SP1Field::generator();
176        let batch_size = 16;
177
178        let p3_dft = Radix2DitParallel;
179
180        for log_d in log_degrees.iter() {
181            let d = 1 << log_d;
182
183            let tensor_h = Tensor::<SP1Field>::rand(&mut rng, [d, batch_size]);
184
185            let tensor_h_sent = tensor_h.clone();
186            let result = run_sync_in_place(|t| {
187                let tensor_raw = DeviceTensor::from_host(&tensor_h_sent, &t).unwrap().into_inner();
188                let tensor = DeviceTensor::from_raw(tensor_raw).transpose().into_inner();
189                let dft = CudaDftKoalaBear::default();
190                let mut dst =
191                    Tensor::<Felt, _>::with_sizes_in([batch_size, d << log_blowup], t.clone());
192                dft.coset_dft_into(
193                    tensor.as_view(),
194                    &mut dst,
195                    shift,
196                    log_blowup,
197                    DftOrdering::BitReversed,
198                    1,
199                )
200                .unwrap();
201
202                let result = DeviceTensor::from_raw(dst).transpose();
203                result.to_host().unwrap()
204            })
205            .unwrap();
206
207            let expected_result = p3_dft
208                .coset_dft(&tensor_h, shift, log_blowup, DftOrdering::BitReversed, 0)
209                .unwrap();
210
211            for (i, (r, e)) in
212                result.as_slice().iter().zip_eq(expected_result.as_slice()).enumerate()
213            {
214                assert_eq!(r, e, "Mismatch for log degree {log_d} at index {i}");
215            }
216        }
217    }
218}