9bow commited on
Commit
b00e259
·
verified ·
1 Parent(s): 363dc91

Add exaone-deep-fp32-8k-state-alias-token-major-v2-v5 (attention bounded by the used context)

Browse files
This view is limited to 50 files because it contains too many changes.   See raw diff
Files changed (50) hide show
  1. exaone-deep-fp32-8k-state-alias-token-major-v2-v5/entrypoints/prefill-16.json +0 -0
  2. exaone-deep-fp32-8k-state-alias-token-major-v2-v5/entrypoints/prefill-4.json +0 -0
  3. exaone-deep-fp32-8k-state-alias-token-major-v2-v5/entrypoints/prefill-64.json +0 -0
  4. exaone-deep-fp32-8k-state-alias-token-major-v2-v5/graph.json +0 -0
  5. exaone-deep-fp32-8k-state-alias-token-major-v2-v5/kernels/00f1518df8d42f8c7c48cbb8d4080ed273fc370066401f5caf6316d3df3a9c72.wgsl +9 -0
  6. exaone-deep-fp32-8k-state-alias-token-major-v2-v5/kernels/010ceb4c6d74413ae14f7a70e539bbcad465036da065b435f5f99436ba183f3a.wgsl +10 -0
  7. exaone-deep-fp32-8k-state-alias-token-major-v2-v5/kernels/011f3693baa378d6f074ec8e475cda2dfee7d2a19dcac8ff041c6eb3569c009a.wgsl +16 -0
  8. exaone-deep-fp32-8k-state-alias-token-major-v2-v5/kernels/03780813af6ed9f38ba73c3512379e4a9db2730eb91497d8dde2d21a3984ae1b.wgsl +9 -0
  9. exaone-deep-fp32-8k-state-alias-token-major-v2-v5/kernels/056a7cde988532d104769fe6246161d0a150f08c4f910918110826ce621bb5a3.wgsl +10 -0
  10. exaone-deep-fp32-8k-state-alias-token-major-v2-v5/kernels/06314129ff03df6da928e3f91acf8b8779caace4b3b722790b5a9215f1129e50.wgsl +17 -0
  11. exaone-deep-fp32-8k-state-alias-token-major-v2-v5/kernels/065de78c2ded88ff8cfca05f54ba48f64761693422c9a04ff2801fac2e85606a.wgsl +10 -0
  12. exaone-deep-fp32-8k-state-alias-token-major-v2-v5/kernels/079286b9ee077ebd54a8fb16dbf2596fd5a558d045e09327698396d6fed19faa.wgsl +25 -0
  13. exaone-deep-fp32-8k-state-alias-token-major-v2-v5/kernels/07ba2a3d7e883e143fe36feb0793c2054ce4384c45117c501cb7837596a9310f.wgsl +25 -0
  14. exaone-deep-fp32-8k-state-alias-token-major-v2-v5/kernels/0b36e6f66583375c2c8a3426a96b1e0c913dfcbb16ab06838d9a186ec06d1289.wgsl +9 -0
  15. exaone-deep-fp32-8k-state-alias-token-major-v2-v5/kernels/0c3c183e8c534c4caf04c61c77ae778d91ea93272df63ad6a5394b67125f284f.wgsl +25 -0
  16. exaone-deep-fp32-8k-state-alias-token-major-v2-v5/kernels/0d83a40babe385a0afb76f881da28aceb4bffa3e8a6d4497523503b935854950.wgsl +59 -0
  17. exaone-deep-fp32-8k-state-alias-token-major-v2-v5/kernels/17a32140ebd06816dfb4d571da93b9fe24e622c069bb563882a65efdc9bd3045.wgsl +12 -0
  18. exaone-deep-fp32-8k-state-alias-token-major-v2-v5/kernels/17cf2e11eabbbfb1c28393ccc0cd4a6f4145566dd8d65f2a87be3e3481dd48e5.wgsl +10 -0
  19. exaone-deep-fp32-8k-state-alias-token-major-v2-v5/kernels/189c65aa14859d6642ba29508ca2150726cc5d94b6cd1e72d8ddb2b14aaf603f.wgsl +9 -0
  20. exaone-deep-fp32-8k-state-alias-token-major-v2-v5/kernels/18a66ec0875c456c9defd45d15e6ef0c148e64a1618e2825c7f96eadbbe8303a.wgsl +9 -0
  21. exaone-deep-fp32-8k-state-alias-token-major-v2-v5/kernels/1b3a44837ef6433e8faa7fb4dae08de71f68861ca914355ed5f8711a0b28719b.wgsl +17 -0
  22. exaone-deep-fp32-8k-state-alias-token-major-v2-v5/kernels/1b9e76d54a728dada6475c0d04d6c5cb5fec9a862cd2e01bd9ce77b16e0d96b6.wgsl +9 -0
  23. exaone-deep-fp32-8k-state-alias-token-major-v2-v5/kernels/1bf469f7b8f1f89afadd837b786ff5e2f9f210b8a163b844343cfed2865d790f.wgsl +9 -0
  24. exaone-deep-fp32-8k-state-alias-token-major-v2-v5/kernels/1e4e50d6a6d5c9380d37a3462475e30fa473e77585e8d36baff974e5c1152782.wgsl +10 -0
  25. exaone-deep-fp32-8k-state-alias-token-major-v2-v5/kernels/1e8ed1fc01e57c0f843a0dfa69589cc0fdcf7dedd2117d6bb8bc2f10a46ded10.wgsl +17 -0
  26. exaone-deep-fp32-8k-state-alias-token-major-v2-v5/kernels/1e9e57e0500e94c2e40db913222f7286a3886aa562c14f234e94361200a5fb2a.wgsl +9 -0
  27. exaone-deep-fp32-8k-state-alias-token-major-v2-v5/kernels/1ec850641f1c68248987e77bb1c9bb8f66bc8ded1c08941ea790ae4979bb7174.wgsl +9 -0
  28. exaone-deep-fp32-8k-state-alias-token-major-v2-v5/kernels/1ece9b553926e5f15825e1e8cb6638808227e8e23c5063f4b6d35ef4ef05a83c.wgsl +10 -0
  29. exaone-deep-fp32-8k-state-alias-token-major-v2-v5/kernels/1ed82ae7c212a7174f1c1a80c3575f118523ff8c1de4ca54fb2d1fe542f4c031.wgsl +10 -0
  30. exaone-deep-fp32-8k-state-alias-token-major-v2-v5/kernels/2024ef81167d6e5e3f3685336fadeaf536fef11d59b778ec9c3ce1c51513a3ea.wgsl +59 -0
  31. exaone-deep-fp32-8k-state-alias-token-major-v2-v5/kernels/265f0d6e72ce92141e27fd00420e8565670ccc37e566aaf34ecfe9ec3d96725a.wgsl +9 -0
  32. exaone-deep-fp32-8k-state-alias-token-major-v2-v5/kernels/27a921d3b5edfdca227a99e1c3328959881c09acb02b12cd09d07d598745b2d3.wgsl +9 -0
  33. exaone-deep-fp32-8k-state-alias-token-major-v2-v5/kernels/2ac9c34882d6464a1170e8f7cfcf32daf8875beea8b6ce7ea0726bf97faf2500.wgsl +16 -0
  34. exaone-deep-fp32-8k-state-alias-token-major-v2-v5/kernels/2b52b6c54ef6a4ea6d83f4bc93fbe576a97374b11e2bc5551a66173b10820075.wgsl +25 -0
  35. exaone-deep-fp32-8k-state-alias-token-major-v2-v5/kernels/2d579253030031996b0e5bc60673be3857dad34cfb4e6d3b5d462286337d2851.wgsl +59 -0
  36. exaone-deep-fp32-8k-state-alias-token-major-v2-v5/kernels/2dc6d073f766830ed46387e14ea0a320bfe8cbbddbad2c450334441885d8d6fe.wgsl +17 -0
  37. exaone-deep-fp32-8k-state-alias-token-major-v2-v5/kernels/2f01373257c48f14505e017a3d4e9236235f64b889afef2d355f75ef376d4453.wgsl +16 -0
  38. exaone-deep-fp32-8k-state-alias-token-major-v2-v5/kernels/307d2f0d7330af4c5fb41a9e52936470ccada38a6fe9bfa4530025d0188b8c3e.wgsl +10 -0
  39. exaone-deep-fp32-8k-state-alias-token-major-v2-v5/kernels/31d236875acba515c98c35ba2b251fe95c92813111f4851a2c817589c5b872e8.wgsl +24 -0
  40. exaone-deep-fp32-8k-state-alias-token-major-v2-v5/kernels/31d59acc55ed479f9268e1ec82ea3975d4972863c865f2e842de5c9aea4902e8.wgsl +96 -0
  41. exaone-deep-fp32-8k-state-alias-token-major-v2-v5/kernels/3203e08f2c99fa4848f553739b7115a7c677840344f17b1d4391d6130517f5a7.wgsl +9 -0
  42. exaone-deep-fp32-8k-state-alias-token-major-v2-v5/kernels/33d5e3ca854722003d2032af6ad94488ac2caaefaec22cd9bed9df8eb51efc11.wgsl +9 -0
  43. exaone-deep-fp32-8k-state-alias-token-major-v2-v5/kernels/34789e50cf72ea181c86db87caa629346be4ae8221e39c85ebbd836bc22d0ca9.wgsl +10 -0
  44. exaone-deep-fp32-8k-state-alias-token-major-v2-v5/kernels/35c0d9ce0c1a0fc678ad4893e9b1508c759402b5649ccaca55d94993875f5c2c.wgsl +11 -0
  45. exaone-deep-fp32-8k-state-alias-token-major-v2-v5/kernels/36eb23461edc6900aa34d9d30acca79f0e3823a35bfbb413e72244886613cf1c.wgsl +10 -0
  46. exaone-deep-fp32-8k-state-alias-token-major-v2-v5/kernels/3714273478bb8a6c9b5f209cc1be287d03b7f3b188b5ed28bf74ff8c1750a5d9.wgsl +17 -0
  47. exaone-deep-fp32-8k-state-alias-token-major-v2-v5/kernels/397c272c70f921a18fa5b359f241c16191dabb59affdb25a5f21e0ac97b6ed39.wgsl +9 -0
  48. exaone-deep-fp32-8k-state-alias-token-major-v2-v5/kernels/39b0e20440b4b41f452d8130d42c31a13a660ab22bd56ebb653bc40eb85d5890.wgsl +59 -0
  49. exaone-deep-fp32-8k-state-alias-token-major-v2-v5/kernels/3c1557fa28d308c0aae510ad9480e3454a28c5f9f409906a5b3ff186abba9ef3.wgsl +9 -0
  50. exaone-deep-fp32-8k-state-alias-token-major-v2-v5/kernels/3c55802ae4ce1a62daf3f5afd38a3ae905812bb4d990d1393cba3636a644994d.wgsl +13 -0
