MCPcopy Create free account
hub / github.com/MegEngine/MegEngine / backward

Method backward

dnn/src/naive/rnn/rnn.cpp:56–102  ·  view source on GitHub ↗

Source from the content-addressed store, hash-verified

54}
55
56void 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
104size_t CellWeightsWrapperBase::backward_workspace_size_in_bytes(
105 Handle* handle, size_t batch_size, size_t hidden_size, size_t input_size,

Callers 2

train_funFunction · 0.45
backward_exec_internalFunction · 0.45

Calls 8

WorkspaceClass · 0.85
dist_byteMethod · 0.80
spanMethod · 0.80
backwardFunction · 0.50
raw_ptrMethod · 0.45
paramMethod · 0.45
execMethod · 0.45
deduce_layoutMethod · 0.45

Tested by

no test coverage detected