MCPcopy Create free account
hub / github.com/DeepRec-AI/DeepRec / Compute

Function Compute

tensorflow/core/kernels/maxpooling_op.cc:1200–1570  ·  view source on GitHub ↗

Source from the content-addressed store, hash-verified

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_;

Callers 15

operator()Method · 0.70
ComputeMethod · 0.70
ComputeMethod · 0.70
operator()Method · 0.70
ComputeMethod · 0.70
ComputeMethod · 0.70
ComputeMethod · 0.70
ComputeMethod · 0.70
ComputeMethod · 0.70
operator()Method · 0.70
ComputeMethod · 0.70
operator()Method · 0.70

Calls 12

ShapeFromFormatFunction · 0.85
InvalidArgumentFunction · 0.85
launchFunction · 0.85
allocate_outputMethod · 0.80
ToStringFunction · 0.70
DataTypeStringFunction · 0.50
NameClass · 0.50
inputMethod · 0.45
shapeMethod · 0.45
okMethod · 0.45
statusMethod · 0.45
DeviceMethod · 0.45

Tested by

no test coverage detected