Skip to main content

msm_webgpu/cuzk/test/
point.rs

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    // Map results back from GPU to CPU.
74    let data = read_from_gpu_test(&device, &queue, encoder, vec![result_sb]).await;
75
76    // Destroy the GPU device object.
77    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
99/// Run WebGPU point op sync
100pub 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
104/// Run WebGPU point op async
105pub 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        // random u32
177        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}