MCPcopy Create free account
hub / github.com/alibaba/MNN / onGrad

Method onGrad

tools/train/source/grad/RasterGrad.cpp:17–70  ·  view source on GitHub ↗

Source from the content-addressed store, hash-verified

15class RasterGrad : public OpGrad {
16public:
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
73static void _create() {

Callers

nothing calls this directly

Calls 11

_RasterRawFunction · 0.85
createFunction · 0.50
sizeMethod · 0.45
getMethod · 0.45
strMethod · 0.45
keyMethod · 0.45
dataMethod · 0.45
iMethod · 0.45
resizeMethod · 0.45
getTensorMethod · 0.45
getInfoMethod · 0.45

Tested by

no test coverage detected