载入中...
搜索中...
未找到
KernelGenResize.cpp
浏览该文件的文档.
1#include <sstream>
3
5void genResize2d(const Graph &g, const FusedGroup &grp, KernelSpec &out) {
6 const GraphNode &rn = g.node(grp.outputNode);
7 const GraphNode &X = g.node(rn.in0);
8 const int H = X.dims[2], W = X.dims[3];
9 const int OH = rn.dims[2], OW = rn.dims[3];
10 const bool nearest = rn.i0 == 0;
11 std::ostringstream os;
12 os << header(kLocalSize);
13 os << bufferDecl(0, "in_");
14 os << bufferDecl(1, "o");
15 os << pushConstant();
16 os << "void main() {\n";
17 os << " uint i_ = gl_GlobalInvocationID.x;\n";
18 os << " if (i_ >= " << rn.size << "u) return;\n";
19 os << " uint ow = i_ % " << OW << "u;\n";
20 os << " uint rem = i_ / " << OW << "u;\n";
21 os << " uint oh = rem % " << OH << "u;\n";
22 os << " uint rem2 = rem / " << OH << "u;\n";
23 os << " uint c = rem2 % " << X.dims[1] << "u;\n";
24 os << " uint n_ = rem2 / " << X.dims[1] << "u;\n";
25 os << " uint base = (n_ * " << X.dims[1] << "u + c) * " << H << "u * " << W << "u;\n";
26 if (nearest) {
27 os << " uint ih = uint(float(oh) * " << scalarStr(float(H) / OH) << ") ;\n";
28 os << " uint iw = uint(float(ow) * " << scalarStr(float(W) / OW) << ") ;\n";
29 os << " ih = min(ih, " << H - 1 << "u); iw = min(iw, " << W - 1 << "u);\n";
30 os << " o[i_] = in_[base + ih * " << W << "u + iw];\n";
31 } else {
32 os << " float fx = (float(ow) + 0.5) * " << scalarStr(float(W) / OW) << " - 0.5;\n";
33 os << " float fy = (float(oh) + 0.5) * " << scalarStr(float(H) / OH) << " - 0.5;\n";
34 os << " fx = clamp(fx, 0.0, " << scalarStr(float(W - 1)) << ");\n";
35 os << " fy = clamp(fy, 0.0, " << scalarStr(float(H - 1)) << ");\n";
36 os << " uint x0 = uint(floor(fx)); uint y0 = uint(floor(fy));\n";
37 os << " uint x1 = min(x0 + 1u, " << W - 1 << "u); uint y1 = min(y0 + 1u, " << H - 1 << "u);\n";
38 os << " float w00 = in_[base + y0 * " << W << "u + x0];\n";
39 os << " float w10 = in_[base + y0 * " << W << "u + x1];\n";
40 os << " float w01 = in_[base + y1 * " << W << "u + x0];\n";
41 os << " float w11 = in_[base + y1 * " << W << "u + x1];\n";
42 os << " float top = w00 + (w10 - w00) * (fx - float(x0));\n";
43 os << " float bot = w01 + (w11 - w01) * (fx - float(x0));\n";
44 os << " o[i_] = top + (bot - top) * (fy - float(y0));\n";
45 }
46 os << "}\n";
47 out.pass1.clear();
48 out.pass2 = os.str();
49 out.groupsX2 = groupsFor(rn.size);
50 out.inputCount = 1;
51 return;
52}
53
54
55} // namespace eve::tensor::glsl_detail
tensor::Graph g
Definition GpuGraph.cpp:7
EVENGINE_API_DOMAINS public API.
Definition Graph.h:111
std::string header(int localX, int localY)
Header.
Definition KernelGen.cpp:13
std::string bufferDecl(int binding, const char *name)
Buffer decl.
Definition KernelGen.cpp:22
int groupsFor(int count)
Groups for.
std::string scalarStr(float v)
Scalar str.
void genResize2d(const Graph &, const FusedGroup &, KernelSpec &)
Gen resize 2 d.
std::string pushConstant()
Pushes constant.
FusedGroup public API.
Definition Optimizer.h:46
GraphNode public API.
Definition Graph.h:80
int dims[Tensor::kMaxRank]
Definition Graph.h:82
KernelSpec public API.
Definition KernelGen.h:27