| 15 | class RasterGrad : 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> result(inputs.size(), nullptr); |
| 21 | auto rasterInfo = expr->get()->main_as_Extra(); |
| 22 | const int32_t* regionData = nullptr; |
| 23 | const int REGION_LENGTH = 11; |
| 24 | std::vector<int32_t> regionDataHolder; |
| 25 | if (nullptr != rasterInfo && nullptr != rasterInfo->attr()) { |
| 26 | for (int i=0; i<rasterInfo->attr()->size(); ++i) { |
| 27 | auto attr = rasterInfo->attr()->GetAs<Attribute>(i); |
| 28 | if (attr->key()->str() == "region") { |
| 29 | regionData = attr->list()->i()->data(); |
| 30 | MNN_ASSERT(inputs.size() * REGION_LENGTH == attr->list()->i()->size()); |
| 31 | break; |
| 32 | } |
| 33 | } |
| 34 | } else { |
| 35 | regionDataHolder.resize(inputs.size() * REGION_LENGTH); |
| 36 | regionData = regionDataHolder.data(); |
| 37 | auto outputTensor = Variable::create(expr)->getTensor(); |
| 38 | auto des = TensorUtils::getDescribe(outputTensor); |
| 39 | MNN_ASSERT(des->regions.size() == inputs.size()); |
| 40 | for (int i=0; i<inputs.size(); ++i) { |
| 41 | auto& r = des->regions[i]; |
| 42 | auto dstPtr = regionDataHolder.data() + REGION_LENGTH * i; |
| 43 | dstPtr[0] = r.src.offset; |
| 44 | dstPtr[1] = r.src.stride[0]; |
| 45 | dstPtr[2] = r.src.stride[1]; |
| 46 | dstPtr[3] = r.src.stride[2]; |
| 47 | |
| 48 | dstPtr[4] = r.dst.offset; |
| 49 | dstPtr[5] = r.dst.stride[0]; |
| 50 | dstPtr[6] = r.dst.stride[1]; |
| 51 | dstPtr[7] = r.dst.stride[2]; |
| 52 | |
| 53 | dstPtr[8] = r.size[0]; |
| 54 | dstPtr[9] = r.size[1]; |
| 55 | dstPtr[10] = r.size[2]; |
| 56 | } |
| 57 | } |
| 58 | for (int i=0; i<inputs.size(); ++i) { |
| 59 | auto regionInfo = regionData + REGION_LENGTH * i; |
| 60 | auto info = inputs[i]->getInfo(); |
| 61 | auto shape = info->dim; |
| 62 | std::vector<int> curRegion(REGION_LENGTH); |
| 63 | ::memcpy(curRegion.data(), regionInfo, REGION_LENGTH * sizeof(int)); |
| 64 | ::memcpy(curRegion.data(), regionInfo + 4, 4 * sizeof(int)); |
| 65 | ::memcpy(curRegion.data() + 4, regionInfo, 4 * sizeof(int)); |
| 66 | auto grad = _RasterRaw({backwardOutput[0]}, curRegion, shape, info->type, info->order); |
| 67 | result[i] = grad; |
| 68 | } |
| 69 | return result; |
| 70 | } |
| 71 | }; |
| 72 | |
| 73 | static void _create() { |