载入中...
搜索中...
未找到
KernelGenWgslReduction.cpp
浏览该文件的文档.
1#include <algorithm>
2#include <cmath>
3#include <sstream>
4#include "common/Exception.h"
6#include "tensor/Quant.h"
7
9void genSoftmax(const Graph &g, const FusedGroup &grp, KernelSpec &out) {
10 const GraphNode &sn = g.node(grp.outputNode);
11 const GraphNode &X = g.node(sn.in0);
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];
15 reduce = X.dims[axis];
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;
19
20 std::ostringstream os1, os2;
21 os1 << header(kLocalSize);
22 os1 << bufferDecl(0, "in_");
23 os1 << bufferDecl(1, "mx");
24 os1 << bufferDecl(2, "sm");
25 os1 << pushConstant();
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>) "
28 "{\n";
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";
36 os1 << " }\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";
40 os1 << " }\n";
41 os1 << " mx[i_] = m;\n";
42 os1 << " sm[i_] = s;\n";
43 os1 << "}\n";
44
45 os2 << header(kLocalSize);
46 os2 << bufferDecl(0, "in_");
47 os2 << bufferDecl(1, "mx");
48 os2 << bufferDecl(2, "sm");
49 os2 << bufferDecl(3, "o");
50 os2 << pushConstant();
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>) "
53 "{\n";
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";
60 if (logMode) {
61 os2 << " o[i_] = (x - m) - log(s);\n";
62 } else {
63 os2 << " o[i_] = exp(x - m) / s;\n";
64 }
65 os2 << "}\n";
66 out.pass1 = os1.str();
67 out.pass2 = os2.str();
68 out.groupsX1 = groupsFor(rows);
69 out.groupsX2 = groupsFor(sn.size);
70 out.inputCount = 1;
71 out.inputsReadPass1 = 1;
72 out.statsCount = 2;
73 out.statsSize = rows;
74 out.twoPass = true;
75 return;
76}
77
78void genNorm(const Graph &g, const FusedGroup &grp, bool rms, KernelSpec &out) {
79 const GraphNode &nn = g.node(grp.outputNode);
80 const GraphNode &X = g.node(nn.in0);
81 const int cols = X.dims[X.rank - 1];
82 const int rows = X.size / cols;
83 const float eps = nn.s0;
84 const bool hasScale = grp.hasScale;
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;
89
90 std::ostringstream os1, os2;
91 os1 << header(kLocalSize);
92 os1 << bufferDecl(0, "in_");
93 os1 << bufferDecl(1, "st0");
94 if (statsCount > 1) os1 << bufferDecl(2, "st1");
95 os1 << pushConstant();
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>) "
98 "{\n";
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";
105 if (rms) {
106 os1 << " s0 += v * v;\n";
107 } else {
108 os1 << " s0 += v; s1 += v * v;\n";
109 }
110 os1 << " }\n";
111 os1 << " st0[i_] = s0;\n";
112 if (statsCount > 1) os1 << " st1[i_] = s1;\n";
113 os1 << "}\n";
114
115 os2 << header(kLocalSize);
116 for (int k = 0; k < inputCount; ++k) os2 << bufferDecl(k, ("a" + std::to_string(k)).c_str());
117 os2 << bufferDecl(inputCount, "st0");
118 if (statsCount > 1) os2 << bufferDecl(inputCount + 1, "st1");
119 os2 << bufferDecl(outBinding, "o");
120 os2 << pushConstant();
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>) "
123 "{\n";
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";
128 if (rms) {
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";
132 } else {
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";
140 }
141 os2 << " o[i_] = y;\n";
142 os2 << "}\n";
143 out.pass1 = os1.str();
144 out.pass2 = os2.str();
145 out.groupsX1 = groupsFor(rows);
146 out.groupsX2 = groupsFor(nn.size);
147 out.inputCount = inputCount;
148 out.inputsReadPass1 = 1;
149 out.statsCount = statsCount;
150 out.statsSize = rows;
151 out.twoPass = true;
152 return;
153}
154
155void genReduceOrArgmax(const Graph &g, const FusedGroup &grp, bool argmax, KernelSpec &out) {
156 const GraphNode &rn = g.node(grp.outputNode);
157 const GraphNode &X = g.node(rn.in0);
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];
161 reduce = X.dims[axis];
162 for (int k = axis + 1; k < X.rank; ++k) inner *= X.dims[k];
163 const int outSize = outer * inner;
164 std::ostringstream os;
165 os << header(kLocalSize);
166 os << bufferDecl(0, "in_");
167 os << bufferDecl(1, "o");
168 os << pushConstant();
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";
175 if (argmax) {
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";
181 os << " }\n";
182 os << " o[i_] = bestJ;\n";
183 } else {
184 switch (grp.op) {
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";
190 if (grp.op == OpType::ReduceMean) os << " acc /= " << scalarStr(float(reduce)) << ";\n";
191 break;
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";
196 break;
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";
201 break;
202 default: throw eve::Exception("Tensor WGSL: unsupported kernel variant or binding count");
203 }
204 os << " o[i_] = acc;\n";
205 }
206 os << "}\n";
207 out.pass1.clear();
208 out.pass2 = os.str();
209 out.groupsX2 = groupsFor(outSize);
210 out.inputCount = 1;
211 return;
212}
213
214
215} // namespace eve::tensor::wgsl_detail
int rows
int cols
tensor::Graph g
Definition GpuGraph.cpp:7
EVENGINE_API_FOUNDATION public API.
Definition Exception.h:13
EVENGINE_API_DOMAINS public API.
Definition Graph.h:111
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.
FusedGroup public API.
Definition Optimizer.h:46
GraphNode public API.
Definition Graph.h:80
KernelSpec public API.
Definition KernelGen.h:27
gpgpu::ComputeShader * reduce