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;
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";
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";
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";
EVENGINE_API_DOMAINS public API.
std::string header(int localX, int localY)
Header.
std::string bufferDecl(int binding, const char *name)
Buffer decl.
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.
int dims[Tensor::kMaxRank]