@@ -69,7 +69,8 @@ void compare_impl(
6969 const size_t out_bind_size = (out_tensor.nbytes + 3 ) & ~size_t (3 );
7070 const uint32_t n_words = (numel + 3u ) / 4u ;
7171
72- uint32_t wg_size = utils::clamp_workgroup_size (device, kCompareWorkgroupSizeX );
72+ uint32_t wg_size =
73+ utils::clamp_workgroup_size (device, kCompareWorkgroupSizeX );
7374 uint32_t workgroup_count =
7475 utils::compute_1d_workgroup_count (device, n_words, wg_size, op_name);
7576
@@ -139,20 +140,21 @@ void compare_impl(
139140 graph.add_dispatch ({pipeline, bind_group, workgroup_count});
140141
141142 WGPUBuffer p_buf = params_buf;
142- auto cmp_resize = [self_id, out_id, mode, scalar, wg_size, dispatch_idx,
143- p_buf, op_name](WebGPUGraph& g) {
144- const auto & d = g.cur_dims (self_id);
145- uint32_t n = 1u ;
146- for (auto x : d) {
147- n *= static_cast <uint32_t >(x);
148- }
149- g.set_cur_dims (out_id, d);
150- CompareParams p = {n, mode, scalar, 0u };
151- wgpuQueueWriteBuffer (g.queue (), p_buf, 0 , &p, sizeof (p));
152- const uint32_t nw = (n + 3u ) / 4u ;
153- g.dispatch_at (dispatch_idx).workgroup_count_x =
154- utils::compute_1d_workgroup_count (g.device (), nw, wg_size, op_name);
155- };
143+ auto cmp_resize =
144+ [self_id, out_id, mode, scalar, wg_size, dispatch_idx, p_buf, op_name](
145+ WebGPUGraph& g) {
146+ const auto & d = g.cur_dims (self_id);
147+ uint32_t n = 1u ;
148+ for (auto x : d) {
149+ n *= static_cast <uint32_t >(x);
150+ }
151+ g.set_cur_dims (out_id, d);
152+ CompareParams p = {n, mode, scalar, 0u };
153+ wgpuQueueWriteBuffer (g.queue (), p_buf, 0 , &p, sizeof (p));
154+ const uint32_t nw = (n + 3u ) / 4u ;
155+ g.dispatch_at (dispatch_idx).workgroup_count_x =
156+ utils::compute_1d_workgroup_count (g.device (), nw, wg_size, op_name);
157+ };
156158 graph.add_tensor_resize_hook (self_id, cmp_resize);
157159
158160 wgpuShaderModuleRelease (shader);
0 commit comments