Switch to Adam optimizer:
WIP
This commit is contained in:
@@ -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]);
|
||||
}
|
||||
@@ -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]);
|
||||
}
|
||||
@@ -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];
|
||||
}
|
||||
@@ -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]);
|
||||
}
|
||||
@@ -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];
|
||||
}
|
||||
@@ -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;
|
||||
}
|
||||
@@ -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;
|
||||
}
|
||||
@@ -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);
|
||||
}
|
||||
@@ -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;
|
||||
}
|
||||
Reference in New Issue
Block a user