载入中...
搜索中...
未找到
KernelGenWgslIndex.cpp
浏览该文件的文档.
1#include <algorithm>
2#include <cmath>
3#include <sstream>
4#include "common/Exception.h"
6#include "tensor/Quant.h"
7
9void genEmbedding(const Graph &g, const FusedGroup &grp, KernelSpec &out) {
10 const GraphNode &en = g.node(grp.outputNode);
11 const GraphNode &T = g.node(en.in0);
12 const bool tQuant = q::isQuantDType(static_cast<DType>(T.dtype)) && !T.constBytes.empty();
13 const bool tInt = T.dtype != static_cast<int>(DType::Fp16);
14 const int vocab = T.dims[0], dim = T.dims[1];
15 std::ostringstream os;
16 os << header(kLocalSize);
17 if (tQuant)
18 os << bufferDeclUint(0, "table");
19 else
20 os << bufferDecl(0, "table");
21 os << bufferDecl(1, "idx");
22 if (tQuant && tInt) os << bufferDecl(2, "bs");
23 os << bufferDecl(tQuant ? 3 : 2, "o");
24 os << pushConstant();
25 if (tQuant) os << emitQuantizedBVal(static_cast<DType>(T.dtype), T.qGroup, "table");
26 os << "@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>) {\n";
28 os << " var i_: u32 = globalId.x;\n";
29 os << " if (i_ >= " << en.size << "u) { return; }\n";
30 os << " var r: u32 = i_ / " << dim << "u;\n";
31 os << " var d: u32 = i_ % " << dim << "u;\n";
32 os << " var ii: i32 = i32(idx[r]);\n";
33 os << " ii = clamp(ii, 0, " << (vocab - 1) << ");\n";
34 os << " o[i_] = "
35 << (tQuant ? "bval(u32(ii) * " + std::to_string(dim) + "u + d)"
36 : "table[u32(ii) * " + std::to_string(dim) + "u + d]")
37 << ";\n";
38 os << "}\n";
39 out.pass1.clear();
40 out.pass2 = os.str();
41 out.groupsX2 = groupsFor(en.size);
42 out.inputCount = 2;
43 out.qDtype = tQuant ? T.dtype : 0;
44 out.qGroup = T.qGroup;
45 out.scalesBinding = tQuant && tInt ? 2 : -1;
46 out.outputBinding = tQuant ? 3 : -1;
47 return;
48}
49
50void genConcat(const Graph &g, const FusedGroup &grp, KernelSpec &out) {
51 const GraphNode &cn = g.node(grp.outputNode);
52 const int axis = cn.i0;
53 const int n = int(grp.inputs.size());
54 if (n < 2 || n > 4) throw eve::Exception("Tensor WGSL: unsupported kernel variant or binding count");
55 int starts[4] = {};
56 int axisTotal = 0;
57 for (int k = 0; k < n; ++k) {
58 starts[k] = axisTotal;
59 axisTotal += g.node(grp.inputs[static_cast<size_t>(k)]).dims[axis];
60 }
61 int inner = 1;
62 for (int k = axis + 1; k < cn.rank; ++k) inner *= cn.dims[k];
63 std::ostringstream os;
64 os << header(kLocalSize);
65 for (int k = 0; k < n; ++k) os << bufferDecl(k, ("a" + std::to_string(k)).c_str());
66 os << bufferDecl(n, "o");
67 os << pushConstant();
68 os << "@compute @workgroup_size(workgroupX, workgroupY)\nfn main(@builtin(global_invocation_id) globalId: "
69 "vec3<u32>, @builtin(local_invocation_id) localId: vec3<u32>, @builtin(workgroup_id) groupId: vec3<u32>) {\n";
70 os << " var i_: u32 = globalId.x;\n";
71 os << " if (i_ >= " << cn.size << "u) { return; }\n";
72 os << " var ax: u32 = (i_ / " << inner << "u) % " << axisTotal << "u;\n";
73 os << " var op: u32 = i_ / (" << axisTotal << "u * " << inner << "u);\n";
74 os << " var ip: u32 = i_ % " << inner << "u;\n";
75 os << " var v: f32 = 0.0;\n";
76 for (int k = 0; k < n; ++k) {
77 const int sz = g.node(grp.inputs[static_cast<size_t>(k)]).dims[axis];
78 const char *cond = k == 0 ? "if" : "else if";
79 os << " " << cond << " (ax >= " << starts[k] << "u && ax < " << starts[k] + sz << "u) {\n";
80 os << " v = a" << k << "[op * " << sz << "u * " << inner << "u + (ax - " << starts[k] << "u) * " << inner
81 << "u + ip];\n";
82 os << " }\n";
83 }
84 os << " o[i_] = v;\n";
85 os << "}\n";
86 out.pass1.clear();
87 out.pass2 = os.str();
88 out.groupsX2 = groupsFor(cn.size);
89 out.inputCount = n;
90 return;
91}
92
93void genSlice(const Graph &g, const FusedGroup &grp, KernelSpec &out) {
94 const GraphNode &sn = g.node(grp.outputNode);
95 const GraphNode &X = g.node(sn.in0);
96 const int axis = sn.i0, begin = sn.i1, end = sn.i2;
97 const int axisSize = end - begin;
98 int inner = 1;
99 for (int k = axis + 1; k < sn.rank; ++k) inner *= sn.dims[k];
100 std::ostringstream os;
101 os << header(kLocalSize);
102 os << bufferDecl(0, "in_");
103 os << bufferDecl(1, "o");
104 os << pushConstant();
105 os << "@compute @workgroup_size(workgroupX, workgroupY)\nfn main(@builtin(global_invocation_id) globalId: "
106 "vec3<u32>, @builtin(local_invocation_id) localId: vec3<u32>, @builtin(workgroup_id) groupId: vec3<u32>) {\n";
107 os << " var i_: u32 = globalId.x;\n";
108 os << " if (i_ >= " << sn.size << "u) { return; }\n";
109 os << " var ax: u32 = (i_ / " << inner << "u) % " << axisSize << "u;\n";
110 os << " var op: u32 = i_ / (" << axisSize << "u * " << inner << "u);\n";
111 os << " var ip: u32 = i_ % " << inner << "u;\n";
112 os << " o[i_] = in_[op * " << X.dims[axis] << "u * " << inner << "u + (ax + " << begin << "u) * " << inner
113 << "u + ip];\n";
114 os << "}\n";
115 out.pass1.clear();
116 out.pass2 = os.str();
117 out.groupsX2 = groupsFor(sn.size);
118 out.inputCount = 1;
119 return;
120}
121
122void genPermute(const Graph &g, const FusedGroup &grp, KernelSpec &out) {
123 const GraphNode &pn = g.node(grp.outputNode);
124 const GraphNode &X = g.node(pn.in0);
125 const int rank = pn.rank;
126 int S[Tensor::kMaxRank] = {};
127 S[rank - 1] = 1;
128 for (int k = rank - 2; k >= 0; --k) S[k] = S[k + 1] * pn.dims[k + 1];
129 int inStride[Tensor::kMaxRank] = {};
130 inStride[rank - 1] = 1;
131 for (int k = rank - 2; k >= 0; --k) inStride[k] = inStride[k + 1] * X.dims[k + 1];
132 std::ostringstream os;
133 os << header(kLocalSize);
134 os << bufferDecl(0, "in_");
135 os << bufferDecl(1, "o");
136 os << pushConstant();
137 os << "@compute @workgroup_size(workgroupX, workgroupY)\nfn main(@builtin(global_invocation_id) globalId: "
138 "vec3<u32>, @builtin(local_invocation_id) localId: vec3<u32>, @builtin(workgroup_id) groupId: vec3<u32>) {\n";
139 os << " var i_: u32 = globalId.x;\n";
140 os << " if (i_ >= " << pn.size << "u) { return; }\n";
141 os << " var idx: u32 = 0u;\n";
142 for (int k = 0; k < rank; ++k) {
143 const int inAxis = pn.perm[k];
144 os << " idx += ((i_ / " << S[k] << "u) % " << pn.dims[k] << "u) * " << inStride[inAxis] << "u;\n";
145 }
146 os << " o[i_] = in_[idx];\n";
147 os << "}\n";
148 out.pass1.clear();
149 out.pass2 = os.str();
150 out.groupsX2 = groupsFor(pn.size);
151 out.inputCount = 1;
152 return;
153}
154
155void genResize2d(const Graph &g, const FusedGroup &grp, KernelSpec &out) {
156 const GraphNode &rn = g.node(grp.outputNode);
157 const GraphNode &X = g.node(rn.in0);
158 const int H = X.dims[2], W = X.dims[3];
159 const int OH = rn.dims[2], OW = rn.dims[3];
160 const bool nearest = rn.i0 == 0;
161 std::ostringstream os;
162 os << header(kLocalSize);
163 os << bufferDecl(0, "in_");
164 os << bufferDecl(1, "o");
165 os << pushConstant();
166 os << "@compute @workgroup_size(workgroupX, workgroupY)\nfn main(@builtin(global_invocation_id) globalId: "
167 "vec3<u32>, @builtin(local_invocation_id) localId: vec3<u32>, @builtin(workgroup_id) groupId: vec3<u32>) {\n";
168 os << " var i_: u32 = globalId.x;\n";
169 os << " if (i_ >= " << rn.size << "u) { return; }\n";
170 os << " var ow: u32 = i_ % " << OW << "u;\n";
171 os << " var rem: u32 = i_ / " << OW << "u;\n";
172 os << " var oh: u32 = rem % " << OH << "u;\n";
173 os << " var rem2: u32 = rem / " << OH << "u;\n";
174 os << " var c: u32 = rem2 % " << X.dims[1] << "u;\n";
175 os << " var n_: u32 = rem2 / " << X.dims[1] << "u;\n";
176 os << " var base: u32 = (n_ * " << X.dims[1] << "u + c) * " << H << "u * " << W << "u;\n";
177 if (nearest) {
178 os << " var ih: u32 = u32(f32(oh) * " << scalarStr(float(H) / OH) << ") ;\n";
179 os << " var iw: u32 = u32(f32(ow) * " << scalarStr(float(W) / OW) << ") ;\n";
180 os << " ih = min(ih, " << H - 1 << "u); iw = min(iw, " << W - 1 << "u);\n";
181 os << " o[i_] = in_[base + ih * " << W << "u + iw];\n";
182 } else {
183 os << " var fx: f32 = (f32(ow) + 0.5) * " << scalarStr(float(W) / OW) << " - 0.5;\n";
184 os << " var fy: f32 = (f32(oh) + 0.5) * " << scalarStr(float(H) / OH) << " - 0.5;\n";
185 os << " fx = clamp(fx, 0.0, " << scalarStr(float(W - 1)) << ");\n";
186 os << " fy = clamp(fy, 0.0, " << scalarStr(float(H - 1)) << ");\n";
187 os << " var x0: u32 = u32(floor(fx)); var y0: u32 = u32(floor(fy));\n";
188 os << " var x1: u32 = min(x0 + 1u, " << W - 1 << "u); var y1: u32 = min(y0 + 1u, " << H - 1 << "u);\n";
189 os << " var w00: f32 = in_[base + y0 * " << W << "u + x0];\n";
190 os << " var w10: f32 = in_[base + y0 * " << W << "u + x1];\n";
191 os << " var w01: f32 = in_[base + y1 * " << W << "u + x0];\n";
192 os << " var w11: f32 = in_[base + y1 * " << W << "u + x1];\n";
193 os << " var top: f32 = w00 + (w10 - w00) * (fx - f32(x0));\n";
194 os << " var bot: f32 = w01 + (w11 - w01) * (fx - f32(x0));\n";
195 os << " o[i_] = top + (bot - top) * (fy - f32(y0));\n";
196 }
197 os << "}\n";
198 out.pass1.clear();
199 out.pass2 = os.str();
200 out.groupsX2 = groupsFor(rn.size);
201 out.inputCount = 1;
202 return;
203}
204
205void genSdpa(const Graph &g, const FusedGroup &grp, KernelSpec &out) {
206 const GraphNode &qn = g.node(grp.outputNode);
207 const GraphNode &Q = g.node(qn.in0);
208 const GraphNode &K = g.node(qn.in1);
209 const int B = Q.dims[0], H = Q.dims[1], T = Q.dims[2], D = Q.dims[3];
210 const int S = K.dims[2];
211 if (S > 2048 || D > 512)
212 throw eve::Exception(
213 "Tensor WGSL: unsupported kernel variant or binding count"); // shared-memory limits -> CPU fallback
214 const float scale = qn.s0;
215 const bool masked = grp.masked;
216 const int bindingOut = masked ? 4 : 3;
217 std::ostringstream os;
218 os << header(128);
219 os << bufferDecl(0, "q");
220 os << bufferDecl(1, "k");
221 os << bufferDecl(2, "v");
222 if (masked) os << bufferDecl(3, "mask");
223 os << bufferDecl(bindingOut, "o");
224 os << "var<workgroup> scores: array<f32, " << S << ">;\n";
225 os << "var<workgroup> maxv: f32;\n";
226 os << "var<workgroup> sumv: f32;\n";
227 os << pushConstant();
228 os << "@compute @workgroup_size(workgroupX, workgroupY)\nfn main(@builtin(global_invocation_id) globalId: "
229 "vec3<u32>, @builtin(local_invocation_id) localId: vec3<u32>, @builtin(workgroup_id) groupId: vec3<u32>) {\n";
230 os << " var tid: u32 = localId.x;\n";
231 os << " var bh: u32 = groupId.x;\n";
232 os << " var t: u32 = groupId.y;\n";
233 os << " var b: u32 = bh / " << H << "u;\n";
234 os << " var h: u32 = bh % " << H << "u;\n";
235 os << " var qbase: u32 = (b * " << H << "u + h) * " << T << "u * " << D << "u + t * " << D << "u;\n";
236 os << " var kbase: u32 = (b * " << H << "u + h) * " << S << "u * " << D << "u;\n";
237 os << " var vbase: u32 = kbase;\n";
238 os << " for (var s: u32 = tid; s < " << S << "u; s += 128u) {\n";
239 os << " var acc: f32 = 0.0;\n";
240 os << " for (var d: u32 = 0u; d < " << D << "u; d++) { acc += q[qbase + d] * k[kbase + s * " << D
241 << "u + d]; }\n";
242 os << " acc *= " << scalarStr(scale) << ";\n";
243 if (masked) {
244 os << " acc += mask[(b * " << H << "u + h) * " << T << "u * " << S << "u + t * " << S << "u + s];\n";
245 }
246 os << " scores[s] = acc;\n";
247 os << " }\n";
248 os << " workgroupBarrier();\n";
249 os << " if (tid == 0u) {\n";
250 os << " var m: f32 = -3.402823e38;\n";
251 os << " for (var s: u32 = 0u; s < " << S << "u; s++) { m = max(m, scores[s]); }\n";
252 os << " var sm: f32 = 0.0;\n";
253 os << " for (var s: u32 = 0u; s < " << S << "u; s++) { sm += exp(scores[s] - m); }\n";
254 os << " maxv = m; sumv = sm;\n";
255 os << " }\n";
256 os << " workgroupBarrier();\n";
257 os << " for (var d: u32 = tid; d < " << D << "u; d += 128u) {\n";
258 os << " var acc: f32 = 0.0;\n";
259 os << " for (var s: u32 = 0u; s < " << S << "u; s++) { acc += exp(scores[s] - maxv) * v[vbase + s * " << D
260 << "u + d]; }\n";
261 os << " o[qbase + d] = acc / sumv;\n";
262 os << " }\n";
263 os << "}\n";
264 out.pass1.clear();
265 out.pass2 = os.str();
266 out.groupsX2 = B * H;
267 out.groupsY2 = T;
268 out.inputCount = masked ? 4 : 3;
269 return;
270}
271
272
273} // namespace eve::tensor::wgsl_detail
tensor::Graph g
Definition GpuGraph.cpp:7
glm::vec3 n
Definition Grass.cpp:63
std::array< float, 3 > scale
float begin
EVENGINE_API_FOUNDATION public API.
Definition Exception.h:13
EVENGINE_API_DOMAINS public API.
Definition Graph.h:111
static constexpr int kMaxRank
Definition Tensor.h:50
bool isQuantDType(DType dt)
True when quant d type.
Definition Quant.h:17
std::string scalarStr(float v)
Scalar str.
void genSlice(const Graph &g, const FusedGroup &grp, KernelSpec &out)
Gen slice.
void genConcat(const Graph &g, const FusedGroup &grp, KernelSpec &out)
Gen concat.
void genResize2d(const Graph &g, const FusedGroup &grp, KernelSpec &out)
Gen resize 2 d.
std::string bufferDeclUint(int binding, const char *name)
Buffer decl uint.
void genPermute(const Graph &g, const FusedGroup &grp, KernelSpec &out)
Gen permute.
void genEmbedding(const Graph &g, const FusedGroup &grp, KernelSpec &out)
Gen embedding.
std::string bufferDecl(int binding, const char *name)
Buffer decl.
std::string pushConstant()
Pushes constant.
int groupsFor(int count)
Groups for.
void genSdpa(const Graph &g, const FusedGroup &grp, KernelSpec &out)
Gen sdpa.
std::string header(int localX, int localY)
Header.
DType
Tensor element types.
Definition Tensor.h:24
FusedGroup public API.
Definition Optimizer.h:46
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