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

Method onGrad

tools/train/source/grad/ReluGrad.cpp:17–43  ·  view source on GitHub ↗

Source from the content-addressed store, hash-verified

15class PReluGrad : public OpGrad {
16public:
17 virtual std::vector<Express::VARP> onGrad(Express::EXPRP expr,
18 const std::vector<Express::VARP>& backwardOutput) override {
19 std::vector<Express::VARP> result(1, nullptr);
20 auto op = expr->get();
21 auto input = expr->inputs()[0];
22 auto mask = _Relu(_Sign(input));
23 auto prelu = op->main_as_PRelu();
24 if (prelu->slope()->size() == 1) {
25 auto slope = prelu->slope()->data()[0];
26 result[0] = (mask + (_Scalar<float>(1.0f) - mask) * _Scalar<float>(slope)) * backwardOutput[0];
27 return result;
28 }
29 auto channel = prelu->slope()->size();
30 std::vector<float> scale(channel);
31 ::memcpy(scale.data(), prelu->slope()->data(), channel * sizeof(float));
32 std::vector<float> bias(channel, 0.0f);
33 auto outputSecond = _Scale(backwardOutput[0], channel, std::move(scale), std::move(bias));
34 result[0] = mask * backwardOutput[0] + (_Scalar<float>(1.0f) - mask) * outputSecond;
35// auto diffInfo = result[0]->getInfo();
36// auto inputInfo = input->getInfo();
37// for (int i=0; i<diffInfo->dim.size(); ++i) {
38// MNN_ASSERT(diffInfo->dim[i] == inputInfo->dim[i]);
39// MNN_PRINT("%s, %d, %d - %d\n", expr->name().c_str(), i, diffInfo->dim[i], inputInfo->dim[i]);
40// }
41// MNN_ASSERT(diffInfo->order == inputInfo->order);
42 return result;
43 }
44
45};
46class ReluGrad : public OpGrad {

Callers

nothing calls this directly

Calls 6

_ReluFunction · 0.85
_SignFunction · 0.85
_ScaleFunction · 0.85
getMethod · 0.45
sizeMethod · 0.45
dataMethod · 0.45

Tested by

no test coverage detected