ggml: add f16 out_prod support for CPU and out_prod op for Vulkan (#23997)
This commit is contained in:
@@ -0,0 +1,59 @@
|
||||
#version 450
|
||||
|
||||
#extension GL_EXT_shader_16bit_storage : require
|
||||
|
||||
layout (push_constant) uniform parameter
|
||||
{
|
||||
uint ne;
|
||||
uint ne00; uint ne01; uint ne02; uint ne03; uint nb00; uint nb01; uint nb02; uint nb03;
|
||||
uint ne10; uint ne11; uint ne12; uint ne13; uint nb10; uint nb11; uint nb12; uint nb13;
|
||||
uint ne20; uint ne21; uint ne22; uint ne23; uint nb20; uint nb21; uint nb22; uint nb23;
|
||||
uint misalign_offsets;
|
||||
float param1; float param2; int param3;
|
||||
} p;
|
||||
|
||||
layout (binding = 0) readonly buffer A {float data_a[];};
|
||||
layout (binding = 1) readonly buffer B {float data_b[];};
|
||||
layout (binding = 2) writeonly buffer D {float data_d[];};
|
||||
|
||||
uint get_idx() {
|
||||
return gl_GlobalInvocationID.z * 262144 + gl_GlobalInvocationID.y * 512 + gl_GlobalInvocationID.x;
|
||||
}
|
||||
|
||||
uint get_aoffset() { return p.misalign_offsets >> 16; }
|
||||
uint get_boffset() { return (p.misalign_offsets >> 8) & 0xFF; }
|
||||
uint get_doffset() { return p.misalign_offsets & 0xFF; }
|
||||
|
||||
layout(local_size_x = 256, local_size_y = 1, local_size_z = 1) in;
|
||||
|
||||
void main() {
|
||||
uint idx = get_idx();
|
||||
if (idx >= p.ne) {
|
||||
return;
|
||||
}
|
||||
|
||||
uint tmp = idx;
|
||||
uint i0 = tmp % p.ne20; tmp /= p.ne20;
|
||||
uint i1 = tmp % p.ne21; tmp /= p.ne21;
|
||||
uint i2 = tmp % p.ne22; tmp /= p.ne22;
|
||||
uint i3 = tmp;
|
||||
|
||||
uint a_i0 = i0 % p.ne00;
|
||||
uint a_i2 = i2 / (p.ne22 / p.ne02);
|
||||
uint a_i3 = i3 / (p.ne23 / p.ne03);
|
||||
|
||||
uint b_i0 = i1 % p.ne10;
|
||||
uint b_i2 = i2;
|
||||
uint b_i3 = i3;
|
||||
|
||||
float sum = 0.0f;
|
||||
uint K = p.ne01;
|
||||
for (uint k = 0; k < K; k++) {
|
||||
uint aoff = get_aoffset() + a_i3*p.nb03 + a_i2*p.nb02 + k*p.nb01 + a_i0*p.nb00;
|
||||
uint boff = get_boffset() + b_i3*p.nb13 + b_i2*p.nb12 + k*p.nb11 + b_i0*p.nb10;
|
||||
sum += data_a[aoff] * data_b[boff];
|
||||
}
|
||||
|
||||
uint doff = get_doffset() + i3*p.nb23 + i2*p.nb22 + i1*p.nb21 + i0*p.nb20;
|
||||
data_d[doff] = sum;
|
||||
}
|
||||
@@ -1036,6 +1036,8 @@ void process_shaders() {
|
||||
}
|
||||
}
|
||||
|
||||
string_to_spv("out_prod_f32", "out_prod.comp", {});
|
||||
|
||||
string_to_spv("timestep_embedding_f32", "timestep_embedding.comp", merge_maps(base_dict, {{"A_TYPE", "float"}, {"D_TYPE", "float"}}));
|
||||
|
||||
string_to_spv("conv_transpose_1d_f32", "conv_transpose_1d.comp", {{"A_TYPE", "float"}, {"B_TYPE", "float"}, {"D_TYPE", "float"}});
|
||||
|
||||
Reference in New Issue
Block a user