MCPcopy Create free account
hub / github.com/MegEngine/MegCC / GetQuaterBcastType

Function GetQuaterBcastType

compiler/lib/KernelGen/Common/ElemwiseCommon.cpp:186–226  ·  view source on GitHub ↗

Source from the content-addressed store, hash-verified

184}
185
186std::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
228std::vector<TensorType> DecodeTernaryBcastType(const BcastType bct_type) {
229 std::vector<TensorType> input_type;

Callers

nothing calls this directly

Calls 1

maxFunction · 0.85

Tested by

no test coverage detected