| 2038 | } |
| 2039 | |
| 2040 | void GetShuffleShape(AxesOrder input_axes_order, AxesOrder output_axes_order, |
| 2041 | std::vector<int>* shuffle) { |
| 2042 | CHECK_EQ(AxesCount(input_axes_order), AxesCount(output_axes_order)); |
| 2043 | shuffle->resize(4); |
| 2044 | for (int i = 0; i < 4; i++) { |
| 2045 | (*shuffle)[i] = i; |
| 2046 | } |
| 2047 | if (input_axes_order == output_axes_order) { |
| 2048 | // nothing to do |
| 2049 | } else if (AxesCount(input_axes_order) == 2) { |
| 2050 | shuffle->resize(2); |
| 2051 | (*shuffle)[0] = 1; |
| 2052 | (*shuffle)[1] = 0; |
| 2053 | } else if (input_axes_order == AxesOrder::kOHWI && |
| 2054 | output_axes_order == AxesOrder::kHWIO) { |
| 2055 | // 3210 <- 3210 |
| 2056 | // HWIO <- OHWI |
| 2057 | *shuffle = {1, 2, 3, 0}; |
| 2058 | } else if (input_axes_order == AxesOrder::kHWIO && |
| 2059 | output_axes_order == AxesOrder::kOHWI) { |
| 2060 | // 3210 <- 3210 |
| 2061 | // OHWI <- HWIO |
| 2062 | *shuffle = {3, 0, 1, 2}; |
| 2063 | } else if (input_axes_order == AxesOrder::kOHWI && |
| 2064 | output_axes_order == AxesOrder::kHWOI) { |
| 2065 | *shuffle = {1, 2, 0, 3}; |
| 2066 | } else { |
| 2067 | LOG(FATAL) << "Bad shuffle"; |
| 2068 | } |
| 2069 | } |
| 2070 | |
| 2071 | void ExtendShuffle(const std::vector<int>& input_shuffle, int newdim, |
| 2072 | std::vector<int>* extended_shuffle) { |
no test coverage detected