| 319 | |
| 320 | template <> |
| 321 | uint64_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 | |
| 341 | template <> |
| 342 | uint64_t opr_footprint_func<opr::LocalShareBackwardFilter>(cg::OperatorNodeBase* opr) { |