MCPcopy Create free account
hub / github.com/MegEngine/MegEngine / get_workspace_bundle

Method get_workspace_bundle

dnn/src/cuda/warp_perspective/forward.cpp:112–144  ·  view source on GitHub ↗

Source from the content-addressed store, hash-verified

110} // namespace warp_perspective
111
112WorkspaceBundle 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
146WorkspaceBundle WarpPerspectiveForwardImpl::get_workspace_bundle(
147 void* ptr, const TensorLayoutArray& srcs, const TensorLayout& mat,

Callers

nothing calls this directly

Calls 9

get_inner_layoutFunction · 0.85
dist_byteMethod · 0.80
spanMethod · 0.80
paramFunction · 0.50
get_workspaceFunction · 0.50
enumvMethod · 0.45
push_backMethod · 0.45
total_nr_elemsMethod · 0.45
sizeMethod · 0.45

Tested by

no test coverage detected