| 199 | } |
| 200 | |
| 201 | __fp16 cactus_max_all_f16(const __fp16* data, size_t num_elements) { |
| 202 | return CactusThreading::parallel_reduce( |
| 203 | num_elements, CactusThreading::Thresholds::ALL_REDUCE, |
| 204 | [&](size_t start, size_t end) -> __fp16 { |
| 205 | const size_t vec_end = start + simd_align(end - start); |
| 206 | float16x8_t acc = vdupq_n_f16(static_cast<__fp16>(-65504.0f)); |
| 207 | |
| 208 | for (size_t i = start; i < vec_end; i += SIMD_F16_WIDTH) { |
| 209 | acc = vmaxq_f16(acc, vld1q_f16(&data[i])); |
| 210 | } |
| 211 | |
| 212 | __fp16 result = static_cast<__fp16>(-65504.0f); |
| 213 | __fp16 arr[8]; |
| 214 | vst1q_f16(arr, acc); |
| 215 | for (int j = 0; j < 8; j++) result = std::max(result, arr[j]); |
| 216 | for (size_t i = vec_end; i < end; ++i) result = std::max(result, data[i]); |
| 217 | return result; |
| 218 | }, |
| 219 | static_cast<__fp16>(-65504.0f), |
| 220 | [](__fp16 a, __fp16 b) { return std::max(a, b); } |
| 221 | ); |
| 222 | } |
| 223 | |
| 224 | void cactus_max_axis_f16(const __fp16* input, __fp16* output, size_t outer_size, size_t axis_size, size_t inner_size) { |
| 225 | axis_reduce_f16_impl(input, output, outer_size, axis_size, inner_size, |
no test coverage detected