@@ -74,7 +74,7 @@ static void CalculateSubgroupArithmetic(
7474 ::benchmark::State &state, ::uvkc::vulkan::Device *device,
7575 const ::uvkc::benchmark::LatencyMeasure *latency_measure,
7676 const uint32_t *code, size_t code_num_words, int num_elements,
77- uint32_t subgroup_size , Arithmetic arith_op) {
77+ uint32_t proposed_subgroup_size , Arithmetic arith_op) {
7878 size_t buffer_num_bytes = num_elements * sizeof (float );
7979
8080 // ===-------------------------------------------------------------------===/
@@ -116,10 +116,11 @@ static void CalculateSubgroupArithmetic(
116116 // ===-------------------------------------------------------------------===/
117117
118118 // +: fill the whole buffer as 1.0f.
119- // *: fill with alternating subgroup_size and (1 / subgroup_size).
119+ // *: fill with alternating values of proposed_subgroup_size and
120+ // (1 / proposed_subgroup_size).
120121 BM_CHECK_OK (::uvkc::benchmark::SetDeviceBufferViaStagingBuffer (
121122 device, src_buffer.get (), buffer_num_bytes,
122- [arith_op, subgroup_size ](void *ptr, size_t num_bytes) {
123+ [arith_op, proposed_subgroup_size ](void *ptr, size_t num_bytes) {
123124 float *src_float_buffer = reinterpret_cast <float *>(ptr);
124125 switch (arith_op) {
125126 case Arithmetic::Add: {
@@ -129,8 +130,8 @@ static void CalculateSubgroupArithmetic(
129130 } break ;
130131 case Arithmetic::Mul: {
131132 for (int i = 0 ; i < num_bytes / sizeof (float ); i += 2 ) {
132- src_float_buffer[i] = subgroup_size ;
133- src_float_buffer[i + 1 ] = 1 .0f / subgroup_size ;
133+ src_float_buffer[i] = proposed_subgroup_size ;
134+ src_float_buffer[i + 1 ] = 1 .0f / proposed_subgroup_size ;
134135 }
135136 } break ;
136137 }
@@ -171,14 +172,17 @@ static void CalculateSubgroupArithmetic(
171172
172173 BM_CHECK_OK (::uvkc::benchmark::GetDeviceBufferViaStagingBuffer (
173174 device, dst_buffer.get (), buffer_num_bytes,
174- [arith_op, subgroup_size](void *ptr, size_t num_bytes) {
175- float *dst_float_buffer = reinterpret_cast <float *>(ptr);
175+ [arith_op, proposed_subgroup_size](void *ptr, size_t num_bytes) {
176+ const uint32_t actual_subgroup_size =
177+ reinterpret_cast <uint32_t *>(ptr)[0 ];
178+ float *dst_float_buffer = reinterpret_cast <float *>(ptr) + 1 ;
179+ const auto num_floats = (num_bytes / sizeof (float )) - 1 ;
176180 switch (arith_op) {
177181 case Arithmetic::Add: {
178- for (int i = 0 ; i < num_bytes / sizeof ( float ) ; ++i) {
182+ for (int i = 0 ; i < num_floats ; ++i) {
179183 float expected_value = 1 .0f ;
180- if (i % subgroup_size == 0 ) {
181- expected_value = subgroup_size ;
184+ if (i % actual_subgroup_size == 0 ) {
185+ expected_value = actual_subgroup_size ;
182186 }
183187
184188 BM_CHECK_EQ (dst_float_buffer[i], expected_value)
@@ -188,14 +192,14 @@ static void CalculateSubgroupArithmetic(
188192 }
189193 } break ;
190194 case Arithmetic::Mul: {
191- for (int i = 0 ; i < num_bytes / sizeof ( float ) ; ++i) {
195+ for (int i = 0 ; i < num_floats ; ++i) {
192196 float expected_value = 0 .0f ;
193- if (i % subgroup_size == 0 ) {
197+ if (i % actual_subgroup_size == 0 ) {
194198 expected_value = 1 .0f ;
195199 } else if (i % 2 == 0 ) {
196- expected_value = subgroup_size ;
200+ expected_value = proposed_subgroup_size ;
197201 } else {
198- expected_value = 1 .0f / subgroup_size ;
202+ expected_value = 1 .0f / proposed_subgroup_size ;
199203 }
200204
201205 BM_CHECK_EQ (dst_float_buffer[i], expected_value)
0 commit comments