| 127 | } |
| 128 | |
| 129 | size_t ConvolutionImpl::get_workspace_in_bytes( |
| 130 | const TensorLayout& src, const TensorLayout& filter, const TensorLayout& dst, |
| 131 | const PreprocessedFilter* preprocessed_filter) { |
| 132 | TensorLayoutArray layouts{src, filter, dst}; |
| 133 | AlgorithmCache::Key key{this->handle(), this->get_opr_type(), |
| 134 | layouts.data(), layouts.size(), |
| 135 | &this->param(), sizeof(this->param())}; |
| 136 | auto rst = AlgorithmCache::instance().get(key); |
| 137 | if (rst.policy.algo.valid()) { |
| 138 | return rst.workspace; |
| 139 | } |
| 140 | |
| 141 | auto fparam = make_ncb_kern_size_param(src, filter, dst, preprocessed_filter); |
| 142 | auto&& algo = get_algorithm(fparam); |
| 143 | if (is_naive_algo(algo)) { |
| 144 | return naive::ConvolutionForwardImpl::get_workspace_in_bytes( |
| 145 | src, filter, dst, preprocessed_filter); |
| 146 | } else { |
| 147 | return NCB_ALGO_FUNC(get_workspace, algo, fparam); |
| 148 | } |
| 149 | } |
| 150 | |
| 151 | size_t ConvolutionImpl::get_preprocess_workspace_in_bytes( |
| 152 | const TensorLayout& src, const TensorLayout& filter, const TensorLayout& dst) { |
no test coverage detected