MCPcopy Create free account
hub / github.com/MegEngine/MegEngine / LocalShareBackwardData>

Method LocalShareBackwardData>

src/plugin/impl/opr_footprint.cpp:321–339  ·  view source on GitHub ↗

Source from the content-addressed store, hash-verified

319
320template <>
321uint64_t opr_footprint_func<opr::LocalShareBackwardData>(cg::OperatorNodeBase* opr) {
322 mgb_assert(
323 opr->input().size() == 3,
324 "LocalShareBackwardData opr should have three inputs");
325 auto&& filter_shape = opr->input()[0]->shape();
326 auto&& diff_shape = opr->input()[1]->shape();
327 auto&& grad_shape = opr->output()[0]->shape();
328 using Param = opr::LocalShareForward::Param;
329 auto&& param = opr->cast_final_safe<opr::LocalShareBackwardData>().param();
330 mgb_assert(param.format == Param::Format::NCHW);
331 size_t groups = 1;
332 size_t kern_spatial_pos = 3;
333 if (param.sparse == Param::Sparse::GROUP) {
334 groups = filter_shape[0];
335 kern_spatial_pos = 4;
336 }
337 size_t fh = filter_shape[kern_spatial_pos], fw = filter_shape[kern_spatial_pos + 1];
338 return diff_shape.total_nr_elems() * fh * fw * grad_shape[1] * 2 / groups;
339}
340
341template <>
342uint64_t opr_footprint_func<opr::LocalShareBackwardFilter>(cg::OperatorNodeBase* opr) {

Callers

nothing calls this directly

Calls 6

sizeMethod · 0.45
inputMethod · 0.45
shapeMethod · 0.45
outputMethod · 0.45
paramMethod · 0.45
total_nr_elemsMethod · 0.45

Tested by

no test coverage detected