exaone-deep-fp32-8k-state-alias-token-major-v2-v5/entrypoints/prefill-16.json ADDED
The diff for this file is too large to render. See raw diff
 
exaone-deep-fp32-8k-state-alias-token-major-v2-v5/entrypoints/prefill-4.json ADDED
The diff for this file is too large to render. See raw diff
 
exaone-deep-fp32-8k-state-alias-token-major-v2-v5/entrypoints/prefill-64.json ADDED
The diff for this file is too large to render. See raw diff
 
exaone-deep-fp32-8k-state-alias-token-major-v2-v5/graph.json ADDED
The diff for this file is too large to render. See raw diff
 
exaone-deep-fp32-8k-state-alias-token-major-v2-v5/kernels/00f1518df8d42f8c7c48cbb8d4080ed273fc370066401f5caf6316d3df3a9c72.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 * 10240u;
7
+ if (i >= 10240u) { return; }
8
+ out[i] = f32(b0[((i / 10240u) % 1u) * 10240u + ((i / 320u) % 32u) * 80u + ((i / 80u) % 4u) * 2560u + ((i / 1u) % 80u) * 1u]);
9
+ }
exaone-deep-fp32-8k-state-alias-token-major-v2-v5/kernels/010ceb4c6d74413ae14f7a70e539bbcad465036da065b435f5f99436ba183f3a.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 * 1280u;
7
+ if (i >= 1280u) { return; }
8
+ let x = f32(b0[i]);
9
+ out[i] = f32(cos(x));
10
+ }
exaone-deep-fp32-8k-state-alias-token-major-v2-v5/kernels/011f3693baa378d6f074ec8e475cda2dfee7d2a19dcac8ff041c6eb3569c009a.wgsl ADDED
@@ -0,0 +1,16 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
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
+ @compute @workgroup_size(64)
8
+ fn main(@builtin(global_invocation_id) gid: vec3<u32>) {
9
+ let i = gid.x + gid.y * 102400u;
10
+ if (i >= 102400u) { return; }
11
+ let coord = (i / 1u) % 102400u;
12
+ if (coord >= 0u && coord < 26214u) { out[i] = f32(b0[(i / 102400u) * 26214u + (coord - 0u) * 1u + i % 1u]); }
13
+ if (coord >= 26214u && coord < 52428u) { out[i] = f32(b1[(i / 102400u) * 26214u + (coord - 26214u) * 1u + i % 1u]); }
14
+ if (coord >= 52428u && coord < 78642u) { out[i] = f32(b2[(i / 102400u) * 26214u + (coord - 52428u) * 1u + i % 1u]); }
15
+ if (coord >= 78642u && coord < 102400u) { out[i] = f32(b3[(i / 102400u) * 23758u + (coord - 78642u) * 1u + i % 1u]); }
16
+ }
exaone-deep-fp32-8k-state-alias-token-major-v2-v5/kernels/03780813af6ed9f38ba73c3512379e4a9db2730eb91497d8dde2d21a3984ae1b.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 * 640u;
7
+ if (i >= 640u) { return; }
8
+ out[i] = f32(b0[((i / 640u) % 1u) * 640u + ((i / 80u) % 8u) * 80u + ((i / 80u) % 1u) * 640u + ((i / 1u) % 80u) * 1u]);
9
+ }
exaone-deep-fp32-8k-state-alias-token-major-v2-v5/kernels/056a7cde988532d104769fe6246161d0a150f08c4f910918110826ce621bb5a3.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 * 40960u;
8
+ if (i >= 40960u) { return; }
9
+ out[i] = f32(f32(b0[i]) * f32(b1[((i / 80u) % 64u) * 80u + ((i / 1u) % 80u) * 1u]));
10
+ }
exaone-deep-fp32-8k-state-alias-token-major-v2-v5/kernels/06314129ff03df6da928e3f91acf8b8779caace4b3b722790b5a9215f1129e50.wgsl ADDED
@@ -0,0 +1,17 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ @group(0) @binding(0) var<storage, read> b0: array<i32>;
2
+ @group(0) @binding(1) var<storage, read> b1: array<u32>;
3
+ @group(0) @binding(2) var<storage, read_write> out: array<f32>;
4
+ fn unpack_bf16_1(index: u32) -> f32 {
5
+ let pair = b1[index / 2u];
6
+ let bits = (pair >> ((index % 2u) * 16u)) & 65535u;
7
+ return bitcast<f32>(bits << 16u);
8
+ }
9
+
10
+ @compute @workgroup_size(64)
11
+ fn main(@builtin(global_invocation_id) gid: vec3<u32>) {
12
+ let i = gid.x + gid.y * 40960u;
13
+ if (i >= 40960u) { return; }
14
+ let token = i32(b0[i / 2560u]);
15
+ if (token < 52428 || token >= 78642) { out[i] = f32(0.0); return; }
16
+ out[i] = f32(unpack_bf16_1(u32(token - 52428) * 2560u + i % 2560u));
17
+ }
exaone-deep-fp32-8k-state-alias-token-major-v2-v5/kernels/065de78c2ded88ff8cfca05f54ba48f64761693422c9a04ff2801fac2e85606a.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 * 5120u;
7
+ if (i >= 5120u) { return; }
8
+ let x = f32(b0[i]);
9
+ out[i] = f32(sin(x));
10
+ }
exaone-deep-fp32-8k-state-alias-token-major-v2-v5/kernels/079286b9ee077ebd54a8fb16dbf2596fd5a558d045e09327698396d6fed19faa.wgsl ADDED
@@ -0,0 +1,25 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ @group(0) @binding(0) var<storage, read> b0: array<f32>;
2
+ @group(0) @binding(1) var<storage, read> b1: array<u32>;
3
+ @group(0) @binding(2) var<storage, read_write> out: array<f32>;
4
+ fn unpack_bf16_1(index: u32) -> f32 {
5
+ let pair = b1[index / 2u];
6
+ let bits = (pair >> ((index % 2u) * 16u)) & 65535u;
7
+ return bitcast<f32>(bits << 16u);
8
+ }
9
+
10
+ var<workgroup> partial: array<f32, 64>;
11
+ @compute @workgroup_size(64)
12
+ fn main(@builtin(workgroup_id) group: vec3<u32>, @builtin(local_invocation_id) local: vec3<u32>) {
13
+ let i = group.x + group.y * 26214u;
14
+ if (i >= 26214u) { return; }
15
+ let lane = local.x;
16
+ var acc = 0.0;
17
+ for (var p = lane; p < 2560u; p += 64u) { acc += f32(b0[(i / 26214u) * 2560u + p]) * unpack_bf16_1((i % 26214u) * 2560u + p); }
18
+ partial[lane] = acc;
19
+ workgroupBarrier();
20
+ for (var stride = 32u; stride > 0u; stride /= 2u) {
21
+ if (lane < stride) { partial[lane] += partial[lane + stride]; }
22
+ workgroupBarrier();
23
+ }
24
+ if (lane == 0u) { out[i] = f32(partial[0] + 0.0); }
25
+ }
exaone-deep-fp32-8k-state-alias-token-major-v2-v5/kernels/07ba2a3d7e883e143fe36feb0793c2054ce4384c45117c501cb7837596a9310f.wgsl ADDED
@@ -0,0 +1,25 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ @group(0) @binding(0) var<storage, read> b0: array<f32>;
2
+ @group(0) @binding(1) var<storage, read> b1: array<u32>;
3
+ @group(0) @binding(2) var<storage, read_write> out: array<f32>;
4
+ fn unpack_bf16_1(index: u32) -> f32 {
5
+ let pair = b1[index / 2u];
6
+ let bits = (pair >> ((index % 2u) * 16u)) & 65535u;
7
+ return bitcast<f32>(bits << 16u);
8
+ }
9
+
10
+ var<workgroup> partial: array<f32, 64>;
11
+ @compute @workgroup_size(64)
12
+ fn main(@builtin(workgroup_id) group: vec3<u32>, @builtin(local_invocation_id) local: vec3<u32>) {
13
+ let i = group.x + group.y * 2560u;
14
+ if (i >= 2560u) { return; }
15
+ let lane = local.x;
16
+ var acc = 0.0;
17
+ for (var p = lane; p < 7168u; p += 64u) { acc += f32(b0[(i / 2560u) * 7168u + p]) * unpack_bf16_1((i % 2560u) * 7168u + p); }
18
+ partial[lane] = acc;
19
+ workgroupBarrier();
20
+ for (var stride = 32u; stride > 0u; stride /= 2u) {
21
+ if (lane < stride) { partial[lane] += partial[lane + stride]; }
22
+ workgroupBarrier();
23
+ }
24
+ if (lane == 0u) { out[i] = f32(partial[0] + 0.0); }
25
+ }
exaone-deep-fp32-8k-state-alias-token-major-v2-v5/kernels/0b36e6f66583375c2c8a3426a96b1e0c913dfcbb16ab06838d9a186ec06d1289.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 * 5120u;
7
+ if (i >= 5120u) { return; }
8
+ out[i] = f32(b0[((i / 5120u) % 1u) * 10240u + ((i / 160u) % 32u) * 320u + ((i / 40u) % 4u) * 80u + (((i / 1u) % 40u) * 1u + 40u) * 1u]);
9
+ }
exaone-deep-fp32-8k-state-alias-token-major-v2-v5/kernels/0c3c183e8c534c4caf04c61c77ae778d91ea93272df63ad6a5394b67125f284f.wgsl ADDED
@@ -0,0 +1,25 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ @group(0) @binding(0) var<storage, read> b0: array<f32>;
2
+ @group(0) @binding(1) var<storage, read> b1: array<u32>;
3
+ @group(0) @binding(2) var<storage, read_write> out: array<f32>;
4
+ fn unpack_bf16_1(index: u32) -> f32 {
5
+ let pair = b1[index / 2u];
6
+ let bits = (pair >> ((index % 2u) * 16u)) & 65535u;
7
+ return bitcast<f32>(bits << 16u);
8
+ }
9
+
10
+ var<workgroup> partial: array<f32, 64>;
11
+ @compute @workgroup_size(64)
12
+ fn main(@builtin(workgroup_id) group: vec3<u32>, @builtin(local_invocation_id) local: vec3<u32>) {
13
+ let i = group.x + group.y * 640u;
14
+ if (i >= 640u) { return; }
15
+ let lane = local.x;
16
+ var acc = 0.0;
17
+ for (var p = lane; p < 2560u; p += 64u) { acc += f32(b0[(i / 640u) * 2560u + p]) * unpack_bf16_1((i % 640u) * 2560u + p); }
18
+ partial[lane] = acc;
19
+ workgroupBarrier();
20
+ for (var stride = 32u; stride > 0u; stride /= 2u) {
21
+ if (lane < stride) { partial[lane] += partial[lane + stride]; }
22
+ workgroupBarrier();
23
+ }
24
+ if (lane == 0u) { out[i] = f32(partial[0] + 0.0); }
25
+ }
exaone-deep-fp32-8k-state-alias-token-major-v2-v5/kernels/0d83a40babe385a0afb76f881da28aceb4bffa3e8a6d4497523503b935854950.wgsl ADDED
@@ -0,0 +1,59 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ @group(0) @binding(0) var<storage, read> b0: array<f32>;
2
+ @group(0) @binding(1) var<storage, read> b1: array<u32>;
3
+ @group(0) @binding(2) var<storage, read_write> out: array<f32>;
4
+ fn unpack_bf16_1(index: u32) -> f32 {
5
+ let pair = b1[index / 2u];
6
+ let bits = (pair >> ((index % 2u) * 16u)) & 65535u;
7
+ return bitcast<f32>(bits << 16u);
8
+ }
9
+
10
+ var<workgroup> tile_a: array<array<f32, 16>, 16>;
11
+ var<workgroup> tile_b: array<array<f32, 64>, 16>;
12
+ @compute @workgroup_size(16, 16)
13
+ fn main(@builtin(workgroup_id) group: vec3<u32>, @builtin(local_invocation_id) local: vec3<u32>) {
14
+ let lane = local.y * 16u + local.x;
15
+ let batch = group.z;
16
+ let tile_n = group.x * 64u;
17
+ let tile_row = group.y * 16u;
18
+ var acc: array<array<f32, 4>, 1>;
19
+ for (var k0 = 0u; k0 < 2560u; k0 += 16u) {
20
+ for (var e = 0u; e < 1u; e++) {
21
+ let flat = lane + e * 256u;
22
+ let m_local = flat / 16u;
23
+ let row = tile_row + m_local;
24
+ let col = k0 + flat % 16u;
25
+ var value = 0.0;
26
+ if (row < 16u && col < 2560u) { value = f32(f32(b0[(batch * 16u + row) * 2560u + col])); }
27
+ tile_a[m_local][flat % 16u] = value;
28
+ }
29
+ for (var e = 0u; e < 4u; e++) {
30
+ let n_local = lane / 4u;
31
+ let k_local = (lane % 4u) * 4u + e;
32
+ let n_index = tile_n + n_local;
33
+ let col = k0 + k_local;
34
+ var value = 0.0;
35
+ if (n_index < 7168u && col < 2560u) { value = f32(unpack_bf16_1(n_index * 2560u + col)); }
36
+ tile_b[k_local][n_local] = value;
37
+ }
38
+ workgroupBarrier();
39
+ for (var kk = 0u; kk < 16u; kk++) {
40
+ var b_values: array<f32, 4>;
41
+ for (var c = 0u; c < 4u; c++) { b_values[c] = tile_b[kk][local.x * 4u + c]; }
42
+ for (var r = 0u; r < 1u; r++) {
43
+ let a_value = tile_a[local.y * 1u + r][kk];
44
+ for (var c = 0u; c < 4u; c++) { acc[r][c] += a_value * b_values[c]; }
45
+ }
46
+ }
47
+ workgroupBarrier();
48
+ }
49
+ for (var r = 0u; r < 1u; r++) {
50
+ let out_row = tile_row + local.y * 1u + r;
51
+ for (var c = 0u; c < 4u; c++) {
52
+ let out_col = tile_n + local.x * 4u + c;
53
+ if (out_row < 16u && out_col < 7168u) {
54
+ let i = (batch * 16u + out_row) * 7168u + out_col;
55
+ out[i] = f32(acc[r][c] + 0.0);
56
+ }
57
+ }
58
+ }
59
+ }
exaone-deep-fp32-8k-state-alias-token-major-v2-v5/kernels/17a32140ebd06816dfb4d571da93b9fe24e622c069bb563882a65efdc9bd3045.wgsl ADDED
@@ -0,0 +1,12 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
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 * 640u;
8
+ if (i >= 640u) { return; }
9
+ let coord = (i / 1u) % 80u;
10
+ if (coord >= 0u && coord < 40u) { out[i] = f32(b0[(i / 80u) * 40u + (coord - 0u) * 1u + i % 1u]); }
11
+ if (coord >= 40u && coord < 80u) { out[i] = f32(b1[(i / 80u) * 40u + (coord - 40u) * 1u + i % 1u]); }
12
+ }
exaone-deep-fp32-8k-state-alias-token-major-v2-v5/kernels/17cf2e11eabbbfb1c28393ccc0cd4a6f4145566dd8d65f2a87be3e3481dd48e5.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 * 81920u;
7
+ if (i >= 81920u) { return; }
8
+ let x = f32(b0[i]);
9
+ out[i] = f32(-x);
10
+ }
exaone-deep-fp32-8k-state-alias-token-major-v2-v5/kernels/189c65aa14859d6642ba29508ca2150726cc5d94b6cd1e72d8ddb2b14aaf603f.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 * 40960u;
7
+ if (i >= 40960u) { return; }
8
+ out[i] = f32(b0[((i / 40960u) % 1u) * 40960u + ((i / 5120u) % 8u) * 80u + ((i / 80u) % 64u) * 640u + ((i / 1u) % 80u) * 1u]);
9
+ }
exaone-deep-fp32-8k-state-alias-token-major-v2-v5/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
+ }
exaone-deep-fp32-8k-state-alias-token-major-v2-v5/kernels/1b3a44837ef6433e8faa7fb4dae08de71f68861ca914355ed5f8711a0b28719b.wgsl ADDED
@@ -0,0 +1,17 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
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 * 64u;
8
+ if (i >= 40u) { return; }
9
+ let batch = i / 40u;
10
+ let row = (i / 1u) % 40u;
11
+ let col = i % 1u;
12
+ var acc = 0.0;
13
+ for (var p = 0u; p < 1u; p++) {
14
+ acc += f32(b0[(0u) + row * 1u + p]) * f32(b1[(0u) + p * 1u + col]);
15
+ }
16
+ out[i] = f32(acc);
17
+ }
exaone-deep-fp32-8k-state-alias-token-major-v2-v5/kernels/1b9e76d54a728dada6475c0d04d6c5cb5fec9a862cd2e01bd9ce77b16e0d96b6.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 * 1280u;
7
+ if (i >= 1280u) { return; }
8
+ out[i] = f32(f32(b0[i]) * 1.0);
9
+ }
exaone-deep-fp32-8k-state-alias-token-major-v2-v5/kernels/1bf469f7b8f1f89afadd837b786ff5e2f9f210b8a163b844343cfed2865d790f.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 * 5120u;
7
+ if (i >= 5120u) { return; }
8
+ out[i] = f32(f32(b0[i]) * 1.0);
9
+ }
exaone-deep-fp32-8k-state-alias-token-major-v2-v5/kernels/1e4e50d6a6d5c9380d37a3462475e30fa473e77585e8d36baff974e5c1152782.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 * 2560u;
8
+ if (i >= 2560u) { return; }
9
+ out[i] = f32(f32(b0[i]) + (f32(b1[i]) * 1.0));
10
+ }
exaone-deep-fp32-8k-state-alias-token-major-v2-v5/kernels/1e8ed1fc01e57c0f843a0dfa69589cc0fdcf7dedd2117d6bb8bc2f10a46ded10.wgsl ADDED
@@ -0,0 +1,17 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ @group(0) @binding(0) var<storage, read> b0: array<i32>;
2
+ @group(0) @binding(1) var<storage, read> b1: array<u32>;
3
+ @group(0) @binding(2) var<storage, read_write> out: array<f32>;
4
+ fn unpack_bf16_1(index: u32) -> f32 {
5
+ let pair = b1[index / 2u];
6
+ let bits = (pair >> ((index % 2u) * 16u)) & 65535u;
7
+ return bitcast<f32>(bits << 16u);
8
+ }
9
+
10
+ @compute @workgroup_size(64)
11
+ fn main(@builtin(global_invocation_id) gid: vec3<u32>) {
12
+ let i = gid.x + gid.y * 2560u;
13
+ if (i >= 2560u) { return; }
14
+ let token = i32(b0[i / 2560u]);
15
+ if (token < 52428 || token >= 78642) { out[i] = f32(0.0); return; }
16
+ out[i] = f32(unpack_bf16_1(u32(token - 52428) * 2560u + i % 2560u));
17
+ }
exaone-deep-fp32-8k-state-alias-token-major-v2-v5/kernels/1e9e57e0500e94c2e40db913222f7286a3886aa562c14f234e94361200a5fb2a.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 * 640u;
7
+ if (i >= 640u) { return; }
8
+ out[i] = f32(b0[((i / 640u) % 1u) * 640u + ((i / 40u) % 16u) * 1u + ((i / 1u) % 40u) * 16u]);
9
+ }
exaone-deep-fp32-8k-state-alias-token-major-v2-v5/kernels/1ec850641f1c68248987e77bb1c9bb8f66bc8ded1c08941ea790ae4979bb7174.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 * 20480u;
7
+ if (i >= 20480u) { return; }
8
+ out[i] = f32(b0[((i / 20480u) % 1u) * 40960u + ((i / 640u) % 32u) * 1280u + ((i / 40u) % 16u) * 80u + (((i / 1u) % 40u) * 1u + 40u) * 1u]);
9
+ }
exaone-deep-fp32-8k-state-alias-token-major-v2-v5/kernels/1ece9b553926e5f15825e1e8cb6638808227e8e23c5063f4b6d35ef4ef05a83c.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 * 2560u;
8
+ if (i >= 2560u) { return; }
9
+ out[i] = f32(f32(b0[i]) * f32(b1[((i / 80u) % 4u) * 80u + ((i / 1u) % 80u) * 1u]));
10
+ }
exaone-deep-fp32-8k-state-alias-token-major-v2-v5/kernels/1ed82ae7c212a7174f1c1a80c3575f118523ff8c1de4ca54fb2d1fe542f4c031.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 * 320u;
7
+ if (i >= 320u) { return; }
8
+ let x = f32(b0[i]);
9
+ out[i] = f32(cos(x));
10
+ }
exaone-deep-fp32-8k-state-alias-token-major-v2-v5/kernels/2024ef81167d6e5e3f3685336fadeaf536fef11d59b778ec9c3ce1c51513a3ea.wgsl ADDED
@@ -0,0 +1,59 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ @group(0) @binding(0) var<storage, read> b0: array<f32>;
2
+ @group(0) @binding(1) var<storage, read> b1: array<u32>;
3
+ @group(0) @binding(2) var<storage, read_write> out: array<f32>;
4
+ fn unpack_bf16_1(index: u32) -> f32 {
5
+ let pair = b1[index / 2u];
6
+ let bits = (pair >> ((index % 2u) * 16u)) & 65535u;
7
+ return bitcast<f32>(bits << 16u);
8
+ }
9
+
10
+ var<workgroup> tile_a: array<array<f32, 16>, 64>;
11
+ var<workgroup> tile_b: array<array<f32, 64>, 16>;
12
+ @compute @workgroup_size(16, 16)
13
+ fn main(@builtin(workgroup_id) group: vec3<u32>, @builtin(local_invocation_id) local: vec3<u32>) {
14
+ let lane = local.y * 16u + local.x;
15
+ let batch = group.z;
16
+ let tile_n = group.x * 64u;
17
+ let tile_row = group.y * 64u;
18
+ var acc: array<array<f32, 4>, 4>;
19
+ for (var k0 = 0u; k0 < 2560u; k0 += 16u) {
20
+ for (var e = 0u; e < 4u; e++) {
21
+ let flat = lane + e * 256u;
22
+ let m_local = flat / 16u;
23
+ let row = tile_row + m_local;
24
+ let col = k0 + flat % 16u;
25
+ var value = 0.0;
26
+ if (row < 64u && col < 2560u) { value = f32(f32(b0[(batch * 64u + row) * 2560u + col])); }
27
+ tile_a[m_local][flat % 16u] = value;
28
+ }
29
+ for (var e = 0u; e < 4u; e++) {
30
+ let n_local = lane / 4u;
31
+ let k_local = (lane % 4u) * 4u + e;
32
+ let n_index = tile_n + n_local;
33
+ let col = k0 + k_local;
34
+ var value = 0.0;
35
+ if (n_index < 23758u && col < 2560u) { value = f32(unpack_bf16_1(n_index * 2560u + col)); }
36
+ tile_b[k_local][n_local] = value;
37
+ }
38
+ workgroupBarrier();
39
+ for (var kk = 0u; kk < 16u; kk++) {
40
+ var b_values: array<f32, 4>;
41
+ for (var c = 0u; c < 4u; c++) { b_values[c] = tile_b[kk][local.x * 4u + c]; }
42
+ for (var r = 0u; r < 4u; r++) {
43
+ let a_value = tile_a[local.y * 4u + r][kk];
44
+ for (var c = 0u; c < 4u; c++) { acc[r][c] += a_value * b_values[c]; }
45
+ }
46
+ }
47
+ workgroupBarrier();
48
+ }
49
+ for (var r = 0u; r < 4u; r++) {
50
+ let out_row = tile_row + local.y * 4u + r;
51
+ for (var c = 0u; c < 4u; c++) {
52
+ let out_col = tile_n + local.x * 4u + c;
53
+ if (out_row < 64u && out_col < 23758u) {
54
+ let i = (batch * 64u + out_row) * 23758u + out_col;
55
+ out[i] = f32(acc[r][c] + 0.0);
56
+ }
57
+ }
58
+ }
59
+ }
exaone-deep-fp32-8k-state-alias-token-major-v2-v5/kernels/265f0d6e72ce92141e27fd00420e8565670ccc37e566aaf34ecfe9ec3d96725a.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 * 20480u;
7
+ if (i >= 20480u) { return; }
8
+ out[i] = f32(b0[((i / 20480u) % 1u) * 40960u + ((i / 2560u) % 8u) * 5120u + ((i / 40u) % 64u) * 80u + (((i / 1u) % 40u) * 1u + 40u) * 1u]);
9
+ }
exaone-deep-fp32-8k-state-alias-token-major-v2-v5/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
+ }
exaone-deep-fp32-8k-state-alias-token-major-v2-v5/kernels/2ac9c34882d6464a1170e8f7cfcf32daf8875beea8b6ce7ea0726bf97faf2500.wgsl ADDED
@@ -0,0 +1,16 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
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
+ @compute @workgroup_size(64)
8
+ fn main(@builtin(global_invocation_id) gid: vec3<u32>) {
9
+ let i = gid.x + gid.y * 1638400u;
10
+ if (i >= 1638400u) { return; }
11
+ let coord = (i / 1u) % 102400u;
12
+ if (coord >= 0u && coord < 26214u) { out[i] = f32(b0[(i / 102400u) * 26214u + (coord - 0u) * 1u + i % 1u]); }
13
+ if (coord >= 26214u && coord < 52428u) { out[i] = f32(b1[(i / 102400u) * 26214u + (coord - 26214u) * 1u + i % 1u]); }
14
+ if (coord >= 52428u && coord < 78642u) { out[i] = f32(b2[(i / 102400u) * 26214u + (coord - 52428u) * 1u + i % 1u]); }
15
+ if (coord >= 78642u && coord < 102400u) { out[i] = f32(b3[(i / 102400u) * 23758u + (coord - 78642u) * 1u + i % 1u]); }
16
+ }
exaone-deep-fp32-8k-state-alias-token-major-v2-v5/kernels/2b52b6c54ef6a4ea6d83f4bc93fbe576a97374b11e2bc5551a66173b10820075.wgsl ADDED
@@ -0,0 +1,25 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ @group(0) @binding(0) var<storage, read> b0: array<f32>;
2
+ @group(0) @binding(1) var<storage, read> b1: array<u32>;
3
+ @group(0) @binding(2) var<storage, read_write> out: array<f32>;
4
+ fn unpack_bf16_1(index: u32) -> f32 {
5
+ let pair = b1[index / 2u];
6
+ let bits = (pair >> ((index % 2u) * 16u)) & 65535u;
7
+ return bitcast<f32>(bits << 16u);
8
+ }
9
+
10
+ var<workgroup> partial: array<f32, 64>;
11
+ @compute @workgroup_size(64)
12
+ fn main(@builtin(workgroup_id) group: vec3<u32>, @builtin(local_invocation_id) local: vec3<u32>) {
13
+ let i = group.x + group.y * 23758u;
14
+ if (i >= 23758u) { return; }
15
+ let lane = local.x;
16
+ var acc = 0.0;
17
+ for (var p = lane; p < 2560u; p += 64u) { acc += f32(b0[(i / 23758u) * 2560u + p]) * unpack_bf16_1((i % 23758u) * 2560u + p); }
18
+ partial[lane] = acc;
19
+ workgroupBarrier();
20
+ for (var stride = 32u; stride > 0u; stride /= 2u) {
21
+ if (lane < stride) { partial[lane] += partial[lane + stride]; }
22
+ workgroupBarrier();
23
+ }
24
+ if (lane == 0u) { out[i] = f32(partial[0] + 0.0); }
25
+ }
exaone-deep-fp32-8k-state-alias-token-major-v2-v5/kernels/2d579253030031996b0e5bc60673be3857dad34cfb4e6d3b5d462286337d2851.wgsl ADDED
@@ -0,0 +1,59 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ @group(0) @binding(0) var<storage, read> b0: array<f32>;
2
+ @group(0) @binding(1) var<storage, read> b1: array<u32>;
3
+ @group(0) @binding(2) var<storage, read_write> out: array<f32>;
4
+ fn unpack_bf16_1(index: u32) -> f32 {
5
+ let pair = b1[index / 2u];
6
+ let bits = (pair >> ((index % 2u) * 16u)) & 65535u;
7
+ return bitcast<f32>(bits << 16u);
8
+ }
9
+
10
+ var<workgroup> tile_a: array<array<f32, 16>, 64>;
11
+ var<workgroup> tile_b: array<array<f32, 64>, 16>;
12
+ @compute @workgroup_size(16, 16)
13
+ fn main(@builtin(workgroup_id) group: vec3<u32>, @builtin(local_invocation_id) local: vec3<u32>) {
14
+ let lane = local.y * 16u + local.x;
15
+ let batch = group.z;
16
+ let tile_n = group.x * 64u;
17
+ let tile_row = group.y * 64u;
18
+ var acc: array<array<f32, 4>, 4>;
19
+ for (var k0 = 0u; k0 < 7168u; k0 += 16u) {
20
+ for (var e = 0u; e < 4u; e++) {
21
+ let flat = lane + e * 256u;
22
+ let m_local = flat / 16u;
23
+ let row = tile_row + m_local;
24
+ let col = k0 + flat % 16u;
25
+ var value = 0.0;
26
+ if (row < 64u && col < 7168u) { value = f32(f32(b0[(batch * 64u + row) * 7168u + col])); }
27
+ tile_a[m_local][flat % 16u] = value;
28
+ }
29
+ for (var e = 0u; e < 4u; e++) {
30
+ let n_local = lane / 4u;
31
+ let k_local = (lane % 4u) * 4u + e;
32
+ let n_index = tile_n + n_local;
33
+ let col = k0 + k_local;
34
+ var value = 0.0;
35
+ if (n_index < 2560u && col < 7168u) { value = f32(unpack_bf16_1(n_index * 7168u + col)); }
36
+ tile_b[k_local][n_local] = value;
37
+ }
38
+ workgroupBarrier();
39
+ for (var kk = 0u; kk < 16u; kk++) {
40
+ var b_values: array<f32, 4>;
41
+ for (var c = 0u; c < 4u; c++) { b_values[c] = tile_b[kk][local.x * 4u + c]; }
42
+ for (var r = 0u; r < 4u; r++) {
43
+ let a_value = tile_a[local.y * 4u + r][kk];
44
+ for (var c = 0u; c < 4u; c++) { acc[r][c] += a_value * b_values[c]; }
45
+ }
46
+ }
47
+ workgroupBarrier();
48
+ }
49
+ for (var r = 0u; r < 4u; r++) {
50
+ let out_row = tile_row + local.y * 4u + r;
51
+ for (var c = 0u; c < 4u; c++) {
52
+ let out_col = tile_n + local.x * 4u + c;
53
+ if (out_row < 64u && out_col < 2560u) {
54
+ let i = (batch * 64u + out_row) * 2560u + out_col;
55
+ out[i] = f32(acc[r][c] + 0.0);
56
+ }
57
+ }
58
+ }
59
+ }
exaone-deep-fp32-8k-state-alias-token-major-v2-v5/kernels/2dc6d073f766830ed46387e14ea0a320bfe8cbbddbad2c450334441885d8d6fe.wgsl ADDED
@@ -0,0 +1,17 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ @group(0) @binding(0) var<storage, read> b0: array<i32>;
2
+ @group(0) @binding(1) var<storage, read> b1: array<u32>;
3
+ @group(0) @binding(2) var<storage, read_write> out: array<f32>;
4
+ fn unpack_bf16_1(index: u32) -> f32 {
5
+ let pair = b1[index / 2u];
6
+ let bits = (pair >> ((index % 2u) * 16u)) & 65535u;
7
+ return bitcast<f32>(bits << 16u);
8
+ }
9
+
10
+ @compute @workgroup_size(64)
11
+ fn main(@builtin(global_invocation_id) gid: vec3<u32>) {
12
+ let i = gid.x + gid.y * 10240u;
13
+ if (i >= 10240u) { return; }
14
+ let token = i32(b0[i / 2560u]);
15
+ if (token < 52428 || token >= 78642) { out[i] = f32(0.0); return; }
16
+ out[i] = f32(unpack_bf16_1(u32(token - 52428) * 2560u + i % 2560u));
17
+ }
exaone-deep-fp32-8k-state-alias-token-major-v2-v5/kernels/2f01373257c48f14505e017a3d4e9236235f64b889afef2d355f75ef376d4453.wgsl ADDED
@@ -0,0 +1,16 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
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
+ @compute @workgroup_size(64)
8
+ fn main(@builtin(global_invocation_id) gid: vec3<u32>) {
9
+ let i = gid.x + gid.y * 4194240u;
10
+ if (i >= 6553600u) { return; }
11
+ let coord = (i / 1u) % 102400u;
12
+ if (coord >= 0u && coord < 26214u) { out[i] = f32(b0[(i / 102400u) * 26214u + (coord - 0u) * 1u + i % 1u]); }
13
+ if (coord >= 26214u && coord < 52428u) { out[i] = f32(b1[(i / 102400u) * 26214u + (coord - 26214u) * 1u + i % 1u]); }
14
+ if (coord >= 52428u && coord < 78642u) { out[i] = f32(b2[(i / 102400u) * 26214u + (coord - 52428u) * 1u + i % 1u]); }
15
+ if (coord >= 78642u && coord < 102400u) { out[i] = f32(b3[(i / 102400u) * 23758u + (coord - 78642u) * 1u + i % 1u]); }
16
+ }
exaone-deep-fp32-8k-state-alias-token-major-v2-v5/kernels/307d2f0d7330af4c5fb41a9e52936470ccada38a6fe9bfa4530025d0188b8c3e.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 * 163840u;
8
+ if (i >= 163840u) { return; }
9
+ out[i] = f32(f32(b0[i]) + (f32(b1[i]) * 1.0));
10
+ }
exaone-deep-fp32-8k-state-alias-token-major-v2-v5/kernels/31d236875acba515c98c35ba2b251fe95c92813111f4851a2c817589c5b872e8.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 * 1u;
9
+ if (row >= 1u) { return; }
10
+ let lane = local.x;
11
+ if (lane == 0u) {
12
+ var total = 0.0;
13
+ for (var j = 0u; j < 2560u; j++) {
14
+ let v = f32(b0[row * 2560u + j]);
15
+ total += v * v;
16
+ }
17
+ factor = inverseSqrt(total / 2560.0 + 1e-05);
18
+ }
19
+ workgroupBarrier();
20
+ for (var p = lane; p < 2560u; p += 64u) {
21
+ let i = row * 2560u + p;
22
+ out[i] = f32(f32(b0[row * 2560u + p]) * factor * (f32(b1[p]) + 0.0));
23
+ }
24
+ }
exaone-deep-fp32-8k-state-alias-token-major-v2-v5/kernels/31d59acc55ed479f9268e1ec82ea3975d4972863c865f2e842de5c9aea4902e8.wgsl ADDED
@@ -0,0 +1,96 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
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, 80>;
8
+ var<workgroup> scores: array<f32, 1024>;
9
+ var<workgroup> partial: array<f32, 64>;
10
+ var<workgroup> limit_shared: u32;
11
+ @compute @workgroup_size(64)
12
+ fn main(@builtin(workgroup_id) group: vec3<u32>, @builtin(local_invocation_id) local: vec3<u32>) {
13
+ let lane = local.x;
14
+ let head = group.x / 1u;
15
+ let row = group.x % 1u;
16
+ let kv_head = head / 4u;
17
+ for (var d = lane; d < 80u; d += 64u) { q_row[d] = f32(b0[(head * 1u + row) * 80u + d]); }
18
+ // Bound the key loop by the last unmasked position so work follows the used context.
19
+ var last = 0.0;
20
+ for (var j = lane; j < 8192u; j += 64u) {
21
+ if (f32(b3[row * 8192u + j]) > -1000.0) { last = max(last, f32(j + 1u)); }
22
+ }
23
+ partial[lane] = last;
24
+ workgroupBarrier();
25
+ for (var stride = 32u; stride > 0u; stride /= 2u) {
26
+ if (lane < stride) { partial[lane] = max(partial[lane], partial[lane + stride]); }
27
+ workgroupBarrier();
28
+ }
29
+ if (lane == 0u) { limit_shared = u32(partial[0]); }
30
+ let limit = workgroupUniformLoad(&limit_shared);
31
+ var running_max = -3.4028234663852886e38;
32
+ var running_sum = 0.0;
33
+ var acc: array<f32, 2>;
34
+ for (var start = 0u; start < limit; start += 1024u) {
35
+ let count = min(1024u, limit - start);
36
+ var tile_max = -3.4028234663852886e38;
37
+ for (var t = lane; t < count; t += 64u) {
38
+ let j = start + t;
39
+ var score = -3.4028234663852886e38;
40
+ if (f32(b3[row * 8192u + j]) > -1000.0) {
41
+ var dot = 0.0;
42
+ for (var d = 0u; d < 80u; d++) { dot += q_row[d] * f32(b1[(((kv_head * 8192u + j) * 80u + d) / 80u) % 8192u * 640u + (((kv_head * 8192u + j) * 80u + d) / 655360u) % 8u * 80u + (((kv_head * 8192u + j) * 80u + d) / 1u) % 80u * 1u]); }
43
+ score = dot * 0.11180339887498948 + f32(b3[row * 8192u + j]);
44
+ }
45
+ scores[t] = score;
46
+ tile_max = max(tile_max, score);
47
+ }
48
+ partial[lane] = tile_max;
49
+ workgroupBarrier();
50
+ for (var stride = 32u; stride > 0u; stride /= 2u) {
51
+ if (lane < stride) { partial[lane] = max(partial[lane], partial[lane + stride]); }
52
+ workgroupBarrier();
53
+ }
54
+ let next_max = max(running_max, partial[0]);
55
+ workgroupBarrier();
56
+ var tile_sum = 0.0;
57
+ for (var t = lane; t < count; t += 64u) {
58
+ var p = 0.0;
59
+ if (scores[t] > -3.0e38) { p = exp(scores[t] - next_max); }
60
+ scores[t] = p;
61
+ tile_sum += p;
62
+ }
63
+ partial[lane] = tile_sum;
64
+ workgroupBarrier();
65
+ for (var stride = 32u; stride > 0u; stride /= 2u) {
66
+ if (lane < stride) { partial[lane] += partial[lane + stride]; }
67
+ workgroupBarrier();
68
+ }
69
+ var correction = 0.0;
70
+ if (running_max > -3.0e38) { correction = exp(running_max - next_max); }
71
+ running_sum = running_sum * correction + partial[0];
72
+ running_max = next_max;
73
+ for (var c = 0u; c < 2u; c++) {
74
+ let d = lane + c * 64u;
75
+ var total = acc[c] * correction;
76
+ if (d < 80u) {
77
+ for (var t = 0u; t < count; t++) {
78
+ let p = scores[t];
79
+ if (p != 0.0) {
80
+ let j = start + t;
81
+ total += p * f32(b2[(((kv_head * 8192u + j) * 80u + d) / 80u) % 8192u * 640u + (((kv_head * 8192u + j) * 80u + d) / 655360u) % 8u * 80u + (((kv_head * 8192u + j) * 80u + d) / 1u) % 80u * 1u]);
82
+ }
83
+ }
84
+ }
85
+ acc[c] = total;
86
+ }
87
+ workgroupBarrier();
88
+ }
89
+ for (var c = 0u; c < 2u; c++) {
90
+ let d = lane + c * 64u;
91
+ if (d < 80u) {
92
+ let i = (head * 1u + row) * 80u + d;
93
+ out[i] = f32(acc[c] / running_sum);
94
+ }
95
+ }
96
+ }
exaone-deep-fp32-8k-state-alias-token-major-v2-v5/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
+ }
exaone-deep-fp32-8k-state-alias-token-major-v2-v5/kernels/33d5e3ca854722003d2032af6ad94488ac2caaefaec22cd9bed9df8eb51efc11.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 * 163840u;
7
+ if (i >= 163840u) { return; }
8
+ out[i] = f32(b0[((i / 163840u) % 1u) * 163840u + ((i / 2560u) % 64u) * 80u + ((i / 80u) % 32u) * 5120u + ((i / 1u) % 80u) * 1u]);
9
+ }
exaone-deep-fp32-8k-state-alias-token-major-v2-v5/kernels/34789e50cf72ea181c86db87caa629346be4ae8221e39c85ebbd836bc22d0ca9.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 * 10240u;
8
+ if (i >= 10240u) { return; }
9
+ out[i] = f32(f32(b0[i]) * f32(b1[((i / 80u) % 16u) * 80u + ((i / 1u) % 80u) * 1u]));
10
+ }
exaone-deep-fp32-8k-state-alias-token-major-v2-v5/kernels/35c0d9ce0c1a0fc678ad4893e9b1508c759402b5649ccaca55d94993875f5c2c.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 * 1280u;
7
+ if (i >= 1280u) { return; }
8
+ let coord = (i / 1u) % 80u;
9
+ if (coord >= 0u && coord < 40u) { out[i] = f32(b0[(i / 80u) * 40u + (coord - 0u) * 1u + i % 1u]); }
10
+ if (coord >= 40u && coord < 80u) { out[i] = f32(b0[(i / 80u) * 40u + (coord - 40u) * 1u + i % 1u]); }
11
+ }
exaone-deep-fp32-8k-state-alias-token-major-v2-v5/kernels/36eb23461edc6900aa34d9d30acca79f0e3823a35bfbb413e72244886613cf1c.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 * 320u;
7
+ if (i >= 320u) { return; }
8
+ let x = f32(b0[i]);
9
+ out[i] = f32(sin(x));
10
+ }
exaone-deep-fp32-8k-state-alias-token-major-v2-v5/kernels/3714273478bb8a6c9b5f209cc1be287d03b7f3b188b5ed28bf74ff8c1750a5d9.wgsl ADDED
@@ -0,0 +1,17 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ @group(0) @binding(0) var<storage, read> b0: array<i32>;
2
+ @group(0) @binding(1) var<storage, read> b1: array<u32>;
3
+ @group(0) @binding(2) var<storage, read_write> out: array<f32>;
4
+ fn unpack_bf16_1(index: u32) -> f32 {
5
+ let pair = b1[index / 2u];
6
+ let bits = (pair >> ((index % 2u) * 16u)) & 65535u;
7
+ return bitcast<f32>(bits << 16u);
8
+ }
9
+
10
+ @compute @workgroup_size(64)
11
+ fn main(@builtin(global_invocation_id) gid: vec3<u32>) {
12
+ let i = gid.x + gid.y * 2560u;
13
+ if (i >= 2560u) { return; }
14
+ let token = i32(b0[i / 2560u]);
15
+ if (token < 0 || token >= 26214) { out[i] = f32(0.0); return; }
16
+ out[i] = f32(unpack_bf16_1(u32(token - 0) * 2560u + i % 2560u));
17
+ }
exaone-deep-fp32-8k-state-alias-token-major-v2-v5/kernels/397c272c70f921a18fa5b359f241c16191dabb59affdb25a5f21e0ac97b6ed39.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 * 10240u;
7
+ if (i >= 10240u) { return; }
8
+ out[i] = f32(b0[((i / 10240u) % 1u) * 10240u + ((i / 1280u) % 8u) * 80u + ((i / 80u) % 16u) * 640u + ((i / 1u) % 80u) * 1u]);
9
+ }
exaone-deep-fp32-8k-state-alias-token-major-v2-v5/kernels/39b0e20440b4b41f452d8130d42c31a13a660ab22bd56ebb653bc40eb85d5890.wgsl ADDED
@@ -0,0 +1,59 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ @group(0) @binding(0) var<storage, read> b0: array<f32>;
2
+ @group(0) @binding(1) var<storage, read> b1: array<u32>;
3
+ @group(0) @binding(2) var<storage, read_write> out: array<f32>;
4
+ fn unpack_bf16_1(index: u32) -> f32 {
5
+ let pair = b1[index / 2u];
6
+ let bits = (pair >> ((index % 2u) * 16u)) & 65535u;
7
+ return bitcast<f32>(bits << 16u);
8
+ }
9
+
10
+ var<workgroup> tile_a: array<array<f32, 16>, 16>;
11
+ var<workgroup> tile_b: array<array<f32, 64>, 16>;
12
+ @compute @workgroup_size(16, 16)
13
+ fn main(@builtin(workgroup_id) group: vec3<u32>, @builtin(local_invocation_id) local: vec3<u32>) {
14
+ let lane = local.y * 16u + local.x;
15
+ let batch = group.z;
16
+ let tile_n = group.x * 64u;
17
+ let tile_row = group.y * 16u;
18
+ var acc: array<array<f32, 4>, 1>;
19
+ for (var k0 = 0u; k0 < 2560u; k0 += 16u) {
20
+ for (var e = 0u; e < 1u; e++) {
21
+ let flat = lane + e * 256u;
22
+ let m_local = flat / 16u;
23
+ let row = tile_row + m_local;
24
+ let col = k0 + flat % 16u;
25
+ var value = 0.0;
26
+ if (row < 4u && col < 2560u) { value = f32(f32(b0[(batch * 4u + row) * 2560u + col])); }
27
+ tile_a[m_local][flat % 16u] = value;
28
+ }
29
+ for (var e = 0u; e < 4u; e++) {
30
+ let n_local = lane / 4u;
31
+ let k_local = (lane % 4u) * 4u + e;
32
+ let n_index = tile_n + n_local;
33
+ let col = k0 + k_local;
34
+ var value = 0.0;
35
+ if (n_index < 23758u && col < 2560u) { value = f32(unpack_bf16_1(n_index * 2560u + col)); }
36
+ tile_b[k_local][n_local] = value;
37
+ }
38
+ workgroupBarrier();
39
+ for (var kk = 0u; kk < 16u; kk++) {
40
+ var b_values: array<f32, 4>;
41
+ for (var c = 0u; c < 4u; c++) { b_values[c] = tile_b[kk][local.x * 4u + c]; }
42
+ for (var r = 0u; r < 1u; r++) {
43
+ let a_value = tile_a[local.y * 1u + r][kk];
44
+ for (var c = 0u; c < 4u; c++) { acc[r][c] += a_value * b_values[c]; }
45
+ }
46
+ }
47
+ workgroupBarrier();
48
+ }
49
+ for (var r = 0u; r < 1u; r++) {
50
+ let out_row = tile_row + local.y * 1u + r;
51
+ for (var c = 0u; c < 4u; c++) {
52
+ let out_col = tile_n + local.x * 4u + c;
53
+ if (out_row < 4u && out_col < 23758u) {
54
+ let i = (batch * 4u + out_row) * 23758u + out_col;
55
+ out[i] = f32(acc[r][c] + 0.0);
56
+ }
57
+ }
58
+ }
59
+ }
exaone-deep-fp32-8k-state-alias-token-major-v2-v5/kernels/3c1557fa28d308c0aae510ad9480e3454a28c5f9f409906a5b3ff186abba9ef3.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 * 1280u;
7
+ if (i >= 1280u) { return; }
8
+ out[i] = f32(b0[((i / 1280u) % 1u) * 2560u + ((i / 40u) % 32u) * 80u + ((i / 40u) % 1u) * 80u + (((i / 1u) % 40u) * 1u + 40u) * 1u]);
9
+ }
exaone-deep-fp32-8k-state-alias-token-major-v2-v5/kernels/3c55802ae4ce1a62daf3f5afd38a3ae905812bb4d990d1393cba3636a644994d.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 * 640u;
8
+ if (i >= 640u) { return; }
9
+ let token = (i / 80u) % 1u;
10
+ let outer = i / 80u;
11
+ let destination = outer * 655360u + u32(b1[token]) * 80u + i % 80u;
12
+ out[((destination) / 80u) % 8192u * 640u + ((destination) / 655360u) % 8u * 80u + ((destination) / 1u) % 80u * 1u] = f32(b0[i]);
13
+ }