| 15 | class PReluGrad : public OpGrad { |
| 16 | public: |
| 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 | }; |
| 46 | class ReluGrad : public OpGrad { |