| 54 | } |
| 55 | |
| 56 | void CellWeightsWrapperBase::backward( |
| 57 | Handle* handle, param::RNNCell::NonlineMode nonlineMode, _megdnn_tensor_in x, |
| 58 | const TensorNDArray& states, _megdnn_tensor_in y, const TensorNDArray& douts, |
| 59 | _megdnn_tensor_out dx, TensorNDArray& dstates, _megdnn_tensor_out dwi, |
| 60 | _megdnn_tensor_out dwh, _megdnn_tensor_out dbias, |
| 61 | _megdnn_workspace workspace) const { |
| 62 | auto dy = douts[0]; |
| 63 | using NonlineMode = param::RNNCell::NonlineMode; |
| 64 | using Mode = Elemwise::Mode; |
| 65 | auto elemwise_opr = handle->create_operator<ElemwiseForward>(); |
| 66 | TensorND tmp = {workspace.raw_ptr, dy.layout}; |
| 67 | auto new_workspace = Workspace( |
| 68 | workspace.raw_ptr + tmp.layout.span().dist_byte(), |
| 69 | workspace.size - tmp.layout.span().dist_byte()); |
| 70 | switch (nonlineMode) { |
| 71 | case (NonlineMode::IDENTITY): |
| 72 | memcpy(tmp.raw_ptr(), dy.raw_ptr(), dy.layout.span().dist_byte()); |
| 73 | break; |
| 74 | case (NonlineMode::TANH): |
| 75 | elemwise_opr->param().mode = Mode::TANH_GRAD; |
| 76 | elemwise_opr->exec({y, dy}, tmp); |
| 77 | break; |
| 78 | case (NonlineMode::RELU): |
| 79 | elemwise_opr->param().mode = Mode::SWITCH_GT0; |
| 80 | elemwise_opr->exec({y, dy}, tmp); |
| 81 | break; |
| 82 | } |
| 83 | auto matrixmul_opr = handle->create_operator<MatrixMulForward>(); |
| 84 | matrixmul_opr->param().transposeA = false; |
| 85 | matrixmul_opr->param().transposeB = false; |
| 86 | // dx |
| 87 | matrixmul_opr->exec(tmp, this->weight_ih, dx, new_workspace); |
| 88 | // dhx |
| 89 | matrixmul_opr->exec(tmp, this->weight_hh, dstates[0], new_workspace); |
| 90 | // dwi |
| 91 | matrixmul_opr->param().transposeA = true; |
| 92 | matrixmul_opr->exec(tmp, x, dwi, new_workspace); |
| 93 | // dwh |
| 94 | matrixmul_opr->exec(tmp, states[0], dwh, new_workspace); |
| 95 | // dbias |
| 96 | auto sum_opr = handle->create_operator<ReduceForward>(); |
| 97 | sum_opr->param().mode = ReduceForward::Mode::SUM; |
| 98 | sum_opr->param().axis = 0; |
| 99 | TensorND dbias_expanded = { |
| 100 | dbias.raw_ptr(), {{1, dbias.layout.shape[0]}, dbias.layout.dtype}}; |
| 101 | sum_opr->exec(tmp, dbias_expanded, new_workspace); |
| 102 | } |
| 103 | |
| 104 | size_t CellWeightsWrapperBase::backward_workspace_size_in_bytes( |
| 105 | Handle* handle, size_t batch_size, size_t hidden_size, size_t input_size, |