| 300 | |
| 301 | template <> |
| 302 | uint64_t opr_footprint_func<opr::LocalShareForward>(cg::OperatorNodeBase* opr) { |
| 303 | mgb_assert(opr->input().size() == 2, "LocalShare opr should have two inputs"); |
| 304 | auto&& out_shape = opr->output()[0]->shape(); |
| 305 | auto&& src_shape = opr->input()[0]->shape(); |
| 306 | auto&& filter_shape = opr->input()[1]->shape(); |
| 307 | using Param = opr::LocalShareForward::Param; |
| 308 | auto&& param = opr->cast_final_safe<opr::LocalShareForward>().param(); |
| 309 | mgb_assert(param.format == Param::Format::NCHW); |
| 310 | size_t groups = 1; |
| 311 | size_t kern_spatial_pos = 3; |
| 312 | if (param.sparse == Param::Sparse::GROUP) { |
| 313 | groups = filter_shape[0]; |
| 314 | kern_spatial_pos = 4; |
| 315 | } |
| 316 | size_t fh = filter_shape[kern_spatial_pos], fw = filter_shape[kern_spatial_pos + 1]; |
| 317 | return out_shape.total_nr_elems() * fh * fw * src_shape[1] * 2 / groups; |
| 318 | } |
| 319 | |
| 320 | template <> |
| 321 | uint64_t opr_footprint_func<opr::LocalShareBackwardData>(cg::OperatorNodeBase* opr) { |