You can not select more than 25 topics Topics must start with a chinese character,a letter or number, can include dashes ('-') and can be up to 35 characters long.

subgraph_detail.cpp 7.3 kB

123456789101112131415161718192021222324252627282930313233343536373839404142434445464748495051525354555657585960616263646566676869707172737475767778798081828384858687888990919293949596979899100101102103104105106107108109110111112113114115116117118119120121122123124125126127128129130131132133134135136137138139140141142143144145146147148149150151152153154155156157158159160161162163164165166167168
  1. /**
  2. * \file imperative/src/impl/subgraph_detail.cpp
  3. * MegEngine is Licensed under the Apache License, Version 2.0 (the "License")
  4. *
  5. * Copyright (c) 2014-2021 Megvii Inc. All rights reserved.
  6. *
  7. * Unless required by applicable law or agreed to in writing,
  8. * software distributed under the License is distributed on an
  9. * "AS IS" BASIS, WITHOUT ARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
  10. */
  11. #include "megbrain/imperative/subgraph_detail.h"
  12. #include "megbrain/imperative/graph_builder.h"
  13. #include "megbrain/imperative/ops/autogen.h"
  14. #include "megbrain/opr/io.h"
  15. #include "./op_trait.h"
  16. namespace mgb {
  17. namespace imperative {
  18. namespace subgraph_detail {
  19. VarNodeArray apply_on_var_node(const OpDef& def, const VarNodeArray& inputs) {
  20. SmallVector<LogicalTensorDesc> input_descs;
  21. for (auto&& input : inputs) {
  22. input_descs.push_back({TensorLayout{input->dtype()}, input->comp_node()});
  23. }
  24. auto apply_functor = [&](const std::shared_ptr<OpDef>& op,
  25. const VarNodeArray& inputs, size_t nr_outputs) {
  26. op->set_scope(def.scope());
  27. return OpDef::apply_on_var_node(*op, inputs);
  28. };
  29. auto const_functor = [&](const TensorPtr& value) {
  30. return opr::ImmutableTensor::make(*inputs[0]->owner_graph(), value->get_value())
  31. .node();
  32. };
  33. auto subgraph = def.trait()->make_forward_graph(def, input_descs);
  34. auto outputs = subgraph.apply<VarNode*>(inputs, apply_functor, const_functor);
  35. return outputs;
  36. }
  37. std::tuple<SmallVector<LogicalTensorDesc>, bool> infer_output_attrs_fallible(
  38. const OpDef& def, const SmallVector<LogicalTensorDesc>& inputs) {
  39. auto subgraph = def.trait()->make_forward_graph(def, inputs);
  40. bool all_validated = true;
  41. auto apply_functor = [&](const std::shared_ptr<OpDef>& op,
  42. const SmallVector<LogicalTensorDesc>& inputs,
  43. size_t nr_outputs) {
  44. auto [outputs, validated] = OpDef::infer_output_attrs_fallible(*op, inputs);
  45. all_validated = all_validated && validated;
  46. return outputs;
  47. };
  48. auto const_functor = [&](const TensorPtr& value) {
  49. return LogicalTensorDesc{
  50. value->layout(), value->comp_node(),
  51. value->get_value().proxy_to_default_cpu()};
  52. };
  53. auto outputs =
  54. subgraph.apply<LogicalTensorDesc>(inputs, apply_functor, const_functor);
  55. return {outputs, all_validated};
  56. }
  57. SmallVector<TensorPtr> apply_on_physical_tensor(
  58. const OpDef& def, SmallVector<TensorPtr> inputs) {
  59. SmallVector<LogicalTensorDesc> input_descs;
  60. for (auto&& input : inputs) {
  61. input_descs.push_back({input->layout(), input->comp_node()});
  62. }
  63. auto subgraph = def.trait()->make_forward_graph(def, input_descs);
  64. auto apply_functor = [](const std::shared_ptr<OpDef>& op,
  65. const SmallVector<TensorPtr>& inputs, size_t nr_outputs) {
  66. return OpDef::apply_on_physical_tensor(*op, inputs);
  67. };
  68. auto const_functor = [&](const TensorPtr& value) { return value; };
  69. auto outputs = subgraph.apply<TensorPtr>(inputs, apply_functor, const_functor);
  70. return outputs;
  71. }
  72. static EncodedSubgraph make_backward_graph_from_forward(
  73. const SmallVector<LogicalTensorDesc>& inputs,
  74. const SmallVector<bool>& input_requires_grad,
  75. const SmallVector<bool>& output_has_grad, EncodedSubgraph forward_graph) {
  76. using namespace std::placeholders;
  77. using var_t = Subgraph::var_t;
  78. using vars_t = Subgraph::vars_t;
  79. Subgraph::Builder<LogicalTensorDesc> builder(
  80. [](auto&& op, auto&& input_descs, size_t nr_outputs) {
  81. auto [descs, _] = OpDef::infer_output_attrs_fallible(*op, input_descs);
  82. return descs;
  83. });
  84. auto accum_grad = [&](var_t lhs, var_t rhs) {
  85. return builder.write_expr(
  86. Elemwise::make(Elemwise::Mode::ADD), {lhs, rhs}, 1)[0];
  87. };
  88. GradContext<var_t> grad_context{accum_grad};
  89. auto input_vars = builder.write_inputs(inputs);
  90. auto outputs = forward_graph.apply<var_t>(
  91. input_vars, std::bind(&decltype(builder)::write_expr, &builder, _1, _2, _3),
  92. [&](TensorPtr constant) {
  93. return builder.write_constant(
  94. constant, {constant->layout(), constant->comp_node()});
  95. });
  96. size_t nr_outputs = outputs.size();
  97. auto apply_mask = [](auto&& values, SmallVector<bool> mask) {
  98. mgb_assert(mask.size() == values.size());
  99. std::decay_t<decltype(values)> results;
  100. for (size_t i = 0; i < mask.size(); ++i) {
  101. if (mask[i]) {
  102. results.push_back(values[i]);
  103. }
  104. }
  105. return results;
  106. };
  107. grad_context.mark_require_grads(apply_mask(input_vars, input_requires_grad));
  108. builder.iterate([&](std::list<Subgraph::expr_t>::iterator iter) {
  109. grad_context.record_expr(iter->op, iter->inputs, iter->outputs);
  110. });
  111. auto output_descs = builder.get_descs(outputs);
  112. auto computed_outputs = builder.write_inputs(output_descs);
  113. auto output_grads = builder.write_inputs(output_descs);
  114. grad_context.backward(
  115. apply_mask(outputs, output_has_grad),
  116. apply_mask(output_grads, output_has_grad),
  117. [&](Subgraph::expr_t expr, vars_t output_grads) {
  118. auto bg = OpDef::make_backward_graph(
  119. *expr.op, builder.get_descs(expr.inputs),
  120. grad_context.get_require_grads(expr.inputs),
  121. grad_context.get_has_grads(expr.outputs));
  122. if (bg.graph.empty()) {
  123. return vars_t(expr.inputs.size(), 0);
  124. }
  125. vars_t grad_inputs;
  126. grad_inputs.insert(
  127. grad_inputs.end(), expr.inputs.begin(), expr.inputs.end());
  128. grad_inputs.insert(
  129. grad_inputs.end(), expr.outputs.begin(), expr.outputs.end());
  130. grad_inputs.insert(
  131. grad_inputs.end(), output_grads.begin(), output_grads.end());
  132. auto apply_functor =
  133. std::bind(&decltype(builder)::write_expr, &builder, _1, _2, _3);
  134. auto const_functor = [&](TensorPtr constant) {
  135. return builder.write_constant(
  136. constant, {constant->layout(), constant->comp_node()});
  137. };
  138. return bg.apply<var_t>(grad_inputs, apply_functor, const_functor);
  139. });
  140. builder.add_outputs(grad_context.get_grads(input_vars));
  141. for (size_t i = 0; i < nr_outputs; ++i) {
  142. builder.replace_var(outputs[i], computed_outputs[i]);
  143. }
  144. auto backward_graph = builder.encode();
  145. return backward_graph;
  146. }
  147. EncodedSubgraph make_backward_graph(
  148. const OpDef& def, const SmallVector<LogicalTensorDesc>& inputs,
  149. const SmallVector<bool>& input_requires_grad,
  150. const SmallVector<bool>& output_has_grad) {
  151. auto forward_graph = OpDef::make_forward_graph(def, inputs);
  152. return make_backward_graph_from_forward(
  153. inputs, input_requires_grad, output_has_grad, forward_graph);
  154. }
  155. } // namespace subgraph_detail
  156. } // namespace imperative
  157. } // namespace mgb