1+ struct Edge {
2+ start : u32 ,
3+ end : u32 ,
4+ weight : f32 ,
5+ }
6+
7+ struct Uniforms {
8+ k : u32 ,
9+ max_distortion : f32 ,
10+ }
11+
12+ @group (0 ) @binding (0 ) var <uniform > uniforms : Uniforms ;
13+ @group (0 ) @binding (1 ) var <storage , read_write > distance_matrix : array <f32 >;
14+ @group (0 ) @binding (2 ) var <storage , read_write > distance_matrix_a : array <f32 >;
15+ @group (0 ) @binding (3 ) var <storage , read > graph_edges : array <Edge >;
16+ @group (0 ) @binding (4 ) var <storage , read_write > spanner_edges : array <u32 >;
17+
18+ override node_count : u32 ;
19+
20+ @compute @workgroup_size (8 , 8 )
21+ fn compute (
22+ @builtin (global_invocation_id ) global_id : vec3 <u32 >,
23+ ) {
24+ let x = global_id . x ;
25+ let y = global_id . y ;
26+ let n = node_count ;
27+
28+ if (x >= n || y >= n ) {
29+ let lol = spanner_edges [uniforms . k ];
30+ let lol2 = distance_matrix_a [0 ];
31+ return ;
32+ }
33+
34+ let edge = graph_edges [uniforms . k ];
35+
36+ var value = min (
37+ distance_matrix_get (x , y ),
38+ distance_matrix_get (x , edge . start ) + edge . weight + distance_matrix_get (edge . end , y )
39+ );
40+
41+ value = min (
42+ value ,
43+ distance_matrix_get (x , edge . end ) + edge . weight + distance_matrix_get (edge . start , y )
44+ );
45+
46+ distance_matrix_set (x , y , value );
47+ }
48+
49+ // Matrix getters and setters
50+
51+ fn distance_matrix_get (x : u32 , y : u32 ) -> f32 {
52+ return distance_matrix [get_matrix_index (x , y )];
53+ }
54+
55+ fn distance_matrix_set (x : u32 , y : u32 , value : f32 ) {
56+ distance_matrix [get_matrix_index (x , y )] = value ;
57+ }
58+
59+ fn get_matrix_index (x : u32 , y : u32 ) -> u32 {
60+ return x * node_count + y ;
61+ }
0 commit comments