| 1198 | } |
| 1199 | |
| 1200 | void Compute(OpKernelContext* context) override { |
| 1201 | const Tensor& tensor_in = context->input(0); |
| 1202 | |
| 1203 | PoolParameters params{context, ksize_, stride_, |
| 1204 | padding_, data_format_, tensor_in.shape()}; |
| 1205 | if (!context->status().ok()) { |
| 1206 | return; |
| 1207 | } |
| 1208 | |
| 1209 | TensorShape out_shape = |
| 1210 | ShapeFromFormat(data_format_, params.tensor_in_batch, params.out_height, |
| 1211 | params.out_width, params.depth); |
| 1212 | |
| 1213 | // Assuming qint8 <--> NCHW_VECT_C (int8x4) here. |
| 1214 | constexpr bool is_int8x4 = std::is_same<T, qint8>::value; |
| 1215 | OP_REQUIRES(context, (is_int8x4 == (data_format_ == FORMAT_NCHW_VECT_C)), |
| 1216 | errors::InvalidArgument( |
| 1217 | "qint8 should be used with data_format NCHW_VECT_C.")); |
| 1218 | |
| 1219 | #if CUDNN_VERSION >= 7300 |
| 1220 | if (use_dnn_) { |
| 1221 | DnnPoolingOp<T>::Compute(context, se::dnn::PoolingMode::kMaximum, ksize_, |
| 1222 | stride_, padding_, data_format_, tensor_in, |
| 1223 | out_shape, propagate_nans_); |
| 1224 | #else |
| 1225 | // These is_int8x4 checks avoid linker errors for missing qint8 kernels. |
| 1226 | if (!is_int8x4 && use_dnn_ && data_format_ == FORMAT_NCHW) { |
| 1227 | DnnPoolingOp<T>::Compute(context, se::dnn::PoolingMode::kMaximum, ksize_, |
| 1228 | stride_, padding_, data_format_, tensor_in, |
| 1229 | out_shape, propagate_nans_); |
| 1230 | #endif |
| 1231 | } else { |
| 1232 | Tensor* output = nullptr; |
| 1233 | OP_REQUIRES_OK(context, context->allocate_output(0, out_shape, &output)); |
| 1234 | if (is_int8x4) { |
| 1235 | LaunchMaxPoolingNoMask_NCHW_VECT_C<Device>::launch(context, params, |
| 1236 | tensor_in, output); |
| 1237 | } else if (data_format_ == FORMAT_NHWC) { |
| 1238 | LaunchMaxPoolingNoMask<Device, T>::launch(context, params, tensor_in, |
| 1239 | output, propagate_nans_); |
| 1240 | } else { |
| 1241 | LOG(FATAL) << "MaxPool currently only supports the following (layout, " |
| 1242 | "type) combinations: (NHWC, non-qint8), " |
| 1243 | "(NCHW, non-qint8) or (NCHW_VECT_C, qint8). The " |
| 1244 | "requested combination (" |
| 1245 | << ToString(data_format_) << ", " |
| 1246 | << DataTypeString(DataTypeToEnum<T>::v()) |
| 1247 | << ") is not supported."; |
| 1248 | } |
| 1249 | } |
| 1250 | } |
| 1251 | |
| 1252 | private: |
| 1253 | std::vector<int32> ksize_; |
| 1254 | std::vector<int32> stride_; |
| 1255 | Padding padding_; |
| 1256 | TensorFormat data_format_; |
| 1257 | bool use_dnn_; |
no test coverage detected