Add exaone4-fp32-8k-state-alias-token-major-v2-v4 (tiled prefill GEMM over the v3 bundle)
Browse filesThis view is limited to 50 files because it contains too many changes. See raw diff
- exaone4-fp32-8k-state-alias-token-major-v2-v4/entrypoints/prefill-16.json +0 -0
- exaone4-fp32-8k-state-alias-token-major-v2-v4/entrypoints/prefill-4.json +0 -0
- exaone4-fp32-8k-state-alias-token-major-v2-v4/entrypoints/prefill-64.json +0 -0
- exaone4-fp32-8k-state-alias-token-major-v2-v4/graph.json +0 -0
- exaone4-fp32-8k-state-alias-token-major-v2-v4/kernels/0112d93bb1aa347d8887ce1b95610a1a21fc8450ff8f59bcd37747fa6bf6fa52.wgsl +10 -0
- exaone4-fp32-8k-state-alias-token-major-v2-v4/kernels/02581e105cf669277edc5f2212ac9b4d31bfbeb797f41f0953679df01f877c67.wgsl +83 -0
- exaone4-fp32-8k-state-alias-token-major-v2-v4/kernels/0337d8a473b6e485bd383533df7d373dbbfe7a5f3a982424004560b7213fbaa3.wgsl +11 -0
- exaone4-fp32-8k-state-alias-token-major-v2-v4/kernels/033ad0e2dfc972b75e5a7861b4f0f237c46b27fd9b1dcd22f01da3a1c2ad3035.wgsl +13 -0
- exaone4-fp32-8k-state-alias-token-major-v2-v4/kernels/03f7b188faa78a1f519a12d8b8866be338dbe7b73d0d10d0b3d2712df698d9ed.wgsl +55 -0
- exaone4-fp32-8k-state-alias-token-major-v2-v4/kernels/07cfabe2e2c2feff22a06c8296600e5ca21efb7444419f0a05e5ff6f48e0292d.wgsl +10 -0
- exaone4-fp32-8k-state-alias-token-major-v2-v4/kernels/09d52a47128189748341738ed18c539c6adb15107224cdf66db5c3f272ede305.wgsl +13 -0
- exaone4-fp32-8k-state-alias-token-major-v2-v4/kernels/0b8815b1efe176d55ab08530c0a9e1913cf6c9ef3e773c574843573d237cf196.wgsl +10 -0
- exaone4-fp32-8k-state-alias-token-major-v2-v4/kernels/0d3892883d1532d173a7dbe85d9ab1daa87f0d9b4923a535b8d2003f24f9500c.wgsl +11 -0
- exaone4-fp32-8k-state-alias-token-major-v2-v4/kernels/0e2ba0d1849cd945e9b533bfc599a5664fa7909515c955681084ab8ad82f8755.wgsl +9 -0
- exaone4-fp32-8k-state-alias-token-major-v2-v4/kernels/11723b648ae4baeb93ebba06009266508b8bb3e9185c049eb4ec6e85199f045b.wgsl +83 -0
- exaone4-fp32-8k-state-alias-token-major-v2-v4/kernels/133b3d1c9ddb9953a9d3cffa5e3bd5e837cf2474ec6641a311a54d911d958f6b.wgsl +9 -0
- exaone4-fp32-8k-state-alias-token-major-v2-v4/kernels/14d88fb2ba88841a93259cafdf9a2cbf86411c82ec5da801fcd8a887de2ec80d.wgsl +13 -0
- exaone4-fp32-8k-state-alias-token-major-v2-v4/kernels/1533ec909b782828ba90ded15ebe230c6f9dbfbeb74335606f979f2b56c58dbf.wgsl +9 -0
- exaone4-fp32-8k-state-alias-token-major-v2-v4/kernels/1701e3705e0212caf2880d0a27ccb420ef83e5c95aeb2e9a0679eacf9c2e0a87.wgsl +21 -0
- exaone4-fp32-8k-state-alias-token-major-v2-v4/kernels/18a66ec0875c456c9defd45d15e6ef0c148e64a1618e2825c7f96eadbbe8303a.wgsl +9 -0
- exaone4-fp32-8k-state-alias-token-major-v2-v4/kernels/1a73b6fceda28f40bba4ae27a986ffa39842db94fa5903ea5bf689a7e854d4f3.wgsl +11 -0
- exaone4-fp32-8k-state-alias-token-major-v2-v4/kernels/1aec53ecf59335904944de2db37f0d3df64e3445556b40324ce611d4c9fcfc2c.wgsl +9 -0
- exaone4-fp32-8k-state-alias-token-major-v2-v4/kernels/1c5b119286250e322ec8f75a46494057ed2485c4ce8ea081b4858aa5fbeadd0d.wgsl +9 -0
- exaone4-fp32-8k-state-alias-token-major-v2-v4/kernels/1d2ff3e2ffac914e03d1495655b136298374a8b07c94134a60bdb6927cbddcd0.wgsl +83 -0
- exaone4-fp32-8k-state-alias-token-major-v2-v4/kernels/1d3137ed94b2dd9113466681ea64a1e9a3086acbadc3c2e2a35f30dc42c67ad5.wgsl +9 -0
- exaone4-fp32-8k-state-alias-token-major-v2-v4/kernels/214ac9825bc5014847277240f7d17b1f0aa4f2272175a7d77e94116dc75f9fff.wgsl +13 -0
- exaone4-fp32-8k-state-alias-token-major-v2-v4/kernels/216a5956ea2a40590a4de1b09152aacde8cac66cbd448be13a81a5ebb828cffa.wgsl +55 -0
- exaone4-fp32-8k-state-alias-token-major-v2-v4/kernels/27a921d3b5edfdca227a99e1c3328959881c09acb02b12cd09d07d598745b2d3.wgsl +9 -0
- exaone4-fp32-8k-state-alias-token-major-v2-v4/kernels/28119aa01122c42bd26ceb37129fe0c95b2a9192163f6dabd90fe058b1e14c54.wgsl +10 -0
- exaone4-fp32-8k-state-alias-token-major-v2-v4/kernels/284513e54dd254e923b7a86bc5703879e3bfb119b6c71aa234ba744c2dc63615.wgsl +9 -0
- exaone4-fp32-8k-state-alias-token-major-v2-v4/kernels/2b9e4b05ef5623d721dbe56b5812e551efcc6fd833d81064320e86fe94d8c670.wgsl +9 -0
- exaone4-fp32-8k-state-alias-token-major-v2-v4/kernels/2c66dcac8fdd3caa9ff1e2441524070c21c24df9526bee44c49e765a8d4c8222.wgsl +9 -0
- exaone4-fp32-8k-state-alias-token-major-v2-v4/kernels/2e1ce6ef01198780289fd4a35d8f86b5172595ed189cf931e8444091d49e8d3a.wgsl +9 -0
- exaone4-fp32-8k-state-alias-token-major-v2-v4/kernels/3203e08f2c99fa4848f553739b7115a7c677840344f17b1d4391d6130517f5a7.wgsl +9 -0
- exaone4-fp32-8k-state-alias-token-major-v2-v4/kernels/3241a118927e109c0dacbb278370abe70d099ee593a07cd023d4530b5fb8f833.wgsl +24 -0
- exaone4-fp32-8k-state-alias-token-major-v2-v4/kernels/33e8ec04c66adb4988f562f0f9d3e92d29dfda91a56d0ed62be598c7c737f5ba.wgsl +9 -0
- exaone4-fp32-8k-state-alias-token-major-v2-v4/kernels/358e73d0a8169937dce5d2c7fd5f4dea41f7c8307696f1698967bc6c4132efc8.wgsl +13 -0
- exaone4-fp32-8k-state-alias-token-major-v2-v4/kernels/37ce78f0b50bc59dc7b6dd62724d6ad3acfbfb73deae12446a4575f5688d8f34.wgsl +9 -0
- exaone4-fp32-8k-state-alias-token-major-v2-v4/kernels/383d0ce2bbe657b6412b8e55e6b0cf06e385cb4efbe45019afff9357f2a77751.wgsl +11 -0
- exaone4-fp32-8k-state-alias-token-major-v2-v4/kernels/3866ae40efed23181bf612efca4329d5348debc12468bf7565a2f13889b73bea.wgsl +13 -0
- exaone4-fp32-8k-state-alias-token-major-v2-v4/kernels/39646bdfafd1ab0cb874474e3f1e9d79b57f10582a1804bc1d3762a492fec851.wgsl +10 -0
- exaone4-fp32-8k-state-alias-token-major-v2-v4/kernels/3b114a6518b88e4d199413a9dffd2ff911ae539546995b683c9aa0f86462044b.wgsl +55 -0
- exaone4-fp32-8k-state-alias-token-major-v2-v4/kernels/3b495259612ab4c2e36bc153ec19119e95ab0d425693674839363ee4339a2770.wgsl +21 -0
- exaone4-fp32-8k-state-alias-token-major-v2-v4/kernels/3e34409407100d4858c4c109cf62f3eaf2f4c6081166b0ec8679105e6dfbe7fe.wgsl +10 -0
- exaone4-fp32-8k-state-alias-token-major-v2-v4/kernels/40f115c8647a238bfb3ccb0bb8446bccdc94a0dfadc89111c4920efb2259f41d.wgsl +10 -0
- exaone4-fp32-8k-state-alias-token-major-v2-v4/kernels/453668cbe741eb27e2cc10194eca66f514a3763d97b46c1733c376986defc322.wgsl +24 -0
- exaone4-fp32-8k-state-alias-token-major-v2-v4/kernels/476f0abdcc4ba14abb8e304788e71549fd6b18fac0f22e03853683751aa205dc.wgsl +55 -0
- exaone4-fp32-8k-state-alias-token-major-v2-v4/kernels/483857d2d9929bc07cd14940dc3bca97abcf9485ce1c805a234d2db4110a7010.wgsl +10 -0
- exaone4-fp32-8k-state-alias-token-major-v2-v4/kernels/4c3a20bb702b43bacecf9c7df989b55517469eba683bb2cb596e1f7a227347c7.wgsl +13 -0
- exaone4-fp32-8k-state-alias-token-major-v2-v4/kernels/4c526eba157bfd1aee20a2637586d6f4cdcdc68dcf238dcef84647649ffc6f9e.wgsl +13 -0
exaone4-fp32-8k-state-alias-token-major-v2-v4/entrypoints/prefill-16.json
ADDED
|
The diff for this file is too large to render.
See raw diff
|
|
|
exaone4-fp32-8k-state-alias-token-major-v2-v4/entrypoints/prefill-4.json
ADDED
|
The diff for this file is too large to render.
See raw diff
|
|
|
exaone4-fp32-8k-state-alias-token-major-v2-v4/entrypoints/prefill-64.json
ADDED
|
The diff for this file is too large to render.
See raw diff
|
|
|
exaone4-fp32-8k-state-alias-token-major-v2-v4/graph.json
ADDED
|
The diff for this file is too large to render.
See raw diff
|
|
|
exaone4-fp32-8k-state-alias-token-major-v2-v4/kernels/0112d93bb1aa347d8887ce1b95610a1a21fc8450ff8f59bcd37747fa6bf6fa52.wgsl
ADDED
|
@@ -0,0 +1,10 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
@group(0) @binding(0) var<storage, read> b0: array<f32>;
|
| 2 |
+
@group(0) @binding(1) var<storage, read> b1: array<f32>;
|
| 3 |
+
@group(0) @binding(2) var<storage, read_write> out: array<f32>;
|
| 4 |
+
|
| 5 |
+
@compute @workgroup_size(64)
|
| 6 |
+
fn main(@builtin(global_invocation_id) gid: vec3<u32>) {
|
| 7 |
+
let i = gid.x + gid.y * 2048u;
|
| 8 |
+
if (i >= 2048u) { return; }
|
| 9 |
+
out[i] = f32(f32(b0[i]) + (f32(b1[i]) * 1.0));
|
| 10 |
+
}
|
exaone4-fp32-8k-state-alias-token-major-v2-v4/kernels/02581e105cf669277edc5f2212ac9b4d31bfbeb797f41f0953679df01f877c67.wgsl
ADDED
|
@@ -0,0 +1,83 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
@group(0) @binding(0) var<storage, read> b0: array<f32>;
|
| 2 |
+
@group(0) @binding(1) var<storage, read> b1: array<f32>;
|
| 3 |
+
@group(0) @binding(2) var<storage, read> b2: array<f32>;
|
| 4 |
+
@group(0) @binding(3) var<storage, read> b3: array<f32>;
|
| 5 |
+
@group(0) @binding(4) var<storage, read_write> out: array<f32>;
|
| 6 |
+
|
| 7 |
+
var<workgroup> q_row: array<f32, 64>;
|
| 8 |
+
var<workgroup> scores: array<f32, 1024>;
|
| 9 |
+
var<workgroup> partial: array<f32, 64>;
|
| 10 |
+
@compute @workgroup_size(64)
|
| 11 |
+
fn main(@builtin(workgroup_id) group: vec3<u32>, @builtin(local_invocation_id) local: vec3<u32>) {
|
| 12 |
+
let lane = local.x;
|
| 13 |
+
let head = group.x / 16u;
|
| 14 |
+
let row = group.x % 16u;
|
| 15 |
+
let kv_head = head / 4u;
|
| 16 |
+
for (var d = lane; d < 64u; d += 64u) { q_row[d] = f32(b0[(head * 16u + row) * 64u + d]); }
|
| 17 |
+
workgroupBarrier();
|
| 18 |
+
var running_max = -3.4028234663852886e38;
|
| 19 |
+
var running_sum = 0.0;
|
| 20 |
+
var acc: array<f32, 1>;
|
| 21 |
+
for (var start = 0u; start < 8192u; start += 1024u) {
|
| 22 |
+
let count = min(1024u, 8192u - start);
|
| 23 |
+
var tile_max = -3.4028234663852886e38;
|
| 24 |
+
for (var t = lane; t < count; t += 64u) {
|
| 25 |
+
let j = start + t;
|
| 26 |
+
var score = -3.4028234663852886e38;
|
| 27 |
+
if (f32(b3[row * 8192u + j]) > -1000.0) {
|
| 28 |
+
var dot = 0.0;
|
| 29 |
+
for (var d = 0u; d < 64u; d++) { dot += q_row[d] * f32(b1[(((kv_head * 8192u + j) * 64u + d) / 64u) % 8192u * 512u + (((kv_head * 8192u + j) * 64u + d) / 524288u) % 8u * 64u + (((kv_head * 8192u + j) * 64u + d) / 1u) % 64u * 1u]); }
|
| 30 |
+
score = dot * 0.125 + f32(b3[row * 8192u + j]);
|
| 31 |
+
}
|
| 32 |
+
scores[t] = score;
|
| 33 |
+
tile_max = max(tile_max, score);
|
| 34 |
+
}
|
| 35 |
+
partial[lane] = tile_max;
|
| 36 |
+
workgroupBarrier();
|
| 37 |
+
for (var stride = 32u; stride > 0u; stride /= 2u) {
|
| 38 |
+
if (lane < stride) { partial[lane] = max(partial[lane], partial[lane + stride]); }
|
| 39 |
+
workgroupBarrier();
|
| 40 |
+
}
|
| 41 |
+
let next_max = max(running_max, partial[0]);
|
| 42 |
+
workgroupBarrier();
|
| 43 |
+
var tile_sum = 0.0;
|
| 44 |
+
for (var t = lane; t < count; t += 64u) {
|
| 45 |
+
var p = 0.0;
|
| 46 |
+
if (scores[t] > -3.0e38) { p = exp(scores[t] - next_max); }
|
| 47 |
+
scores[t] = p;
|
| 48 |
+
tile_sum += p;
|
| 49 |
+
}
|
| 50 |
+
partial[lane] = tile_sum;
|
| 51 |
+
workgroupBarrier();
|
| 52 |
+
for (var stride = 32u; stride > 0u; stride /= 2u) {
|
| 53 |
+
if (lane < stride) { partial[lane] += partial[lane + stride]; }
|
| 54 |
+
workgroupBarrier();
|
| 55 |
+
}
|
| 56 |
+
var correction = 0.0;
|
| 57 |
+
if (running_max > -3.0e38) { correction = exp(running_max - next_max); }
|
| 58 |
+
running_sum = running_sum * correction + partial[0];
|
| 59 |
+
running_max = next_max;
|
| 60 |
+
for (var c = 0u; c < 1u; c++) {
|
| 61 |
+
let d = lane + c * 64u;
|
| 62 |
+
var total = acc[c] * correction;
|
| 63 |
+
if (d < 64u) {
|
| 64 |
+
for (var t = 0u; t < count; t++) {
|
| 65 |
+
let p = scores[t];
|
| 66 |
+
if (p != 0.0) {
|
| 67 |
+
let j = start + t;
|
| 68 |
+
total += p * f32(b2[(((kv_head * 8192u + j) * 64u + d) / 64u) % 8192u * 512u + (((kv_head * 8192u + j) * 64u + d) / 524288u) % 8u * 64u + (((kv_head * 8192u + j) * 64u + d) / 1u) % 64u * 1u]);
|
| 69 |
+
}
|
| 70 |
+
}
|
| 71 |
+
}
|
| 72 |
+
acc[c] = total;
|
| 73 |
+
}
|
| 74 |
+
workgroupBarrier();
|
| 75 |
+
}
|
| 76 |
+
for (var c = 0u; c < 1u; c++) {
|
| 77 |
+
let d = lane + c * 64u;
|
| 78 |
+
if (d < 64u) {
|
| 79 |
+
let i = (head * 16u + row) * 64u + d;
|
| 80 |
+
out[i] = f32(acc[c] / running_sum);
|
| 81 |
+
}
|
| 82 |
+
}
|
| 83 |
+
}
|
exaone4-fp32-8k-state-alias-token-major-v2-v4/kernels/0337d8a473b6e485bd383533df7d373dbbfe7a5f3a982424004560b7213fbaa3.wgsl
ADDED
|
@@ -0,0 +1,11 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
@group(0) @binding(0) var<storage, read> b0: array<f32>;
|
| 2 |
+
@group(0) @binding(1) var<storage, read_write> out: array<f32>;
|
| 3 |
+
|
| 4 |
+
@compute @workgroup_size(64)
|
| 5 |
+
fn main(@builtin(global_invocation_id) gid: vec3<u32>) {
|
| 6 |
+
let i = gid.x + gid.y * 4096u;
|
| 7 |
+
if (i >= 4096u) { return; }
|
| 8 |
+
let coord = (i / 1u) % 64u;
|
| 9 |
+
if (coord >= 0u && coord < 32u) { out[i] = f32(b0[(i / 64u) * 32u + (coord - 0u) * 1u + i % 1u]); }
|
| 10 |
+
if (coord >= 32u && coord < 64u) { out[i] = f32(b0[(i / 64u) * 32u + (coord - 32u) * 1u + i % 1u]); }
|
| 11 |
+
}
|
exaone4-fp32-8k-state-alias-token-major-v2-v4/kernels/033ad0e2dfc972b75e5a7861b4f0f237c46b27fd9b1dcd22f01da3a1c2ad3035.wgsl
ADDED
|
@@ -0,0 +1,13 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
enable f16;
|
| 2 |
+
@group(0) @binding(0) var<storage, read> b0: array<i32>;
|
| 3 |
+
@group(0) @binding(1) var<storage, read> b1: array<f16>;
|
| 4 |
+
@group(0) @binding(2) var<storage, read_write> out: array<f32>;
|
| 5 |
+
|
| 6 |
+
@compute @workgroup_size(64)
|
| 7 |
+
fn main(@builtin(global_invocation_id) gid: vec3<u32>) {
|
| 8 |
+
let i = gid.x + gid.y * 8192u;
|
| 9 |
+
if (i >= 8192u) { return; }
|
| 10 |
+
let token = i32(b0[i / 2048u]);
|
| 11 |
+
if (token < 98304 || token >= 102400) { out[i] = f32(0.0); return; }
|
| 12 |
+
out[i] = f32(b1[u32(token - 98304) * 2048u + i % 2048u]);
|
| 13 |
+
}
|
exaone4-fp32-8k-state-alias-token-major-v2-v4/kernels/03f7b188faa78a1f519a12d8b8866be338dbe7b73d0d10d0b3d2712df698d9ed.wgsl
ADDED
|
@@ -0,0 +1,55 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
enable f16;
|
| 2 |
+
@group(0) @binding(0) var<storage, read> b0: array<f32>;
|
| 3 |
+
@group(0) @binding(1) var<storage, read> b1: array<f16>;
|
| 4 |
+
@group(0) @binding(2) var<storage, read_write> out: array<f32>;
|
| 5 |
+
|
| 6 |
+
var<workgroup> tile_a: array<array<f32, 16>, 16>;
|
| 7 |
+
var<workgroup> tile_b: array<array<f32, 64>, 16>;
|
| 8 |
+
@compute @workgroup_size(16, 16)
|
| 9 |
+
fn main(@builtin(workgroup_id) group: vec3<u32>, @builtin(local_invocation_id) local: vec3<u32>) {
|
| 10 |
+
let lane = local.y * 16u + local.x;
|
| 11 |
+
let batch = group.z;
|
| 12 |
+
let tile_n = group.x * 64u;
|
| 13 |
+
let tile_row = group.y * 16u;
|
| 14 |
+
var acc: array<array<f32, 4>, 1>;
|
| 15 |
+
for (var k0 = 0u; k0 < 4096u; k0 += 16u) {
|
| 16 |
+
for (var e = 0u; e < 1u; e++) {
|
| 17 |
+
let flat = lane + e * 256u;
|
| 18 |
+
let m_local = flat / 16u;
|
| 19 |
+
let row = tile_row + m_local;
|
| 20 |
+
let col = k0 + flat % 16u;
|
| 21 |
+
var value = 0.0;
|
| 22 |
+
if (row < 4u && col < 4096u) { value = f32(f32(b0[(batch * 4u + row) * 4096u + col])); }
|
| 23 |
+
tile_a[m_local][flat % 16u] = value;
|
| 24 |
+
}
|
| 25 |
+
for (var e = 0u; e < 4u; e++) {
|
| 26 |
+
let n_local = lane / 4u;
|
| 27 |
+
let k_local = (lane % 4u) * 4u + e;
|
| 28 |
+
let n_index = tile_n + n_local;
|
| 29 |
+
let col = k0 + k_local;
|
| 30 |
+
var value = 0.0;
|
| 31 |
+
if (n_index < 2048u && col < 4096u) { value = f32(f32(b1[n_index * 4096u + col])); }
|
| 32 |
+
tile_b[k_local][n_local] = value;
|
| 33 |
+
}
|
| 34 |
+
workgroupBarrier();
|
| 35 |
+
for (var kk = 0u; kk < 16u; kk++) {
|
| 36 |
+
var b_values: array<f32, 4>;
|
| 37 |
+
for (var c = 0u; c < 4u; c++) { b_values[c] = tile_b[kk][local.x * 4u + c]; }
|
| 38 |
+
for (var r = 0u; r < 1u; r++) {
|
| 39 |
+
let a_value = tile_a[local.y * 1u + r][kk];
|
| 40 |
+
for (var c = 0u; c < 4u; c++) { acc[r][c] += a_value * b_values[c]; }
|
| 41 |
+
}
|
| 42 |
+
}
|
| 43 |
+
workgroupBarrier();
|
| 44 |
+
}
|
| 45 |
+
for (var r = 0u; r < 1u; r++) {
|
| 46 |
+
let out_row = tile_row + local.y * 1u + r;
|
| 47 |
+
for (var c = 0u; c < 4u; c++) {
|
| 48 |
+
let out_col = tile_n + local.x * 4u + c;
|
| 49 |
+
if (out_row < 4u && out_col < 2048u) {
|
| 50 |
+
let i = (batch * 4u + out_row) * 2048u + out_col;
|
| 51 |
+
out[i] = f32(acc[r][c] + 0.0);
|
| 52 |
+
}
|
| 53 |
+
}
|
| 54 |
+
}
|
| 55 |
+
}
|
exaone4-fp32-8k-state-alias-token-major-v2-v4/kernels/07cfabe2e2c2feff22a06c8296600e5ca21efb7444419f0a05e5ff6f48e0292d.wgsl
ADDED
|
@@ -0,0 +1,10 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
@group(0) @binding(0) var<storage, read> b0: array<f32>;
|
| 2 |
+
@group(0) @binding(1) var<storage, read_write> out: array<f32>;
|
| 3 |
+
|
| 4 |
+
@compute @workgroup_size(64)
|
| 5 |
+
fn main(@builtin(global_invocation_id) gid: vec3<u32>) {
|
| 6 |
+
let i = gid.x + gid.y * 1024u;
|
| 7 |
+
if (i >= 1024u) { return; }
|
| 8 |
+
let x = f32(b0[i]);
|
| 9 |
+
out[i] = f32(-x);
|
| 10 |
+
}
|
exaone4-fp32-8k-state-alias-token-major-v2-v4/kernels/09d52a47128189748341738ed18c539c6adb15107224cdf66db5c3f272ede305.wgsl
ADDED
|
@@ -0,0 +1,13 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
@group(0) @binding(0) var<storage, read> b0: array<f32>;
|
| 2 |
+
@group(0) @binding(1) var<storage, read> b1: array<i32>;
|
| 3 |
+
@group(0) @binding(2) var<storage, read_write> out: array<f32>;
|
| 4 |
+
|
| 5 |
+
@compute @workgroup_size(64)
|
| 6 |
+
fn main(@builtin(global_invocation_id) gid: vec3<u32>) {
|
| 7 |
+
let i = gid.x + gid.y * 2048u;
|
| 8 |
+
if (i >= 2048u) { return; }
|
| 9 |
+
let token = (i / 64u) % 4u;
|
| 10 |
+
let outer = i / 256u;
|
| 11 |
+
let destination = outer * 524288u + u32(b1[token]) * 64u + i % 64u;
|
| 12 |
+
out[((destination) / 64u) % 8192u * 512u + ((destination) / 524288u) % 8u * 64u + ((destination) / 1u) % 64u * 1u] = f32(b0[i]);
|
| 13 |
+
}
|
exaone4-fp32-8k-state-alias-token-major-v2-v4/kernels/0b8815b1efe176d55ab08530c0a9e1913cf6c9ef3e773c574843573d237cf196.wgsl
ADDED
|
@@ -0,0 +1,10 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
@group(0) @binding(0) var<storage, read> b0: array<f32>;
|
| 2 |
+
@group(0) @binding(1) var<storage, read_write> out: array<f32>;
|
| 3 |
+
|
| 4 |
+
@compute @workgroup_size(64)
|
| 5 |
+
fn main(@builtin(global_invocation_id) gid: vec3<u32>) {
|
| 6 |
+
let i = gid.x + gid.y * 256u;
|
| 7 |
+
if (i >= 256u) { return; }
|
| 8 |
+
let x = f32(b0[i]);
|
| 9 |
+
out[i] = f32(sin(x));
|
| 10 |
+
}
|
exaone4-fp32-8k-state-alias-token-major-v2-v4/kernels/0d3892883d1532d173a7dbe85d9ab1daa87f0d9b4923a535b8d2003f24f9500c.wgsl
ADDED
|
@@ -0,0 +1,11 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
@group(0) @binding(0) var<storage, read> b0: array<f32>;
|
| 2 |
+
@group(0) @binding(1) var<storage, read_write> out: array<f32>;
|
| 3 |
+
|
| 4 |
+
@compute @workgroup_size(64)
|
| 5 |
+
fn main(@builtin(global_invocation_id) gid: vec3<u32>) {
|
| 6 |
+
let i = gid.x + gid.y * 256u;
|
| 7 |
+
if (i >= 256u) { return; }
|
| 8 |
+
let coord = (i / 1u) % 64u;
|
| 9 |
+
if (coord >= 0u && coord < 32u) { out[i] = f32(b0[(i / 64u) * 32u + (coord - 0u) * 1u + i % 1u]); }
|
| 10 |
+
if (coord >= 32u && coord < 64u) { out[i] = f32(b0[(i / 64u) * 32u + (coord - 32u) * 1u + i % 1u]); }
|
| 11 |
+
}
|
exaone4-fp32-8k-state-alias-token-major-v2-v4/kernels/0e2ba0d1849cd945e9b533bfc599a5664fa7909515c955681084ab8ad82f8755.wgsl
ADDED
|
@@ -0,0 +1,9 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
@group(0) @binding(0) var<storage, read> b0: array<f32>;
|
| 2 |
+
@group(0) @binding(1) var<storage, read_write> out: array<f32>;
|
| 3 |
+
|
| 4 |
+
@compute @workgroup_size(64)
|
| 5 |
+
fn main(@builtin(global_invocation_id) gid: vec3<u32>) {
|
| 6 |
+
let i = gid.x + gid.y * 512u;
|
| 7 |
+
if (i >= 512u) { return; }
|
| 8 |
+
out[i] = f32(b0[((i / 512u) % 1u) * 512u + ((i / 64u) % 8u) * 64u + ((i / 64u) % 1u) * 512u + ((i / 1u) % 64u) * 1u]);
|
| 9 |
+
}
|
exaone4-fp32-8k-state-alias-token-major-v2-v4/kernels/11723b648ae4baeb93ebba06009266508b8bb3e9185c049eb4ec6e85199f045b.wgsl
ADDED
|
@@ -0,0 +1,83 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
@group(0) @binding(0) var<storage, read> b0: array<f32>;
|
| 2 |
+
@group(0) @binding(1) var<storage, read> b1: array<f32>;
|
| 3 |
+
@group(0) @binding(2) var<storage, read> b2: array<f32>;
|
| 4 |
+
@group(0) @binding(3) var<storage, read> b3: array<f32>;
|
| 5 |
+
@group(0) @binding(4) var<storage, read_write> out: array<f32>;
|
| 6 |
+
|
| 7 |
+
var<workgroup> q_row: array<f32, 64>;
|
| 8 |
+
var<workgroup> scores: array<f32, 1024>;
|
| 9 |
+
var<workgroup> partial: array<f32, 64>;
|
| 10 |
+
@compute @workgroup_size(64)
|
| 11 |
+
fn main(@builtin(workgroup_id) group: vec3<u32>, @builtin(local_invocation_id) local: vec3<u32>) {
|
| 12 |
+
let lane = local.x;
|
| 13 |
+
let head = group.x / 1u;
|
| 14 |
+
let row = group.x % 1u;
|
| 15 |
+
let kv_head = head / 4u;
|
| 16 |
+
for (var d = lane; d < 64u; d += 64u) { q_row[d] = f32(b0[(head * 1u + row) * 64u + d]); }
|
| 17 |
+
workgroupBarrier();
|
| 18 |
+
var running_max = -3.4028234663852886e38;
|
| 19 |
+
var running_sum = 0.0;
|
| 20 |
+
var acc: array<f32, 1>;
|
| 21 |
+
for (var start = 0u; start < 8192u; start += 1024u) {
|
| 22 |
+
let count = min(1024u, 8192u - start);
|
| 23 |
+
var tile_max = -3.4028234663852886e38;
|
| 24 |
+
for (var t = lane; t < count; t += 64u) {
|
| 25 |
+
let j = start + t;
|
| 26 |
+
var score = -3.4028234663852886e38;
|
| 27 |
+
if (f32(b3[row * 8192u + j]) > -1000.0) {
|
| 28 |
+
var dot = 0.0;
|
| 29 |
+
for (var d = 0u; d < 64u; d++) { dot += q_row[d] * f32(b1[(((kv_head * 8192u + j) * 64u + d) / 64u) % 8192u * 512u + (((kv_head * 8192u + j) * 64u + d) / 524288u) % 8u * 64u + (((kv_head * 8192u + j) * 64u + d) / 1u) % 64u * 1u]); }
|
| 30 |
+
score = dot * 0.125 + f32(b3[row * 8192u + j]);
|
| 31 |
+
}
|
| 32 |
+
scores[t] = score;
|
| 33 |
+
tile_max = max(tile_max, score);
|
| 34 |
+
}
|
| 35 |
+
partial[lane] = tile_max;
|
| 36 |
+
workgroupBarrier();
|
| 37 |
+
for (var stride = 32u; stride > 0u; stride /= 2u) {
|
| 38 |
+
if (lane < stride) { partial[lane] = max(partial[lane], partial[lane + stride]); }
|
| 39 |
+
workgroupBarrier();
|
| 40 |
+
}
|
| 41 |
+
let next_max = max(running_max, partial[0]);
|
| 42 |
+
workgroupBarrier();
|
| 43 |
+
var tile_sum = 0.0;
|
| 44 |
+
for (var t = lane; t < count; t += 64u) {
|
| 45 |
+
var p = 0.0;
|
| 46 |
+
if (scores[t] > -3.0e38) { p = exp(scores[t] - next_max); }
|
| 47 |
+
scores[t] = p;
|
| 48 |
+
tile_sum += p;
|
| 49 |
+
}
|
| 50 |
+
partial[lane] = tile_sum;
|
| 51 |
+
workgroupBarrier();
|
| 52 |
+
for (var stride = 32u; stride > 0u; stride /= 2u) {
|
| 53 |
+
if (lane < stride) { partial[lane] += partial[lane + stride]; }
|
| 54 |
+
workgroupBarrier();
|
| 55 |
+
}
|
| 56 |
+
var correction = 0.0;
|
| 57 |
+
if (running_max > -3.0e38) { correction = exp(running_max - next_max); }
|
| 58 |
+
running_sum = running_sum * correction + partial[0];
|
| 59 |
+
running_max = next_max;
|
| 60 |
+
for (var c = 0u; c < 1u; c++) {
|
| 61 |
+
let d = lane + c * 64u;
|
| 62 |
+
var total = acc[c] * correction;
|
| 63 |
+
if (d < 64u) {
|
| 64 |
+
for (var t = 0u; t < count; t++) {
|
| 65 |
+
let p = scores[t];
|
| 66 |
+
if (p != 0.0) {
|
| 67 |
+
let j = start + t;
|
| 68 |
+
total += p * f32(b2[(((kv_head * 8192u + j) * 64u + d) / 64u) % 8192u * 512u + (((kv_head * 8192u + j) * 64u + d) / 524288u) % 8u * 64u + (((kv_head * 8192u + j) * 64u + d) / 1u) % 64u * 1u]);
|
| 69 |
+
}
|
| 70 |
+
}
|
| 71 |
+
}
|
| 72 |
+
acc[c] = total;
|
| 73 |
+
}
|
| 74 |
+
workgroupBarrier();
|
| 75 |
+
}
|
| 76 |
+
for (var c = 0u; c < 1u; c++) {
|
| 77 |
+
let d = lane + c * 64u;
|
| 78 |
+
if (d < 64u) {
|
| 79 |
+
let i = (head * 1u + row) * 64u + d;
|
| 80 |
+
out[i] = f32(acc[c] / running_sum);
|
| 81 |
+
}
|
| 82 |
+
}
|
| 83 |
+
}
|
exaone4-fp32-8k-state-alias-token-major-v2-v4/kernels/133b3d1c9ddb9953a9d3cffa5e3bd5e837cf2474ec6641a311a54d911d958f6b.wgsl
ADDED
|
@@ -0,0 +1,9 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
@group(0) @binding(0) var<storage, read> b0: array<f32>;
|
| 2 |
+
@group(0) @binding(1) var<storage, read_write> out: array<f32>;
|
| 3 |
+
|
| 4 |
+
@compute @workgroup_size(64)
|
| 5 |
+
fn main(@builtin(global_invocation_id) gid: vec3<u32>) {
|
| 6 |
+
let i = gid.x + gid.y * 65536u;
|
| 7 |
+
if (i >= 65536u) { return; }
|
| 8 |
+
out[i] = f32(b0[((i / 65536u) % 1u) * 131072u + ((i / 2048u) % 32u) * 4096u + ((i / 32u) % 64u) * 64u + (((i / 1u) % 32u) * 1u + 32u) * 1u]);
|
| 9 |
+
}
|
exaone4-fp32-8k-state-alias-token-major-v2-v4/kernels/14d88fb2ba88841a93259cafdf9a2cbf86411c82ec5da801fcd8a887de2ec80d.wgsl
ADDED
|
@@ -0,0 +1,13 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
enable f16;
|
| 2 |
+
@group(0) @binding(0) var<storage, read> b0: array<i32>;
|
| 3 |
+
@group(0) @binding(1) var<storage, read> b1: array<f16>;
|
| 4 |
+
@group(0) @binding(2) var<storage, read_write> out: array<f32>;
|
| 5 |
+
|
| 6 |
+
@compute @workgroup_size(64)
|
| 7 |
+
fn main(@builtin(global_invocation_id) gid: vec3<u32>) {
|
| 8 |
+
let i = gid.x + gid.y * 32768u;
|
| 9 |
+
if (i >= 32768u) { return; }
|
| 10 |
+
let token = i32(b0[i / 2048u]);
|
| 11 |
+
if (token < 32768 || token >= 65536) { out[i] = f32(0.0); return; }
|
| 12 |
+
out[i] = f32(b1[u32(token - 32768) * 2048u + i % 2048u]);
|
| 13 |
+
}
|
exaone4-fp32-8k-state-alias-token-major-v2-v4/kernels/1533ec909b782828ba90ded15ebe230c6f9dbfbeb74335606f979f2b56c58dbf.wgsl
ADDED
|
@@ -0,0 +1,9 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
@group(0) @binding(0) var<storage, read> b0: array<f32>;
|
| 2 |
+
@group(0) @binding(1) var<storage, read_write> out: array<f32>;
|
| 3 |
+
|
| 4 |
+
@compute @workgroup_size(64)
|
| 5 |
+
fn main(@builtin(global_invocation_id) gid: vec3<u32>) {
|
| 6 |
+
let i = gid.x + gid.y * 1024u;
|
| 7 |
+
if (i >= 1024u) { return; }
|
| 8 |
+
out[i] = f32(b0[((i / 1024u) % 1u) * 2048u + ((i / 32u) % 32u) * 64u + ((i / 32u) % 1u) * 64u + (((i / 1u) % 32u) * 1u + 0u) * 1u]);
|
| 9 |
+
}
|
exaone4-fp32-8k-state-alias-token-major-v2-v4/kernels/1701e3705e0212caf2880d0a27ccb420ef83e5c95aeb2e9a0679eacf9c2e0a87.wgsl
ADDED
|
@@ -0,0 +1,21 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
enable f16;
|
| 2 |
+
@group(0) @binding(0) var<storage, read> b0: array<f32>;
|
| 3 |
+
@group(0) @binding(1) var<storage, read> b1: array<f16>;
|
| 4 |
+
@group(0) @binding(2) var<storage, read_write> out: array<f32>;
|
| 5 |
+
|
| 6 |
+
var<workgroup> partial: array<f32, 64>;
|
| 7 |
+
@compute @workgroup_size(64)
|
| 8 |
+
fn main(@builtin(workgroup_id) group: vec3<u32>, @builtin(local_invocation_id) local: vec3<u32>) {
|
| 9 |
+
let i = group.x + group.y * 2048u;
|
| 10 |
+
if (i >= 2048u) { return; }
|
| 11 |
+
let lane = local.x;
|
| 12 |
+
var acc = 0.0;
|
| 13 |
+
for (var p = lane; p < 4096u; p += 64u) { acc += f32(b0[(i / 2048u) * 4096u + p]) * f32(b1[(i % 2048u) * 4096u + p]); }
|
| 14 |
+
partial[lane] = acc;
|
| 15 |
+
workgroupBarrier();
|
| 16 |
+
for (var stride = 32u; stride > 0u; stride /= 2u) {
|
| 17 |
+
if (lane < stride) { partial[lane] += partial[lane + stride]; }
|
| 18 |
+
workgroupBarrier();
|
| 19 |
+
}
|
| 20 |
+
if (lane == 0u) { out[i] = f32(partial[0] + 0.0); }
|
| 21 |
+
}
|
exaone4-fp32-8k-state-alias-token-major-v2-v4/kernels/18a66ec0875c456c9defd45d15e6ef0c148e64a1618e2825c7f96eadbbe8303a.wgsl
ADDED
|
@@ -0,0 +1,9 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
@group(0) @binding(0) var<storage, read> b0: array<i32>;
|
| 2 |
+
@group(0) @binding(1) var<storage, read_write> out: array<i32>;
|
| 3 |
+
|
| 4 |
+
@compute @workgroup_size(64)
|
| 5 |
+
fn main(@builtin(global_invocation_id) gid: vec3<u32>) {
|
| 6 |
+
let i = gid.x + gid.y * 64u;
|
| 7 |
+
if (i >= 16u) { return; }
|
| 8 |
+
out[i] = i32(b0[i]);
|
| 9 |
+
}
|
exaone4-fp32-8k-state-alias-token-major-v2-v4/kernels/1a73b6fceda28f40bba4ae27a986ffa39842db94fa5903ea5bf689a7e854d4f3.wgsl
ADDED
|
@@ -0,0 +1,11 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
@group(0) @binding(0) var<storage, read> b0: array<f32>;
|
| 2 |
+
@group(0) @binding(1) var<storage, read_write> out: array<f32>;
|
| 3 |
+
|
| 4 |
+
@compute @workgroup_size(64)
|
| 5 |
+
fn main(@builtin(global_invocation_id) gid: vec3<u32>) {
|
| 6 |
+
let i = gid.x + gid.y * 1024u;
|
| 7 |
+
if (i >= 1024u) { return; }
|
| 8 |
+
let coord = (i / 1u) % 64u;
|
| 9 |
+
if (coord >= 0u && coord < 32u) { out[i] = f32(b0[(i / 64u) * 32u + (coord - 0u) * 1u + i % 1u]); }
|
| 10 |
+
if (coord >= 32u && coord < 64u) { out[i] = f32(b0[(i / 64u) * 32u + (coord - 32u) * 1u + i % 1u]); }
|
| 11 |
+
}
|
exaone4-fp32-8k-state-alias-token-major-v2-v4/kernels/1aec53ecf59335904944de2db37f0d3df64e3445556b40324ce611d4c9fcfc2c.wgsl
ADDED
|
@@ -0,0 +1,9 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
@group(0) @binding(0) var<storage, read> b0: array<f32>;
|
| 2 |
+
@group(0) @binding(1) var<storage, read_write> out: array<f32>;
|
| 3 |
+
|
| 4 |
+
@compute @workgroup_size(64)
|
| 5 |
+
fn main(@builtin(global_invocation_id) gid: vec3<u32>) {
|
| 6 |
+
let i = gid.x + gid.y * 1024u;
|
| 7 |
+
if (i >= 1024u) { return; }
|
| 8 |
+
out[i] = f32(b0[((i / 1024u) % 1u) * 2048u + ((i / 32u) % 32u) * 64u + ((i / 32u) % 1u) * 64u + (((i / 1u) % 32u) * 1u + 32u) * 1u]);
|
| 9 |
+
}
|
exaone4-fp32-8k-state-alias-token-major-v2-v4/kernels/1c5b119286250e322ec8f75a46494057ed2485c4ce8ea081b4858aa5fbeadd0d.wgsl
ADDED
|
@@ -0,0 +1,9 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
@group(0) @binding(0) var<storage, read> b0: array<f32>;
|
| 2 |
+
@group(0) @binding(1) var<storage, read_write> out: array<f32>;
|
| 3 |
+
|
| 4 |
+
@compute @workgroup_size(64)
|
| 5 |
+
fn main(@builtin(global_invocation_id) gid: vec3<u32>) {
|
| 6 |
+
let i = gid.x + gid.y * 4096u;
|
| 7 |
+
if (i >= 4096u) { return; }
|
| 8 |
+
out[i] = f32(b0[((i / 4096u) % 1u) * 8192u + ((i / 128u) % 32u) * 256u + ((i / 32u) % 4u) * 64u + (((i / 1u) % 32u) * 1u + 32u) * 1u]);
|
| 9 |
+
}
|
exaone4-fp32-8k-state-alias-token-major-v2-v4/kernels/1d2ff3e2ffac914e03d1495655b136298374a8b07c94134a60bdb6927cbddcd0.wgsl
ADDED
|
@@ -0,0 +1,83 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
@group(0) @binding(0) var<storage, read> b0: array<f32>;
|
| 2 |
+
@group(0) @binding(1) var<storage, read> b1: array<f32>;
|
| 3 |
+
@group(0) @binding(2) var<storage, read> b2: array<f32>;
|
| 4 |
+
@group(0) @binding(3) var<storage, read> b3: array<f32>;
|
| 5 |
+
@group(0) @binding(4) var<storage, read_write> out: array<f32>;
|
| 6 |
+
|
| 7 |
+
var<workgroup> q_row: array<f32, 64>;
|
| 8 |
+
var<workgroup> scores: array<f32, 1024>;
|
| 9 |
+
var<workgroup> partial: array<f32, 64>;
|
| 10 |
+
@compute @workgroup_size(64)
|
| 11 |
+
fn main(@builtin(workgroup_id) group: vec3<u32>, @builtin(local_invocation_id) local: vec3<u32>) {
|
| 12 |
+
let lane = local.x;
|
| 13 |
+
let head = group.x / 64u;
|
| 14 |
+
let row = group.x % 64u;
|
| 15 |
+
let kv_head = head / 4u;
|
| 16 |
+
for (var d = lane; d < 64u; d += 64u) { q_row[d] = f32(b0[(head * 64u + row) * 64u + d]); }
|
| 17 |
+
workgroupBarrier();
|
| 18 |
+
var running_max = -3.4028234663852886e38;
|
| 19 |
+
var running_sum = 0.0;
|
| 20 |
+
var acc: array<f32, 1>;
|
| 21 |
+
for (var start = 0u; start < 8192u; start += 1024u) {
|
| 22 |
+
let count = min(1024u, 8192u - start);
|
| 23 |
+
var tile_max = -3.4028234663852886e38;
|
| 24 |
+
for (var t = lane; t < count; t += 64u) {
|
| 25 |
+
let j = start + t;
|
| 26 |
+
var score = -3.4028234663852886e38;
|
| 27 |
+
if (f32(b3[row * 8192u + j]) > -1000.0) {
|
| 28 |
+
var dot = 0.0;
|
| 29 |
+
for (var d = 0u; d < 64u; d++) { dot += q_row[d] * f32(b1[(((kv_head * 8192u + j) * 64u + d) / 64u) % 8192u * 512u + (((kv_head * 8192u + j) * 64u + d) / 524288u) % 8u * 64u + (((kv_head * 8192u + j) * 64u + d) / 1u) % 64u * 1u]); }
|
| 30 |
+
score = dot * 0.125 + f32(b3[row * 8192u + j]);
|
| 31 |
+
}
|
| 32 |
+
scores[t] = score;
|
| 33 |
+
tile_max = max(tile_max, score);
|
| 34 |
+
}
|
| 35 |
+
partial[lane] = tile_max;
|
| 36 |
+
workgroupBarrier();
|
| 37 |
+
for (var stride = 32u; stride > 0u; stride /= 2u) {
|
| 38 |
+
if (lane < stride) { partial[lane] = max(partial[lane], partial[lane + stride]); }
|
| 39 |
+
workgroupBarrier();
|
| 40 |
+
}
|
| 41 |
+
let next_max = max(running_max, partial[0]);
|
| 42 |
+
workgroupBarrier();
|
| 43 |
+
var tile_sum = 0.0;
|
| 44 |
+
for (var t = lane; t < count; t += 64u) {
|
| 45 |
+
var p = 0.0;
|
| 46 |
+
if (scores[t] > -3.0e38) { p = exp(scores[t] - next_max); }
|
| 47 |
+
scores[t] = p;
|
| 48 |
+
tile_sum += p;
|
| 49 |
+
}
|
| 50 |
+
partial[lane] = tile_sum;
|
| 51 |
+
workgroupBarrier();
|
| 52 |
+
for (var stride = 32u; stride > 0u; stride /= 2u) {
|
| 53 |
+
if (lane < stride) { partial[lane] += partial[lane + stride]; }
|
| 54 |
+
workgroupBarrier();
|
| 55 |
+
}
|
| 56 |
+
var correction = 0.0;
|
| 57 |
+
if (running_max > -3.0e38) { correction = exp(running_max - next_max); }
|
| 58 |
+
running_sum = running_sum * correction + partial[0];
|
| 59 |
+
running_max = next_max;
|
| 60 |
+
for (var c = 0u; c < 1u; c++) {
|
| 61 |
+
let d = lane + c * 64u;
|
| 62 |
+
var total = acc[c] * correction;
|
| 63 |
+
if (d < 64u) {
|
| 64 |
+
for (var t = 0u; t < count; t++) {
|
| 65 |
+
let p = scores[t];
|
| 66 |
+
if (p != 0.0) {
|
| 67 |
+
let j = start + t;
|
| 68 |
+
total += p * f32(b2[(((kv_head * 8192u + j) * 64u + d) / 64u) % 8192u * 512u + (((kv_head * 8192u + j) * 64u + d) / 524288u) % 8u * 64u + (((kv_head * 8192u + j) * 64u + d) / 1u) % 64u * 1u]);
|
| 69 |
+
}
|
| 70 |
+
}
|
| 71 |
+
}
|
| 72 |
+
acc[c] = total;
|
| 73 |
+
}
|
| 74 |
+
workgroupBarrier();
|
| 75 |
+
}
|
| 76 |
+
for (var c = 0u; c < 1u; c++) {
|
| 77 |
+
let d = lane + c * 64u;
|
| 78 |
+
if (d < 64u) {
|
| 79 |
+
let i = (head * 64u + row) * 64u + d;
|
| 80 |
+
out[i] = f32(acc[c] / running_sum);
|
| 81 |
+
}
|
| 82 |
+
}
|
| 83 |
+
}
|
exaone4-fp32-8k-state-alias-token-major-v2-v4/kernels/1d3137ed94b2dd9113466681ea64a1e9a3086acbadc3c2e2a35f30dc42c67ad5.wgsl
ADDED
|
@@ -0,0 +1,9 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
@group(0) @binding(0) var<storage, read> b0: array<f32>;
|
| 2 |
+
@group(0) @binding(1) var<storage, read_write> out: array<f32>;
|
| 3 |
+
|
| 4 |
+
@compute @workgroup_size(64)
|
| 5 |
+
fn main(@builtin(global_invocation_id) gid: vec3<u32>) {
|
| 6 |
+
let i = gid.x + gid.y * 4096u;
|
| 7 |
+
if (i >= 4096u) { return; }
|
| 8 |
+
out[i] = f32(b0[((i / 4096u) % 1u) * 8192u + ((i / 128u) % 32u) * 256u + ((i / 32u) % 4u) * 64u + (((i / 1u) % 32u) * 1u + 0u) * 1u]);
|
| 9 |
+
}
|
exaone4-fp32-8k-state-alias-token-major-v2-v4/kernels/214ac9825bc5014847277240f7d17b1f0aa4f2272175a7d77e94116dc75f9fff.wgsl
ADDED
|
@@ -0,0 +1,13 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
enable f16;
|
| 2 |
+
@group(0) @binding(0) var<storage, read> b0: array<i32>;
|
| 3 |
+
@group(0) @binding(1) var<storage, read> b1: array<f16>;
|
| 4 |
+
@group(0) @binding(2) var<storage, read_write> out: array<f32>;
|
| 5 |
+
|
| 6 |
+
@compute @workgroup_size(64)
|
| 7 |
+
fn main(@builtin(global_invocation_id) gid: vec3<u32>) {
|
| 8 |
+
let i = gid.x + gid.y * 131072u;
|
| 9 |
+
if (i >= 131072u) { return; }
|
| 10 |
+
let token = i32(b0[i / 2048u]);
|
| 11 |
+
if (token < 0 || token >= 32768) { out[i] = f32(0.0); return; }
|
| 12 |
+
out[i] = f32(b1[u32(token - 0) * 2048u + i % 2048u]);
|
| 13 |
+
}
|
exaone4-fp32-8k-state-alias-token-major-v2-v4/kernels/216a5956ea2a40590a4de1b09152aacde8cac66cbd448be13a81a5ebb828cffa.wgsl
ADDED
|
@@ -0,0 +1,55 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
enable f16;
|
| 2 |
+
@group(0) @binding(0) var<storage, read> b0: array<f32>;
|
| 3 |
+
@group(0) @binding(1) var<storage, read> b1: array<f16>;
|
| 4 |
+
@group(0) @binding(2) var<storage, read_write> out: array<f32>;
|
| 5 |
+
|
| 6 |
+
var<workgroup> tile_a: array<array<f32, 16>, 16>;
|
| 7 |
+
var<workgroup> tile_b: array<array<f32, 64>, 16>;
|
| 8 |
+
@compute @workgroup_size(16, 16)
|
| 9 |
+
fn main(@builtin(workgroup_id) group: vec3<u32>, @builtin(local_invocation_id) local: vec3<u32>) {
|
| 10 |
+
let lane = local.y * 16u + local.x;
|
| 11 |
+
let batch = group.z;
|
| 12 |
+
let tile_n = group.x * 64u;
|
| 13 |
+
let tile_row = group.y * 16u;
|
| 14 |
+
var acc: array<array<f32, 4>, 1>;
|
| 15 |
+
for (var k0 = 0u; k0 < 2048u; k0 += 16u) {
|
| 16 |
+
for (var e = 0u; e < 1u; e++) {
|
| 17 |
+
let flat = lane + e * 256u;
|
| 18 |
+
let m_local = flat / 16u;
|
| 19 |
+
let row = tile_row + m_local;
|
| 20 |
+
let col = k0 + flat % 16u;
|
| 21 |
+
var value = 0.0;
|
| 22 |
+
if (row < 16u && col < 2048u) { value = f32(f32(b0[(batch * 16u + row) * 2048u + col])); }
|
| 23 |
+
tile_a[m_local][flat % 16u] = value;
|
| 24 |
+
}
|
| 25 |
+
for (var e = 0u; e < 4u; e++) {
|
| 26 |
+
let n_local = lane / 4u;
|
| 27 |
+
let k_local = (lane % 4u) * 4u + e;
|
| 28 |
+
let n_index = tile_n + n_local;
|
| 29 |
+
let col = k0 + k_local;
|
| 30 |
+
var value = 0.0;
|
| 31 |
+
if (n_index < 32768u && col < 2048u) { value = f32(f32(b1[n_index * 2048u + col])); }
|
| 32 |
+
tile_b[k_local][n_local] = value;
|
| 33 |
+
}
|
| 34 |
+
workgroupBarrier();
|
| 35 |
+
for (var kk = 0u; kk < 16u; kk++) {
|
| 36 |
+
var b_values: array<f32, 4>;
|
| 37 |
+
for (var c = 0u; c < 4u; c++) { b_values[c] = tile_b[kk][local.x * 4u + c]; }
|
| 38 |
+
for (var r = 0u; r < 1u; r++) {
|
| 39 |
+
let a_value = tile_a[local.y * 1u + r][kk];
|
| 40 |
+
for (var c = 0u; c < 4u; c++) { acc[r][c] += a_value * b_values[c]; }
|
| 41 |
+
}
|
| 42 |
+
}
|
| 43 |
+
workgroupBarrier();
|
| 44 |
+
}
|
| 45 |
+
for (var r = 0u; r < 1u; r++) {
|
| 46 |
+
let out_row = tile_row + local.y * 1u + r;
|
| 47 |
+
for (var c = 0u; c < 4u; c++) {
|
| 48 |
+
let out_col = tile_n + local.x * 4u + c;
|
| 49 |
+
if (out_row < 16u && out_col < 32768u) {
|
| 50 |
+
let i = (batch * 16u + out_row) * 32768u + out_col;
|
| 51 |
+
out[i] = f32(acc[r][c] + 0.0);
|
| 52 |
+
}
|
| 53 |
+
}
|
| 54 |
+
}
|
| 55 |
+
}
|
exaone4-fp32-8k-state-alias-token-major-v2-v4/kernels/27a921d3b5edfdca227a99e1c3328959881c09acb02b12cd09d07d598745b2d3.wgsl
ADDED
|
@@ -0,0 +1,9 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
@group(0) @binding(0) var<storage, read> b0: array<i32>;
|
| 2 |
+
@group(0) @binding(1) var<storage, read_write> out: array<f32>;
|
| 3 |
+
|
| 4 |
+
@compute @workgroup_size(64)
|
| 5 |
+
fn main(@builtin(global_invocation_id) gid: vec3<u32>) {
|
| 6 |
+
let i = gid.x + gid.y * 64u;
|
| 7 |
+
if (i >= 1u) { return; }
|
| 8 |
+
out[i] = f32(b0[i]);
|
| 9 |
+
}
|
exaone4-fp32-8k-state-alias-token-major-v2-v4/kernels/28119aa01122c42bd26ceb37129fe0c95b2a9192163f6dabd90fe058b1e14c54.wgsl
ADDED
|
@@ -0,0 +1,10 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
@group(0) @binding(0) var<storage, read> b0: array<f32>;
|
| 2 |
+
@group(0) @binding(1) var<storage, read_write> out: array<f32>;
|
| 3 |
+
|
| 4 |
+
@compute @workgroup_size(64)
|
| 5 |
+
fn main(@builtin(global_invocation_id) gid: vec3<u32>) {
|
| 6 |
+
let i = gid.x + gid.y * 4096u;
|
| 7 |
+
if (i >= 4096u) { return; }
|
| 8 |
+
let x = f32(b0[i]);
|
| 9 |
+
out[i] = f32(sin(x));
|
| 10 |
+
}
|
exaone4-fp32-8k-state-alias-token-major-v2-v4/kernels/284513e54dd254e923b7a86bc5703879e3bfb119b6c71aa234ba744c2dc63615.wgsl
ADDED
|
@@ -0,0 +1,9 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
@group(0) @binding(0) var<storage, read> b0: array<f32>;
|
| 2 |
+
@group(0) @binding(1) var<storage, read_write> out: array<f32>;
|
| 3 |
+
|
| 4 |
+
@compute @workgroup_size(64)
|
| 5 |
+
fn main(@builtin(global_invocation_id) gid: vec3<u32>) {
|
| 6 |
+
let i = gid.x + gid.y * 2048u;
|
| 7 |
+
if (i >= 2048u) { return; }
|
| 8 |
+
out[i] = f32(b0[((i / 2048u) % 1u) * 2048u + ((i / 32u) % 64u) * 1u + ((i / 1u) % 32u) * 64u]);
|
| 9 |
+
}
|
exaone4-fp32-8k-state-alias-token-major-v2-v4/kernels/2b9e4b05ef5623d721dbe56b5812e551efcc6fd833d81064320e86fe94d8c670.wgsl
ADDED
|
@@ -0,0 +1,9 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
@group(0) @binding(0) var<storage, read> b0: array<f32>;
|
| 2 |
+
@group(0) @binding(1) var<storage, read_write> out: array<f32>;
|
| 3 |
+
|
| 4 |
+
@compute @workgroup_size(64)
|
| 5 |
+
fn main(@builtin(global_invocation_id) gid: vec3<u32>) {
|
| 6 |
+
let i = gid.x + gid.y * 4096u;
|
| 7 |
+
if (i >= 4096u) { return; }
|
| 8 |
+
out[i] = f32(f32(b0[i]) * 1.0);
|
| 9 |
+
}
|
exaone4-fp32-8k-state-alias-token-major-v2-v4/kernels/2c66dcac8fdd3caa9ff1e2441524070c21c24df9526bee44c49e765a8d4c8222.wgsl
ADDED
|
@@ -0,0 +1,9 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
@group(0) @binding(0) var<storage, read> b0: array<f32>;
|
| 2 |
+
@group(0) @binding(1) var<storage, read_write> out: array<f32>;
|
| 3 |
+
|
| 4 |
+
@compute @workgroup_size(64)
|
| 5 |
+
fn main(@builtin(global_invocation_id) gid: vec3<u32>) {
|
| 6 |
+
let i = gid.x + gid.y * 16384u;
|
| 7 |
+
if (i >= 16384u) { return; }
|
| 8 |
+
out[i] = f32(b0[((i / 16384u) % 1u) * 32768u + ((i / 2048u) % 8u) * 4096u + ((i / 32u) % 64u) * 64u + (((i / 1u) % 32u) * 1u + 0u) * 1u]);
|
| 9 |
+
}
|
exaone4-fp32-8k-state-alias-token-major-v2-v4/kernels/2e1ce6ef01198780289fd4a35d8f86b5172595ed189cf931e8444091d49e8d3a.wgsl
ADDED
|
@@ -0,0 +1,9 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
@group(0) @binding(0) var<storage, read> b0: array<f32>;
|
| 2 |
+
@group(0) @binding(1) var<storage, read_write> out: array<f32>;
|
| 3 |
+
|
| 4 |
+
@compute @workgroup_size(64)
|
| 5 |
+
fn main(@builtin(global_invocation_id) gid: vec3<u32>) {
|
| 6 |
+
let i = gid.x + gid.y * 64u;
|
| 7 |
+
if (i >= 64u) { return; }
|
| 8 |
+
out[i] = f32(f32(b0[i]) * 1.0);
|
| 9 |
+
}
|
exaone4-fp32-8k-state-alias-token-major-v2-v4/kernels/3203e08f2c99fa4848f553739b7115a7c677840344f17b1d4391d6130517f5a7.wgsl
ADDED
|
@@ -0,0 +1,9 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
@group(0) @binding(0) var<storage, read> b0: array<i32>;
|
| 2 |
+
@group(0) @binding(1) var<storage, read_write> out: array<i32>;
|
| 3 |
+
|
| 4 |
+
@compute @workgroup_size(64)
|
| 5 |
+
fn main(@builtin(global_invocation_id) gid: vec3<u32>) {
|
| 6 |
+
let i = gid.x + gid.y * 64u;
|
| 7 |
+
if (i >= 4u) { return; }
|
| 8 |
+
out[i] = i32(b0[i]);
|
| 9 |
+
}
|
exaone4-fp32-8k-state-alias-token-major-v2-v4/kernels/3241a118927e109c0dacbb278370abe70d099ee593a07cd023d4530b5fb8f833.wgsl
ADDED
|
@@ -0,0 +1,24 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
@group(0) @binding(0) var<storage, read> b0: array<f32>;
|
| 2 |
+
@group(0) @binding(1) var<storage, read> b1: array<f32>;
|
| 3 |
+
@group(0) @binding(2) var<storage, read_write> out: array<f32>;
|
| 4 |
+
|
| 5 |
+
var<workgroup> factor: f32;
|
| 6 |
+
@compute @workgroup_size(64)
|
| 7 |
+
fn main(@builtin(workgroup_id) group: vec3<u32>, @builtin(local_invocation_id) local: vec3<u32>) {
|
| 8 |
+
let row = group.x + group.y * 8u;
|
| 9 |
+
if (row >= 8u) { return; }
|
| 10 |
+
let lane = local.x;
|
| 11 |
+
if (lane == 0u) {
|
| 12 |
+
var total = 0.0;
|
| 13 |
+
for (var j = 0u; j < 64u; j++) {
|
| 14 |
+
let v = f32(b0[row * 64u + j]);
|
| 15 |
+
total += v * v;
|
| 16 |
+
}
|
| 17 |
+
factor = inverseSqrt(total / 64.0 + 1e-05);
|
| 18 |
+
}
|
| 19 |
+
workgroupBarrier();
|
| 20 |
+
for (var p = lane; p < 64u; p += 64u) {
|
| 21 |
+
let i = row * 64u + p;
|
| 22 |
+
out[i] = f32(f32(b0[row * 64u + p]) * factor * (f32(b1[p]) + 0.0));
|
| 23 |
+
}
|
| 24 |
+
}
|
exaone4-fp32-8k-state-alias-token-major-v2-v4/kernels/33e8ec04c66adb4988f562f0f9d3e92d29dfda91a56d0ed62be598c7c737f5ba.wgsl
ADDED
|
@@ -0,0 +1,9 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
@group(0) @binding(0) var<storage, read> b0: array<f32>;
|
| 2 |
+
@group(0) @binding(1) var<storage, read_write> out: array<f32>;
|
| 3 |
+
|
| 4 |
+
@compute @workgroup_size(64)
|
| 5 |
+
fn main(@builtin(global_invocation_id) gid: vec3<u32>) {
|
| 6 |
+
let i = gid.x + gid.y * 16384u;
|
| 7 |
+
if (i >= 16384u) { return; }
|
| 8 |
+
out[i] = f32(b0[((i / 16384u) % 1u) * 32768u + ((i / 2048u) % 8u) * 4096u + ((i / 32u) % 64u) * 64u + (((i / 1u) % 32u) * 1u + 32u) * 1u]);
|
| 9 |
+
}
|
exaone4-fp32-8k-state-alias-token-major-v2-v4/kernels/358e73d0a8169937dce5d2c7fd5f4dea41f7c8307696f1698967bc6c4132efc8.wgsl
ADDED
|
@@ -0,0 +1,13 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
enable f16;
|
| 2 |
+
@group(0) @binding(0) var<storage, read> b0: array<i32>;
|
| 3 |
+
@group(0) @binding(1) var<storage, read> b1: array<f16>;
|
| 4 |
+
@group(0) @binding(2) var<storage, read_write> out: array<f32>;
|
| 5 |
+
|
| 6 |
+
@compute @workgroup_size(64)
|
| 7 |
+
fn main(@builtin(global_invocation_id) gid: vec3<u32>) {
|
| 8 |
+
let i = gid.x + gid.y * 8192u;
|
| 9 |
+
if (i >= 8192u) { return; }
|
| 10 |
+
let token = i32(b0[i / 2048u]);
|
| 11 |
+
if (token < 32768 || token >= 65536) { out[i] = f32(0.0); return; }
|
| 12 |
+
out[i] = f32(b1[u32(token - 32768) * 2048u + i % 2048u]);
|
| 13 |
+
}
|
exaone4-fp32-8k-state-alias-token-major-v2-v4/kernels/37ce78f0b50bc59dc7b6dd62724d6ad3acfbfb73deae12446a4575f5688d8f34.wgsl
ADDED
|
@@ -0,0 +1,9 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
@group(0) @binding(0) var<storage, read> b0: array<f32>;
|
| 2 |
+
@group(0) @binding(1) var<storage, read_write> out: array<f32>;
|
| 3 |
+
|
| 4 |
+
@compute @workgroup_size(64)
|
| 5 |
+
fn main(@builtin(global_invocation_id) gid: vec3<u32>) {
|
| 6 |
+
let i = gid.x + gid.y * 128u;
|
| 7 |
+
if (i >= 128u) { return; }
|
| 8 |
+
out[i] = f32(b0[((i / 128u) % 1u) * 128u + ((i / 32u) % 4u) * 1u + ((i / 1u) % 32u) * 4u]);
|
| 9 |
+
}
|
exaone4-fp32-8k-state-alias-token-major-v2-v4/kernels/383d0ce2bbe657b6412b8e55e6b0cf06e385cb4efbe45019afff9357f2a77751.wgsl
ADDED
|
@@ -0,0 +1,11 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
@group(0) @binding(0) var<storage, read> b0: array<f32>;
|
| 2 |
+
@group(0) @binding(1) var<storage, read> b1: array<f32>;
|
| 3 |
+
@group(0) @binding(2) var<storage, read_write> out: array<f32>;
|
| 4 |
+
|
| 5 |
+
@compute @workgroup_size(64)
|
| 6 |
+
fn main(@builtin(global_invocation_id) gid: vec3<u32>) {
|
| 7 |
+
let i = gid.x + gid.y * 16384u;
|
| 8 |
+
if (i >= 16384u) { return; }
|
| 9 |
+
let x = f32(b0[i]); let s = f32(f32(x / (1.0 + exp(-x))));
|
| 10 |
+
out[i] = f32(s * f32(b1[i]));
|
| 11 |
+
}
|
exaone4-fp32-8k-state-alias-token-major-v2-v4/kernels/3866ae40efed23181bf612efca4329d5348debc12468bf7565a2f13889b73bea.wgsl
ADDED
|
@@ -0,0 +1,13 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
enable f16;
|
| 2 |
+
@group(0) @binding(0) var<storage, read> b0: array<i32>;
|
| 3 |
+
@group(0) @binding(1) var<storage, read> b1: array<f16>;
|
| 4 |
+
@group(0) @binding(2) var<storage, read_write> out: array<f32>;
|
| 5 |
+
|
| 6 |
+
@compute @workgroup_size(64)
|
| 7 |
+
fn main(@builtin(global_invocation_id) gid: vec3<u32>) {
|
| 8 |
+
let i = gid.x + gid.y * 8192u;
|
| 9 |
+
if (i >= 8192u) { return; }
|
| 10 |
+
let token = i32(b0[i / 2048u]);
|
| 11 |
+
if (token < 65536 || token >= 98304) { out[i] = f32(0.0); return; }
|
| 12 |
+
out[i] = f32(b1[u32(token - 65536) * 2048u + i % 2048u]);
|
| 13 |
+
}
|
exaone4-fp32-8k-state-alias-token-major-v2-v4/kernels/39646bdfafd1ab0cb874474e3f1e9d79b57f10582a1804bc1d3762a492fec851.wgsl
ADDED
|
@@ -0,0 +1,10 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
@group(0) @binding(0) var<storage, read> b0: array<f32>;
|
| 2 |
+
@group(0) @binding(1) var<storage, read> b1: array<f32>;
|
| 3 |
+
@group(0) @binding(2) var<storage, read_write> out: array<f32>;
|
| 4 |
+
|
| 5 |
+
@compute @workgroup_size(64)
|
| 6 |
+
fn main(@builtin(global_invocation_id) gid: vec3<u32>) {
|
| 7 |
+
let i = gid.x + gid.y * 32768u;
|
| 8 |
+
if (i >= 32768u) { return; }
|
| 9 |
+
out[i] = f32(f32(b0[i]) * f32(b1[((i / 64u) % 64u) * 64u + ((i / 1u) % 64u) * 1u]));
|
| 10 |
+
}
|
exaone4-fp32-8k-state-alias-token-major-v2-v4/kernels/3b114a6518b88e4d199413a9dffd2ff911ae539546995b683c9aa0f86462044b.wgsl
ADDED
|
@@ -0,0 +1,55 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
enable f16;
|
| 2 |
+
@group(0) @binding(0) var<storage, read> b0: array<f32>;
|
| 3 |
+
@group(0) @binding(1) var<storage, read> b1: array<f16>;
|
| 4 |
+
@group(0) @binding(2) var<storage, read_write> out: array<f32>;
|
| 5 |
+
|
| 6 |
+
var<workgroup> tile_a: array<array<f32, 16>, 64>;
|
| 7 |
+
var<workgroup> tile_b: array<array<f32, 64>, 16>;
|
| 8 |
+
@compute @workgroup_size(16, 16)
|
| 9 |
+
fn main(@builtin(workgroup_id) group: vec3<u32>, @builtin(local_invocation_id) local: vec3<u32>) {
|
| 10 |
+
let lane = local.y * 16u + local.x;
|
| 11 |
+
let batch = group.z;
|
| 12 |
+
let tile_n = group.x * 64u;
|
| 13 |
+
let tile_row = group.y * 64u;
|
| 14 |
+
var acc: array<array<f32, 4>, 4>;
|
| 15 |
+
for (var k0 = 0u; k0 < 2048u; k0 += 16u) {
|
| 16 |
+
for (var e = 0u; e < 4u; e++) {
|
| 17 |
+
let flat = lane + e * 256u;
|
| 18 |
+
let m_local = flat / 16u;
|
| 19 |
+
let row = tile_row + m_local;
|
| 20 |
+
let col = k0 + flat % 16u;
|
| 21 |
+
var value = 0.0;
|
| 22 |
+
if (row < 64u && col < 2048u) { value = f32(f32(b0[(batch * 64u + row) * 2048u + col])); }
|
| 23 |
+
tile_a[m_local][flat % 16u] = value;
|
| 24 |
+
}
|
| 25 |
+
for (var e = 0u; e < 4u; e++) {
|
| 26 |
+
let n_local = lane / 4u;
|
| 27 |
+
let k_local = (lane % 4u) * 4u + e;
|
| 28 |
+
let n_index = tile_n + n_local;
|
| 29 |
+
let col = k0 + k_local;
|
| 30 |
+
var value = 0.0;
|
| 31 |
+
if (n_index < 512u && col < 2048u) { value = f32(f32(b1[n_index * 2048u + col])); }
|
| 32 |
+
tile_b[k_local][n_local] = value;
|
| 33 |
+
}
|
| 34 |
+
workgroupBarrier();
|
| 35 |
+
for (var kk = 0u; kk < 16u; kk++) {
|
| 36 |
+
var b_values: array<f32, 4>;
|
| 37 |
+
for (var c = 0u; c < 4u; c++) { b_values[c] = tile_b[kk][local.x * 4u + c]; }
|
| 38 |
+
for (var r = 0u; r < 4u; r++) {
|
| 39 |
+
let a_value = tile_a[local.y * 4u + r][kk];
|
| 40 |
+
for (var c = 0u; c < 4u; c++) { acc[r][c] += a_value * b_values[c]; }
|
| 41 |
+
}
|
| 42 |
+
}
|
| 43 |
+
workgroupBarrier();
|
| 44 |
+
}
|
| 45 |
+
for (var r = 0u; r < 4u; r++) {
|
| 46 |
+
let out_row = tile_row + local.y * 4u + r;
|
| 47 |
+
for (var c = 0u; c < 4u; c++) {
|
| 48 |
+
let out_col = tile_n + local.x * 4u + c;
|
| 49 |
+
if (out_row < 64u && out_col < 512u) {
|
| 50 |
+
let i = (batch * 64u + out_row) * 512u + out_col;
|
| 51 |
+
out[i] = f32(acc[r][c] + 0.0);
|
| 52 |
+
}
|
| 53 |
+
}
|
| 54 |
+
}
|
| 55 |
+
}
|
exaone4-fp32-8k-state-alias-token-major-v2-v4/kernels/3b495259612ab4c2e36bc153ec19119e95ab0d425693674839363ee4339a2770.wgsl
ADDED
|
@@ -0,0 +1,21 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
enable f16;
|
| 2 |
+
@group(0) @binding(0) var<storage, read> b0: array<f32>;
|
| 3 |
+
@group(0) @binding(1) var<storage, read> b1: array<f16>;
|
| 4 |
+
@group(0) @binding(2) var<storage, read_write> out: array<f32>;
|
| 5 |
+
|
| 6 |
+
var<workgroup> partial: array<f32, 64>;
|
| 7 |
+
@compute @workgroup_size(64)
|
| 8 |
+
fn main(@builtin(workgroup_id) group: vec3<u32>, @builtin(local_invocation_id) local: vec3<u32>) {
|
| 9 |
+
let i = group.x + group.y * 4096u;
|
| 10 |
+
if (i >= 4096u) { return; }
|
| 11 |
+
let lane = local.x;
|
| 12 |
+
var acc = 0.0;
|
| 13 |
+
for (var p = lane; p < 2048u; p += 64u) { acc += f32(b0[(i / 4096u) * 2048u + p]) * f32(b1[(i % 4096u) * 2048u + p]); }
|
| 14 |
+
partial[lane] = acc;
|
| 15 |
+
workgroupBarrier();
|
| 16 |
+
for (var stride = 32u; stride > 0u; stride /= 2u) {
|
| 17 |
+
if (lane < stride) { partial[lane] += partial[lane + stride]; }
|
| 18 |
+
workgroupBarrier();
|
| 19 |
+
}
|
| 20 |
+
if (lane == 0u) { out[i] = f32(partial[0] + 0.0); }
|
| 21 |
+
}
|
exaone4-fp32-8k-state-alias-token-major-v2-v4/kernels/3e34409407100d4858c4c109cf62f3eaf2f4c6081166b0ec8679105e6dfbe7fe.wgsl
ADDED
|
@@ -0,0 +1,10 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
@group(0) @binding(0) var<storage, read> b0: array<f32>;
|
| 2 |
+
@group(0) @binding(1) var<storage, read_write> out: array<f32>;
|
| 3 |
+
|
| 4 |
+
@compute @workgroup_size(64)
|
| 5 |
+
fn main(@builtin(global_invocation_id) gid: vec3<u32>) {
|
| 6 |
+
let i = gid.x + gid.y * 16384u;
|
| 7 |
+
if (i >= 16384u) { return; }
|
| 8 |
+
let x = f32(b0[i]);
|
| 9 |
+
out[i] = f32(-x);
|
| 10 |
+
}
|
exaone4-fp32-8k-state-alias-token-major-v2-v4/kernels/40f115c8647a238bfb3ccb0bb8446bccdc94a0dfadc89111c4920efb2259f41d.wgsl
ADDED
|
@@ -0,0 +1,10 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
@group(0) @binding(0) var<storage, read> b0: array<f32>;
|
| 2 |
+
@group(0) @binding(1) var<storage, read> b1: array<f32>;
|
| 3 |
+
@group(0) @binding(2) var<storage, read_write> out: array<f32>;
|
| 4 |
+
|
| 5 |
+
@compute @workgroup_size(64)
|
| 6 |
+
fn main(@builtin(global_invocation_id) gid: vec3<u32>) {
|
| 7 |
+
let i = gid.x + gid.y * 2048u;
|
| 8 |
+
if (i >= 2048u) { return; }
|
| 9 |
+
out[i] = f32(f32(b0[i]) * f32(b1[((i / 64u) % 4u) * 64u + ((i / 1u) % 64u) * 1u]));
|
| 10 |
+
}
|
exaone4-fp32-8k-state-alias-token-major-v2-v4/kernels/453668cbe741eb27e2cc10194eca66f514a3763d97b46c1733c376986defc322.wgsl
ADDED
|
@@ -0,0 +1,24 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
@group(0) @binding(0) var<storage, read> b0: array<f32>;
|
| 2 |
+
@group(0) @binding(1) var<storage, read> b1: array<f32>;
|
| 3 |
+
@group(0) @binding(2) var<storage, read_write> out: array<f32>;
|
| 4 |
+
|
| 5 |
+
var<workgroup> factor: f32;
|
| 6 |
+
@compute @workgroup_size(64)
|
| 7 |
+
fn main(@builtin(workgroup_id) group: vec3<u32>, @builtin(local_invocation_id) local: vec3<u32>) {
|
| 8 |
+
let row = group.x + group.y * 4u;
|
| 9 |
+
if (row >= 4u) { return; }
|
| 10 |
+
let lane = local.x;
|
| 11 |
+
if (lane == 0u) {
|
| 12 |
+
var total = 0.0;
|
| 13 |
+
for (var j = 0u; j < 2048u; j++) {
|
| 14 |
+
let v = f32(b0[row * 2048u + j]);
|
| 15 |
+
total += v * v;
|
| 16 |
+
}
|
| 17 |
+
factor = inverseSqrt(total / 2048.0 + 1e-05);
|
| 18 |
+
}
|
| 19 |
+
workgroupBarrier();
|
| 20 |
+
for (var p = lane; p < 2048u; p += 64u) {
|
| 21 |
+
let i = row * 2048u + p;
|
| 22 |
+
out[i] = f32(f32(b0[row * 2048u + p]) * factor * (f32(b1[p]) + 0.0));
|
| 23 |
+
}
|
| 24 |
+
}
|
exaone4-fp32-8k-state-alias-token-major-v2-v4/kernels/476f0abdcc4ba14abb8e304788e71549fd6b18fac0f22e03853683751aa205dc.wgsl
ADDED
|
@@ -0,0 +1,55 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
enable f16;
|
| 2 |
+
@group(0) @binding(0) var<storage, read> b0: array<f32>;
|
| 3 |
+
@group(0) @binding(1) var<storage, read> b1: array<f16>;
|
| 4 |
+
@group(0) @binding(2) var<storage, read_write> out: array<f32>;
|
| 5 |
+
|
| 6 |
+
var<workgroup> tile_a: array<array<f32, 16>, 16>;
|
| 7 |
+
var<workgroup> tile_b: array<array<f32, 64>, 16>;
|
| 8 |
+
@compute @workgroup_size(16, 16)
|
| 9 |
+
fn main(@builtin(workgroup_id) group: vec3<u32>, @builtin(local_invocation_id) local: vec3<u32>) {
|
| 10 |
+
let lane = local.y * 16u + local.x;
|
| 11 |
+
let batch = group.z;
|
| 12 |
+
let tile_n = group.x * 64u;
|
| 13 |
+
let tile_row = group.y * 16u;
|
| 14 |
+
var acc: array<array<f32, 4>, 1>;
|
| 15 |
+
for (var k0 = 0u; k0 < 2048u; k0 += 16u) {
|
| 16 |
+
for (var e = 0u; e < 1u; e++) {
|
| 17 |
+
let flat = lane + e * 256u;
|
| 18 |
+
let m_local = flat / 16u;
|
| 19 |
+
let row = tile_row + m_local;
|
| 20 |
+
let col = k0 + flat % 16u;
|
| 21 |
+
var value = 0.0;
|
| 22 |
+
if (row < 16u && col < 2048u) { value = f32(f32(b0[(batch * 16u + row) * 2048u + col])); }
|
| 23 |
+
tile_a[m_local][flat % 16u] = value;
|
| 24 |
+
}
|
| 25 |
+
for (var e = 0u; e < 4u; e++) {
|
| 26 |
+
let n_local = lane / 4u;
|
| 27 |
+
let k_local = (lane % 4u) * 4u + e;
|
| 28 |
+
let n_index = tile_n + n_local;
|
| 29 |
+
let col = k0 + k_local;
|
| 30 |
+
var value = 0.0;
|
| 31 |
+
if (n_index < 2048u && col < 2048u) { value = f32(f32(b1[n_index * 2048u + col])); }
|
| 32 |
+
tile_b[k_local][n_local] = value;
|
| 33 |
+
}
|
| 34 |
+
workgroupBarrier();
|
| 35 |
+
for (var kk = 0u; kk < 16u; kk++) {
|
| 36 |
+
var b_values: array<f32, 4>;
|
| 37 |
+
for (var c = 0u; c < 4u; c++) { b_values[c] = tile_b[kk][local.x * 4u + c]; }
|
| 38 |
+
for (var r = 0u; r < 1u; r++) {
|
| 39 |
+
let a_value = tile_a[local.y * 1u + r][kk];
|
| 40 |
+
for (var c = 0u; c < 4u; c++) { acc[r][c] += a_value * b_values[c]; }
|
| 41 |
+
}
|
| 42 |
+
}
|
| 43 |
+
workgroupBarrier();
|
| 44 |
+
}
|
| 45 |
+
for (var r = 0u; r < 1u; r++) {
|
| 46 |
+
let out_row = tile_row + local.y * 1u + r;
|
| 47 |
+
for (var c = 0u; c < 4u; c++) {
|
| 48 |
+
let out_col = tile_n + local.x * 4u + c;
|
| 49 |
+
if (out_row < 16u && out_col < 2048u) {
|
| 50 |
+
let i = (batch * 16u + out_row) * 2048u + out_col;
|
| 51 |
+
out[i] = f32(acc[r][c] + 0.0);
|
| 52 |
+
}
|
| 53 |
+
}
|
| 54 |
+
}
|
| 55 |
+
}
|
exaone4-fp32-8k-state-alias-token-major-v2-v4/kernels/483857d2d9929bc07cd14940dc3bca97abcf9485ce1c805a234d2db4110a7010.wgsl
ADDED
|
@@ -0,0 +1,10 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
@group(0) @binding(0) var<storage, read> b0: array<f32>;
|
| 2 |
+
@group(0) @binding(1) var<storage, read> b1: array<f32>;
|
| 3 |
+
@group(0) @binding(2) var<storage, read_write> out: array<f32>;
|
| 4 |
+
|
| 5 |
+
@compute @workgroup_size(64)
|
| 6 |
+
fn main(@builtin(global_invocation_id) gid: vec3<u32>) {
|
| 7 |
+
let i = gid.x + gid.y * 32768u;
|
| 8 |
+
if (i >= 32768u) { return; }
|
| 9 |
+
out[i] = f32(f32(b0[i]) * f32(b1[((i / 64u) % 16u) * 64u + ((i / 1u) % 64u) * 1u]));
|
| 10 |
+
}
|
exaone4-fp32-8k-state-alias-token-major-v2-v4/kernels/4c3a20bb702b43bacecf9c7df989b55517469eba683bb2cb596e1f7a227347c7.wgsl
ADDED
|
@@ -0,0 +1,13 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
enable f16;
|
| 2 |
+
@group(0) @binding(0) var<storage, read> b0: array<i32>;
|
| 3 |
+
@group(0) @binding(1) var<storage, read> b1: array<f16>;
|
| 4 |
+
@group(0) @binding(2) var<storage, read_write> out: array<f32>;
|
| 5 |
+
|
| 6 |
+
@compute @workgroup_size(64)
|
| 7 |
+
fn main(@builtin(global_invocation_id) gid: vec3<u32>) {
|
| 8 |
+
let i = gid.x + gid.y * 131072u;
|
| 9 |
+
if (i >= 131072u) { return; }
|
| 10 |
+
let token = i32(b0[i / 2048u]);
|
| 11 |
+
if (token < 32768 || token >= 65536) { out[i] = f32(0.0); return; }
|
| 12 |
+
out[i] = f32(b1[u32(token - 32768) * 2048u + i % 2048u]);
|
| 13 |
+
}
|
exaone4-fp32-8k-state-alias-token-major-v2-v4/kernels/4c526eba157bfd1aee20a2637586d6f4cdcdc68dcf238dcef84647649ffc6f9e.wgsl
ADDED
|
@@ -0,0 +1,13 @@
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1 |
+
enable f16;
|
| 2 |
+
@group(0) @binding(0) var<storage, read> b0: array<i32>;
|
| 3 |
+
@group(0) @binding(1) var<storage, read> b1: array<f16>;
|
| 4 |
+
@group(0) @binding(2) var<storage, read_write> out: array<f32>;
|
| 5 |
+
|
| 6 |
+
@compute @workgroup_size(64)
|
| 7 |
+
fn main(@builtin(global_invocation_id) gid: vec3<u32>) {
|
| 8 |
+
let i = gid.x + gid.y * 131072u;
|
| 9 |
+
if (i >= 131072u) { return; }
|
| 10 |
+
let token = i32(b0[i / 2048u]);
|
| 11 |
+
if (token < 98304 || token >= 102400) { out[i] = f32(0.0); return; }
|
| 12 |
+
out[i] = f32(b1[u32(token - 98304) * 2048u + i % 2048u]);
|
| 13 |
+
}
|