载入中...
搜索中...
未找到
KernelGen.cpp
浏览该文件的文档.
1#include "tensor/CpuKernels.h"
3#include "tensor/Quant.h"
4
5#include <algorithm>
6#include <cmath>
7#include <sstream>
8
9namespace eve::tensor {
10namespace glsl_detail {
11
12
13std::string header(int localX, int localY) {
14 std::ostringstream os;
15 os << "#version 450\n";
16 os << "layout(local_size_x = " << localX;
17 if (localY > 1) os << ", local_size_y = " << localY;
18 os << ") in;\n";
19 return os.str();
20}
21
22std::string bufferDecl(int binding, const char *name) {
23 std::ostringstream os;
24 os << "layout(set = 0, binding = " << binding << ") buffer B" << binding << " { float "
25 << name << "[]; };\n";
26 return os.str();
27}
28
29std::string bufferDeclUint(int binding, const char *name) {
30 std::ostringstream os;
31 os << "layout(set = 0, binding = " << binding << ") buffer B" << binding << " { uint "
32 << name << "[]; };\n";
33 return os.str();
34}
35
41std::string emitQuantizedBVal(DType dt, int group, const char *bufName = "b") {
42 std::ostringstream os;
43 const int g = group > 0 ? group : 1;
44 switch (dt) {
45 case DType::Int8:
46 os << "float bval(uint idx) {\n"
47 << " uint word = " << bufName << "[idx >> 2u];\n"
48 << " uint byte = (word >> ((idx & 3u) * 8u)) & 0xFFu;\n"
49 << " int v = int(byte);\n"
50 << " if (v >= 128) v -= 256;\n"
51 << " return float(v) * bs[idx / " << g << "u];\n"
52 << "}\n";
53 break;
54 case DType::Int4:
55 os << "float bval(uint idx) {\n"
56 << " uint word = " << bufName << "[idx >> 3u];\n"
57 << " uint byte = (word >> (((idx >> 1u) & 3u) * 8u)) & 0xFFu;\n"
58 << " uint nib = (idx & 1u) == 0u ? (byte & 0xFu) : (byte >> 4u);\n"
59 << " int v = int(nib);\n"
60 << " if (v >= 8) v -= 16;\n"
61 << " return float(v) * bs[idx / " << g << "u];\n"
62 << "}\n";
63 break;
64 case DType::Fp16:
65 os << "float bval(uint idx) {\n"
66 << " uint word = " << bufName << "[idx >> 1u];\n"
67 << " uint hb = (word >> ((idx & 1u) * 16u)) & 0xFFFFu;\n"
68 << " return unpackHalf2x16(hb).x;\n"
69 << "}\n";
70 break;
71 case DType::Fp8E4M3:
72 os << "float bval(uint idx) {\n"
73 << " uint word = " << bufName << "[idx >> 2u];\n"
74 << " uint byte = (word >> ((idx & 3u) * 8u)) & 0xFFu;\n"
75 << " int s = (int(byte & 0x80u) != 0) ? -1 : 1;\n"
76 << " int e = int((byte >> 3u) & 0xFu);\n"
77 << " int m = int(byte & 0x7u);\n"
78 << " float v = (e == 0) ? exp2(-6.0) * float(m) / 8.0\n"
79 << " : exp2(float(e - 7)) * (1.0 + float(m) / 8.0);\n"
80 << " return float(s) * v * bs[idx / " << g << "u];\n"
81 << "}\n";
82 break;
83 case DType::Fp4E2M1:
84 os << "float bval(uint idx) {\n"
85 << " uint word = " << bufName << "[idx >> 3u];\n"
86 << " uint byte = (word >> (((idx >> 1u) & 3u) * 8u)) & 0xFFu;\n"
87 << " uint nib = (idx & 1u) == 0u ? (byte & 0xFu) : (byte >> 4u);\n"
88 << " int s = (int(nib & 8u) != 0) ? -1 : 1;\n"
89 << " int e = int((nib >> 1u) & 3u);\n"
90 << " int m = int(nib & 1u);\n"
91 << " float v = (e == 0) ? 0.5 * float(m)\n"
92 << " : exp2(float(e - 1)) * (1.0 + 0.5 * float(m));\n"
93 << " return float(s) * v * bs[idx / " << g << "u];\n"
94 << "}\n";
95 break;
96 default: break;
97 }
98 return os.str();
99}
100
101std::string pushConstant() {
102 return "layout(push_constant) uniform PC { float data[32]; } pc;\n";
103}
104
105int groupsFor(int count) { return (count + kLocalSize - 1) / kLocalSize; }
106
107std::string scalarStr(float v) {
108 if (v == int(v) && std::fabs(v) < 1e9f) return std::to_string(int(v)) + ".0";
109 std::ostringstream os;
110 os << v << "f";
111 return os.str();
112}
113
120 const Graph &graph;
122 const std::vector<std::string> &indexExprs; // per group input
123 std::string rootVar;
124 std::ostringstream &os;
125 int temp = 0;
126 std::string biasIndexExpr; // e.g. "jj" (matmul) / "f" (conv), empty = none
127
128 std::string operand(int nodeId) {
129 if (nodeId == group.biasNode && !biasIndexExpr.empty())
130 return "bias[" + biasIndexExpr + "]";
131 const auto it = std::find(group.inputs.begin(), group.inputs.end(), nodeId);
132 if (it != group.inputs.end()) {
133 const size_t idx = static_cast<size_t>(it - group.inputs.begin());
134 return "a" + std::to_string(idx) + "[" + indexExprs[idx] + "]";
135 }
136 return vars[static_cast<size_t>(nodeId)];
137 }
138
139 std::vector<std::string> vars;
140
141 std::string emit(int nodeId) {
142 const GraphNode &nd = graph.node(nodeId);
143 const std::string x = (nd.in0 == group.nodes.front() && rootVar != "") ? rootVar
144 : (nd.in0 >= 0 ? operand(nd.in0) : rootVar);
145 std::string expr;
146 switch (nd.type) {
147 case OpType::Add: expr = "(" + x + " + " + operand(nd.in1) + ")"; break;
148 case OpType::Sub: expr = "(" + x + " - " + operand(nd.in1) + ")"; break;
149 case OpType::Multiply: expr = "(" + x + " * " + operand(nd.in1) + ")"; break;
150 case OpType::Divide: expr = "(" + x + " / " + operand(nd.in1) + ")"; break;
151 case OpType::AddScalar: expr = "(" + x + " + " + scalarStr(nd.s0) + ")"; break;
152 case OpType::SubScalar: expr = "(" + x + " - " + scalarStr(nd.s0) + ")"; break;
153 case OpType::MulScalar: expr = "(" + x + " * " + scalarStr(nd.s0) + ")"; break;
154 case OpType::DivScalar: expr = "(" + x + " / " + scalarStr(nd.s0) + ")"; break;
155 case OpType::PowScalar: expr = "pow(" + x + ", " + scalarStr(nd.s0) + ")"; break;
156 case OpType::Neg: expr = "(-" + x + ")"; break;
157 case OpType::Abs: expr = "abs(" + x + ")"; break;
158 case OpType::Sqrt: expr = "sqrt(" + x + ")"; break;
159 case OpType::Exp: expr = "exp(" + x + ")"; break;
160 case OpType::Log: expr = "log(" + x + ")"; break;
161 case OpType::Sin: expr = "sin(" + x + ")"; break;
162 case OpType::Cos: expr = "cos(" + x + ")"; break;
163 case OpType::Tanh: expr = "tanh(" + x + ")"; break;
164 case OpType::Relu: expr = "max(" + x + ", 0.0)"; break;
165 case OpType::Sigmoid: expr = "(1.0 / (1.0 + exp(-" + x + ")))"; break;
166 case OpType::Gelu:
167 expr = "(0.5 * " + x +
168 " * (1.0 + tanh(0.7978845608028654 * (" + x +
169 " + 0.044715 * " + x + " * " + x + " * " + x + "))))";
170 break;
171 case OpType::Silu: expr = "(" + x + " / (1.0 + exp(-" + x + ")))"; break;
172 case OpType::Clamp: {
173 float lo = nd.s0, hi = nd.s1;
174 if (lo > hi) std::swap(lo, hi);
175 expr = "clamp(" + x + ", " + scalarStr(lo) + ", " + scalarStr(hi) + ")";
176 break;
177 }
178 case OpType::MaximumScalar: expr = "max(" + x + ", " + scalarStr(nd.s0) + ")"; break;
179 case OpType::MinimumScalar: expr = "min(" + x + ", " + scalarStr(nd.s0) + ")"; break;
180 case OpType::Where:
181 expr = "(" + operand(nd.in0) + " > 0.5 ? " + operand(nd.in1) + " : " +
182 operand(nd.in2) + ")";
183 break;
184 default: return "";
185 }
186 const std::string var = "t" + std::to_string(temp++);
187 os << " float " << var << " = " << expr << ";\n";
188 if (vars.size() <= static_cast<size_t>(nodeId)) vars.resize(static_cast<size_t>(nodeId) + 1);
189 vars[static_cast<size_t>(nodeId)] = var;
190 return var;
191 }
192};
193
199void emitInputIndexExprs(std::ostringstream &os, const Graph &g, const FusedGroup &grp,
200 std::vector<std::string> &indexExprs) {
201 const GraphNode &on = g.node(grp.outputNode);
202 const int rank = on.rank;
203 std::vector<int> S(static_cast<size_t>(rank), 1);
204 if (rank > 0) {
205 S[static_cast<size_t>(rank - 1)] = 1;
206 for (int k = rank - 2; k >= 0; --k) S[static_cast<size_t>(k)] =
207 S[static_cast<size_t>(k + 1)] * on.dims[k + 1];
208 }
209 indexExprs.resize(grp.inputs.size());
210 for (size_t k = 0; k < grp.inputs.size(); ++k) {
211 const GraphNode &inN = g.node(grp.inputs[k]);
212 bool identical = inN.rank == rank;
213 if (identical) {
214 for (int d = 0; d < rank; ++d)
215 if (inN.dims[d] != on.dims[d]) {
216 identical = false;
217 break;
218 }
219 }
220 if (identical) {
221 indexExprs[k] = "i_";
222 continue;
223 }
224 const int pad = rank - inN.rank;
225 std::vector<int> stride(static_cast<size_t>(rank), 0);
226 for (int d = 0; d < rank; ++d) {
227 const int dim = d < pad ? 1 : inN.dims[d - pad];
228 if (dim == 1) continue;
229 int s = 1;
230 for (int t = d + 1; t < rank; ++t) {
231 const int dt = t < pad ? 1 : inN.dims[t - pad];
232 if (dt != 1) s *= dt;
233 }
234 stride[static_cast<size_t>(d)] = s;
235 }
236 const std::string var = "idx" + std::to_string(k);
237 os << " uint " << var << " = 0u;\n";
238 for (int d = 0; d < rank; ++d) {
239 if (stride[static_cast<size_t>(d)] == 0) continue;
240 os << " " << var << " += ((i_ / " << S[static_cast<size_t>(d)] << "u) % "
241 << on.dims[d] << "u) * " << stride[static_cast<size_t>(d)] << "u;\n";
242 }
243 indexExprs[k] = var;
244 }
245}
246
247bool genElementwise(const Graph &g, const FusedGroup &grp, KernelSpec &out) {
248 if (grp.inputs.size() > size_t(kMaxKernelBindings - 1)) return false;
249 const GraphNode &on = g.node(grp.outputNode);
250 std::ostringstream os;
251 os << header(kLocalSize);
252 for (size_t k = 0; k < grp.inputs.size(); ++k)
253 os << bufferDecl(int(k), ("a" + std::to_string(k)).c_str());
254 os << bufferDecl(int(grp.inputs.size()), "o");
255 os << pushConstant();
256 os << "void main() {\n";
257 os << " uint i_ = gl_GlobalInvocationID.x;\n";
258 os << " if (i_ >= " << on.size << "u) return;\n";
259 std::vector<std::string> indexExprs;
260 emitInputIndexExprs(os, g, grp, indexExprs);
261 ChainContext ctx{g, grp, indexExprs, "", os, 0, ""};
262 for (int u : grp.nodes) ctx.emit(u);
263 os << " o[i_] = " << ctx.vars[static_cast<size_t>(grp.outputNode)] << ";\n";
264 os << "}\n";
265 out.pass1.clear();
266 out.pass2 = os.str();
267 out.groupsX2 = groupsFor(on.size);
268 out.inputCount = int(grp.inputs.size());
269 return true;
270}
271
273bool genMatMul(const Graph &g, const FusedGroup &grp, bool tiled, KernelSpec &out) {
274 const GraphNode &mm = g.node(grp.nodes.front());
275 const GraphNode &A = g.node(mm.in0);
276 const GraphNode &B = g.node(mm.in1);
277 const bool batched = mm.rank == 3;
278 const int batch = batched ? A.dims[0] : 1;
279 const int m = A.dims[batched ? 1 : 0];
280 const int k = A.dims[batched ? 2 : 1];
281 const int n = B.dims[batched ? 2 : 1];
282 const bool hasBias = grp.biasNode >= 0;
283 const bool bQuant = q::isQuantDType(static_cast<DType>(B.dtype)) && !B.constBytes.empty();
284 if (bQuant && tiled) return false; // quantized weights use the naive variant only
285 const int scalesBinding = hasBias ? 3 : 2;
286 const int bindingOut = bQuant ? (hasBias ? 4 : 3) : (hasBias ? 3 : 2);
287
288 std::ostringstream os;
289 if (!tiled) {
290 os << header(kLocalSize);
291 } else {
292 os << header(16, 16);
293 }
294 os << bufferDecl(0, "a");
295 if (bQuant) os << bufferDeclUint(1, "b");
296 else os << bufferDecl(1, "b");
297 if (hasBias) os << bufferDecl(2, "bias");
298 if (bQuant && B.dtype != static_cast<int>(DType::Fp16))
299 os << bufferDecl(scalesBinding, "bs");
300 os << bufferDecl(bindingOut, "o");
301 os << pushConstant();
302 if (bQuant) os << emitQuantizedBVal(static_cast<DType>(B.dtype), B.qGroup);
303
304 std::vector<std::string> indexExprs(grp.inputs.size(), "i_");
305 ChainContext ctx{g, grp, indexExprs, "r", os, 0, ""};
306 ctx.biasIndexExpr = "jj";
307
308 if (!tiled) {
309 os << "void main() {\n";
310 os << " uint i_ = gl_GlobalInvocationID.x;\n";
311 os << " if (i_ >= " << batch * m * n << "u) return;\n";
312 if (batched) {
313 os << " uint bb = i_ / " << m * n << "u;\n";
314 os << " uint rem = i_ % " << m * n << "u;\n";
315 os << " uint ii = rem / " << n << "u;\n";
316 os << " uint jj = rem % " << n << "u;\n";
317 } else {
318 os << " uint ii = i_ / " << n << "u;\n";
319 os << " uint jj = i_ % " << n << "u;\n";
320 }
321 os << " float r = 0.0;\n";
322 os << " for (uint t = 0u; t < " << k << "u; ++t) {\n";
323 if (batched) {
324 os << " r += a[bb * " << m * k << "u + ii * " << k << "u + t] * "
325 << (bQuant ? "bval(bb * " + std::to_string(k * n) + "u + t * " +
326 std::to_string(n) + "u + jj)"
327 : "b[bb * " + std::to_string(k * n) + "u + t * " +
328 std::to_string(n) + "u + jj]")
329 << ";\n";
330 } else {
331 os << " r += a[ii * " << k << "u + t] * "
332 << (bQuant ? "bval(t * " + std::to_string(n) + "u + jj)"
333 : "b[t * " + std::to_string(n) + "u + jj]")
334 << ";\n";
335 }
336 os << " }\n";
337 for (int u : grp.epilogue) ctx.emit(u);
338 const std::string finalExpr =
339 grp.epilogue.empty() ? std::string("r") : ctx.vars[static_cast<size_t>(grp.outputNode)];
340 os << " o[i_] = " << finalExpr << ";\n";
341 os << "}\n";
342 out.groupsX2 = groupsFor(batch * m * n);
343 } else {
344 ctx.biasIndexExpr = "col";
345 const int gx = (n + 15) / 16;
346 const int gy = (m + 15) / 16;
347 os << "shared float As[16][17];\n";
348 os << "shared float Bs[16][17];\n";
349 os << "void main() {\n";
350 os << " uint tx = gl_LocalInvocationID.x;\n";
351 os << " uint ty = gl_LocalInvocationID.y;\n";
352 os << " uint row = gl_GlobalInvocationID.y;\n";
353 os << " uint col = gl_GlobalInvocationID.x;\n";
354 os << " float acc = 0.0;\n";
355 os << " for (uint tile = 0u; tile < " << (k + 15) / 16 << "u; ++tile) {\n";
356 os << " uint aCol = tile * 16u + tx;\n";
357 os << " uint bRow = tile * 16u + ty;\n";
358 os << " As[ty][tx] = (aCol < " << k << "u && row < " << m << "u) ? a[row * "
359 << k << "u + aCol] : 0.0;\n";
360 os << " Bs[ty][tx] = (bRow < " << k << "u && col < " << n << "u) ? b[bRow * "
361 << n << "u + col] : 0.0;\n";
362 os << " barrier();\n";
363 os << " for (uint t = 0u; t < 16u; ++t) acc += As[ty][t] * Bs[t][tx];\n";
364 os << " barrier();\n";
365 os << " }\n";
366 os << " float r = acc;\n";
367 for (int u : grp.epilogue) ctx.emit(u);
368 const std::string finalExpr =
369 grp.epilogue.empty() ? std::string("r") : ctx.vars[static_cast<size_t>(grp.outputNode)];
370 os << " if (row < " << m << "u && col < " << n << "u) o[row * " << n << "u + col] = "
371 << finalExpr << ";\n";
372 os << "}\n";
373 out.groupsX2 = gx;
374 out.groupsY2 = gy;
375 }
376 out.pass1.clear();
377 out.pass2 = os.str();
378 out.inputCount = hasBias ? 3 : 2;
379 out.qDtype = bQuant ? B.dtype : 0;
380 out.qGroup = B.qGroup;
381 out.scalesBinding = bQuant && B.dtype != static_cast<int>(DType::Fp16) ? scalesBinding : -1;
382 out.outputBinding = bQuant ? bindingOut : -1;
383 if (tiled && batched) return false; // batched matmul uses the naive variant
384 return true;
385}
386
387bool genConv(const Graph &g, const FusedGroup &grp, KernelSpec &out) {
388 const GraphNode &cn = g.node(grp.nodes.front());
389 const GraphNode &X = g.node(cn.in0);
390 const GraphNode &Wt = g.node(cn.in1);
391 const bool is1d = cn.type == OpType::Conv1d;
392 const int stride = cn.i0, pad = cn.i1;
393 const bool hasBias = grp.biasNode >= 0;
394 const bool hasConvBias = cn.in2 >= 0;
395 const int bindingOut = 2 + int(hasConvBias) + int(hasBias);
396
397 std::ostringstream os;
398 os << header(kLocalSize);
399 os << bufferDecl(0, "x");
400 os << bufferDecl(1, "w");
401 if (hasConvBias) os << bufferDecl(2, "convBias");
402 if (hasBias) os << bufferDecl(hasConvBias ? 3 : 2, "bias");
403 os << bufferDecl(bindingOut, "o");
404 os << pushConstant();
405 std::vector<std::string> indexExprs(grp.inputs.size(), "i_");
406 ChainContext ctx{g, grp, indexExprs, "r", os, 0, ""};
407 ctx.biasIndexExpr = is1d ? "ol" : "ow";
408
409 os << "void main() {\n";
410 os << " uint i_ = gl_GlobalInvocationID.x;\n";
411 os << " if (i_ >= " << g.node(grp.outputNode).size << "u) return;\n";
412 if (is1d) {
413 const int C = X.dims[1], L = X.dims[2];
414 const int F = Wt.dims[0], K = Wt.dims[2];
415 os << " uint n_ = i_ / (" << F << "u * " << cn.dims[2] << "u);\n";
416 os << " uint rem = i_ % (" << F << "u * " << cn.dims[2] << "u);\n";
417 os << " uint f = rem / " << cn.dims[2] << "u;\n";
418 os << " uint ol = rem % " << cn.dims[2] << "u;\n";
419 os << " float r = " << (hasConvBias ? "convBias[f]" : "0.0") << ";\n";
420 os << " for (uint c = 0u; c < " << C << "u; ++c) {\n";
421 os << " for (uint kk = 0u; kk < " << K << "u; ++kk) {\n";
422 os << " int il = int(ol) * " << stride << " + int(kk) - " << pad << ";\n";
423 os << " if (il < 0 || il >= " << L << ") continue;\n";
424 os << " r += x[(n_ * " << C << "u + c) * " << L << "u + uint(il)] * w[(f * "
425 << C << "u + c) * " << K << "u + kk];\n";
426 os << " }\n";
427 os << " }\n";
428 } else {
429 const int C = X.dims[1], H = X.dims[2], W = X.dims[3];
430 const int F = Wt.dims[0], KH = Wt.dims[2], KW = Wt.dims[3];
431 os << " uint n_ = i_ / (" << F << "u * " << cn.dims[2] << "u * " << cn.dims[3] << "u);\n";
432 os << " uint rem = i_ % (" << F << "u * " << cn.dims[2] << "u * " << cn.dims[3] << "u);\n";
433 os << " uint f = rem / (" << cn.dims[2] << "u * " << cn.dims[3] << "u);\n";
434 os << " uint rem2 = rem % (" << cn.dims[2] << "u * " << cn.dims[3] << "u);\n";
435 os << " uint oh = rem2 / " << cn.dims[3] << "u;\n";
436 os << " uint ow = rem2 % " << cn.dims[3] << "u;\n";
437 os << " float r = " << (hasConvBias ? "convBias[f]" : "0.0") << ";\n";
438 os << " for (uint c = 0u; c < " << C << "u; ++c) {\n";
439 os << " for (uint kh = 0u; kh < " << KH << "u; ++kh) {\n";
440 os << " int ih = int(oh) * " << stride << " + int(kh) - " << pad << ";\n";
441 os << " if (ih < 0 || ih >= " << H << ") continue;\n";
442 os << " for (uint kw = 0u; kw < " << KW << "u; ++kw) {\n";
443 os << " int iw = int(ow) * " << stride << " + int(kw) - " << pad << ";\n";
444 os << " if (iw < 0 || iw >= " << W << ") continue;\n";
445 os << " r += x[((n_ * " << C << "u + c) * " << H << "u + uint(ih)) * " << W
446 << "u + uint(iw)] * w[((f * " << C << "u + c) * " << KH << "u + kh) * " << KW
447 << "u + kw];\n";
448 os << " }\n";
449 os << " }\n";
450 os << " }\n";
451 }
452 for (int u : grp.epilogue) ctx.emit(u);
453 const std::string finalExpr =
454 grp.epilogue.empty() ? std::string("r") : ctx.vars[static_cast<size_t>(grp.outputNode)];
455 os << " o[i_] = " << finalExpr << ";\n";
456 os << "}\n";
457 out.pass1.clear();
458 out.pass2 = os.str();
459 out.groupsX2 = groupsFor(g.node(grp.outputNode).size);
460 out.inputCount = 2 + int(hasConvBias) + int(hasBias);
461 return true;
462}
463
464bool genPool(const Graph &g, const FusedGroup &grp, KernelSpec &out) {
465 const GraphNode &pn = g.node(grp.outputNode);
466 const GraphNode &X = g.node(pn.in0);
467 const int ksize = pn.i0, stride = pn.i1, pad = pn.i2;
468 const int C = X.dims[1], H = X.dims[2], W = X.dims[3];
469 const int OH = pn.dims[2], OW = pn.dims[3];
470 const bool maxPool = pn.type == OpType::MaxPool2d;
471 std::ostringstream os;
472 os << header(kLocalSize);
473 os << bufferDecl(0, "in_");
474 os << bufferDecl(1, "o");
475 os << pushConstant();
476 os << "void main() {\n";
477 os << " uint i_ = gl_GlobalInvocationID.x;\n";
478 os << " if (i_ >= " << pn.size << "u) return;\n";
479 os << " uint n_ = i_ / (" << C << "u * " << OH << "u * " << OW << "u);\n";
480 os << " uint rem = i_ % (" << C << "u * " << OH << "u * " << OW << "u);\n";
481 os << " uint c = rem / (" << OH << "u * " << OW << "u);\n";
482 os << " uint rem2 = rem % (" << OH << "u * " << OW << "u);\n";
483 os << " uint oh = rem2 / " << OW << "u;\n";
484 os << " uint ow = rem2 % " << OW << "u;\n";
485 os << " float acc = " << (maxPool ? "-3.402823e38" : "0.0") << ";\n";
486 os << " int valid = 0;\n";
487 os << " for (int kh = 0; kh < " << ksize << "; ++kh) {\n";
488 os << " int ih = int(oh) * " << stride << " + kh - " << pad << ";\n";
489 os << " if (ih < 0 || ih >= " << H << ") continue;\n";
490 os << " for (int kw = 0; kw < " << ksize << "; ++kw) {\n";
491 os << " int iw = int(ow) * " << stride << " + kw - " << pad << ";\n";
492 os << " if (iw < 0 || iw >= " << W << ") continue;\n";
493 os << " float v = in_[((n_ * " << C << "u + c) * " << H << "u + uint(ih)) * " << W
494 << "u + uint(iw)];\n";
495 if (maxPool) {
496 os << " acc = max(acc, v);\n";
497 } else {
498 os << " acc += v; ++valid;\n";
499 }
500 os << " }\n";
501 os << " }\n";
502 if (!maxPool) os << " acc = valid > 0 ? acc / float(valid) : 0.0;\n";
503 os << " o[i_] = acc;\n";
504 os << "}\n";
505 out.pass1.clear();
506 out.pass2 = os.str();
507 out.groupsX2 = groupsFor(pn.size);
508 out.inputCount = 1;
509 return true;
510}
511
512bool genSoftmax(const Graph &g, const FusedGroup &grp, KernelSpec &out) {
513 const GraphNode &sn = g.node(grp.outputNode);
514 const GraphNode &X = g.node(sn.in0);
515 const int axis = sn.i0;
516 int outer = 1, reduce = 1, inner = 1;
517 for (int k = 0; k < axis; ++k) outer *= X.dims[k];
518 reduce = X.dims[axis];
519 for (int k = axis + 1; k < X.rank; ++k) inner *= X.dims[k];
520 const int rows = outer * inner;
521 const bool logMode = grp.logMode;
522
523 std::ostringstream os1, os2;
524 os1 << header(kLocalSize);
525 os1 << bufferDecl(0, "in_");
526 os1 << bufferDecl(1, "mx");
527 os1 << bufferDecl(2, "sm");
528 os1 << pushConstant();
529 os1 << "void main() {\n";
530 os1 << " uint i_ = gl_GlobalInvocationID.x;\n";
531 os1 << " if (i_ >= " << rows << "u) return;\n";
532 os1 << " uint o_ = i_ / " << inner << "u;\n";
533 os1 << " uint ii = i_ % " << inner << "u;\n";
534 os1 << " float m = -3.402823e38;\n";
535 os1 << " for (uint j = 0u; j < " << reduce << "u; ++j) {\n";
536 os1 << " m = max(m, in_[(o_ * " << reduce << "u + j) * " << inner << "u + ii]);\n";
537 os1 << " }\n";
538 os1 << " float s = 0.0;\n";
539 os1 << " for (uint j = 0u; j < " << reduce << "u; ++j) {\n";
540 os1 << " s += exp(in_[(o_ * " << reduce << "u + j) * " << inner << "u + ii] - m);\n";
541 os1 << " }\n";
542 os1 << " mx[i_] = m;\n";
543 os1 << " sm[i_] = s;\n";
544 os1 << "}\n";
545
546 os2 << header(kLocalSize);
547 os2 << bufferDecl(0, "in_");
548 os2 << bufferDecl(1, "mx");
549 os2 << bufferDecl(2, "sm");
550 os2 << bufferDecl(3, "o");
551 os2 << pushConstant();
552 os2 << "void main() {\n";
553 os2 << " uint i_ = gl_GlobalInvocationID.x;\n";
554 os2 << " if (i_ >= " << sn.size << "u) return;\n";
555 os2 << " uint o_ = (i_ / (" << reduce * inner << "u)) * " << inner << "u + (i_ % "
556 << inner << "u);\n";
557 os2 << " float m = mx[o_];\n";
558 os2 << " float s = sm[o_];\n";
559 os2 << " float x = in_[i_];\n";
560 if (logMode) {
561 os2 << " o[i_] = (x - m) - log(s);\n";
562 } else {
563 os2 << " o[i_] = exp(x - m) / s;\n";
564 }
565 os2 << "}\n";
566 out.pass1 = os1.str();
567 out.pass2 = os2.str();
568 out.groupsX1 = groupsFor(rows);
569 out.groupsX2 = groupsFor(sn.size);
570 out.inputCount = 1;
571 out.inputsReadPass1 = 1;
572 out.statsCount = 2;
573 out.statsSize = rows;
574 out.twoPass = true;
575 return true;
576}
577
578bool genNorm(const Graph &g, const FusedGroup &grp, bool rms, KernelSpec &out) {
579 const GraphNode &nn = g.node(grp.outputNode);
580 const GraphNode &X = g.node(nn.in0);
581 const int cols = X.dims[X.rank - 1];
582 const int rows = X.size / cols;
583 const float eps = nn.s0;
584 const bool hasScale = grp.hasScale;
585 const bool hasBias = !rms && grp.hasBias;
586 const int inputCount = 1 + (hasScale ? 1 : 0) + (hasBias ? 1 : 0);
587 const int statsCount = rms ? 1 : 2;
588 const int outBinding = inputCount + statsCount;
589
590 std::ostringstream os1, os2;
591 os1 << header(kLocalSize);
592 os1 << bufferDecl(0, "in_");
593 os1 << bufferDecl(1, "st0");
594 if (statsCount > 1) os1 << bufferDecl(2, "st1");
595 os1 << pushConstant();
596 os1 << "void main() {\n";
597 os1 << " uint i_ = gl_GlobalInvocationID.x;\n";
598 os1 << " if (i_ >= " << rows << "u) return;\n";
599 os1 << " float s0 = 0.0;\n";
600 if (statsCount > 1) os1 << " float s1 = 0.0;\n";
601 os1 << " for (uint j = 0u; j < " << cols << "u; ++j) {\n";
602 os1 << " float v = in_[i_ * " << cols << "u + j];\n";
603 if (rms) {
604 os1 << " s0 += v * v;\n";
605 } else {
606 os1 << " s0 += v; s1 += v * v;\n";
607 }
608 os1 << " }\n";
609 os1 << " st0[i_] = s0;\n";
610 if (statsCount > 1) os1 << " st1[i_] = s1;\n";
611 os1 << "}\n";
612
613 os2 << header(kLocalSize);
614 for (int k = 0; k < inputCount; ++k)
615 os2 << bufferDecl(k, ("a" + std::to_string(k)).c_str());
616 os2 << bufferDecl(inputCount, "st0");
617 if (statsCount > 1) os2 << bufferDecl(inputCount + 1, "st1");
618 os2 << bufferDecl(outBinding, "o");
619 os2 << pushConstant();
620 os2 << "void main() {\n";
621 os2 << " uint i_ = gl_GlobalInvocationID.x;\n";
622 os2 << " if (i_ >= " << nn.size << "u) return;\n";
623 os2 << " uint r = i_ / " << cols << "u;\n";
624 os2 << " uint c = i_ % " << cols << "u;\n";
625 if (rms) {
626 os2 << " float inv = 1.0 / sqrt(st0[r] / " << cols << ".0 + " << scalarStr(eps)
627 << ");\n";
628 os2 << " float y = a0[i_] * inv;\n";
629 if (hasScale) os2 << " y *= a1[c];\n";
630 } else {
631 os2 << " float mean = st0[r] / " << cols << ".0;\n";
632 os2 << " float var = st1[r] / " << cols << ".0 - mean * mean;\n";
633 os2 << " var = max(var, 0.0);\n";
634 os2 << " float inv = 1.0 / sqrt(var + " << scalarStr(eps) << ");\n";
635 os2 << " float y = (a0[i_] - mean) * inv;\n";
636 if (hasScale) os2 << " y *= a1[c];\n";
637 if (hasBias) os2 << " y += a" << (hasScale ? 2 : 1) << "[c];\n";
638 }
639 os2 << " o[i_] = y;\n";
640 os2 << "}\n";
641 out.pass1 = os1.str();
642 out.pass2 = os2.str();
643 out.groupsX1 = groupsFor(rows);
644 out.groupsX2 = groupsFor(nn.size);
645 out.inputCount = inputCount;
646 out.inputsReadPass1 = 1;
647 out.statsCount = statsCount;
648 out.statsSize = rows;
649 out.twoPass = true;
650 return true;
651}
652
653bool genReduceOrArgmax(const Graph &g, const FusedGroup &grp, bool argmax, KernelSpec &out) {
654 const GraphNode &rn = g.node(grp.outputNode);
655 const GraphNode &X = g.node(rn.in0);
656 const int axis = rn.i0;
657 int outer = 1, reduce = 1, inner = 1;
658 for (int k = 0; k < axis; ++k) outer *= X.dims[k];
659 reduce = X.dims[axis];
660 for (int k = axis + 1; k < X.rank; ++k) inner *= X.dims[k];
661 const int outSize = outer * inner;
662 std::ostringstream os;
663 os << header(kLocalSize);
664 os << bufferDecl(0, "in_");
665 os << bufferDecl(1, "o");
666 os << pushConstant();
667 os << "void main() {\n";
668 os << " uint i_ = gl_GlobalInvocationID.x;\n";
669 os << " if (i_ >= " << outSize << "u) return;\n";
670 os << " uint o_ = i_ / " << inner << "u;\n";
671 os << " uint ii = i_ % " << inner << "u;\n";
672 if (argmax) {
673 os << " float best = -3.402823e38;\n";
674 os << " float bestJ = 0.0;\n";
675 os << " for (uint j = 0u; j < " << reduce << "u; ++j) {\n";
676 os << " float v = in_[(o_ * " << reduce << "u + j) * " << inner << "u + ii];\n";
677 os << " if (v > best) { best = v; bestJ = float(j); }\n";
678 os << " }\n";
679 os << " o[i_] = bestJ;\n";
680 } else {
681 switch (grp.op) {
684 os << " float acc = 0.0;\n";
685 os << " for (uint j = 0u; j < " << reduce << "u; ++j) acc += in_[(o_ * "
686 << reduce << "u + j) * " << inner << "u + ii];\n";
687 if (grp.op == OpType::ReduceMean)
688 os << " acc /= " << scalarStr(float(reduce)) << ";\n";
689 break;
691 os << " float acc = 3.402823e38;\n";
692 os << " for (uint j = 0u; j < " << reduce << "u; ++j) acc = min(acc, in_[(o_ * "
693 << reduce << "u + j) * " << inner << "u + ii]);\n";
694 break;
696 os << " float acc = -3.402823e38;\n";
697 os << " for (uint j = 0u; j < " << reduce << "u; ++j) acc = max(acc, in_[(o_ * "
698 << reduce << "u + j) * " << inner << "u + ii]);\n";
699 break;
700 default: return false;
701 }
702 os << " o[i_] = acc;\n";
703 }
704 os << "}\n";
705 out.pass1.clear();
706 out.pass2 = os.str();
707 out.groupsX2 = groupsFor(outSize);
708 out.inputCount = 1;
709 return true;
710}
711
712bool genEmbedding(const Graph &g, const FusedGroup &grp, KernelSpec &out) {
713 const GraphNode &en = g.node(grp.outputNode);
714 const GraphNode &T = g.node(en.in0);
715 const bool tQuant = q::isQuantDType(static_cast<DType>(T.dtype)) && !T.constBytes.empty();
716 const bool tInt = T.dtype != static_cast<int>(DType::Fp16);
717 const int vocab = T.dims[0], dim = T.dims[1];
718 std::ostringstream os;
719 os << header(kLocalSize);
720 if (tQuant) os << bufferDeclUint(0, "table");
721 else os << bufferDecl(0, "table");
722 os << bufferDecl(1, "idx");
723 if (tQuant && tInt) os << bufferDecl(2, "bs");
724 os << bufferDecl(tQuant ? 3 : 2, "o");
725 os << pushConstant();
726 if (tQuant) os << emitQuantizedBVal(static_cast<DType>(T.dtype), T.qGroup, "table");
727 os << "void main() {\n";
728 os << " uint i_ = gl_GlobalInvocationID.x;\n";
729 os << " if (i_ >= " << en.size << "u) return;\n";
730 os << " uint r = i_ / " << dim << "u;\n";
731 os << " uint d = i_ % " << dim << "u;\n";
732 os << " int ii = int(idx[r]);\n";
733 os << " ii = clamp(ii, 0, " << (vocab - 1) << ");\n";
734 os << " o[i_] = " << (tQuant ? "bval(uint(ii) * " + std::to_string(dim) + "u + d)"
735 : "table[uint(ii) * " + std::to_string(dim) + "u + d]")
736 << ";\n";
737 os << "}\n";
738 out.pass1.clear();
739 out.pass2 = os.str();
740 out.groupsX2 = groupsFor(en.size);
741 out.inputCount = 2;
742 out.qDtype = tQuant ? T.dtype : 0;
743 out.qGroup = T.qGroup;
744 out.scalesBinding = tQuant && tInt ? 2 : -1;
745 out.outputBinding = tQuant ? 3 : -1;
746 return true;
747}
748
749bool genConcat(const Graph &g, const FusedGroup &grp, KernelSpec &out) {
750 const GraphNode &cn = g.node(grp.outputNode);
751 const int axis = cn.i0;
752 const int n = int(grp.inputs.size());
753 if (n < 2 || n > 4) return false;
754 int starts[4] = {};
755 int axisTotal = 0;
756 for (int k = 0; k < n; ++k) {
757 starts[k] = axisTotal;
758 axisTotal += g.node(grp.inputs[static_cast<size_t>(k)]).dims[axis];
759 }
760 int inner = 1;
761 for (int k = axis + 1; k < cn.rank; ++k) inner *= cn.dims[k];
762 std::ostringstream os;
763 os << header(kLocalSize);
764 for (int k = 0; k < n; ++k) os << bufferDecl(k, ("a" + std::to_string(k)).c_str());
765 os << bufferDecl(n, "o");
766 os << pushConstant();
767 os << "void main() {\n";
768 os << " uint i_ = gl_GlobalInvocationID.x;\n";
769 os << " if (i_ >= " << cn.size << "u) return;\n";
770 os << " uint ax = (i_ / " << inner << "u) % " << axisTotal << "u;\n";
771 os << " uint op = i_ / (" << axisTotal << "u * " << inner << "u);\n";
772 os << " uint ip = i_ % " << inner << "u;\n";
773 os << " float v = 0.0;\n";
774 for (int k = 0; k < n; ++k) {
775 const int sz = g.node(grp.inputs[static_cast<size_t>(k)]).dims[axis];
776 const char *cond = k == 0 ? "if" : "else if";
777 os << " " << cond << " (ax >= " << starts[k] << "u && ax < " << starts[k] + sz
778 << "u) {\n";
779 os << " v = a" << k << "[op * " << sz << "u * " << inner << "u + (ax - " << starts[k]
780 << "u) * " << inner << "u + ip];\n";
781 os << " }\n";
782 }
783 os << " o[i_] = v;\n";
784 os << "}\n";
785 out.pass1.clear();
786 out.pass2 = os.str();
787 out.groupsX2 = groupsFor(cn.size);
788 out.inputCount = n;
789 return true;
790}
791
792bool genSlice(const Graph &g, const FusedGroup &grp, KernelSpec &out) {
793 const GraphNode &sn = g.node(grp.outputNode);
794 const GraphNode &X = g.node(sn.in0);
795 const int axis = sn.i0, begin = sn.i1, end = sn.i2;
796 const int axisSize = end - begin;
797 int inner = 1;
798 for (int k = axis + 1; k < sn.rank; ++k) inner *= sn.dims[k];
799 std::ostringstream os;
800 os << header(kLocalSize);
801 os << bufferDecl(0, "in_");
802 os << bufferDecl(1, "o");
803 os << pushConstant();
804 os << "void main() {\n";
805 os << " uint i_ = gl_GlobalInvocationID.x;\n";
806 os << " if (i_ >= " << sn.size << "u) return;\n";
807 os << " uint ax = (i_ / " << inner << "u) % " << axisSize << "u;\n";
808 os << " uint op = i_ / (" << axisSize << "u * " << inner << "u);\n";
809 os << " uint ip = i_ % " << inner << "u;\n";
810 os << " o[i_] = in_[op * " << X.dims[axis] << "u * " << inner << "u + (ax + " << begin
811 << "u) * " << inner << "u + ip];\n";
812 os << "}\n";
813 out.pass1.clear();
814 out.pass2 = os.str();
815 out.groupsX2 = groupsFor(sn.size);
816 out.inputCount = 1;
817 return true;
818}
819
820bool genPermute(const Graph &g, const FusedGroup &grp, KernelSpec &out) {
821 const GraphNode &pn = g.node(grp.outputNode);
822 const GraphNode &X = g.node(pn.in0);
823 const int rank = pn.rank;
824 int S[Tensor::kMaxRank] = {};
825 S[rank - 1] = 1;
826 for (int k = rank - 2; k >= 0; --k) S[k] = S[k + 1] * pn.dims[k + 1];
827 int inStride[Tensor::kMaxRank] = {};
828 inStride[rank - 1] = 1;
829 for (int k = rank - 2; k >= 0; --k) inStride[k] = inStride[k + 1] * X.dims[k + 1];
830 std::ostringstream os;
831 os << header(kLocalSize);
832 os << bufferDecl(0, "in_");
833 os << bufferDecl(1, "o");
834 os << pushConstant();
835 os << "void main() {\n";
836 os << " uint i_ = gl_GlobalInvocationID.x;\n";
837 os << " if (i_ >= " << pn.size << "u) return;\n";
838 os << " uint idx = 0u;\n";
839 for (int k = 0; k < rank; ++k) {
840 const int inAxis = pn.perm[k];
841 os << " idx += ((i_ / " << S[k] << "u) % " << pn.dims[k] << "u) * "
842 << inStride[inAxis] << "u;\n";
843 }
844 os << " o[i_] = in_[idx];\n";
845 os << "}\n";
846 out.pass1.clear();
847 out.pass2 = os.str();
848 out.groupsX2 = groupsFor(pn.size);
849 out.inputCount = 1;
850 return true;
851}
852
853bool genSdpa(const Graph &g, const FusedGroup &grp, KernelSpec &out) {
854 const GraphNode &qn = g.node(grp.outputNode);
855 const GraphNode &Q = g.node(qn.in0);
856 const GraphNode &K = g.node(qn.in1);
857 const int B = Q.dims[0], H = Q.dims[1], T = Q.dims[2], D = Q.dims[3];
858 const int S = K.dims[2];
859 if (S > 2048 || D > 512) return false; // shared-memory limits -> CPU fallback
860 const float scale = qn.s0;
861 const bool masked = grp.masked;
862 const int bindingOut = masked ? 4 : 3;
863 std::ostringstream os;
864 os << header(128);
865 os << bufferDecl(0, "q");
866 os << bufferDecl(1, "k");
867 os << bufferDecl(2, "v");
868 if (masked) os << bufferDecl(3, "mask");
869 os << bufferDecl(bindingOut, "o");
870 os << "shared float scores[" << S << "];\n";
871 os << "shared float maxv;\n";
872 os << "shared float sumv;\n";
873 os << pushConstant();
874 os << "void main() {\n";
875 os << " uint tid = gl_LocalInvocationID.x;\n";
876 os << " uint bh = gl_WorkGroupID.x;\n";
877 os << " uint t = gl_WorkGroupID.y;\n";
878 os << " uint b = bh / " << H << "u;\n";
879 os << " uint h = bh % " << H << "u;\n";
880 os << " uint qbase = (b * " << H << "u + h) * " << T << "u * " << D << "u + t * " << D
881 << "u;\n";
882 os << " uint kbase = (b * " << H << "u + h) * " << S << "u * " << D << "u;\n";
883 os << " uint vbase = kbase;\n";
884 os << " for (uint s = tid; s < " << S << "u; s += 128u) {\n";
885 os << " float acc = 0.0;\n";
886 os << " for (uint d = 0u; d < " << D << "u; ++d) acc += q[qbase + d] * k[kbase + s * "
887 << D << "u + d];\n";
888 os << " acc *= " << scalarStr(scale) << ";\n";
889 if (masked) {
890 os << " acc += mask[(b * " << H << "u + h) * " << T << "u * " << S << "u + t * " << S
891 << "u + s];\n";
892 }
893 os << " scores[s] = acc;\n";
894 os << " }\n";
895 os << " barrier();\n";
896 os << " if (tid == 0u) {\n";
897 os << " float m = -3.402823e38;\n";
898 os << " for (uint s = 0u; s < " << S << "u; ++s) m = max(m, scores[s]);\n";
899 os << " float sm = 0.0;\n";
900 os << " for (uint s = 0u; s < " << S << "u; ++s) sm += exp(scores[s] - m);\n";
901 os << " maxv = m; sumv = sm;\n";
902 os << " }\n";
903 os << " barrier();\n";
904 os << " for (uint d = tid; d < " << D << "u; d += 128u) {\n";
905 os << " float acc = 0.0;\n";
906 os << " for (uint s = 0u; s < " << S << "u; ++s) acc += exp(scores[s] - maxv) * v[vbase + s * "
907 << D << "u + d];\n";
908 os << " o[qbase + d] = acc / sumv;\n";
909 os << " }\n";
910 os << "}\n";
911 out.pass1.clear();
912 out.pass2 = os.str();
913 out.groupsX2 = B * H;
914 out.groupsY2 = T;
915 out.inputCount = masked ? 4 : 3;
916 return true;
917}
918
919} // namespace glsl_detail
920
922 using namespace glsl_detail;
923 out = KernelSpec{};
924 switch (group.kind) {
925 case GroupKind::Elementwise: return genElementwise(graph, group, out);
927 // naive variant first; the runtime autotunes between naive and tiled
928 return genMatMul(graph, group, false, out);
930 case GroupKind::Conv2d: return genConv(graph, group, out);
932 case GroupKind::AvgPool2d: return genPool(graph, group, out);
933 case GroupKind::Softmax: return genSoftmax(graph, group, out);
934 case GroupKind::LayerNorm: return genNorm(graph, group, false, out);
935 case GroupKind::RMSNorm: return genNorm(graph, group, true, out);
938 return genReduceOrArgmax(graph, group, group.kind == GroupKind::ArgMax, out);
939 case GroupKind::Embedding: return genEmbedding(graph, group, out);
940 case GroupKind::Concat: return genConcat(graph, group, out);
941 case GroupKind::Slice: return genSlice(graph, group, out);
942 case GroupKind::Permute: return genPermute(graph, group, out);
943 case GroupKind::Resize2d: genResize2d(graph, group, out); return true;
944 case GroupKind::Sdpa: return genSdpa(graph, group, out);
945 case GroupKind::Alias: return true; // no kernel; pure buffer alias
946 }
947 return false;
948}
949
950bool generateMatMulVariant(const Graph &graph, const FusedGroup &group, bool tiled,
951 KernelSpec &out) {
952 using namespace glsl_detail;
953 out = KernelSpec{};
954 if (group.kind != GroupKind::MatMul) return false;
955 return genMatMul(graph, group, tiled, out);
956}
957
958} // namespace eve::tensor
float x
Definition AnimClip.cpp:738
const std::string & s
building::EdgeCurveGroup group
std::string nodeId
int rows
int cols
tensor::Graph g
Definition GpuGraph.cpp:7
float u
Definition Grass.cpp:233
glm::vec3 n
Definition Grass.cpp:63
float v
std::array< float, 3 > scale
std::string name
std::map< std::string, std::vector< std::string > > graph
Definition Package.cpp:59
eve::action::ActionVfxBinding binding
int idx
float begin
float d
float t
std::uint32_t count
float m[16]
EVENGINE_API_DOMAINS public API.
Definition Graph.h:111
const GraphNode & node(int id) const
Node.
Definition Graph.h:116
static constexpr int kMaxRank
Definition Tensor.h:50
bool genConv(const Graph &g, const FusedGroup &grp, KernelSpec &out)
std::string emitQuantizedBVal(DType dt, int group, const char *bufName="b")
Definition KernelGen.cpp:41
bool genSdpa(const Graph &g, const FusedGroup &grp, KernelSpec &out)
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
bool genEmbedding(const Graph &g, const FusedGroup &grp, KernelSpec &out)
bool genSoftmax(const Graph &g, const FusedGroup &grp, KernelSpec &out)
bool genConcat(const Graph &g, const FusedGroup &grp, KernelSpec &out)
int groupsFor(int count)
Groups for.
bool genNorm(const Graph &g, const FusedGroup &grp, bool rms, KernelSpec &out)
std::string bufferDeclUint(int binding, const char *name)
Definition KernelGen.cpp:29
std::string scalarStr(float v)
Scalar str.
bool genSlice(const Graph &g, const FusedGroup &grp, KernelSpec &out)
bool genReduceOrArgmax(const Graph &g, const FusedGroup &grp, bool argmax, KernelSpec &out)
std::string pushConstant()
Pushes constant.
bool genMatMul(const Graph &g, const FusedGroup &grp, bool tiled, KernelSpec &out)
bool genPermute(const Graph &g, const FusedGroup &grp, KernelSpec &out)
bool genPool(const Graph &g, const FusedGroup &grp, KernelSpec &out)
void emitInputIndexExprs(std::ostringstream &os, const Graph &g, const FusedGroup &grp, std::vector< std::string > &indexExprs)
bool genElementwise(const Graph &g, const FusedGroup &grp, KernelSpec &out)
bool isQuantDType(DType dt)
True when quant d type.
Definition Quant.h:17
DType
Tensor element types.
Definition Tensor.h:24
bool generateKernel(const Graph &graph, const FusedGroup &group, KernelSpec &out)
Generate kernel.
constexpr int kMaxKernelBindings
Definition KernelGen.h:69
bool generateMatMulVariant(const Graph &graph, const FusedGroup &group, bool tiled, KernelSpec &out)
Generate mat mul variant.
FusedGroup public API.
Definition Optimizer.h:46
std::vector< int > epilogue
Definition Optimizer.h:65
std::vector< int > nodes
Definition Optimizer.h:51
std::vector< int > inputs
Definition Optimizer.h:53
GraphNode public API.
Definition Graph.h:80
int perm[Tensor::kMaxRank]
Definition Graph.h:99
int dims[Tensor::kMaxRank]
Definition Graph.h:82
std::vector< uint8_t > constBytes
Definition Graph.h:105
KernelSpec public API.
Definition KernelGen.h:27
std::vector< std::string > vars
const std::vector< std::string > & indexExprs
gpgpu::ComputeShader * reduce
uint32_t pad[2]