| 44 | using ctype_to_hl_type_t = typename ctype_to_hl_type<T>::hl_type; |
| 45 | |
| 46 | Halide::Expr dispatch_elemwise_mode( |
| 47 | opr::Elemwise::Mode mode, const AstNodeArray& inputs, DType out_dtype, |
| 48 | const std::vector<std::vector<Halide::Expr>>& exprs_of_inps) { |
| 49 | using Mode = opr::Elemwise::Mode; |
| 50 | auto cv = [&](Halide::Expr a) { return hl_type_cast(a, out_dtype); }; |
| 51 | |
| 52 | #define inp(i) (inputs[(i)]->m_func(exprs_of_inps[(i)])) |
| 53 | switch (mode) { |
| 54 | // unary |
| 55 | case Mode::RELU: |
| 56 | return Halide::select(inp(0) <= 0, cv(0), inp(0)); |
| 57 | case Mode::ABS: |
| 58 | return Halide::abs(inp(0)); |
| 59 | case Mode::ACOS: |
| 60 | return Halide::acos(inp(0)); |
| 61 | case Mode::ASIN: |
| 62 | return Halide::asin(inp(0)); |
| 63 | case Mode::CEIL: |
| 64 | return Halide::ceil(inp(0)); |
| 65 | case Mode::COS: |
| 66 | return Halide::cos(inp(0)); |
| 67 | case Mode::EXP: |
| 68 | return Halide::exp(inp(0)); |
| 69 | case Mode::EXPM1: |
| 70 | return Halide::exp(inp(0)) - cv(1); |
| 71 | case Mode::FLOOR: |
| 72 | return Halide::floor(inp(0)); |
| 73 | case Mode::LOG: |
| 74 | return Halide::log(inp(0)); |
| 75 | case Mode::LOG1P: |
| 76 | return Halide::log(inp(0) + cv(1)); |
| 77 | case Mode::NEGATE: |
| 78 | return -inp(0); |
| 79 | case Mode::SIGMOID: |
| 80 | return cv(1) / (cv(1) + Halide::exp(-inp(0))); |
| 81 | case Mode::SIN: |
| 82 | return Halide::sin(inp(0)); |
| 83 | case Mode::TANH: |
| 84 | return Halide::tanh(inp(0)); |
| 85 | case Mode::ERF: |
| 86 | return Halide::erf(inp(0)); |
| 87 | case Mode::ERFC: |
| 88 | return cv(1) - Halide::erf(inp(0)); |
| 89 | case Mode::H_SWISH: |
| 90 | return inp(0) * Halide::max(Halide::min(inp(0) + cv(3), cv(6)), cv(0)) / |
| 91 | cv(6); |
| 92 | |
| 93 | // binary |
| 94 | case Mode::ABS_GRAD: |
| 95 | return Halide::select(inp(0) > 0, inp(1), -inp(1)); |
| 96 | case Mode::ADD: |
| 97 | return inp(0) + inp(1); |
| 98 | case Mode::FLOOR_DIV: |
| 99 | return Halide::floor(inp(0) / inp(1)); |
| 100 | case Mode::MAX: |
| 101 | return Halide::max(inp(0), inp(1)); |
| 102 | case Mode::MIN: |
| 103 | return Halide::min(inp(0), inp(1)); |
no test coverage detected