| 40 | } // namespace |
| 41 | |
| 42 | void GroupLocalForwardImpl::exec( |
| 43 | _megdnn_tensor_in src, _megdnn_tensor_in filter, _megdnn_tensor_out dst, |
| 44 | _megdnn_workspace workspace) { |
| 45 | megdnn_assert( |
| 46 | src.layout.dtype == dtype::Float32(), |
| 47 | "cuda do not support fp16 group local operator"); |
| 48 | check_exec(src.layout, filter.layout, dst.layout, workspace.size); |
| 49 | |
| 50 | auto handle = concrete_handle(this->handle()); |
| 51 | auto G = filter.layout[0]; |
| 52 | auto IH = src.layout.shape[2], IW = src.layout.shape[3], OH = dst.layout.shape[2], |
| 53 | OW = dst.layout.shape[3]; |
| 54 | if (prefer_inference_kernel(src.layout, filter.layout, dst.layout)) { |
| 55 | auto N = src.layout.shape[0], ICg = src.layout.shape[1] / G, |
| 56 | OCg = dst.layout.shape[1] / G; |
| 57 | auto FH = filter.layout.shape[4], FW = filter.layout.shape[5]; |
| 58 | auto PH = param().pad_h, PW = param().pad_w; |
| 59 | auto SH = param().stride_h, SW = param().stride_w; |
| 60 | const float* sptr = src.ptr<dt_float32>(); |
| 61 | const float* fptr = filter.ptr<dt_float32>(); |
| 62 | float* dptr = dst.ptr<dt_float32>(); |
| 63 | float* wptr = workspace.ptr<dt_float32>(); |
| 64 | auto stream = cuda_stream(this->handle()); |
| 65 | |
| 66 | group_local::exec( |
| 67 | sptr, fptr, dptr, wptr, N, ICg, IH, IW, OCg, OH, OW, FH, FW, G, PH, PW, |
| 68 | SH, SW, stream); |
| 69 | } else { |
| 70 | auto&& opr = get_opr(handle, param()); |
| 71 | TensorND src_g = {src.raw_ptr(), prepare_src_dst(src.layout, G)}; |
| 72 | TensorND dst_g = {dst.raw_ptr(), prepare_src_dst(dst.layout, G)}; |
| 73 | TensorND filter_g = {filter.raw_ptr(), prepare_filter(filter.layout)}; |
| 74 | for (size_t g = 0; g < G; ++g) { |
| 75 | opr->exec(src_g, filter_g, dst_g, workspace); |
| 76 | incr_refp( |
| 77 | src_g.get_ref_ptr(), src_g.layout.stride[1] * |
| 78 | src_g.layout.shape[1] * |
| 79 | src_g.layout.dtype.size()); |
| 80 | incr_refp( |
| 81 | dst_g.get_ref_ptr(), dst_g.layout.stride[1] * |
| 82 | dst_g.layout.shape[1] * |
| 83 | dst_g.layout.dtype.size()); |
| 84 | incr_refp(filter_g.get_ref_ptr(), filter_g.layout.span().dist_byte()); |
| 85 | } |
| 86 | } |
| 87 | } |
| 88 | |
| 89 | GroupLocalForwardImpl::GroupLocalForwardImpl(Handle* handle) |
| 90 | : GroupLocalForward(handle) {} |
nothing calls this directly
no test coverage detected