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

Method PredEvaluator

src/opr/impl/cond.cpp:110–142  ·  view source on GitHub ↗

Source from the content-addressed store, hash-verified

108};
109
110CondExecPred::PredEvaluator::PredEvaluator(
111 const CondExecPred& opr, const DeviceTensorND& pred) {
112 pre_check(pred);
113 switch (pred.dtype().enumv()) {
114#define cbf(dt) \
115 case DTypeTrait<dt>::enumv: { \
116 using ct = DTypeTrait<dt>::ctype; \
117 m_compare = [eps = opr.m_param.eps, \
118 p = pred.ptr<ct>()[0]](const DeviceTensorND& key) { \
119 ct k = key.ptr<ct>()[0]; \
120 return std::abs(p - k) < eps ? EQ : (p < k ? LT : GT); \
121 }; \
122 break; \
123 }
124#define cbi(dt) \
125 case DTypeTrait<dt>::enumv: { \
126 using ct = DTypeTrait<dt>::ctype; \
127 m_compare = [p = pred.ptr<ct>()[0]](const DeviceTensorND& key) { \
128 ct k = key.ptr<ct>()[0]; \
129 return p == k ? EQ : (p < k ? LT : GT); \
130 }; \
131 break; \
132 }
133
134 MEGDNN_FOREACH_COMPUTING_DTYPE_FLOAT(cbf);
135 MEGDNN_FOREACH_COMPUTING_DTYPE_INT(cbi)
136#undef cbf
137#undef cbi
138
139 default:
140 mgb_throw(GraphError, "unsupported pred dtype: %s", pred.dtype().name());
141 }
142}
143
144class CondExecPred::GlobalRegistry final : public UserDataContainer::UserData {
145 MGB_TYPEINFO_OBJ_DECL;

Callers

nothing calls this directly

Calls 3

enumvMethod · 0.45
dtypeMethod · 0.45
nameMethod · 0.45

Tested by

no test coverage detected