| 184 | } |
| 185 | |
| 186 | std::vector<TensorType> GetQuaterBcastType( |
| 187 | const CCOperand& operand0, const CCOperand& operand1, const CCOperand& operand2, |
| 188 | const CCOperand& operand3) { |
| 189 | auto shape0 = operand0.shape; |
| 190 | auto shape1 = operand1.shape; |
| 191 | auto shape2 = operand2.shape; |
| 192 | auto shape3 = operand3.shape; |
| 193 | |
| 194 | auto get_nr_elem = [](const std::vector<size_t>& shape) { |
| 195 | size_t nr_elem = 1; |
| 196 | for (size_t i = 0; i < shape.size(); i++) { |
| 197 | nr_elem *= shape[i]; |
| 198 | } |
| 199 | return nr_elem; |
| 200 | }; |
| 201 | size_t nr_elem0 = get_nr_elem(shape0); |
| 202 | size_t nr_elem1 = get_nr_elem(shape1); |
| 203 | size_t nr_elem2 = get_nr_elem(shape2); |
| 204 | size_t nr_elem3 = get_nr_elem(shape3); |
| 205 | size_t max_elemwise = |
| 206 | std::max(std::max(nr_elem0, nr_elem1), std::max(nr_elem2, nr_elem3)); |
| 207 | auto get_tensor_type = [&](size_t nr_elem, std::vector<size_t>& shape) { |
| 208 | if (nr_elem == 1) { |
| 209 | return SCALAR; |
| 210 | } else if (nr_elem == max_elemwise) { |
| 211 | return VECTOR; |
| 212 | } else { |
| 213 | if (shape[shape.size() - 1] != 4 && shape[shape.size() - 1] != 8) { |
| 214 | return BCAST101; |
| 215 | } else { |
| 216 | return BCAST101xX; |
| 217 | } |
| 218 | } |
| 219 | }; |
| 220 | std::vector<TensorType> ret; |
| 221 | ret.push_back(get_tensor_type(nr_elem0, shape0)); |
| 222 | ret.push_back(get_tensor_type(nr_elem1, shape1)); |
| 223 | ret.push_back(get_tensor_type(nr_elem2, shape2)); |
| 224 | ret.push_back(get_tensor_type(nr_elem3, shape3)); |
| 225 | return ret; |
| 226 | } |
| 227 | |
| 228 | std::vector<TensorType> DecodeTernaryBcastType(const BcastType bct_type) { |
| 229 | std::vector<TensorType> input_type; |