| 110 | } // namespace warp_perspective |
| 111 | |
| 112 | WorkspaceBundle WarpPerspectiveForwardImpl::get_workspace_bundle( |
| 113 | void* ptr, const TensorLayout& src, const TensorLayout& mat, |
| 114 | const TensorLayout& mat_idx, const TensorLayout& dst) const { |
| 115 | MEGDNN_MARK_USED_VAR(mat_idx); |
| 116 | SmallVector<size_t> sizes; |
| 117 | TensorLayout fsrc = src; |
| 118 | TensorLayout fmat = mat; |
| 119 | TensorLayout fdst = dst; |
| 120 | if ((src.dtype.enumv() == DTypeEnum::QuantizedS4 || |
| 121 | src.dtype.enumv() == DTypeEnum::Quantized4Asymm) && |
| 122 | param().format == param::WarpPerspective::Format::NCHW) { |
| 123 | get_inner_layout(src, dst, fsrc, fdst, handle(), param().format); |
| 124 | sizes.push_back(fsrc.span().dist_byte()); |
| 125 | sizes.push_back(fdst.span().dist_byte()); |
| 126 | } else { |
| 127 | auto get_workspace = [&sizes](TensorLayout& layout) { |
| 128 | if (layout.dtype == dtype::BFloat16()) { |
| 129 | layout.dtype = dtype::Float32(); |
| 130 | sizes.push_back(layout.span().dist_byte()); |
| 131 | } |
| 132 | }; |
| 133 | get_workspace(fsrc); |
| 134 | get_workspace(fmat); |
| 135 | get_workspace(fdst); |
| 136 | } |
| 137 | if (param().format == param::WarpPerspective::Format::NHWC) { |
| 138 | //! use double for the workspace dtype as float may cause |
| 139 | //! accuracy problems |
| 140 | sizes.push_back(mat.total_nr_elems() * sizeof(double)); |
| 141 | } |
| 142 | |
| 143 | return {ptr, std::move(sizes)}; |
| 144 | } |
| 145 | |
| 146 | WorkspaceBundle WarpPerspectiveForwardImpl::get_workspace_bundle( |
| 147 | void* ptr, const TensorLayoutArray& srcs, const TensorLayout& mat, |
nothing calls this directly
no test coverage detected