Skip to main content

sp1_gpu_cudart/tensor/
dot.rs

1use slop_tensor::{Tensor, TensorView};
2use sp1_gpu_sys::{
3    reduce::{
4        dot_along_short_dimension_kernel_koala_bear_base_base,
5        dot_along_short_dimension_kernel_koala_bear_base_extension,
6        dot_along_short_dimension_kernel_koala_bear_extension_extension,
7        partial_dot_koala_bear_base_extension_kernel, partial_dot_koala_bear_extension_kernel,
8        partial_dot_koala_bear_kernel,
9    },
10    runtime::KernelPtr,
11};
12use sp1_primitives::{SP1ExtensionField, SP1Field};
13
14use crate::{args, reduce::partial_sum_reduction_into, DeviceCopy, DeviceTensor, TaskScope};
15
16use super::reduce::DeviceSumKernel;
17
18/// # Safety
19///
20pub unsafe trait DotKernel<T: DeviceCopy, U: DeviceCopy>: DeviceSumKernel<U> {
21    fn partial_dot_kernel_last_dim() -> KernelPtr;
22
23    fn dot_along_short_dimension_kernel() -> KernelPtr;
24}
25
26pub fn dot_along_dim_view<'a, T: DeviceCopy, U: DeviceCopy>(
27    src: TensorView<'a, T, TaskScope>,
28    scalars: TensorView<'a, U, TaskScope>,
29    dim: usize,
30) -> Tensor<U, TaskScope>
31where
32    TaskScope: DotKernel<T, U>,
33{
34    let mut sizes = src.sizes().to_vec();
35    sizes.remove(dim);
36    let mut dst = Tensor::with_sizes_in(sizes, src.backend().clone());
37    assert_eq!(src.sizes().len(), 2, "Dot product only supported for 2D tensors",);
38    let max_scalar_dim = *scalars.sizes().iter().max().unwrap();
39    assert_eq!(max_scalar_dim, scalars.total_len(), "The scalar tensor must be a 1D tensor");
40    // The kernels read `scalars[0..src.sizes()[dim]]`: scalars beyond the dimension length
41    // are simply unused (callers dotting zero-padded data rely on reading a prefix), but a
42    // scalar tensor shorter than the dimension would be an out-of-bounds read on device.
43    assert!(
44        src.sizes()[dim] <= scalars.total_len(),
45        "dot along dimension {} of a tensor with sizes {:?} reads {} scalars, but only {} were \
46         provided",
47        dim,
48        src.sizes(),
49        src.sizes()[dim],
50        scalars.total_len()
51    );
52    match dim {
53        dim if dim == src.sizes().len() - 1 => {
54            let height = src.sizes()[dim];
55            let width = src.total_len() / height;
56
57            let null_ptr = std::ptr::null::<std::ffi::c_void>();
58            let partial_args = args!(null_ptr, src.as_ptr(), scalars.as_ptr(), width, height);
59            const BLOCK_SIZE: usize = 256;
60            const INTIAL_STRIDE: usize = 4;
61            dst.storage.write_bytes(0, dst.total_len() * std::mem::size_of::<U>()).unwrap();
62            unsafe {
63                partial_sum_reduction_into::<U, BLOCK_SIZE, INTIAL_STRIDE, 5>(
64                    dst.as_view_mut(),
65                    TaskScope::partial_dot_kernel_last_dim(),
66                    partial_args,
67                    0,
68                    src.shape(),
69                    dim,
70                    src.backend(),
71                );
72            }
73        }
74        0 => {
75            let height = src.sizes()[1];
76            let width = src.total_len() / height;
77
78            const BLOCK_SIZE: usize = 256;
79            let args = args!(dst.as_mut_ptr(), src.as_ptr(), scalars.as_ptr(), width, height);
80            let grid_dim = height.div_ceil(BLOCK_SIZE);
81            unsafe {
82                dst.assume_init();
83                src.backend()
84                    .launch_kernel(
85                        TaskScope::dot_along_short_dimension_kernel(),
86                        grid_dim,
87                        BLOCK_SIZE,
88                        &args,
89                        0,
90                    )
91                    .unwrap();
92            }
93        }
94        _ => panic!(
95            "Dot product is not supported along dimension {} for tensor of sizes {:?}",
96            dim,
97            src.sizes()
98        ),
99    }
100    dst
101}
102
103impl<T: DeviceCopy> DeviceTensor<T> {
104    pub fn dot_along_dim<U: DeviceCopy>(
105        &self,
106        scalars: &DeviceTensor<U>,
107        dim: usize,
108    ) -> DeviceTensor<U>
109    where
110        TaskScope: DotKernel<T, U>,
111    {
112        let raw = dot_along_dim_view(self.raw.as_view(), scalars.raw.as_view(), dim);
113        DeviceTensor { raw }
114    }
115}
116
117unsafe impl DotKernel<SP1Field, SP1Field> for TaskScope {
118    fn partial_dot_kernel_last_dim() -> KernelPtr {
119        unsafe { partial_dot_koala_bear_kernel() }
120    }
121
122    fn dot_along_short_dimension_kernel() -> KernelPtr {
123        unsafe { dot_along_short_dimension_kernel_koala_bear_base_base() }
124    }
125}
126
127unsafe impl DotKernel<SP1ExtensionField, SP1ExtensionField> for TaskScope {
128    fn partial_dot_kernel_last_dim() -> KernelPtr {
129        unsafe { partial_dot_koala_bear_extension_kernel() }
130    }
131
132    fn dot_along_short_dimension_kernel() -> KernelPtr {
133        unsafe { dot_along_short_dimension_kernel_koala_bear_extension_extension() }
134    }
135}
136
137unsafe impl DotKernel<SP1Field, SP1ExtensionField> for TaskScope {
138    fn partial_dot_kernel_last_dim() -> KernelPtr {
139        unsafe { partial_dot_koala_bear_base_extension_kernel() }
140    }
141
142    fn dot_along_short_dimension_kernel() -> KernelPtr {
143        unsafe { dot_along_short_dimension_kernel_koala_bear_base_extension() }
144    }
145}
146
147#[cfg(test)]
148mod tests {
149    use itertools::Itertools;
150    use slop_algebra::AbstractField;
151    use slop_tensor::Tensor;
152    use sp1_primitives::{SP1ExtensionField, SP1Field};
153
154    use super::DeviceTensor;
155
156    type SP1FieldExt = SP1ExtensionField;
157
158    #[test]
159    fn test_koala_bear_dot() {
160        let num_summands = 100;
161        let mut rng = rand::thread_rng();
162
163        for size in [10, 100, 1 << 16] {
164            let tensor = Tensor::<SP1Field>::rand(&mut rng, [num_summands, size]);
165            let scalars = Tensor::<SP1Field>::rand(&mut rng, [size]);
166
167            let inner_product = crate::run_sync_in_place(|t| {
168                let device_tensor = DeviceTensor::from_host(&tensor, &t).unwrap();
169                let device_scalars = DeviceTensor::from_host(&scalars, &t).unwrap();
170                let inner_product = device_tensor.dot_along_dim(&device_scalars, 1);
171                inner_product.to_host().unwrap()
172            })
173            .unwrap();
174
175            assert_eq!(inner_product.sizes(), [num_summands]);
176            for i in 0..num_summands {
177                let expected_inner_product: SP1Field = tensor
178                    .get(i)
179                    .unwrap()
180                    .as_slice()
181                    .iter()
182                    .copied()
183                    .zip_eq(scalars.as_buffer().iter().copied())
184                    .map(|(a, b)| a * b)
185                    .sum();
186                assert_eq!(expected_inner_product, *inner_product[[i]]);
187            }
188        }
189    }
190
191    #[test]
192    fn test_koala_bear_extension_dot() {
193        let num_summands = 100;
194        let mut rng = rand::thread_rng();
195
196        type EF = SP1ExtensionField;
197
198        for size in [10, 100, 1 << 16] {
199            let tensor = Tensor::<EF>::rand(&mut rng, [num_summands, size]);
200            let scalars = Tensor::<EF>::rand(&mut rng, [size]);
201
202            let inner_product = crate::run_sync_in_place(|t| {
203                let device_tensor = DeviceTensor::from_host(&tensor, &t).unwrap();
204                let device_scalars = DeviceTensor::from_host(&scalars, &t).unwrap();
205                let inner_product = device_tensor.dot_along_dim(&device_scalars, 1);
206                inner_product.to_host().unwrap()
207            })
208            .unwrap();
209
210            assert_eq!(inner_product.sizes(), [num_summands]);
211            for i in 0..num_summands {
212                let expected_inner_product: EF = tensor
213                    .get(i)
214                    .unwrap()
215                    .as_slice()
216                    .iter()
217                    .copied()
218                    .zip_eq(scalars.as_buffer().iter().copied())
219                    .map(|(a, b)| a * b)
220                    .sum();
221                assert_eq!(expected_inner_product, *inner_product[[i]]);
222            }
223        }
224    }
225
226    #[test]
227    fn test_koala_bear_base_extension_dot() {
228        let mut rng = rand::thread_rng();
229
230        type F = SP1Field;
231        type EF = SP1ExtensionField;
232
233        for size in [10, 100, 1 << 10, 1 << 12, 1 << 16] {
234            for num_summands in [64, 128] {
235                let tensor = Tensor::<F>::rand(&mut rng, [num_summands, size]);
236                let scalars = Tensor::<EF>::rand(&mut rng, [size]);
237
238                let inner_product = crate::run_sync_in_place(|t| {
239                    let device_tensor = DeviceTensor::from_host(&tensor, &t).unwrap();
240                    let device_scalars = DeviceTensor::from_host(&scalars, &t).unwrap();
241                    t.synchronize_blocking().unwrap();
242                    let time = std::time::Instant::now();
243                    let inner_product = device_tensor.dot_along_dim(&device_scalars, 1);
244                    t.synchronize_blocking().unwrap();
245                    tracing::info!(
246                        "Dot time for size {}, num_summands: {}, time: {:?}",
247                        size,
248                        num_summands,
249                        time.elapsed()
250                    );
251                    inner_product.to_host().unwrap()
252                })
253                .unwrap();
254
255                assert_eq!(inner_product.sizes(), [num_summands]);
256                for i in 0..num_summands {
257                    let expected_inner_product: EF = tensor
258                        .get(i)
259                        .unwrap()
260                        .as_slice()
261                        .iter()
262                        .copied()
263                        .zip_eq(scalars.as_buffer().iter().copied())
264                        .map(|(a, b)| b * a)
265                        .sum();
266                    assert_eq!(expected_inner_product, *inner_product[[i]]);
267                }
268            }
269        }
270    }
271
272    #[test]
273    fn test_dot_along_dim_0_base_base() {
274        let mut rng = rand::thread_rng();
275
276        let width = 10;
277        let height = 1500;
278
279        let host_tensor = Tensor::<SP1Field>::rand(&mut rng, [width, height]);
280        let host_scalars = Tensor::<SP1Field>::rand(&mut rng, [width]);
281
282        let dot = crate::run_sync_in_place(|t| {
283            let tensor = DeviceTensor::from_host(&host_tensor, &t).unwrap();
284            let scalars = DeviceTensor::from_host(&host_scalars, &t).unwrap();
285            let dot = tensor.dot_along_dim(&scalars, 0);
286            dot.to_host().unwrap()
287        })
288        .unwrap();
289
290        assert_eq!(dot.sizes(), [height]);
291        for i in 0..height {
292            let mut dot_product = SP1Field::zero();
293            for j in 0..width {
294                dot_product += *host_scalars[[j]] * *host_tensor[[j, i]];
295            }
296            assert_eq!(*dot[[i]], dot_product, "Dot product at index {i} is incorrect");
297        }
298    }
299
300    #[test]
301    fn test_dot_along_dim_0_base_ext() {
302        let mut rng = rand::thread_rng();
303
304        let width = 10;
305        let height = 1500;
306
307        let host_tensor = Tensor::<SP1Field>::rand(&mut rng, [width, height]);
308        let host_scalars = Tensor::<SP1FieldExt>::rand(&mut rng, [width]);
309
310        let dot = crate::run_sync_in_place(|t| {
311            let tensor = DeviceTensor::from_host(&host_tensor, &t).unwrap();
312            let scalars = DeviceTensor::from_host(&host_scalars, &t).unwrap();
313            let dot = tensor.dot_along_dim(&scalars, 0);
314            dot.to_host().unwrap()
315        })
316        .unwrap();
317
318        assert_eq!(dot.sizes(), [height]);
319        for i in 0..height {
320            let mut dot_product = SP1FieldExt::zero();
321            for j in 0..width {
322                dot_product += *host_scalars[[j]] * *host_tensor[[j, i]];
323            }
324            assert_eq!(*dot[[i]], dot_product, "Dot product at index {i} is incorrect");
325        }
326    }
327
328    #[test]
329    fn test_dot_along_dim_0_ext_ext() {
330        let mut rng = rand::thread_rng();
331
332        let width = 10;
333        let height = 1500;
334
335        let host_tensor = Tensor::<SP1FieldExt>::rand(&mut rng, [width, height]);
336        let host_scalars = Tensor::<SP1FieldExt>::rand(&mut rng, [width]);
337
338        let dot = crate::run_sync_in_place(|t| {
339            let tensor = DeviceTensor::from_host(&host_tensor, &t).unwrap();
340            let scalars = DeviceTensor::from_host(&host_scalars, &t).unwrap();
341            let dot = tensor.dot_along_dim(&scalars, 0);
342            dot.to_host().unwrap()
343        })
344        .unwrap();
345
346        assert_eq!(dot.sizes(), [height]);
347        for i in 0..height {
348            let mut dot_product = SP1FieldExt::zero();
349            for j in 0..width {
350                dot_product += *host_scalars[[j]] * *host_tensor[[j, i]];
351            }
352            assert_eq!(*dot[[i]], dot_product, "Dot product at index {i} is incorrect");
353        }
354    }
355}