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

Method backward

imperative/src/impl/transformations/grad.cpp:206–263  ·  view source on GitHub ↗

Source from the content-addressed store, hash-verified

204}
205
206void GradKey::backward() {
207 mgb_assert(m_frozen);
208 auto& tape = m_frozen_tape;
209 for (std::ptrdiff_t k = tape.size() - 1; k >= 0; --k) {
210 auto& [grad_fn, op] = tape[k];
211 std::string scope_name = op ? op->make_name() + ".Backward" : "CustomBackward";
212 Transformation::push_scope(scope_name);
213 auto grad_receiver = [&, grad_fn = grad_fn](size_t i, ValueRef grad) {
214 auto& dest = grad_fn->m_dests[i];
215 if (dest) {
216 auto& existing_grad = dest->m_grad;
217 if (!existing_grad) {
218 existing_grad = grad;
219 } else {
220 existing_grad = imperative::apply(
221 ApplyOp(*Elemwise::make(Elemwise::Mode::ADD)),
222 existing_grad, grad)[0];
223 }
224 }
225 };
226 // clang-format off
227 std::visit([&, grad_fn = grad_fn, op = op](auto&& backward) {
228 using T = std::decay_t<decltype(backward)>;
229 if constexpr (std::is_same_v<T, std::monostate>) {
230 mgb_throw(AssertionError, "invalid backward");
231 } else {
232 // mgb_assert(grad_fn->m_slots.size() > 0);
233 SmallVector<ValueRef> grads (grad_fn->m_slots.size());
234 auto iter = grads.begin();
235 for (auto&& slot : grad_fn->m_slots) {
236 *iter++ = slot.m_grad;
237 }
238 if (Profiler::is_profiling()) {
239 imperative::apply(PushScope(scope_name, ScopeType::BACKWARD), Span<ValueRef>(nullptr, nullptr));
240 }
241 backward(grads, grad_receiver);
242 if (Profiler::is_profiling()) {
243 imperative::apply(PopScope(scope_name, ScopeType::BACKWARD), Span<ValueRef>(nullptr, nullptr));
244 }
245 }
246 }, grad_fn->m_backward);
247 // clang-format on
248 for (auto&& dest : grad_fn->m_dests) {
249 if (!dest) {
250 continue;
251 }
252 if (!dest.m_producer_record.next && dest->callback) {
253 // I'm the last grad producer, invoke callback
254 if (dest->m_grad) {
255 dest->callback(dest->m_grad);
256 }
257 }
258 }
259 grad_fn->clear();
260 Transformation::pop_scope(scope_name);
261 }
262 tape.clear();
263}

Callers 3

make_backward_graphMethod · 0.45
make_backward_closureMethod · 0.45

Calls 12

is_profilingFunction · 0.85
applyFunction · 0.50
ApplyOpClass · 0.50
makeFunction · 0.50
PushScopeClass · 0.50
backwardFunction · 0.50
PopScopeClass · 0.50
sizeMethod · 0.45
make_nameMethod · 0.45
beginMethod · 0.45
callbackMethod · 0.45
clearMethod · 0.45

Tested by

no test coverage detected