| 3257 | } |
| 3258 | |
| 3259 | port::StatusOr<std::unique_ptr<cudnn_frontend::OperationGraph>> |
| 3260 | GetCudnnOperationGraph( |
| 3261 | dnn::ConvolutionKind kind, dnn::DataType element_type, Stream* stream, |
| 3262 | const dnn::BatchDescriptor& input_descriptor, |
| 3263 | const dnn::FilterDescriptor& filter_descriptor, |
| 3264 | const dnn::BatchDescriptor& output_descriptor, |
| 3265 | const dnn::ConvolutionDescriptor& convolution_descriptor, |
| 3266 | CudnnHandle &cudnn) { |
| 3267 | cudnnBackendDescriptorType_t conv_mode = GetCudnnConvolutionType(kind); |
| 3268 | cudnnDataType_t cudnn_type = ToCudnnDataType(element_type); |
| 3269 | |
| 3270 | // x tensor. |
| 3271 | std::vector<int64> input_strides64 = input_descriptor.full_strides( |
| 3272 | dnn::DataLayout::kBatchDepthYX); |
| 3273 | std::vector<int64> input_dims64 = input_descriptor.full_dims( |
| 3274 | dnn::DataLayout::kBatchDepthYX); |
| 3275 | std::vector<int64_t> input_strides(input_strides64.cbegin(), |
| 3276 | input_strides64.cend()); |
| 3277 | std::vector<int64_t> input_dims(input_dims64.cbegin(), input_dims64.cend()); |
| 3278 | auto tensor_x = cudnn_frontend::TensorBuilder() |
| 3279 | .setDim(input_dims.size(), &input_dims[0]) |
| 3280 | .setStrides(input_dims.size(), &input_strides[0]) |
| 3281 | .setId('x') |
| 3282 | .setAlignment(32) |
| 3283 | .setDataType(cudnn_type) |
| 3284 | .build(); |
| 3285 | RETURN_MSG_IF_CUDNN_ERROR(tensor_x); |
| 3286 | |
| 3287 | // y tensor. |
| 3288 | std::vector<int64> output_strides64 = output_descriptor.full_strides( |
| 3289 | dnn::DataLayout::kBatchDepthYX); |
| 3290 | std::vector<int64> output_dims64 = output_descriptor.full_dims( |
| 3291 | dnn::DataLayout::kBatchDepthYX); |
| 3292 | std::vector<int64_t> output_strides(output_strides64.cbegin(), |
| 3293 | output_strides64.cend()); |
| 3294 | std::vector<int64_t> output_dims(output_dims64.cbegin(), |
| 3295 | output_dims64.cend()); |
| 3296 | auto tensor_y = cudnn_frontend::TensorBuilder() |
| 3297 | .setDim(output_dims.size(), &output_dims[0]) |
| 3298 | .setStrides(output_dims.size(), &output_strides[0]) |
| 3299 | .setId('y') |
| 3300 | .setAlignment(32) |
| 3301 | .setDataType(cudnn_type) |
| 3302 | .build(); |
| 3303 | RETURN_MSG_IF_CUDNN_ERROR(tensor_y); |
| 3304 | |
| 3305 | // w tensor: Transform HWNC (XYIO) format to NCHW/NHWC. |
| 3306 | std::vector<int64> filter_dims64(2 + filter_descriptor.ndims()); |
| 3307 | filter_dims64[0] = filter_descriptor.output_feature_map_count(); |
| 3308 | filter_dims64[1] = filter_descriptor.input_feature_map_count(); |
| 3309 | auto spatial_dims64 = filter_descriptor.input_filter_dims(); |
| 3310 | std::copy(spatial_dims64.begin(), spatial_dims64.end(), |
| 3311 | filter_dims64.begin() + 2); |
| 3312 | cudnnTensorFormat_t format; |
| 3313 | dnn::DataLayout tensor_format; |
| 3314 | switch (filter_descriptor.layout()) { |
| 3315 | case dnn::FilterLayout::kOutputInputYX: |
| 3316 | format = CUDNN_TENSOR_NCHW; |
no test coverage detected