Map MklTensorFormat to OneDNN format tag @input: MklTensorFormat i.e. TensorFlow data format @return: OneDNN's memory format tag corresponding to MklTensorFormat. Fails with an error if invalid data format.
| 1008 | // @return: OneDNN's memory format tag corresponding to MklTensorFormat. |
| 1009 | // Fails with an error if invalid data format. |
| 1010 | inline memory::format_tag MklTensorFormatToMklDnnDataFormat( |
| 1011 | MklTensorFormat format) { |
| 1012 | if (format == MklTensorFormat::FORMAT_NHWC) return memory::format_tag::nhwc; |
| 1013 | if (format == MklTensorFormat::FORMAT_NCHW) return memory::format_tag::nchw; |
| 1014 | if (format == MklTensorFormat::FORMAT_NDHWC) return memory::format_tag::ndhwc; |
| 1015 | if (format == MklTensorFormat::FORMAT_NCDHW) return memory::format_tag::ncdhw; |
| 1016 | if (format == MklTensorFormat::FORMAT_X) return memory::format_tag::x; |
| 1017 | if (format == MklTensorFormat::FORMAT_NC) return memory::format_tag::nc; |
| 1018 | if (format == MklTensorFormat::FORMAT_TNC) return memory::format_tag::tnc; |
| 1019 | return memory::format_tag::undef; |
| 1020 | } |
| 1021 | |
| 1022 | /// Map TensorFlow data format into OneDNN 3D data format |
| 1023 | /// @input: TensorFlow data format |
no outgoing calls
no test coverage detected