| 10 | namespace cuda { |
| 11 | |
| 12 | void LocalForwardImpl::exec( |
| 13 | _megdnn_tensor_in src, _megdnn_tensor_in filter, _megdnn_tensor_out dst, |
| 14 | _megdnn_workspace workspace) { |
| 15 | megdnn_assert( |
| 16 | src.layout.dtype == dtype::Float32(), |
| 17 | "cuda do not support fp16 local operator"); |
| 18 | check_exec(src.layout, filter.layout, dst.layout, workspace.size); |
| 19 | bool is_xcorr = param().mode == Mode::CROSS_CORRELATION; |
| 20 | auto N = src.layout.shape[0], IC = src.layout.shape[1], IH = src.layout.shape[2], |
| 21 | IW = src.layout.shape[3]; |
| 22 | auto OC = dst.layout.shape[1], OH = dst.layout.shape[2], OW = dst.layout.shape[3]; |
| 23 | auto FH = filter.layout.shape[3], FW = filter.layout.shape[4]; |
| 24 | auto handle = concrete_handle(this->handle()); |
| 25 | auto stream = cuda_stream(this->handle()); |
| 26 | auto cublas = cublas_handle(this->handle()); |
| 27 | auto one = handle->one_device(); |
| 28 | auto zero = handle->zero_device(); |
| 29 | size_t src_batch_strd = src.layout.stride[0]; |
| 30 | size_t dst_batch_strd = dst.layout.stride[0]; |
| 31 | if (use_cuda_convnet(src.layout, filter.layout, dst.layout)) { |
| 32 | local::forward_proxy_convnet( |
| 33 | src.ptr<dt_float32>(), filter.ptr<dt_float32>(), dst.ptr<dt_float32>(), |
| 34 | reinterpret_cast<float*>(workspace.raw_ptr), N, IC, IH, IW, OC, OH, OW, |
| 35 | FH, FW, src_batch_strd, dst_batch_strd, param().pad_h, param().pad_w, |
| 36 | param().stride_h, param().stride_w, cublas, stream, one, zero); |
| 37 | } else if ( |
| 38 | local::forward_proxy_default_share_mem_in_bytes(IH, IW) <= |
| 39 | handle->device_prop().sharedMemPerBlock) { |
| 40 | local::forward_proxy_default( |
| 41 | src.ptr<dt_float32>(), filter.ptr<dt_float32>(), dst.ptr<dt_float32>(), |
| 42 | N, IC, IH, IW, OC, OH, OW, FH, FW, src_batch_strd, dst_batch_strd, |
| 43 | param().pad_h, param().pad_w, param().stride_h, param().stride_w, |
| 44 | is_xcorr, stream); |
| 45 | } else { |
| 46 | megdnn_throw(ssprintf( |
| 47 | "No usable kernel for local conv, src: %s filter: %s \n", |
| 48 | src.layout.to_string().c_str(), filter.layout.to_string().c_str())); |
| 49 | } |
| 50 | } |
| 51 | |
| 52 | size_t LocalForwardImpl::get_workspace_in_bytes( |
| 53 | const TensorLayout& src, const TensorLayout& filter, const TensorLayout& dst) { |
nothing calls this directly
no test coverage detected