=================== MakeShapeEmitter ====================*/
| 66 | |
| 67 | // =================== MakeShapeEmitter ====================*/ |
| 68 | MakeShapeEmitter::EmitResult MakeShapeEmitter::emit() const { |
| 69 | auto pattern = mixin_analyze(); |
| 70 | auto builder = [pattern](const VarNodeArray& input) { |
| 71 | mgb_assert( |
| 72 | input.size() == 1, |
| 73 | "number of input of MakeShapeBuilder should be 1(got:%zu)", |
| 74 | input.size()); |
| 75 | auto sym_var = SymbolVar(input.front()); |
| 76 | auto shp = opr::GetVarShape::make(sym_var); |
| 77 | auto cv = [&sym_var](int c) { return sym_var.make_scalar(c); }; |
| 78 | auto sub = [&shp, &cv](int ax) { |
| 79 | return opr::IndexAt::make(shp, {{0, cv(ax)}}); |
| 80 | }; |
| 81 | SymbolVarArray axs; |
| 82 | for (auto&& i : pattern) { |
| 83 | int axis, factor; |
| 84 | bool mul; |
| 85 | std::tie(axis, factor, mul) = i; |
| 86 | if (axis >= 0) { |
| 87 | if (mul) |
| 88 | axs.emplace_back(sub(axis) * factor); |
| 89 | else |
| 90 | axs.emplace_back(sub(axis) / factor); |
| 91 | } else { |
| 92 | axs.emplace_back(cv(factor)); |
| 93 | } |
| 94 | } |
| 95 | auto tshp = opr::Concat::make(axs, 0); |
| 96 | return tshp.node(); |
| 97 | }; |
| 98 | auto checker = mixin_emit_checker(pattern); |
| 99 | return std::make_tuple(builder, checker); |
| 100 | } |
| 101 | |
| 102 | // =================== ReshapeEmitter ====================*/ |
| 103 | ReshapeEmitter::EmitResult ReshapeEmitter::emit() const { |