12 const int axis = sn.
i0;
13 int outer = 1,
reduce = 1, inner = 1;
14 for (
int k = 0; k < axis; ++k) outer *= X.dims[k];
16 for (
int k = axis + 1; k < X.rank; ++k) inner *= X.dims[k];
17 const int rows = outer * inner;
18 const bool logMode = grp.
logMode;
20 std::ostringstream os1, os2;
26 os1 <<
"@compute @workgroup_size(workgroupX, workgroupY)\nfn main(@builtin(global_invocation_id) globalId: "
27 "vec3<u32>, @builtin(local_invocation_id) localId: vec3<u32>, @builtin(workgroup_id) groupId: vec3<u32>) "
29 os1 <<
" var i_: u32 = globalId.x;\n";
30 os1 <<
" if (i_ >= " <<
rows <<
"u) { return; }\n";
31 os1 <<
" var o_: u32 = i_ / " << inner <<
"u;\n";
32 os1 <<
" var ii: u32 = i_ % " << inner <<
"u;\n";
33 os1 <<
" var m: f32 = -3.402823e38;\n";
34 os1 <<
" for (var j: u32 = 0u; j < " <<
reduce <<
"u; j++) {\n";
35 os1 <<
" m = max(m, in_[(o_ * " <<
reduce <<
"u + j) * " << inner <<
"u + ii]);\n";
37 os1 <<
" var s: f32 = 0.0;\n";
38 os1 <<
" for (var j: u32 = 0u; j < " <<
reduce <<
"u; j++) {\n";
39 os1 <<
" s += exp(in_[(o_ * " <<
reduce <<
"u + j) * " << inner <<
"u + ii] - m);\n";
41 os1 <<
" mx[i_] = m;\n";
42 os1 <<
" sm[i_] = s;\n";
51 os2 <<
"@compute @workgroup_size(workgroupX, workgroupY)\nfn main(@builtin(global_invocation_id) globalId: "
52 "vec3<u32>, @builtin(local_invocation_id) localId: vec3<u32>, @builtin(workgroup_id) groupId: vec3<u32>) "
54 os2 <<
" var i_: u32 = globalId.x;\n";
55 os2 <<
" if (i_ >= " << sn.
size <<
"u) { return; }\n";
56 os2 <<
" var o_: u32 = (i_ / (" <<
reduce * inner <<
"u)) * " << inner <<
"u + (i_ % " << inner <<
"u);\n";
57 os2 <<
" var m: f32 = mx[o_];\n";
58 os2 <<
" var s: f32 = sm[o_];\n";
59 os2 <<
" var x: f32 = in_[i_];\n";
61 os2 <<
" o[i_] = (x - m) - log(s);\n";
63 os2 <<
" o[i_] = exp(x - m) / s;\n";
66 out.
pass1 = os1.str();
67 out.
pass2 = os2.str();
81 const int cols = X.dims[X.rank - 1];
83 const float eps = nn.
s0;
85 const bool hasBias = !rms && grp.
hasBias;
86 const int inputCount = 1 + (hasScale ? 1 : 0) + (hasBias ? 1 : 0);
87 const int statsCount = rms ? 1 : 2;
88 const int outBinding = inputCount + statsCount;
90 std::ostringstream os1, os2;
94 if (statsCount > 1) os1 <<
bufferDecl(2,
"st1");
96 os1 <<
"@compute @workgroup_size(workgroupX, workgroupY)\nfn main(@builtin(global_invocation_id) globalId: "
97 "vec3<u32>, @builtin(local_invocation_id) localId: vec3<u32>, @builtin(workgroup_id) groupId: vec3<u32>) "
99 os1 <<
" var i_: u32 = globalId.x;\n";
100 os1 <<
" if (i_ >= " <<
rows <<
"u) { return; }\n";
101 os1 <<
" var s0: f32 = 0.0;\n";
102 if (statsCount > 1) os1 <<
" var s1: f32 = 0.0;\n";
103 os1 <<
" for (var j: u32 = 0u; j < " <<
cols <<
"u; j++) {\n";
104 os1 <<
" var v: f32 = in_[i_ * " <<
cols <<
"u + j];\n";
106 os1 <<
" s0 += v * v;\n";
108 os1 <<
" s0 += v; s1 += v * v;\n";
111 os1 <<
" st0[i_] = s0;\n";
112 if (statsCount > 1) os1 <<
" st1[i_] = s1;\n";
116 for (
int k = 0; k < inputCount; ++k) os2 <<
bufferDecl(k, (
"a" + std::to_string(k)).c_str());
118 if (statsCount > 1) os2 <<
bufferDecl(inputCount + 1,
"st1");
121 os2 <<
"@compute @workgroup_size(workgroupX, workgroupY)\nfn main(@builtin(global_invocation_id) globalId: "
122 "vec3<u32>, @builtin(local_invocation_id) localId: vec3<u32>, @builtin(workgroup_id) groupId: vec3<u32>) "
124 os2 <<
" var i_: u32 = globalId.x;\n";
125 os2 <<
" if (i_ >= " << nn.
size <<
"u) { return; }\n";
126 os2 <<
" var r: u32 = i_ / " <<
cols <<
"u;\n";
127 os2 <<
" var c: u32 = i_ % " <<
cols <<
"u;\n";
129 os2 <<
" var inv: f32 = 1.0 / sqrt(st0[r] / " <<
cols <<
".0 + " <<
scalarStr(eps) <<
");\n";
130 os2 <<
" var y: f32 = a0[i_] * inv;\n";
131 if (hasScale) os2 <<
" y *= a1[c];\n";
133 os2 <<
" var mean: f32 = st0[r] / " <<
cols <<
".0;\n";
134 os2 <<
" var variance: f32 = st1[r] / " <<
cols <<
".0 - mean * mean;\n";
135 os2 <<
" variance = max(variance, 0.0);\n";
136 os2 <<
" var inv: f32 = 1.0 / sqrt(variance + " <<
scalarStr(eps) <<
");\n";
137 os2 <<
" var y: f32 = (a0[i_] - mean) * inv;\n";
138 if (hasScale) os2 <<
" y *= a1[c];\n";
139 if (hasBias) os2 <<
" y += a" << (hasScale ? 2 : 1) <<
"[c];\n";
141 os2 <<
" o[i_] = y;\n";
143 out.
pass1 = os1.str();
144 out.
pass2 = os2.str();
158 const int axis = rn.
i0;
159 int outer = 1,
reduce = 1, inner = 1;
160 for (
int k = 0; k < axis; ++k) outer *= X.dims[k];
162 for (
int k = axis + 1; k < X.rank; ++k) inner *= X.dims[k];
163 const int outSize = outer * inner;
164 std::ostringstream os;
169 os <<
"@compute @workgroup_size(workgroupX, workgroupY)\nfn main(@builtin(global_invocation_id) globalId: "
170 "vec3<u32>, @builtin(local_invocation_id) localId: vec3<u32>, @builtin(workgroup_id) groupId: vec3<u32>) {\n";
171 os <<
" var i_: u32 = globalId.x;\n";
172 os <<
" if (i_ >= " << outSize <<
"u) { return; }\n";
173 os <<
" var o_: u32 = i_ / " << inner <<
"u;\n";
174 os <<
" var ii: u32 = i_ % " << inner <<
"u;\n";
176 os <<
" var best: f32 = -3.402823e38;\n";
177 os <<
" var bestJ: f32 = 0.0;\n";
178 os <<
" for (var j: u32 = 0u; j < " <<
reduce <<
"u; j++) {\n";
179 os <<
" var v: f32 = in_[(o_ * " <<
reduce <<
"u + j) * " << inner <<
"u + ii];\n";
180 os <<
" if (v > best) { best = v; bestJ = f32(j); }\n";
182 os <<
" o[i_] = bestJ;\n";
187 os <<
" var acc: f32 = 0.0;\n";
188 os <<
" for (var j: u32 = 0u; j < " <<
reduce <<
"u; j++) { acc += in_[(o_ * " <<
reduce <<
"u + j) * "
189 << inner <<
"u + ii]; }\n";
193 os <<
" var acc: f32 = 3.402823e38;\n";
194 os <<
" for (var j: u32 = 0u; j < " <<
reduce <<
"u; j++) { acc = min(acc, in_[(o_ * " <<
reduce
195 <<
"u + j) * " << inner <<
"u + ii]); }\n";
198 os <<
" var acc: f32 = -3.402823e38;\n";
199 os <<
" for (var j: u32 = 0u; j < " <<
reduce <<
"u; j++) { acc = max(acc, in_[(o_ * " <<
reduce
200 <<
"u + j) * " << inner <<
"u + ii]); }\n";
202 default:
throw eve::Exception(
"Tensor WGSL: unsupported kernel variant or binding count");
204 os <<
" o[i_] = acc;\n";
208 out.
pass2 = os.str();
EVENGINE_API_FOUNDATION public API.
EVENGINE_API_DOMAINS public API.
std::string scalarStr(float v)
Scalar str.
std::string bufferDecl(int binding, const char *name)
Buffer decl.
std::string pushConstant()
Pushes constant.
void genSoftmax(const Graph &, const FusedGroup &, KernelSpec &)
Gen softmax.
int groupsFor(int count)
Groups for.
void genReduceOrArgmax(const Graph &, const FusedGroup &, bool argmax, KernelSpec &)
Gen reduce or argmax.
void genNorm(const Graph &, const FusedGroup &, bool rms, KernelSpec &)
Gen norm.
std::string header(int localX, int localY)
Header.
gpgpu::ComputeShader * reduce