volvoxai 0.1.0
This diff represents the content of publicly available package versions that have been released to one of the supported registries. The information contained in this diff is provided for informational purposes only and reflects changes between package versions as they appear in their respective public registries.
- package/LICENSE +21 -0
- package/README.md +145 -0
- package/bin/volvox.js +72 -0
- package/dist/v0.1.0/volvoxai.js +4664 -0
- package/dist/v0.1.0/volvoxai.min.js +1848 -0
- package/dist/v0.1.0/volvoxai.wasm +0 -0
- package/dist/volvoxai.js +4664 -0
- package/dist/volvoxai.min.js +1848 -0
- package/dist/volvoxai.wasm +0 -0
- package/docs/README.md +22 -0
- package/docs/browser-runtime.md +87 -0
- package/docs/efficientdet_tflite_vs_volvoxai.md +445 -0
- package/docs/microkernel_optimization_guide.md +153 -0
- package/docs/model-format.md +108 -0
- package/docs/models.md +103 -0
- package/docs/native-runtime.md +189 -0
- package/docs/operation_list.md +232 -0
- package/docs/operator_fusion_patterns.md +58 -0
- package/docs/quickstart.md +115 -0
- package/docs/roadmap.md +19 -0
- package/docs/testing.md +97 -0
- package/docs/textbook/01-foundations.md +233 -0
- package/docs/textbook/02-tinystories-language-model.md +300 -0
- package/docs/textbook/03-efficientdet-vision-model.md +281 -0
- package/docs/textbook/04-precision-and-quantization.md +208 -0
- package/docs/textbook/05-inside-the-engine.md +155 -0
- package/docs/textbook/06-native-engine-architecture.md +338 -0
- package/docs/textbook/07-glossary-and-next-steps.md +258 -0
- package/docs/textbook/README.md +85 -0
- package/docs/textbook/ko/01-foundations.md +231 -0
- package/docs/textbook/ko/02-tinystories-language-model.md +300 -0
- package/docs/textbook/ko/03-efficientdet-vision-model.md +277 -0
- package/docs/textbook/ko/04-precision-and-quantization.md +206 -0
- package/docs/textbook/ko/05-inside-the-engine.md +154 -0
- package/docs/textbook/ko/06-native-engine-architecture.md +333 -0
- package/docs/textbook/ko/07-glossary-and-next-steps.md +253 -0
- package/docs/textbook/ko/README.md +83 -0
- package/docs/xnnpack_optimization_guide.md +197 -0
- package/js/CPUEngine.js +241 -0
- package/js/Graph.js +49 -0
- package/js/GraphExecutor.js +1020 -0
- package/js/GraphLoader.js +282 -0
- package/js/ShaderLibrary.js +236 -0
- package/js/Tensor.js +25 -0
- package/js/Tokenizer.js +266 -0
- package/js/VolvoxAI.js +130 -0
- package/js/WasmEngine.js +378 -0
- package/js/WebNNEngine.js +169 -0
- package/js/index.js +11 -0
- package/js/ops/add.js +31 -0
- package/js/ops/argMax.js +33 -0
- package/js/ops/averagePool2D.js +38 -0
- package/js/ops/batchNorm2D.js +28 -0
- package/js/ops/cast.js +19 -0
- package/js/ops/clip.js +15 -0
- package/js/ops/concat2.js +18 -0
- package/js/ops/conv1D.js +35 -0
- package/js/ops/conv2D.js +70 -0
- package/js/ops/convTranspose2D.js +45 -0
- package/js/ops/crossAttention.js +69 -0
- package/js/ops/crossSDPA.js +41 -0
- package/js/ops/dequantizeLinear.js +9 -0
- package/js/ops/div.js +15 -0
- package/js/ops/embedding.js +14 -0
- package/js/ops/expand.js +24 -0
- package/js/ops/gELU.js +9 -0
- package/js/ops/gather.js +51 -0
- package/js/ops/gatherElements.js +33 -0
- package/js/ops/globalAveragePool.js +21 -0
- package/js/ops/hardSigmoid.js +12 -0
- package/js/ops/hardSwish.js +12 -0
- package/js/ops/interp1D.js +25 -0
- package/js/ops/layerNorm.js +25 -0
- package/js/ops/leakyReLU.js +10 -0
- package/js/ops/logSoftmax.js +15 -0
- package/js/ops/matMul.js +35 -0
- package/js/ops/maxPool2D.js +36 -0
- package/js/ops/meanHeight.js +17 -0
- package/js/ops/mul.js +31 -0
- package/js/ops/nonMaxSuppression.js +72 -0
- package/js/ops/pReLU.js +11 -0
- package/js/ops/pad.js +35 -0
- package/js/ops/profileX.js +22 -0
- package/js/ops/profileY.js +22 -0
- package/js/ops/rMSNorm.js +14 -0
- package/js/ops/reLU.js +8 -0
- package/js/ops/reduceMean.js +17 -0
- package/js/ops/reduceSum.js +19 -0
- package/js/ops/reshape.js +6 -0
- package/js/ops/resize.js +44 -0
- package/js/ops/sDPA.js +44 -0
- package/js/ops/siLU.js +8 -0
- package/js/ops/sigmoid.js +6 -0
- package/js/ops/slice.js +36 -0
- package/js/ops/softmax.js +18 -0
- package/js/ops/spatialSoftargmaxY.js +28 -0
- package/js/ops/split.js +24 -0
- package/js/ops/sub.js +11 -0
- package/js/ops/tanh.js +7 -0
- package/js/ops/transpose.js +34 -0
- package/js/ops/upsample2x.js +23 -0
- package/js/ops/where.js +15 -0
- package/package.json +33 -0
- package/shaders/add.wgsl +13 -0
- package/shaders/add3Relu.wgsl +23 -0
- package/shaders/addRelu.wgsl +22 -0
- package/shaders/averagePool2D.wgsl +24 -0
- package/shaders/batchNorm2D.wgsl +21 -0
- package/shaders/binaryBroadcast.wgsl +34 -0
- package/shaders/broadcastBinary.wgsl +26 -0
- package/shaders/clip.wgsl +10 -0
- package/shaders/concat2.wgsl +16 -0
- package/shaders/concatCopy.wgsl +10 -0
- package/shaders/concatSigmoidCopy.wgsl +16 -0
- package/shaders/conv1D.wgsl +37 -0
- package/shaders/conv2D.wgsl +80 -0
- package/shaders/conv2DDepthwise4.wgsl +74 -0
- package/shaders/conv2DDepthwise8.wgsl +66 -0
- package/shaders/conv2DPointwise16.wgsl +67 -0
- package/shaders/conv2DPointwise16Tile.wgsl +86 -0
- package/shaders/conv2DPointwise8.wgsl +85 -0
- package/shaders/conv2DPointwise8Vec2.wgsl +70 -0
- package/shaders/conv2DPointwise8Vec4.wgsl +65 -0
- package/shaders/conv2DRegularC3Out16.wgsl +75 -0
- package/shaders/convTranspose2D.wgsl +33 -0
- package/shaders/copy.wgsl +13 -0
- package/shaders/crossAttention.wgsl +140 -0
- package/shaders/crossAttentionF32.wgsl +98 -0
- package/shaders/crossSDPA.wgsl +74 -0
- package/shaders/dequantizeLinear.wgsl +14 -0
- package/shaders/div.wgsl +34 -0
- package/shaders/elementwise.wgsl +13 -0
- package/shaders/embedding.wgsl +22 -0
- package/shaders/expand.wgsl +18 -0
- package/shaders/gELU.wgsl +13 -0
- package/shaders/gather.wgsl +17 -0
- package/shaders/generalTranspose.wgsl +19 -0
- package/shaders/globalAveragePool.wgsl +19 -0
- package/shaders/hardSigmoid.wgsl +13 -0
- package/shaders/hardSwish.wgsl +13 -0
- package/shaders/interp1D.wgsl +28 -0
- package/shaders/layerNorm.wgsl +33 -0
- package/shaders/leakyReLU.wgsl +11 -0
- package/shaders/linearF32.wgsl +33 -0
- package/shaders/linearF32RowMajor.wgsl +24 -0
- package/shaders/linearInt8.wgsl +42 -0
- package/shaders/logSoftmax.wgsl +22 -0
- package/shaders/maxPool2D.wgsl +37 -0
- package/shaders/meanHeight.wgsl +18 -0
- package/shaders/mul.wgsl +32 -0
- package/shaders/nonMaxSuppression.wgsl +92 -0
- package/shaders/pReLU.wgsl +14 -0
- package/shaders/pad.wgsl +19 -0
- package/shaders/profileX.wgsl +28 -0
- package/shaders/profileY.wgsl +28 -0
- package/shaders/quantizeLinear.wgsl +69 -0
- package/shaders/rMSNorm.wgsl +21 -0
- package/shaders/reLU.wgsl +13 -0
- package/shaders/reduce.wgsl +17 -0
- package/shaders/resize.wgsl +52 -0
- package/shaders/sDPA.wgsl +71 -0
- package/shaders/siLU.wgsl +13 -0
- package/shaders/sigmoid.wgsl +13 -0
- package/shaders/slice.wgsl +26 -0
- package/shaders/softmax.wgsl +23 -0
- package/shaders/spatialSoftargmaxY.wgsl +32 -0
- package/shaders/split.wgsl +15 -0
- package/shaders/sub.wgsl +34 -0
- package/shaders/tanh.wgsl +13 -0
- package/shaders/upsample2x.wgsl +24 -0
- package/shaders/where.wgsl +12 -0
- package/volvoxai.wasm +0 -0
|
@@ -0,0 +1,86 @@
|
|
|
1
|
+
@group(0) @binding(0) var<storage, read> input : array<f32>;
|
|
2
|
+
@group(0) @binding(1) var<storage, read> weight : array<vec4<f32>>;
|
|
3
|
+
@group(0) @binding(2) var<storage, read> bias : array<vec4<f32>>;
|
|
4
|
+
@group(0) @binding(3) var<storage, read_write> output : array<vec4<f32>>;
|
|
5
|
+
|
|
6
|
+
struct Params {
|
|
7
|
+
n : u32,
|
|
8
|
+
in_h : u32,
|
|
9
|
+
in_w : u32,
|
|
10
|
+
in_c : u32,
|
|
11
|
+
out_c : u32,
|
|
12
|
+
out_h : u32,
|
|
13
|
+
out_w : u32,
|
|
14
|
+
kh : u32,
|
|
15
|
+
kw : u32,
|
|
16
|
+
sy : u32,
|
|
17
|
+
sx : u32,
|
|
18
|
+
pt : u32,
|
|
19
|
+
pl : u32,
|
|
20
|
+
groups : u32,
|
|
21
|
+
relu : u32,
|
|
22
|
+
dy : u32,
|
|
23
|
+
dx : u32,
|
|
24
|
+
}
|
|
25
|
+
@group(0) @binding(4) var<uniform> params : Params;
|
|
26
|
+
|
|
27
|
+
const TILE_C : u32 = 64u;
|
|
28
|
+
var<workgroup> tile_weight : array<vec4<f32>, 256>;
|
|
29
|
+
|
|
30
|
+
fn apply_relu4(v_in : vec4<f32>) -> vec4<f32> {
|
|
31
|
+
var v = v_in;
|
|
32
|
+
if (params.relu == 1u) {
|
|
33
|
+
v = max(v, vec4<f32>(0.0));
|
|
34
|
+
} else if (params.relu >= 2u) {
|
|
35
|
+
v = min(max(v, vec4<f32>(0.0)), vec4<f32>(6.0));
|
|
36
|
+
}
|
|
37
|
+
return v;
|
|
38
|
+
}
|
|
39
|
+
|
|
40
|
+
@compute @workgroup_size(8, 8, 1)
|
|
41
|
+
fn main(@builtin(global_invocation_id) gid : vec3<u32>,
|
|
42
|
+
@builtin(local_invocation_id) lid3 : vec3<u32>) {
|
|
43
|
+
let ox = gid.x;
|
|
44
|
+
let oy = gid.y;
|
|
45
|
+
let oc16_count = params.out_c / 16u;
|
|
46
|
+
let nb = gid.z / oc16_count;
|
|
47
|
+
let oc0 = (gid.z - nb * oc16_count) * 16u;
|
|
48
|
+
let lid = lid3.y * 8u + lid3.x;
|
|
49
|
+
let in_bounds = nb < params.n && ox < params.out_w && oy < params.out_h;
|
|
50
|
+
|
|
51
|
+
var sum0 = vec4<f32>(0.0);
|
|
52
|
+
var sum1 = vec4<f32>(0.0);
|
|
53
|
+
var sum2 = vec4<f32>(0.0);
|
|
54
|
+
var sum3 = vec4<f32>(0.0);
|
|
55
|
+
let input_base = ((nb * params.in_h + oy) * params.in_w + ox) * params.in_c;
|
|
56
|
+
|
|
57
|
+
for (var tile_start = 0u; tile_start < params.in_c; tile_start = tile_start + TILE_C) {
|
|
58
|
+
let tile_count = min(TILE_C, params.in_c - tile_start);
|
|
59
|
+
for (var wi = lid; wi < tile_count * 4u; wi = wi + 64u) {
|
|
60
|
+
let ic = tile_start + wi / 4u;
|
|
61
|
+
let oc_vec = wi - (wi / 4u) * 4u;
|
|
62
|
+
tile_weight[wi] = weight[(ic * params.out_c + oc0) / 4u + oc_vec];
|
|
63
|
+
}
|
|
64
|
+
workgroupBarrier();
|
|
65
|
+
|
|
66
|
+
if (in_bounds) {
|
|
67
|
+
for (var ti = 0u; ti < tile_count; ti = ti + 1u) {
|
|
68
|
+
let x = input[input_base + tile_start + ti];
|
|
69
|
+
let wbase = ti * 4u;
|
|
70
|
+
sum0 = sum0 + x * tile_weight[wbase];
|
|
71
|
+
sum1 = sum1 + x * tile_weight[wbase + 1u];
|
|
72
|
+
sum2 = sum2 + x * tile_weight[wbase + 2u];
|
|
73
|
+
sum3 = sum3 + x * tile_weight[wbase + 3u];
|
|
74
|
+
}
|
|
75
|
+
}
|
|
76
|
+
workgroupBarrier();
|
|
77
|
+
}
|
|
78
|
+
|
|
79
|
+
if (!in_bounds) { return; }
|
|
80
|
+
let output_base = (((nb * params.out_h + oy) * params.out_w + ox) * params.out_c + oc0) / 4u;
|
|
81
|
+
let bias_base = oc0 / 4u;
|
|
82
|
+
output[output_base] = apply_relu4(sum0 + bias[bias_base]);
|
|
83
|
+
output[output_base + 1u] = apply_relu4(sum1 + bias[bias_base + 1u]);
|
|
84
|
+
output[output_base + 2u] = apply_relu4(sum2 + bias[bias_base + 2u]);
|
|
85
|
+
output[output_base + 3u] = apply_relu4(sum3 + bias[bias_base + 3u]);
|
|
86
|
+
}
|
|
@@ -0,0 +1,85 @@
|
|
|
1
|
+
@group(0) @binding(0) var<storage, read> input : array<f32>;
|
|
2
|
+
@group(0) @binding(1) var<storage, read> weight : array<f32>;
|
|
3
|
+
@group(0) @binding(2) var<storage, read> bias : array<f32>;
|
|
4
|
+
@group(0) @binding(3) var<storage, read_write> output : array<f32>;
|
|
5
|
+
|
|
6
|
+
struct Params {
|
|
7
|
+
n : u32,
|
|
8
|
+
in_h : u32,
|
|
9
|
+
in_w : u32,
|
|
10
|
+
in_c : u32,
|
|
11
|
+
out_c : u32,
|
|
12
|
+
out_h : u32,
|
|
13
|
+
out_w : u32,
|
|
14
|
+
kh : u32,
|
|
15
|
+
kw : u32,
|
|
16
|
+
sy : u32,
|
|
17
|
+
sx : u32,
|
|
18
|
+
pt : u32,
|
|
19
|
+
pl : u32,
|
|
20
|
+
groups : u32,
|
|
21
|
+
relu : u32,
|
|
22
|
+
dy : u32,
|
|
23
|
+
dx : u32,
|
|
24
|
+
}
|
|
25
|
+
@group(0) @binding(4) var<uniform> params : Params;
|
|
26
|
+
|
|
27
|
+
fn apply_relu(v_in : f32) -> f32 {
|
|
28
|
+
var v = v_in;
|
|
29
|
+
if (params.relu == 1u) {
|
|
30
|
+
v = max(v, 0.0);
|
|
31
|
+
} else if (params.relu >= 2u) {
|
|
32
|
+
v = min(max(v, 0.0), 6.0);
|
|
33
|
+
}
|
|
34
|
+
return v;
|
|
35
|
+
}
|
|
36
|
+
|
|
37
|
+
@compute @workgroup_size(8, 8, 1)
|
|
38
|
+
fn main(@builtin(global_invocation_id) gid : vec3<u32>) {
|
|
39
|
+
let ox = gid.x;
|
|
40
|
+
let oy = gid.y;
|
|
41
|
+
let oc8_count = (params.out_c + 7u) / 8u;
|
|
42
|
+
let nb = gid.z / oc8_count;
|
|
43
|
+
let oc0 = (gid.z - nb * oc8_count) * 8u;
|
|
44
|
+
if (nb >= params.n || ox >= params.out_w || oy >= params.out_h || oc0 >= params.out_c) { return; }
|
|
45
|
+
|
|
46
|
+
let has1 = oc0 + 1u < params.out_c;
|
|
47
|
+
let has2 = oc0 + 2u < params.out_c;
|
|
48
|
+
let has3 = oc0 + 3u < params.out_c;
|
|
49
|
+
let has4 = oc0 + 4u < params.out_c;
|
|
50
|
+
let has5 = oc0 + 5u < params.out_c;
|
|
51
|
+
let has6 = oc0 + 6u < params.out_c;
|
|
52
|
+
let has7 = oc0 + 7u < params.out_c;
|
|
53
|
+
var sum0 = 0.0;
|
|
54
|
+
var sum1 = 0.0;
|
|
55
|
+
var sum2 = 0.0;
|
|
56
|
+
var sum3 = 0.0;
|
|
57
|
+
var sum4 = 0.0;
|
|
58
|
+
var sum5 = 0.0;
|
|
59
|
+
var sum6 = 0.0;
|
|
60
|
+
var sum7 = 0.0;
|
|
61
|
+
|
|
62
|
+
let input_base = ((nb * params.in_h + oy) * params.in_w + ox) * params.in_c;
|
|
63
|
+
for (var ic = 0u; ic < params.in_c; ic = ic + 1u) {
|
|
64
|
+
let x = input[input_base + ic];
|
|
65
|
+
let wbase = ic * params.out_c + oc0;
|
|
66
|
+
sum0 = sum0 + x * weight[wbase];
|
|
67
|
+
if (has1) { sum1 = sum1 + x * weight[wbase + 1u]; }
|
|
68
|
+
if (has2) { sum2 = sum2 + x * weight[wbase + 2u]; }
|
|
69
|
+
if (has3) { sum3 = sum3 + x * weight[wbase + 3u]; }
|
|
70
|
+
if (has4) { sum4 = sum4 + x * weight[wbase + 4u]; }
|
|
71
|
+
if (has5) { sum5 = sum5 + x * weight[wbase + 5u]; }
|
|
72
|
+
if (has6) { sum6 = sum6 + x * weight[wbase + 6u]; }
|
|
73
|
+
if (has7) { sum7 = sum7 + x * weight[wbase + 7u]; }
|
|
74
|
+
}
|
|
75
|
+
|
|
76
|
+
let output_base = ((nb * params.out_h + oy) * params.out_w + ox) * params.out_c + oc0;
|
|
77
|
+
output[output_base] = apply_relu(sum0 + bias[oc0]);
|
|
78
|
+
if (has1) { output[output_base + 1u] = apply_relu(sum1 + bias[oc0 + 1u]); }
|
|
79
|
+
if (has2) { output[output_base + 2u] = apply_relu(sum2 + bias[oc0 + 2u]); }
|
|
80
|
+
if (has3) { output[output_base + 3u] = apply_relu(sum3 + bias[oc0 + 3u]); }
|
|
81
|
+
if (has4) { output[output_base + 4u] = apply_relu(sum4 + bias[oc0 + 4u]); }
|
|
82
|
+
if (has5) { output[output_base + 5u] = apply_relu(sum5 + bias[oc0 + 5u]); }
|
|
83
|
+
if (has6) { output[output_base + 6u] = apply_relu(sum6 + bias[oc0 + 6u]); }
|
|
84
|
+
if (has7) { output[output_base + 7u] = apply_relu(sum7 + bias[oc0 + 7u]); }
|
|
85
|
+
}
|
|
@@ -0,0 +1,70 @@
|
|
|
1
|
+
@group(0) @binding(0) var<storage, read> input : array<f32>;
|
|
2
|
+
@group(0) @binding(1) var<storage, read> weight : array<vec2<f32>>;
|
|
3
|
+
@group(0) @binding(2) var<storage, read> bias : array<vec2<f32>>;
|
|
4
|
+
@group(0) @binding(3) var<storage, read_write> output : array<vec2<f32>>;
|
|
5
|
+
|
|
6
|
+
struct Params {
|
|
7
|
+
n : u32,
|
|
8
|
+
in_h : u32,
|
|
9
|
+
in_w : u32,
|
|
10
|
+
in_c : u32,
|
|
11
|
+
out_c : u32,
|
|
12
|
+
out_h : u32,
|
|
13
|
+
out_w : u32,
|
|
14
|
+
kh : u32,
|
|
15
|
+
kw : u32,
|
|
16
|
+
sy : u32,
|
|
17
|
+
sx : u32,
|
|
18
|
+
pt : u32,
|
|
19
|
+
pl : u32,
|
|
20
|
+
groups : u32,
|
|
21
|
+
relu : u32,
|
|
22
|
+
dy : u32,
|
|
23
|
+
dx : u32,
|
|
24
|
+
}
|
|
25
|
+
@group(0) @binding(4) var<uniform> params : Params;
|
|
26
|
+
|
|
27
|
+
fn apply_relu2(v_in : vec2<f32>) -> vec2<f32> {
|
|
28
|
+
var v = v_in;
|
|
29
|
+
if (params.relu == 1u) {
|
|
30
|
+
v = max(v, vec2<f32>(0.0));
|
|
31
|
+
} else if (params.relu >= 2u) {
|
|
32
|
+
v = min(max(v, vec2<f32>(0.0)), vec2<f32>(6.0));
|
|
33
|
+
}
|
|
34
|
+
return v;
|
|
35
|
+
}
|
|
36
|
+
|
|
37
|
+
@compute @workgroup_size(8, 8, 1)
|
|
38
|
+
fn main(@builtin(global_invocation_id) gid : vec3<u32>) {
|
|
39
|
+
let ox = gid.x;
|
|
40
|
+
let oy = gid.y;
|
|
41
|
+
let oc8_count = (params.out_c + 7u) / 8u;
|
|
42
|
+
let nb = gid.z / oc8_count;
|
|
43
|
+
let oc0 = (gid.z - nb * oc8_count) * 8u;
|
|
44
|
+
if (nb >= params.n || ox >= params.out_w || oy >= params.out_h || oc0 >= params.out_c) { return; }
|
|
45
|
+
|
|
46
|
+
let has1 = oc0 + 2u < params.out_c;
|
|
47
|
+
let has2 = oc0 + 4u < params.out_c;
|
|
48
|
+
let has3 = oc0 + 6u < params.out_c;
|
|
49
|
+
var sum0 = vec2<f32>(0.0);
|
|
50
|
+
var sum1 = vec2<f32>(0.0);
|
|
51
|
+
var sum2 = vec2<f32>(0.0);
|
|
52
|
+
var sum3 = vec2<f32>(0.0);
|
|
53
|
+
|
|
54
|
+
let input_base = ((nb * params.in_h + oy) * params.in_w + ox) * params.in_c;
|
|
55
|
+
for (var ic = 0u; ic < params.in_c; ic = ic + 1u) {
|
|
56
|
+
let x = input[input_base + ic];
|
|
57
|
+
let wbase = (ic * params.out_c + oc0) / 2u;
|
|
58
|
+
sum0 = sum0 + x * weight[wbase];
|
|
59
|
+
if (has1) { sum1 = sum1 + x * weight[wbase + 1u]; }
|
|
60
|
+
if (has2) { sum2 = sum2 + x * weight[wbase + 2u]; }
|
|
61
|
+
if (has3) { sum3 = sum3 + x * weight[wbase + 3u]; }
|
|
62
|
+
}
|
|
63
|
+
|
|
64
|
+
let output_base = (((nb * params.out_h + oy) * params.out_w + ox) * params.out_c + oc0) / 2u;
|
|
65
|
+
let bias_base = oc0 / 2u;
|
|
66
|
+
output[output_base] = apply_relu2(sum0 + bias[bias_base]);
|
|
67
|
+
if (has1) { output[output_base + 1u] = apply_relu2(sum1 + bias[bias_base + 1u]); }
|
|
68
|
+
if (has2) { output[output_base + 2u] = apply_relu2(sum2 + bias[bias_base + 2u]); }
|
|
69
|
+
if (has3) { output[output_base + 3u] = apply_relu2(sum3 + bias[bias_base + 3u]); }
|
|
70
|
+
}
|
|
@@ -0,0 +1,65 @@
|
|
|
1
|
+
@group(0) @binding(0) var<storage, read> input : array<f32>;
|
|
2
|
+
@group(0) @binding(1) var<storage, read> weight : array<vec4<f32>>;
|
|
3
|
+
@group(0) @binding(2) var<storage, read> bias : array<vec4<f32>>;
|
|
4
|
+
@group(0) @binding(3) var<storage, read_write> output : array<vec4<f32>>;
|
|
5
|
+
|
|
6
|
+
struct Params {
|
|
7
|
+
n : u32,
|
|
8
|
+
in_h : u32,
|
|
9
|
+
in_w : u32,
|
|
10
|
+
in_c : u32,
|
|
11
|
+
out_c : u32,
|
|
12
|
+
out_h : u32,
|
|
13
|
+
out_w : u32,
|
|
14
|
+
kh : u32,
|
|
15
|
+
kw : u32,
|
|
16
|
+
sy : u32,
|
|
17
|
+
sx : u32,
|
|
18
|
+
pt : u32,
|
|
19
|
+
pl : u32,
|
|
20
|
+
groups : u32,
|
|
21
|
+
relu : u32,
|
|
22
|
+
dy : u32,
|
|
23
|
+
dx : u32,
|
|
24
|
+
}
|
|
25
|
+
@group(0) @binding(4) var<uniform> params : Params;
|
|
26
|
+
|
|
27
|
+
fn apply_relu4(v_in : vec4<f32>) -> vec4<f32> {
|
|
28
|
+
var v = v_in;
|
|
29
|
+
if (params.relu == 1u) {
|
|
30
|
+
v = max(v, vec4<f32>(0.0));
|
|
31
|
+
} else if (params.relu >= 2u) {
|
|
32
|
+
v = min(max(v, vec4<f32>(0.0)), vec4<f32>(6.0));
|
|
33
|
+
}
|
|
34
|
+
return v;
|
|
35
|
+
}
|
|
36
|
+
|
|
37
|
+
@compute @workgroup_size(8, 8, 1)
|
|
38
|
+
fn main(@builtin(global_invocation_id) gid : vec3<u32>) {
|
|
39
|
+
let ox = gid.x;
|
|
40
|
+
let oy = gid.y;
|
|
41
|
+
let oc8_count = (params.out_c + 7u) / 8u;
|
|
42
|
+
let nb = gid.z / oc8_count;
|
|
43
|
+
let oc0 = (gid.z - nb * oc8_count) * 8u;
|
|
44
|
+
if (nb >= params.n || ox >= params.out_w || oy >= params.out_h || oc0 >= params.out_c) { return; }
|
|
45
|
+
|
|
46
|
+
var sum0 = vec4<f32>(0.0);
|
|
47
|
+
var sum1 = vec4<f32>(0.0);
|
|
48
|
+
let input_base = ((nb * params.in_h + oy) * params.in_w + ox) * params.in_c;
|
|
49
|
+
let has_second = oc0 + 4u < params.out_c;
|
|
50
|
+
for (var ic = 0u; ic < params.in_c; ic = ic + 1u) {
|
|
51
|
+
let x = input[input_base + ic];
|
|
52
|
+
let wbase = (ic * params.out_c + oc0) / 4u;
|
|
53
|
+
sum0 = sum0 + x * weight[wbase];
|
|
54
|
+
if (has_second) {
|
|
55
|
+
sum1 = sum1 + x * weight[wbase + 1u];
|
|
56
|
+
}
|
|
57
|
+
}
|
|
58
|
+
|
|
59
|
+
let output_base = (((nb * params.out_h + oy) * params.out_w + ox) * params.out_c + oc0) / 4u;
|
|
60
|
+
let bias_base = oc0 / 4u;
|
|
61
|
+
output[output_base] = apply_relu4(sum0 + bias[bias_base]);
|
|
62
|
+
if (has_second) {
|
|
63
|
+
output[output_base + 1u] = apply_relu4(sum1 + bias[bias_base + 1u]);
|
|
64
|
+
}
|
|
65
|
+
}
|
|
@@ -0,0 +1,75 @@
|
|
|
1
|
+
@group(0) @binding(0) var<storage, read> input : array<f32>;
|
|
2
|
+
@group(0) @binding(1) var<storage, read> weight : array<vec4<f32>>;
|
|
3
|
+
@group(0) @binding(2) var<storage, read> bias : array<vec4<f32>>;
|
|
4
|
+
@group(0) @binding(3) var<storage, read_write> output : array<vec4<f32>>;
|
|
5
|
+
|
|
6
|
+
struct Params {
|
|
7
|
+
n : u32,
|
|
8
|
+
in_h : u32,
|
|
9
|
+
in_w : u32,
|
|
10
|
+
in_c : u32,
|
|
11
|
+
out_c : u32,
|
|
12
|
+
out_h : u32,
|
|
13
|
+
out_w : u32,
|
|
14
|
+
kh : u32,
|
|
15
|
+
kw : u32,
|
|
16
|
+
sy : u32,
|
|
17
|
+
sx : u32,
|
|
18
|
+
pt : u32,
|
|
19
|
+
pl : u32,
|
|
20
|
+
groups : u32,
|
|
21
|
+
relu : u32,
|
|
22
|
+
dy : u32,
|
|
23
|
+
dx : u32,
|
|
24
|
+
}
|
|
25
|
+
@group(0) @binding(4) var<uniform> params : Params;
|
|
26
|
+
|
|
27
|
+
fn apply_relu4(v_in : vec4<f32>) -> vec4<f32> {
|
|
28
|
+
var v = v_in;
|
|
29
|
+
if (params.relu == 1u) {
|
|
30
|
+
v = max(v, vec4<f32>(0.0));
|
|
31
|
+
} else if (params.relu >= 2u) {
|
|
32
|
+
v = min(max(v, vec4<f32>(0.0)), vec4<f32>(6.0));
|
|
33
|
+
}
|
|
34
|
+
return v;
|
|
35
|
+
}
|
|
36
|
+
|
|
37
|
+
@compute @workgroup_size(8, 8, 1)
|
|
38
|
+
fn main(@builtin(global_invocation_id) gid : vec3<u32>) {
|
|
39
|
+
let ox = gid.x;
|
|
40
|
+
let oy = gid.y;
|
|
41
|
+
let oc16_count = params.out_c / 16u;
|
|
42
|
+
let nb = gid.z / oc16_count;
|
|
43
|
+
let oc0 = (gid.z - nb * oc16_count) * 16u;
|
|
44
|
+
if (nb >= params.n || ox >= params.out_w || oy >= params.out_h) { return; }
|
|
45
|
+
|
|
46
|
+
var sum0 = vec4<f32>(0.0);
|
|
47
|
+
var sum1 = vec4<f32>(0.0);
|
|
48
|
+
var sum2 = vec4<f32>(0.0);
|
|
49
|
+
var sum3 = vec4<f32>(0.0);
|
|
50
|
+
for (var yy = 0u; yy < params.kh; yy = yy + 1u) {
|
|
51
|
+
let iy = i32(oy * params.sy + yy * params.dy) - i32(params.pt);
|
|
52
|
+
if (iy < 0 || iy >= i32(params.in_h)) { continue; }
|
|
53
|
+
for (var xx = 0u; xx < params.kw; xx = xx + 1u) {
|
|
54
|
+
let ix = i32(ox * params.sx + xx * params.dx) - i32(params.pl);
|
|
55
|
+
if (ix < 0 || ix >= i32(params.in_w)) { continue; }
|
|
56
|
+
let input_base = ((nb * params.in_h + u32(iy)) * params.in_w + u32(ix)) * 3u;
|
|
57
|
+
let weight_base = (((yy * params.kw + xx) * 3u) * params.out_c + oc0) / 4u;
|
|
58
|
+
for (var ic = 0u; ic < 3u; ic = ic + 1u) {
|
|
59
|
+
let x = input[input_base + ic];
|
|
60
|
+
let wbase = weight_base + ic * (params.out_c / 4u);
|
|
61
|
+
sum0 = sum0 + x * weight[wbase];
|
|
62
|
+
sum1 = sum1 + x * weight[wbase + 1u];
|
|
63
|
+
sum2 = sum2 + x * weight[wbase + 2u];
|
|
64
|
+
sum3 = sum3 + x * weight[wbase + 3u];
|
|
65
|
+
}
|
|
66
|
+
}
|
|
67
|
+
}
|
|
68
|
+
|
|
69
|
+
let output_base = (((nb * params.out_h + oy) * params.out_w + ox) * params.out_c + oc0) / 4u;
|
|
70
|
+
let bias_base = oc0 / 4u;
|
|
71
|
+
output[output_base] = apply_relu4(sum0 + bias[bias_base]);
|
|
72
|
+
output[output_base + 1u] = apply_relu4(sum1 + bias[bias_base + 1u]);
|
|
73
|
+
output[output_base + 2u] = apply_relu4(sum2 + bias[bias_base + 2u]);
|
|
74
|
+
output[output_base + 3u] = apply_relu4(sum3 + bias[bias_base + 3u]);
|
|
75
|
+
}
|
|
@@ -0,0 +1,33 @@
|
|
|
1
|
+
@group(0) @binding(0) var<storage, read> input : array<f32>;
|
|
2
|
+
@group(0) @binding(1) var<storage, read> weight : array<f32>;
|
|
3
|
+
@group(0) @binding(2) var<storage, read> bias : array<f32>;
|
|
4
|
+
@group(0) @binding(3) var<storage, read_write> output : array<f32>;
|
|
5
|
+
struct Params { b: u32, in_h: u32, in_w: u32, in_c: u32, out_h: u32, out_w: u32, out_c: u32, kh: u32, kw: u32, sh: u32, sw: u32, ph: u32, pw: u32, has_bias: u32 }
|
|
6
|
+
@group(0) @binding(4) var<uniform> params : Params;
|
|
7
|
+
@compute @workgroup_size(8, 8, 1)
|
|
8
|
+
fn main(@builtin(global_invocation_id) global_id : vec3<u32>) {
|
|
9
|
+
let x = global_id.x; let y = global_id.y; let batch_c = global_id.z;
|
|
10
|
+
if (x >= params.out_w || y >= params.out_h || batch_c >= (params.b * params.out_c)) { return; }
|
|
11
|
+
let ob = batch_c / params.out_c;
|
|
12
|
+
let oc = batch_c % params.out_c;
|
|
13
|
+
var sum = 0.0;
|
|
14
|
+
if (params.has_bias == 1u) { sum = bias[oc]; }
|
|
15
|
+
for (var ic = 0u; ic < params.in_c; ic = ic + 1u) {
|
|
16
|
+
for (var ky = 0u; ky < params.kh; ky = ky + 1u) {
|
|
17
|
+
for (var kx = 0u; kx < params.kw; kx = kx + 1u) {
|
|
18
|
+
let oy_shifted = i32(y) + i32(params.ph) - i32(ky);
|
|
19
|
+
let ox_shifted = i32(x) + i32(params.pw) - i32(kx);
|
|
20
|
+
if (oy_shifted % i32(params.sh) == 0 && ox_shifted % i32(params.sw) == 0) {
|
|
21
|
+
let iy = oy_shifted / i32(params.sh);
|
|
22
|
+
let ix = ox_shifted / i32(params.sw);
|
|
23
|
+
if (iy >= 0 && iy < i32(params.in_h) && ix >= 0 && ix < i32(params.in_w)) {
|
|
24
|
+
let in_val = input[((ob * params.in_h + u32(iy)) * params.in_w + u32(ix)) * params.in_c + ic];
|
|
25
|
+
let w_val = weight[ic * (params.out_c * params.kh * params.kw) + oc * (params.kh * params.kw) + ky * params.kw + kx];
|
|
26
|
+
sum = sum + in_val * w_val;
|
|
27
|
+
}
|
|
28
|
+
}
|
|
29
|
+
}
|
|
30
|
+
}
|
|
31
|
+
}
|
|
32
|
+
output[((ob * params.out_h + y) * params.out_w + x) * params.out_c + oc] = sum;
|
|
33
|
+
}
|
|
@@ -0,0 +1,13 @@
|
|
|
1
|
+
@group(0) @binding(0) var<storage, read> input : array<f32>;
|
|
2
|
+
@group(0) @binding(1) var<storage, read_write> output : array<f32>;
|
|
3
|
+
struct Params { size : u32 }
|
|
4
|
+
@group(0) @binding(2) var<uniform> params : Params;
|
|
5
|
+
@compute @workgroup_size(64)
|
|
6
|
+
fn main(@builtin(global_invocation_id) global_id : vec3<u32>) {
|
|
7
|
+
let idx = global_id.x;
|
|
8
|
+
if (idx >= params.size) { return; }
|
|
9
|
+
let x = input[idx];
|
|
10
|
+
var out_val = x;
|
|
11
|
+
out_val = x;
|
|
12
|
+
output[idx] = out_val;
|
|
13
|
+
}
|
|
@@ -0,0 +1,140 @@
|
|
|
1
|
+
@group(0) @binding(0) var<storage, read> q_in : array<f32>;
|
|
2
|
+
@group(0) @binding(1) var<storage, read> kv_in : array<f32>;
|
|
3
|
+
@group(0) @binding(2) var<storage, read> weight_int8_packed : array<u32>;
|
|
4
|
+
@group(0) @binding(3) var<storage, read> w_scale : array<f32>;
|
|
5
|
+
@group(0) @binding(4) var<storage, read> bias : array<f32>;
|
|
6
|
+
@group(0) @binding(5) var<storage, read_write> output : array<f32>;
|
|
7
|
+
|
|
8
|
+
struct Params {
|
|
9
|
+
seq_len_q : u32,
|
|
10
|
+
seq_len_kv : u32,
|
|
11
|
+
d_model : u32,
|
|
12
|
+
num_heads : u32,
|
|
13
|
+
head_dim : u32,
|
|
14
|
+
scale_factor : f32,
|
|
15
|
+
has_scale : u32,
|
|
16
|
+
has_bias : u32,
|
|
17
|
+
}
|
|
18
|
+
@group(0) @binding(6) var<uniform> params : Params;
|
|
19
|
+
|
|
20
|
+
@compute @workgroup_size(64, 1, 1)
|
|
21
|
+
fn main(@builtin(global_invocation_id) global_id : vec3<u32>) {
|
|
22
|
+
let q_idx = global_id.x;
|
|
23
|
+
let h_idx = global_id.y;
|
|
24
|
+
|
|
25
|
+
if (q_idx >= params.seq_len_q || h_idx >= params.num_heads) { return; }
|
|
26
|
+
|
|
27
|
+
let head_dim = params.head_dim;
|
|
28
|
+
let d_model = params.d_model;
|
|
29
|
+
let d_model_4 = d_model / 4u;
|
|
30
|
+
|
|
31
|
+
// Cache Q for this head and this q_idx
|
|
32
|
+
var q_cache : array<f32, 64>; // max head_dim 64
|
|
33
|
+
for (var d = 0u; d < head_dim; d = d + 1u) {
|
|
34
|
+
var sum = 0.0;
|
|
35
|
+
let out_col = h_idx * head_dim + d;
|
|
36
|
+
for (var i = 0u; i < d_model_4; i = i + 1u) {
|
|
37
|
+
let w_packed = weight_int8_packed[out_col * d_model_4 + i];
|
|
38
|
+
let in_base = q_idx * d_model + i * 4u;
|
|
39
|
+
sum = sum + q_in[in_base + 0u] * f32(extractBits(i32(w_packed), 0u, 8u));
|
|
40
|
+
sum = sum + q_in[in_base + 1u] * f32(extractBits(i32(w_packed), 8u, 8u));
|
|
41
|
+
sum = sum + q_in[in_base + 2u] * f32(extractBits(i32(w_packed), 16u, 8u));
|
|
42
|
+
sum = sum + q_in[in_base + 3u] * f32(extractBits(i32(w_packed), 24u, 8u));
|
|
43
|
+
}
|
|
44
|
+
if (params.has_scale > 0u) { sum = sum * w_scale[out_col]; }
|
|
45
|
+
if (params.has_bias > 0u) { sum = sum + bias[out_col]; }
|
|
46
|
+
q_cache[d] = sum;
|
|
47
|
+
}
|
|
48
|
+
|
|
49
|
+
var max_logit : f32 = -1e38;
|
|
50
|
+
|
|
51
|
+
// Pass 1: find max
|
|
52
|
+
for (var k_idx = 0u; k_idx < params.seq_len_kv; k_idx = k_idx + 1u) {
|
|
53
|
+
var score : f32 = 0.0;
|
|
54
|
+
for (var d = 0u; d < head_dim; d = d + 1u) {
|
|
55
|
+
var k_val = 0.0;
|
|
56
|
+
let out_col = d_model + h_idx * head_dim + d;
|
|
57
|
+
for (var i = 0u; i < d_model_4; i = i + 1u) {
|
|
58
|
+
let w_packed = weight_int8_packed[out_col * d_model_4 + i];
|
|
59
|
+
let in_base = k_idx * d_model + i * 4u;
|
|
60
|
+
k_val = k_val + kv_in[in_base + 0u] * f32(extractBits(i32(w_packed), 0u, 8u));
|
|
61
|
+
k_val = k_val + kv_in[in_base + 1u] * f32(extractBits(i32(w_packed), 8u, 8u));
|
|
62
|
+
k_val = k_val + kv_in[in_base + 2u] * f32(extractBits(i32(w_packed), 16u, 8u));
|
|
63
|
+
k_val = k_val + kv_in[in_base + 3u] * f32(extractBits(i32(w_packed), 24u, 8u));
|
|
64
|
+
}
|
|
65
|
+
if (params.has_scale > 0u) { k_val = k_val * w_scale[out_col]; }
|
|
66
|
+
if (params.has_bias > 0u) { k_val = k_val + bias[out_col]; }
|
|
67
|
+
|
|
68
|
+
score = score + (q_cache[d] * k_val);
|
|
69
|
+
}
|
|
70
|
+
score = score * params.scale_factor;
|
|
71
|
+
if (score > max_logit) { max_logit = score; }
|
|
72
|
+
}
|
|
73
|
+
|
|
74
|
+
// Pass 2: sum exp
|
|
75
|
+
var sum_exp : f32 = 0.0;
|
|
76
|
+
for (var k_idx = 0u; k_idx < params.seq_len_kv; k_idx = k_idx + 1u) {
|
|
77
|
+
var score : f32 = 0.0;
|
|
78
|
+
for (var d = 0u; d < head_dim; d = d + 1u) {
|
|
79
|
+
var k_val = 0.0;
|
|
80
|
+
let out_col = d_model + h_idx * head_dim + d;
|
|
81
|
+
for (var i = 0u; i < d_model_4; i = i + 1u) {
|
|
82
|
+
let w_packed = weight_int8_packed[out_col * d_model_4 + i];
|
|
83
|
+
let in_base = k_idx * d_model + i * 4u;
|
|
84
|
+
k_val = k_val + kv_in[in_base + 0u] * f32(extractBits(i32(w_packed), 0u, 8u));
|
|
85
|
+
k_val = k_val + kv_in[in_base + 1u] * f32(extractBits(i32(w_packed), 8u, 8u));
|
|
86
|
+
k_val = k_val + kv_in[in_base + 2u] * f32(extractBits(i32(w_packed), 16u, 8u));
|
|
87
|
+
k_val = k_val + kv_in[in_base + 3u] * f32(extractBits(i32(w_packed), 24u, 8u));
|
|
88
|
+
}
|
|
89
|
+
if (params.has_scale > 0u) { k_val = k_val * w_scale[out_col]; }
|
|
90
|
+
if (params.has_bias > 0u) { k_val = k_val + bias[out_col]; }
|
|
91
|
+
|
|
92
|
+
score = score + (q_cache[d] * k_val);
|
|
93
|
+
}
|
|
94
|
+
score = score * params.scale_factor;
|
|
95
|
+
sum_exp = sum_exp + exp(score - max_logit);
|
|
96
|
+
}
|
|
97
|
+
|
|
98
|
+
// Pass 3: output
|
|
99
|
+
for (var d = 0u; d < head_dim; d = d + 1u) {
|
|
100
|
+
var out_val : f32 = 0.0;
|
|
101
|
+
for (var k_idx = 0u; k_idx < params.seq_len_kv; k_idx = k_idx + 1u) {
|
|
102
|
+
var score : f32 = 0.0;
|
|
103
|
+
for (var kd = 0u; kd < head_dim; kd = kd + 1u) {
|
|
104
|
+
var k_val = 0.0;
|
|
105
|
+
let out_col = d_model + h_idx * head_dim + kd;
|
|
106
|
+
for (var i = 0u; i < d_model_4; i = i + 1u) {
|
|
107
|
+
let w_packed = weight_int8_packed[out_col * d_model_4 + i];
|
|
108
|
+
let in_base = k_idx * d_model + i * 4u;
|
|
109
|
+
k_val = k_val + kv_in[in_base + 0u] * f32(extractBits(i32(w_packed), 0u, 8u));
|
|
110
|
+
k_val = k_val + kv_in[in_base + 1u] * f32(extractBits(i32(w_packed), 8u, 8u));
|
|
111
|
+
k_val = k_val + kv_in[in_base + 2u] * f32(extractBits(i32(w_packed), 16u, 8u));
|
|
112
|
+
k_val = k_val + kv_in[in_base + 3u] * f32(extractBits(i32(w_packed), 24u, 8u));
|
|
113
|
+
}
|
|
114
|
+
if (params.has_scale > 0u) { k_val = k_val * w_scale[out_col]; }
|
|
115
|
+
if (params.has_bias > 0u) { k_val = k_val + bias[out_col]; }
|
|
116
|
+
|
|
117
|
+
score = score + (q_cache[kd] * k_val);
|
|
118
|
+
}
|
|
119
|
+
score = score * params.scale_factor;
|
|
120
|
+
|
|
121
|
+
let w = exp(score - max_logit) / sum_exp;
|
|
122
|
+
|
|
123
|
+
var v_val = 0.0;
|
|
124
|
+
let v_col = d_model * 2u + h_idx * head_dim + d;
|
|
125
|
+
for (var i = 0u; i < d_model_4; i = i + 1u) {
|
|
126
|
+
let w_packed = weight_int8_packed[v_col * d_model_4 + i];
|
|
127
|
+
let in_base = k_idx * d_model + i * 4u;
|
|
128
|
+
v_val = v_val + kv_in[in_base + 0u] * f32(extractBits(i32(w_packed), 0u, 8u));
|
|
129
|
+
v_val = v_val + kv_in[in_base + 1u] * f32(extractBits(i32(w_packed), 8u, 8u));
|
|
130
|
+
v_val = v_val + kv_in[in_base + 2u] * f32(extractBits(i32(w_packed), 16u, 8u));
|
|
131
|
+
v_val = v_val + kv_in[in_base + 3u] * f32(extractBits(i32(w_packed), 24u, 8u));
|
|
132
|
+
}
|
|
133
|
+
if (params.has_scale > 0u) { v_val = v_val * w_scale[v_col]; }
|
|
134
|
+
if (params.has_bias > 0u) { v_val = v_val + bias[v_col]; }
|
|
135
|
+
|
|
136
|
+
out_val = out_val + (w * v_val);
|
|
137
|
+
}
|
|
138
|
+
output[q_idx * d_model + (h_idx * head_dim) + d] = out_val;
|
|
139
|
+
}
|
|
140
|
+
}
|
|
@@ -0,0 +1,98 @@
|
|
|
1
|
+
@group(0) @binding(0) var<storage, read> q_in : array<f32>;
|
|
2
|
+
@group(0) @binding(1) var<storage, read> kv_in : array<f32>;
|
|
3
|
+
@group(0) @binding(2) var<storage, read> weight : array<f32>;
|
|
4
|
+
@group(0) @binding(3) var<storage, read> w_scale : array<f32>;
|
|
5
|
+
@group(0) @binding(4) var<storage, read> bias : array<f32>;
|
|
6
|
+
@group(0) @binding(5) var<storage, read_write> output : array<f32>;
|
|
7
|
+
|
|
8
|
+
struct Params {
|
|
9
|
+
seq_len_q : u32,
|
|
10
|
+
seq_len_kv : u32,
|
|
11
|
+
d_model : u32,
|
|
12
|
+
num_heads : u32,
|
|
13
|
+
head_dim : u32,
|
|
14
|
+
scale_factor : f32,
|
|
15
|
+
has_scale : u32,
|
|
16
|
+
has_bias : u32,
|
|
17
|
+
}
|
|
18
|
+
@group(0) @binding(6) var<uniform> params : Params;
|
|
19
|
+
|
|
20
|
+
fn project(src : ptr<function, array<f32, 64>>, row : u32, out_col : u32) -> f32 {
|
|
21
|
+
var sum = 0.0;
|
|
22
|
+
let w_base = out_col * params.d_model;
|
|
23
|
+
for (var i = 0u; i < params.d_model; i = i + 1u) {
|
|
24
|
+
sum = sum + (*src)[i] * weight[w_base + i];
|
|
25
|
+
}
|
|
26
|
+
if (params.has_scale > 0u) { sum = sum * w_scale[out_col]; }
|
|
27
|
+
if (params.has_bias > 0u) { sum = sum + bias[out_col]; }
|
|
28
|
+
return sum;
|
|
29
|
+
}
|
|
30
|
+
|
|
31
|
+
@compute @workgroup_size(64, 1, 1)
|
|
32
|
+
fn main(@builtin(global_invocation_id) gid : vec3<u32>) {
|
|
33
|
+
let q_idx = gid.x;
|
|
34
|
+
let h_idx = gid.y;
|
|
35
|
+
if (q_idx >= params.seq_len_q || h_idx >= params.num_heads) { return; }
|
|
36
|
+
if (params.d_model > 64u || params.head_dim > 64u) { return; }
|
|
37
|
+
|
|
38
|
+
var q_src : array<f32, 64>;
|
|
39
|
+
for (var i = 0u; i < params.d_model; i = i + 1u) {
|
|
40
|
+
q_src[i] = q_in[q_idx * params.d_model + i];
|
|
41
|
+
}
|
|
42
|
+
|
|
43
|
+
var q_proj : array<f32, 64>;
|
|
44
|
+
for (var d = 0u; d < params.head_dim; d = d + 1u) {
|
|
45
|
+
let out_col = h_idx * params.head_dim + d;
|
|
46
|
+
q_proj[d] = project(&q_src, q_idx, out_col);
|
|
47
|
+
}
|
|
48
|
+
|
|
49
|
+
var max_logit = -1e38;
|
|
50
|
+
for (var k_idx = 0u; k_idx < params.seq_len_kv; k_idx = k_idx + 1u) {
|
|
51
|
+
var kv_src : array<f32, 64>;
|
|
52
|
+
for (var i = 0u; i < params.d_model; i = i + 1u) {
|
|
53
|
+
kv_src[i] = kv_in[k_idx * params.d_model + i];
|
|
54
|
+
}
|
|
55
|
+
var score = 0.0;
|
|
56
|
+
for (var d = 0u; d < params.head_dim; d = d + 1u) {
|
|
57
|
+
let out_col = params.d_model + h_idx * params.head_dim + d;
|
|
58
|
+
score = score + q_proj[d] * project(&kv_src, k_idx, out_col);
|
|
59
|
+
}
|
|
60
|
+
score = score * params.scale_factor;
|
|
61
|
+
if (score > max_logit) { max_logit = score; }
|
|
62
|
+
}
|
|
63
|
+
|
|
64
|
+
var sum_exp = 0.0;
|
|
65
|
+
for (var k_idx = 0u; k_idx < params.seq_len_kv; k_idx = k_idx + 1u) {
|
|
66
|
+
var kv_src : array<f32, 64>;
|
|
67
|
+
for (var i = 0u; i < params.d_model; i = i + 1u) {
|
|
68
|
+
kv_src[i] = kv_in[k_idx * params.d_model + i];
|
|
69
|
+
}
|
|
70
|
+
var score = 0.0;
|
|
71
|
+
for (var d = 0u; d < params.head_dim; d = d + 1u) {
|
|
72
|
+
let out_col = params.d_model + h_idx * params.head_dim + d;
|
|
73
|
+
score = score + q_proj[d] * project(&kv_src, k_idx, out_col);
|
|
74
|
+
}
|
|
75
|
+
score = score * params.scale_factor;
|
|
76
|
+
sum_exp = sum_exp + exp(score - max_logit);
|
|
77
|
+
}
|
|
78
|
+
|
|
79
|
+
for (var d = 0u; d < params.head_dim; d = d + 1u) {
|
|
80
|
+
var out_val = 0.0;
|
|
81
|
+
for (var k_idx = 0u; k_idx < params.seq_len_kv; k_idx = k_idx + 1u) {
|
|
82
|
+
var kv_src : array<f32, 64>;
|
|
83
|
+
for (var i = 0u; i < params.d_model; i = i + 1u) {
|
|
84
|
+
kv_src[i] = kv_in[k_idx * params.d_model + i];
|
|
85
|
+
}
|
|
86
|
+
var score = 0.0;
|
|
87
|
+
for (var kd = 0u; kd < params.head_dim; kd = kd + 1u) {
|
|
88
|
+
let k_col = params.d_model + h_idx * params.head_dim + kd;
|
|
89
|
+
score = score + q_proj[kd] * project(&kv_src, k_idx, k_col);
|
|
90
|
+
}
|
|
91
|
+
score = score * params.scale_factor;
|
|
92
|
+
let attn = exp(score - max_logit) / sum_exp;
|
|
93
|
+
let v_col = params.d_model * 2u + h_idx * params.head_dim + d;
|
|
94
|
+
out_val = out_val + attn * project(&kv_src, k_idx, v_col);
|
|
95
|
+
}
|
|
96
|
+
output[q_idx * params.d_model + h_idx * params.head_dim + d] = out_val;
|
|
97
|
+
}
|
|
98
|
+
}
|