14 const int vocab = T.
dims[0], dim = T.
dims[1];
15 std::ostringstream os;
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";
35 << (tQuant ?
"bval(u32(ii) * " + std::to_string(dim) +
"u + d)"
36 :
"table[u32(ii) * " + std::to_string(dim) +
"u + d]")
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");
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];
62 for (
int k = axis + 1; k < cn.
rank; ++k) inner *= cn.
dims[k];
63 std::ostringstream os;
65 for (
int k = 0; k <
n; ++k) os <<
bufferDecl(k, (
"a" + std::to_string(k)).c_str());
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
84 os <<
" o[i_] = v;\n";
99 for (
int k = axis + 1; k < sn.
rank; ++k) inner *= sn.
dims[k];
100 std::ostringstream os;
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
116 out.
pass2 = os.str();
125 const int rank = pn.
rank;
128 for (
int k = rank - 2; k >= 0; --k) S[k] = S[k + 1] * pn.
dims[k + 1];
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;
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";
146 os <<
" o[i_] = in_[idx];\n";
149 out.
pass2 = os.str();
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;
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";
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";
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";
199 out.
pass2 = os.str();
210 const int S = K.
dims[2];
211 if (S > 2048 || D > 512)
213 "Tensor WGSL: unsupported kernel variant or binding count");
215 const bool masked = grp.
masked;
216 const int bindingOut = masked ? 4 : 3;
217 std::ostringstream os;
224 os <<
"var<workgroup> scores: array<f32, " << S <<
">;\n";
225 os <<
"var<workgroup> maxv: f32;\n";
226 os <<
"var<workgroup> sumv: f32;\n";
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
244 os <<
" acc += mask[(b * " << H <<
"u + h) * " << T <<
"u * " << S <<
"u + t * " << S <<
"u + s];\n";
246 os <<
" scores[s] = acc;\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";
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
261 os <<
" o[qbase + d] = acc / sumv;\n";
265 out.
pass2 = os.str();
std::array< float, 3 > scale
EVENGINE_API_FOUNDATION public API.
EVENGINE_API_DOMAINS public API.
static constexpr int kMaxRank
bool isQuantDType(DType dt)
True when quant d type.
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.
std::vector< int > inputs
int perm[Tensor::kMaxRank]
int dims[Tensor::kMaxRank]
std::vector< uint8_t > constBytes