| 15 | class StridedSliceGrad : public OpGrad { |
| 16 | public: |
| 17 | virtual std::vector<Express::VARP> onGrad(Express::EXPRP expr, |
| 18 | const std::vector<Express::VARP>& backwardOutput) override { |
| 19 | auto inputs = expr->inputs(); |
| 20 | std::vector<VARP> res(inputs.size(), nullptr); |
| 21 | auto outputDiff = backwardOutput[0]; |
| 22 | |
| 23 | std::unique_ptr<OpT> forwardOp(expr->get()->UnPack()); |
| 24 | int beginMask = forwardOp->main.AsStridedSliceParam()->beginMask; |
| 25 | int endMask = forwardOp->main.AsStridedSliceParam()->endMask; |
| 26 | int ellipsisMask = forwardOp->main.AsStridedSliceParam()->ellipsisMask; |
| 27 | int newAxisMask = forwardOp->main.AsStridedSliceParam()->newAxisMask; |
| 28 | int shrinkAxisMask = forwardOp->main.AsStridedSliceParam()->shrinkAxisMask; |
| 29 | |
| 30 | auto input = inputs[0]; |
| 31 | auto begin = inputs[1]; |
| 32 | auto end = inputs[2]; |
| 33 | auto stride = inputs[3]; |
| 34 | |
| 35 | auto zeros = _ZerosLike(input); |
| 36 | res[0] = _StridedSliceWrite(zeros, begin, end, stride, outputDiff, beginMask, endMask, ellipsisMask, newAxisMask, shrinkAxisMask); |
| 37 | |
| 38 | return res; |
| 39 | } |
| 40 | }; |
| 41 | |
| 42 | static void _create() { |
nothing calls this directly
no test coverage detected