| 189 | |
| 190 | template<typename T> |
| 191 | static void ReluGrad(const T* dy_ptr, const int32_t* mask_ptr, T* relu_dx_ptr, |
| 192 | const int64_t elem_cnt) { |
| 193 | const int32_t step = 32; |
| 194 | const int64_t outer_loop = elem_cnt / step; |
| 195 | const int64_t remain_loop_start_idx = outer_loop * step; |
| 196 | |
| 197 | for (int64_t outer = 0; outer < outer_loop; ++outer) { |
| 198 | const int32_t mask_val = mask_ptr[outer]; |
| 199 | for (int32_t s = 0; s < step; ++s) { |
| 200 | bool is_positive = mask_val & (1 << s); |
| 201 | relu_dx_ptr[s] = static_cast<T>(is_positive) * dy_ptr[s]; |
| 202 | } |
| 203 | relu_dx_ptr += step; |
| 204 | dy_ptr += step; |
| 205 | } |
| 206 | |
| 207 | if (remain_loop_start_idx < elem_cnt) { |
| 208 | const int32_t mask_val = mask_ptr[outer_loop]; |
| 209 | const int32_t remain = elem_cnt - remain_loop_start_idx; |
| 210 | for (int32_t i = 0; i < remain; ++i) { |
| 211 | bool is_positive = mask_val & (1 << i); |
| 212 | relu_dx_ptr[i] = static_cast<T>(is_positive) * dy_ptr[i]; |
| 213 | } |
| 214 | } |
| 215 | } |
| 216 | |
| 217 | static size_t InferGradTmpSizeForCpuKernel(user_op::InferContext* ctx) { |
| 218 | const auto& dy = ctx->InputTensorDesc("dy", 0); |