1use std::time::Instant;
2
3use halo2curves::{CurveAffine, CurveExt};
4use wgpu::CommandEncoderDescriptor;
5
6use crate::cuzk::{
7 gpu::{
8 create_and_write_storage_buffer, create_and_write_uniform_buffer, create_bind_group,
9 create_bind_group_layout, create_compute_pipeline, create_storage_buffer, execute_pipeline,
10 get_adapter, get_device, read_from_gpu_test,
11 },
12 msm::{P, PARAMS, WORD_SIZE},
13 shader_manager::ShaderManager,
14 utils::{bytes_to_field, points_to_bytes_for_gpu, to_biguint_le},
15};
16
17async fn point_op<C: CurveAffine>(op: &str, a: C, b: C, scalar: u32) -> C::Curve {
18 let a_bytes = points_to_bytes_for_gpu(&[a], PARAMS.num_words, WORD_SIZE);
19 let b_bytes = points_to_bytes_for_gpu(&[b], PARAMS.num_words, WORD_SIZE);
20 let scalar_bytes = scalar.to_le_bytes();
21 let input_size = 1;
22 let chunk_size = if input_size >= 65536 { 16 } else { 4 };
23 let num_words = PARAMS.num_words;
24 println!("Input size: {input_size}");
25 println!("Chunk size: {chunk_size}");
26 println!("Num words: {num_words}");
27 println!("Word size: {WORD_SIZE}");
28 println!("Params: {PARAMS:?}");
29
30 let shader_manager = ShaderManager::new(WORD_SIZE, chunk_size, input_size);
31
32 let adapter = get_adapter().await;
33 let (device, queue) = get_device(&adapter).await;
34 let mut encoder = device.create_command_encoder(&CommandEncoderDescriptor {
35 label: Some("Point Encoder"),
36 });
37
38 let shader_code = shader_manager.gen_test_point_shader();
39
40 let a_sb = create_and_write_storage_buffer(Some("A buffer"), &device, &a_bytes);
41 let b_sb = create_and_write_storage_buffer(Some("B buffer"), &device, &b_bytes);
42
43 let result_sb = create_storage_buffer(Some("Result buffer"), &device, 240);
44
45 let scalar_sb =
46 create_and_write_uniform_buffer(Some("Scalar buffer"), &device, &queue, &scalar_bytes);
47 let bind_group_layout = create_bind_group_layout(
48 Some("Bind group layout"),
49 &device,
50 vec![],
51 vec![&a_sb, &b_sb, &result_sb],
52 vec![&scalar_sb],
53 );
54
55 let bind_group = create_bind_group(
56 Some("Bind group"),
57 &device,
58 &bind_group_layout,
59 vec![&a_sb, &b_sb, &result_sb, &scalar_sb],
60 );
61
62 let compute_pipeline = create_compute_pipeline(
63 Some("Point shader"),
64 &device,
65 &bind_group_layout,
66 &shader_code,
67 op,
68 )
69 .await;
70
71 execute_pipeline(&mut encoder, compute_pipeline, bind_group, 1, 1, 1).await;
72
73 let data = read_from_gpu_test(&device, &queue, encoder, vec![result_sb]).await;
75
76 device.destroy();
78
79 let data_u32 = bytemuck::cast_slice::<u8, u32>(&data[0]);
80 println!("Data u32: {data_u32:?}");
81 println!("Data length: {:?}", data_u32.len());
82
83 let results = data_u32
84 .chunks(20)
85 .map(|chunk| {
86 let biguint_montgomery = to_biguint_le(chunk, num_words, WORD_SIZE as u32);
87 let biguint = biguint_montgomery * &PARAMS.rinv % P.clone();
88 let field: <<C as CurveAffine>::CurveExt as CurveExt>::Base =
89 bytes_to_field(&biguint.to_bytes_le());
90 field
91 })
92 .collect::<Vec<_>>();
93
94 println!("Results: {results:?}");
95
96 C::Curve::new_jacobian(results[0], results[1], results[2]).unwrap()
97}
98
99pub fn run_webgpu_point_op<C: CurveAffine>(op: &str, a: C, b: C, scalar: u32) -> C::Curve {
101 pollster::block_on(run_webgpu_point_op_async(op, a, b, scalar))
102}
103
104pub async fn run_webgpu_point_op_async<C: CurveAffine>(
106 op: &str,
107 a: C,
108 b: C,
109 scalar: u32,
110) -> C::Curve {
111 let now = Instant::now();
112 let result = point_op::<C>(op, a, b, scalar).await;
113 println!("Point op time: {:?}", now.elapsed());
114 result
115}
116
117#[cfg(test)]
118mod tests {
119 use super::*;
120 use group::Curve;
121 use group::cofactor::CofactorCurveAffine;
122 use halo2curves::bn256::{Fr, G1Affine};
123 use rand::{Rng, thread_rng};
124
125 #[test]
126 fn test_webgpu_point_add() {
127 let mut rng = thread_rng();
128 let a = G1Affine::random(&mut rng);
129 println!("a: {:?}", a);
130 let b = G1Affine::random(&mut rng);
131 println!("b: {:?}", b);
132
133 let fast = a + b;
134
135 let result = run_webgpu_point_op::<G1Affine>("test_point_add", a, b, 0);
136
137 println!("Result: {:?}", result);
138 assert_eq!(fast, result);
139 }
140
141 #[test]
142 fn test_webgpu_point_add_identity() {
143 let mut rng = thread_rng();
144 let a = G1Affine::random(&mut rng);
145 println!("a: {:?}", a);
146 let b = G1Affine::identity();
147 println!("b: {:?}", b);
148
149 let fast = a + b;
150
151 let result = run_webgpu_point_op::<G1Affine>("test_point_add_identity", a, b, 0);
152
153 println!("Result: {:?}", result);
154 assert_eq!(fast, result);
155 }
156
157 #[test]
158 fn test_webgpu_point_negate() {
159 let mut rng = thread_rng();
160 let a = G1Affine::random(&mut rng);
161 println!("a: {:?}", a);
162
163 let fast = -a;
164
165 let result = run_webgpu_point_op::<G1Affine>("test_negate_point", a, a, 0);
166
167 println!("Result: {:?}", result);
168 assert_eq!(fast, result.to_affine());
169 }
170
171 #[test]
172 fn test_webgpu_point_double_and_add() {
173 let mut rng = thread_rng();
174 let a = G1Affine::random(&mut rng);
175 println!("a: {:?}", a);
176 let scalar = rng.gen_range(0..u32::MAX);
178 println!("scalar: {:?}", scalar);
179
180 let fast = a * Fr::from(scalar as u64);
181
182 let result = run_webgpu_point_op::<G1Affine>("test_double_and_add", a, a, scalar);
183
184 println!("Result: {:?}", result);
185 assert_eq!(fast, result);
186 }
187}