Switch to Adam optimizer:

WIP
This commit is contained in:
2024-07-27 14:39:08 +00:00
parent dda7369ee7
commit ad8d7a8920
19 changed files with 806 additions and 308050 deletions
+28
View File
@@ -0,0 +1,28 @@
struct ScalarBatch {
size: vec4<u32>,
data: array<f32>,
}
struct Scalar {
size: vec4<u32>,
data: array<f32>,
}
struct MatrixBatch {
size: vec4<u32>,
data: array<vec4<f32>>,
}
@group(0) @binding(0)
var<storage, read> norm: ScalarBatch;
@group(0) @binding(1)
var<storage, read> threshold: Scalar;
@group(0) @binding(2)
var<storage, read_write> m: MatrixBatch;
@compute @workgroup_size(1)
fn main(@builtin(global_invocation_id) global_id: vec3<u32>) {
m.size = m.size;
m.data[(global_id.z * m.size.y + global_id.y) * m.size.x / 4 + global_id.x] *= min(1f, threshold.data[0] / norm.data[global_id.z]);
}
+21
View File
@@ -0,0 +1,21 @@
alias number = f32;
struct MatrixBatch {
size: vec4<u32>,
data: array<vec4<number>>,
}
@group(0) @binding(0)
var<storage, read> l: MatrixBatch;
@group(0) @binding(1)
var<storage, read> r_or_result: MatrixBatch;
@group(0) @binding(2)
var<storage, read_write> result: MatrixBatch;
@compute @workgroup_size(1)
fn main(@builtin(global_invocation_id) global_id: vec3<u32>) {
result.size = l.size;
var resi = (global_id.z * result.size.y + global_id.y) * result.size.x / 4 + global_id.x;
result.data[resi] = pow(l.data[resi], r_or_result.data[resi]);
}
+19
View File
@@ -0,0 +1,19 @@
alias number = f32;
struct MatrixBatch {
size: vec4<u32>,
data: array<vec4<number>>,
}
@group(0) @binding(0)
var<storage, read> v_or_result: MatrixBatch;
@group(0) @binding(1)
var<storage, read_write> result: MatrixBatch;
@compute @workgroup_size(1)
fn main(@builtin(global_invocation_id) global_id: vec3<u32>) {
result.size = v_or_result.size;
var resi = (global_id.z * result.size.y + global_id.y) * result.size.x / 4 + global_id.x;
result.data[resi] = v_or_result.data[resi] * v_or_result.data[resi];
}
+19
View File
@@ -0,0 +1,19 @@
alias number = f32;
struct MatrixBatch {
size: vec4<u32>,
data: array<vec4<number>>,
}
@group(0) @binding(0)
var<storage, read> v_or_result: MatrixBatch;
@group(0) @binding(1)
var<storage, read_write> result: MatrixBatch;
@compute @workgroup_size(1)
fn main(@builtin(global_invocation_id) global_id: vec3<u32>) {
result.size = v_or_result.size;
var resi = (global_id.z * result.size.y + global_id.y) * result.size.x / 4 + global_id.x;
result.data[resi] = sqrt(v_or_result.data[resi]);
}
+23
View File
@@ -0,0 +1,23 @@
alias number = f32;
struct MatrixBatch {
size: vec4<u32>,
data: array<vec4<number>>,
}
@group(0) @binding(0)
var<storage, read> l: MatrixBatch;
@group(0) @binding(1)
var<storage, read> m: MatrixBatch;
@group(0) @binding(2)
var<storage, read> r_or_result: MatrixBatch;
@group(0) @binding(3)
var<storage, read_write> result: MatrixBatch;
@compute @workgroup_size(1)
fn main(@builtin(global_invocation_id) global_id: vec3<u32>) {
result.size = l.size;
var resi = (global_id.z * result.size.y + global_id.y) * result.size.x / 4 + global_id.x;
result.data[resi] = l.data[resi] + m.data[resi] + r_or_result.data[resi];
}
+32
View File
@@ -0,0 +1,32 @@
alias number = f32;
struct MatrixBatch {
size: vec4<u32>,
data: array<number>,
}
@group(0) @binding(0)
var<storage, read> l: MatrixBatch;
@group(0) @binding(1)
var<storage, read> r: MatrixBatch;
@group(0) @binding(2)
var<storage, read_write> result: MatrixBatch;
@compute @workgroup_size(1)
fn main(@builtin(global_invocation_id) global_id: vec3<u32>) {
result.size.x = l.size.y;
result.size.y = r.size.x;
result.size.z = l.size.z;
var s = 0f;
var li = global_id.z * l.size.x * l.size.y + global_id.x * l.size.x;
var ri = global_id.z * r.size.x * r.size.y + global_id.y;
for (var i = 0u; i < l.size.x; i++) {
s += l.data[li] * r.data[ri];
li++;
ri += r.size.x;
}
let l = global_id.z * result.size.y * result.size.x;
result.data[l + global_id.y * result.size.x + global_id.x] = s;
}
+30
View File
@@ -0,0 +1,30 @@
alias number = f32;
struct MatrixBatch {
size: vec4<u32>,
data: array<number>,
}
struct VectorBatch {
size: vec4<u32>,
data: array<number>,
}
@group(0) @binding(0)
var<storage, read> m: MatrixBatch;
@group(0) @binding(1)
var<storage, read_write> result: VectorBatch;
@compute @workgroup_size(1)
fn main(@builtin(global_invocation_id) global_id: vec3<u32>) {
result.size.x = m.size.y;
result.size.z = m.size.z;
var s = 0f;
let lo = global_id.z * m.size.x * m.size.y + global_id.x * m.size.x;
for (var i = 0u; i < m.size.x; i++) {
s += m.data[lo + i] * m.data[lo + i];
}
result.data[global_id.z * result.size.x + global_id.x] = s;
}
+29
View File
@@ -0,0 +1,29 @@
alias number = f32;
struct VectorBatch {
size: vec4<u32>,
data: array<number>,
}
struct ScalarBatch {
size: vec4<u32>,
data: array<number>,
}
@group(0) @binding(0)
var<storage, read> v: VectorBatch;
@group(0) @binding(1)
var<storage, read_write> result: ScalarBatch;
@compute @workgroup_size(1)
fn main(@builtin(global_invocation_id) global_id: vec3<u32>) {
result.size.z = v.size.z;
var s = 0f;
let lo = global_id.z * v.size.x;
for (var i = 0u; i < v.size.x; i++) {
s += v.data[lo + i];
}
result.data[global_id.z] = sqrt(s);
}
+31
View File
@@ -0,0 +1,31 @@
alias number = f32;
struct MatrixBatch {
size: vec4<u32>,
data: array<number>,
}
@group(0) @binding(0)
var<storage, read> l: MatrixBatch;
@group(0) @binding(1)
var<storage, read> r: MatrixBatch;
@group(0) @binding(2)
var<storage, read_write> result: MatrixBatch;
@compute @workgroup_size(1)
fn main(@builtin(global_invocation_id) global_id: vec3<u32>) {
result.size.x = r.size.x;
result.size.y = l.size.x;
result.size.z = l.size.z;
var s = 0f;
var li = global_id.z * l.size.y * l.size.x + global_id.y;
var ri = global_id.z * r.size.y * r.size.x + global_id.x;
for (var i = 0u; i < l.size.y; i++) {
s += l.data[li] * r.data[ri];
li += l.size.x;
ri += r.size.x;
}
result.data[global_id.z * result.size.y * result.size.x + global_id.y * result.size.x + global_id.x] = s;
}