| 64 | |
| 65 | MGB_DYN_TYPE_OBJ_FINAL_IMPL(Elemwise); |
| 66 | Elemwise::Elemwise( |
| 67 | const ModeTrait& mode_trait, const VarNodeArrayView& inputs, Param param, |
| 68 | const OperatorNodeConfig& config) |
| 69 | : Super{inputs.at(0)->owner_graph(), config, mode_trait.name, inputs} { |
| 70 | init_megdnn_opr(*this, param); |
| 71 | output(0)->add_flag(VarNode::Flag::ALLOW_EMPTY_SHAPE); |
| 72 | if (mode_trait.commutable) { |
| 73 | mgb_assert(inputs.size() == 2); |
| 74 | add_input({inputs[0], inputs[1]}, AddInputSortType::CUR_ADDED); |
| 75 | } else { |
| 76 | if (param.mode == Mode::FUSE_MUL_ADD3) { |
| 77 | add_input({inputs[0], inputs[1]}, AddInputSortType::CUR_ADDED); |
| 78 | add_input({inputs[2]}); |
| 79 | } else if (param.mode == Mode::FUSE_MUL_ADD4) { |
| 80 | auto i0 = inputs[0], i1 = inputs[1], i2 = inputs[2], i3 = inputs[3]; |
| 81 | if (i0->id() > i1->id()) |
| 82 | std::swap(i0, i1); |
| 83 | if (i2->id() > i3->id()) |
| 84 | std::swap(i2, i3); |
| 85 | if (i0->id() > i2->id()) { |
| 86 | std::swap(i0, i2); |
| 87 | std::swap(i1, i3); |
| 88 | } |
| 89 | add_input({i0, i1, i2, i3}); |
| 90 | } else { |
| 91 | for (auto i : inputs) |
| 92 | add_input({i}); |
| 93 | } |
| 94 | } |
| 95 | |
| 96 | mgb_assert(m_input_broadcastable.size() >= inputs.size()); |
| 97 | for (size_t i = 0; i < inputs.size(); ++i) { |
| 98 | if (input()[i]->owner_opr()->same_type<opr::MarkNoBroadcastElemwise>()) { |
| 99 | m_input_broadcastable[i] = false; |
| 100 | } else { |
| 101 | m_input_broadcastable[i] = true; |
| 102 | } |
| 103 | } |
| 104 | if (inputs.size() == 1) { |
| 105 | m_input_broadcastable[0] = false; |
| 106 | } else { |
| 107 | Maybe<size_t> non_scalar; |
| 108 | using namespace cg::static_infer; |
| 109 | auto&& mgr = owner_graph()->static_infer_manager(); |
| 110 | for (size_t i = 0; i < input().size(); ++i) { |
| 111 | auto it = mgr.get_infer_type(input(i)); |
| 112 | if (!((it.shape & InferType::CONST) && |
| 113 | mgr.infer_shape(input(i)).is_scalar())) { |
| 114 | if (non_scalar.valid()) { |
| 115 | non_scalar.invalidate(); |
| 116 | break; |
| 117 | } |
| 118 | non_scalar = i; |
| 119 | } |
| 120 | } |
| 121 | if (non_scalar.valid()) { |
| 122 | // exactly one input is non-scalar |
| 123 | m_input_broadcastable[non_scalar.val()] = false; |