| 108 | }; |
| 109 | |
| 110 | CondExecPred::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 | |
| 144 | class CondExecPred::GlobalRegistry final : public UserDataContainer::UserData { |
| 145 | MGB_TYPEINFO_OBJ_DECL; |