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
18pub 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 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}