MCPcopy Create free account
hub / github.com/MegEngine/MegEngine / Elemwise

Method Elemwise

src/opr/impl/basic_arith.cpp:66–138  ·  view source on GitHub ↗

Source from the content-addressed store, hash-verified

64
65MGB_DYN_TYPE_OBJ_FINAL_IMPL(Elemwise);
66Elemwise::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;

Callers 9

_elwise_applyFunction · 0.80
utils.pyFile · 0.80
diag_plane_subgraphFunction · 0.80
__init__Method · 0.80
forwardMethod · 0.80
__init__Method · 0.80
__init__Method · 0.80
test_opdef_serializationFunction · 0.80
test_elemwiseFunction · 0.80

Calls 13

swapFunction · 0.85
invalidateMethod · 0.80
categoryMethod · 0.80
owner_graphMethod · 0.45
atMethod · 0.45
sizeMethod · 0.45
idMethod · 0.45
owner_oprMethod · 0.45
get_infer_typeMethod · 0.45
is_scalarMethod · 0.45
infer_shapeMethod · 0.45
validMethod · 0.45

Tested by 6

__init__Method · 0.64
forwardMethod · 0.64
__init__Method · 0.64
__init__Method · 0.64
test_opdef_serializationFunction · 0.64
test_elemwiseFunction · 0.64