ggml: add f16 out_prod support for CPU and out_prod op for Vulkan (#23997)
This commit is contained in:
@@ -2859,7 +2859,8 @@ struct ggml_cplan ggml_graph_plan(
|
|||||||
} break;
|
} break;
|
||||||
case GGML_OP_OUT_PROD:
|
case GGML_OP_OUT_PROD:
|
||||||
{
|
{
|
||||||
if (ggml_is_quantized(node->src[0]->type)) {
|
if (ggml_is_quantized(node->src[0]->type) ||
|
||||||
|
node->src[0]->type == GGML_TYPE_F16) {
|
||||||
cur = ggml_type_size(GGML_TYPE_F32) * node->src[0]->ne[0] * n_tasks;
|
cur = ggml_type_size(GGML_TYPE_F32) * node->src[0]->ne[0] * n_tasks;
|
||||||
}
|
}
|
||||||
} break;
|
} break;
|
||||||
|
|||||||
@@ -462,11 +462,12 @@ static bool ggml_backend_cpu_device_supports_op(ggml_backend_dev_t dev, const st
|
|||||||
return max_bias == 0.0f;
|
return max_bias == 0.0f;
|
||||||
}
|
}
|
||||||
case GGML_OP_IM2COL_BACK:
|
case GGML_OP_IM2COL_BACK:
|
||||||
return src0->type == GGML_TYPE_F32 && src1->type == GGML_TYPE_F32;
|
return src0->type == GGML_TYPE_F32 && (src1->type == GGML_TYPE_F32 || src1->type == GGML_TYPE_F16);
|
||||||
case GGML_OP_GET_ROWS_BACK:
|
case GGML_OP_GET_ROWS_BACK:
|
||||||
return src0->type == GGML_TYPE_F32 || src0->type == GGML_TYPE_F16;
|
return src0->type == GGML_TYPE_F32 || src0->type == GGML_TYPE_F16;
|
||||||
case GGML_OP_OUT_PROD:
|
case GGML_OP_OUT_PROD:
|
||||||
return (src0->type == GGML_TYPE_F32 || (ggml_is_quantized(src0->type) && src0->ne[2] == src1->ne[2] && src0->ne[3] == src1->ne[3])) &&
|
return (src0->type == GGML_TYPE_F32 ||
|
||||||
|
((src0->type == GGML_TYPE_F16 || ggml_is_quantized(src0->type)) && src0->ne[2] == src1->ne[2] && src0->ne[3] == src1->ne[3])) &&
|
||||||
src1->type == GGML_TYPE_F32 && op->type == GGML_TYPE_F32;
|
src1->type == GGML_TYPE_F32 && op->type == GGML_TYPE_F32;
|
||||||
default:
|
default:
|
||||||
return true;
|
return true;
|
||||||
|
|||||||
@@ -4449,6 +4449,70 @@ static void ggml_compute_forward_out_prod_q_f32(
|
|||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
|
static void ggml_compute_forward_out_prod_f16_f32(
|
||||||
|
const ggml_compute_params * params,
|
||||||
|
ggml_tensor * dst) {
|
||||||
|
|
||||||
|
const ggml_tensor * src0 = dst->src[0];
|
||||||
|
const ggml_tensor * src1 = dst->src[1];
|
||||||
|
|
||||||
|
GGML_TENSOR_BINARY_OP_LOCALS;
|
||||||
|
|
||||||
|
const int ith = params->ith;
|
||||||
|
const int nth = params->nth;
|
||||||
|
|
||||||
|
GGML_ASSERT(src0->type == GGML_TYPE_F16);
|
||||||
|
GGML_ASSERT(src1->type == GGML_TYPE_F32);
|
||||||
|
GGML_ASSERT(dst->type == GGML_TYPE_F32);
|
||||||
|
|
||||||
|
GGML_ASSERT(ne02 == ne12);
|
||||||
|
GGML_ASSERT(ne03 == ne13);
|
||||||
|
GGML_ASSERT(ne2 == ne12);
|
||||||
|
GGML_ASSERT(ne3 == ne13);
|
||||||
|
|
||||||
|
GGML_ASSERT(nb00 == sizeof(ggml_fp16_t));
|
||||||
|
GGML_ASSERT(nb0 == sizeof(float));
|
||||||
|
|
||||||
|
GGML_ASSERT(ne0 == ne00);
|
||||||
|
GGML_ASSERT(ne1 == ne10);
|
||||||
|
GGML_ASSERT(ne2 == ne02);
|
||||||
|
GGML_ASSERT(ne3 == ne03);
|
||||||
|
|
||||||
|
if (ith == 0) {
|
||||||
|
ggml_vec_set_f32(ne0*ne1*ne2*ne3, (float *)dst->data, 0);
|
||||||
|
}
|
||||||
|
ggml_barrier(params->threadpool);
|
||||||
|
|
||||||
|
const int64_t nr = ne1*ne2*ne3;
|
||||||
|
const int64_t dr = (nr + nth - 1)/nth;
|
||||||
|
const int64_t ir0 = dr*ith;
|
||||||
|
const int64_t ir1 = MIN(ir0 + dr, nr);
|
||||||
|
|
||||||
|
float * wdata = (float *) params->wdata + (ne0 + CACHE_LINE_SIZE_F32) * ith;
|
||||||
|
|
||||||
|
for (int64_t ir = ir0; ir < ir1; ++ir) {
|
||||||
|
const int64_t i3 = ir/(ne2*ne1);
|
||||||
|
const int64_t i2 = (ir - i3*ne2*ne1)/ne1;
|
||||||
|
const int64_t i1 = (ir - i3*ne2*ne1 - i2*ne1);
|
||||||
|
|
||||||
|
const int64_t i02 = i2;
|
||||||
|
const int64_t i03 = i3;
|
||||||
|
|
||||||
|
const int64_t i12 = i2;
|
||||||
|
const int64_t i13 = i3;
|
||||||
|
|
||||||
|
float * d = (float *) ((char *) dst->data + (i1*nb1 + i2*nb2 + i3*nb3));
|
||||||
|
|
||||||
|
for (int64_t i01 = 0; i01 < ne01; ++i01) {
|
||||||
|
const int64_t i11 = i01;
|
||||||
|
ggml_fp16_t * s0 = (ggml_fp16_t *) ((char *) src0->data + (i01*nb01 + i02*nb02 + i03*nb03));
|
||||||
|
float * s1 = (float *) ((char *) src1->data + (i1*nb10 + i11*nb11 + i12*nb12 + i13*nb13));
|
||||||
|
ggml_fp16_to_fp32_row(s0, wdata, ne0);
|
||||||
|
ggml_vec_mad_f32(ne0, d, wdata, *s1);
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
void ggml_compute_forward_out_prod(
|
void ggml_compute_forward_out_prod(
|
||||||
const ggml_compute_params * params,
|
const ggml_compute_params * params,
|
||||||
ggml_tensor * dst) {
|
ggml_tensor * dst) {
|
||||||
@@ -4486,9 +4550,8 @@ void ggml_compute_forward_out_prod(
|
|||||||
} break;
|
} break;
|
||||||
case GGML_TYPE_F16:
|
case GGML_TYPE_F16:
|
||||||
{
|
{
|
||||||
GGML_ABORT("fatal error"); // todo
|
ggml_compute_forward_out_prod_f16_f32(params, dst);
|
||||||
// ggml_compute_forward_out_prod_f16_f32(params, dst);
|
} break;
|
||||||
}
|
|
||||||
case GGML_TYPE_F32:
|
case GGML_TYPE_F32:
|
||||||
{
|
{
|
||||||
ggml_compute_forward_out_prod_f32(params, dst);
|
ggml_compute_forward_out_prod_f32(params, dst);
|
||||||
@@ -6469,7 +6532,7 @@ void ggml_compute_forward_im2col_back_f32(
|
|||||||
const ggml_tensor * src1 = dst->src[1]; // convolution kernel
|
const ggml_tensor * src1 = dst->src[1]; // convolution kernel
|
||||||
|
|
||||||
GGML_ASSERT(src0->type == GGML_TYPE_F32);
|
GGML_ASSERT(src0->type == GGML_TYPE_F32);
|
||||||
GGML_ASSERT(src1->type == GGML_TYPE_F32);
|
GGML_ASSERT(src1->type == GGML_TYPE_F32 || src1->type == GGML_TYPE_F16);
|
||||||
GGML_ASSERT( dst->type == GGML_TYPE_F32);
|
GGML_ASSERT( dst->type == GGML_TYPE_F32);
|
||||||
|
|
||||||
GGML_TENSOR_BINARY_OP_LOCALS;
|
GGML_TENSOR_BINARY_OP_LOCALS;
|
||||||
|
|||||||
@@ -961,6 +961,7 @@ struct vk_device_struct {
|
|||||||
vk_pipeline pipeline_col2im_1d_f32;
|
vk_pipeline pipeline_col2im_1d_f32;
|
||||||
vk_pipeline pipeline_col2im_1d_f16;
|
vk_pipeline pipeline_col2im_1d_f16;
|
||||||
vk_pipeline pipeline_col2im_1d_bf16;
|
vk_pipeline pipeline_col2im_1d_bf16;
|
||||||
|
vk_pipeline pipeline_out_prod_f32;
|
||||||
vk_pipeline pipeline_snake_f32;
|
vk_pipeline pipeline_snake_f32;
|
||||||
vk_pipeline pipeline_snake_f16;
|
vk_pipeline pipeline_snake_f16;
|
||||||
vk_pipeline pipeline_snake_bf16;
|
vk_pipeline pipeline_snake_bf16;
|
||||||
@@ -5479,6 +5480,8 @@ static void ggml_vk_load_shaders(vk_device& device, vk_pipeline requested) {
|
|||||||
ggml_vk_create_pipeline(device, device->pipeline_col2im_1d_f16, "col2im_1d_f16", col2im_1d_f16_len, col2im_1d_f16_data, "main", 2, sizeof(vk_op_col2im_1d_push_constants), {256, 1, 1}, {}, 1, true);
|
ggml_vk_create_pipeline(device, device->pipeline_col2im_1d_f16, "col2im_1d_f16", col2im_1d_f16_len, col2im_1d_f16_data, "main", 2, sizeof(vk_op_col2im_1d_push_constants), {256, 1, 1}, {}, 1, true);
|
||||||
ggml_vk_create_pipeline(device, device->pipeline_col2im_1d_bf16, "col2im_1d_bf16", col2im_1d_bf16_len, col2im_1d_bf16_data, "main", 2, sizeof(vk_op_col2im_1d_push_constants), {256, 1, 1}, {}, 1, true);
|
ggml_vk_create_pipeline(device, device->pipeline_col2im_1d_bf16, "col2im_1d_bf16", col2im_1d_bf16_len, col2im_1d_bf16_data, "main", 2, sizeof(vk_op_col2im_1d_push_constants), {256, 1, 1}, {}, 1, true);
|
||||||
|
|
||||||
|
ggml_vk_create_pipeline(device, device->pipeline_out_prod_f32, "out_prod_f32", out_prod_f32_len, out_prod_f32_data, "main", 3, sizeof(vk_op_binary_push_constants), {256, 1, 1}, {}, 1);
|
||||||
|
|
||||||
ggml_vk_create_pipeline(device, device->pipeline_snake_f32, "snake_f32", snake_f32_len, snake_f32_data, "main", 4, sizeof(vk_op_snake_push_constants), {256, 1, 1}, {}, 1);
|
ggml_vk_create_pipeline(device, device->pipeline_snake_f32, "snake_f32", snake_f32_len, snake_f32_data, "main", 4, sizeof(vk_op_snake_push_constants), {256, 1, 1}, {}, 1);
|
||||||
ggml_vk_create_pipeline(device, device->pipeline_snake_f16, "snake_f16", snake_f16_len, snake_f16_data, "main", 4, sizeof(vk_op_snake_push_constants), {256, 1, 1}, {}, 1);
|
ggml_vk_create_pipeline(device, device->pipeline_snake_f16, "snake_f16", snake_f16_len, snake_f16_data, "main", 4, sizeof(vk_op_snake_push_constants), {256, 1, 1}, {}, 1);
|
||||||
ggml_vk_create_pipeline(device, device->pipeline_snake_bf16, "snake_bf16", snake_bf16_len, snake_bf16_data, "main", 4, sizeof(vk_op_snake_push_constants), {256, 1, 1}, {}, 1);
|
ggml_vk_create_pipeline(device, device->pipeline_snake_bf16, "snake_bf16", snake_bf16_len, snake_bf16_data, "main", 4, sizeof(vk_op_snake_push_constants), {256, 1, 1}, {}, 1);
|
||||||
@@ -10745,6 +10748,11 @@ static vk_pipeline ggml_vk_op_get_pipeline(ggml_backend_vk_context * ctx, const
|
|||||||
return ctx->device->pipeline_add_id_f32;
|
return ctx->device->pipeline_add_id_f32;
|
||||||
}
|
}
|
||||||
return nullptr;
|
return nullptr;
|
||||||
|
case GGML_OP_OUT_PROD:
|
||||||
|
if (src0->type == GGML_TYPE_F32 && src1->type == GGML_TYPE_F32 && dst->type == GGML_TYPE_F32) {
|
||||||
|
return ctx->device->pipeline_out_prod_f32;
|
||||||
|
}
|
||||||
|
return nullptr;
|
||||||
case GGML_OP_CONCAT: {
|
case GGML_OP_CONCAT: {
|
||||||
if (src0->type != src1->type || src0->type != dst->type) {
|
if (src0->type != src1->type || src0->type != dst->type) {
|
||||||
return nullptr;
|
return nullptr;
|
||||||
@@ -11701,6 +11709,7 @@ static void ggml_vk_op_f32(ggml_backend_vk_context * ctx, vk_context& subctx, co
|
|||||||
case GGML_OP_DIV:
|
case GGML_OP_DIV:
|
||||||
case GGML_OP_MUL:
|
case GGML_OP_MUL:
|
||||||
case GGML_OP_ADD1:
|
case GGML_OP_ADD1:
|
||||||
|
case GGML_OP_OUT_PROD:
|
||||||
case GGML_OP_ARANGE:
|
case GGML_OP_ARANGE:
|
||||||
case GGML_OP_FILL:
|
case GGML_OP_FILL:
|
||||||
case GGML_OP_SCALE:
|
case GGML_OP_SCALE:
|
||||||
@@ -12014,6 +12023,24 @@ static void ggml_vk_add(ggml_backend_vk_context * ctx, vk_context& subctx, const
|
|||||||
});
|
});
|
||||||
}
|
}
|
||||||
|
|
||||||
|
static void ggml_vk_out_prod(ggml_backend_vk_context * ctx, vk_context& subctx, const ggml_tensor * src0, const ggml_tensor * src1, ggml_tensor * dst) {
|
||||||
|
const uint32_t src0_type_size = ggml_type_size(src0->type);
|
||||||
|
const uint32_t src1_type_size = ggml_type_size(src1->type);
|
||||||
|
const uint32_t dst_type_size = ggml_type_size(dst->type);
|
||||||
|
|
||||||
|
ggml_vk_op_f32<vk_op_binary_push_constants>(ctx, subctx, src0, src1, nullptr, nullptr, dst, GGML_OP_OUT_PROD, {
|
||||||
|
(uint32_t)ggml_nelements(dst),
|
||||||
|
(uint32_t)src0->ne[0], (uint32_t)src0->ne[1], (uint32_t)src0->ne[2],(uint32_t)src0->ne[3],
|
||||||
|
(uint32_t)src0->nb[0] / src0_type_size, (uint32_t)src0->nb[1] / src0_type_size, (uint32_t)src0->nb[2] / src0_type_size, (uint32_t)src0->nb[3] / src0_type_size,
|
||||||
|
(uint32_t)src1->ne[0], (uint32_t)src1->ne[1], (uint32_t)src1->ne[2],(uint32_t)src1->ne[3],
|
||||||
|
(uint32_t)src1->nb[0] / src1_type_size, (uint32_t)src1->nb[1] / src1_type_size, (uint32_t)src1->nb[2] / src1_type_size, (uint32_t)src1->nb[3] / src1_type_size,
|
||||||
|
(uint32_t) dst->ne[0], (uint32_t) dst->ne[1], (uint32_t) dst->ne[2],(uint32_t) dst->ne[3],
|
||||||
|
(uint32_t) dst->nb[0] / dst_type_size, (uint32_t) dst->nb[1] / dst_type_size, (uint32_t) dst->nb[2] / dst_type_size, (uint32_t) dst->nb[3] / dst_type_size,
|
||||||
|
0,
|
||||||
|
0.0f, 0.0f, 0,
|
||||||
|
});
|
||||||
|
}
|
||||||
|
|
||||||
static void ggml_vk_sub(ggml_backend_vk_context * ctx, vk_context& subctx, const ggml_tensor * src0, const ggml_tensor * src1, ggml_tensor * dst) {
|
static void ggml_vk_sub(ggml_backend_vk_context * ctx, vk_context& subctx, const ggml_tensor * src0, const ggml_tensor * src1, ggml_tensor * dst) {
|
||||||
const uint32_t src0_type_size = ggml_type_size(src0->type);
|
const uint32_t src0_type_size = ggml_type_size(src0->type);
|
||||||
const uint32_t src1_type_size = ggml_type_size(src1->type);
|
const uint32_t src1_type_size = ggml_type_size(src1->type);
|
||||||
@@ -14801,6 +14828,9 @@ static bool ggml_vk_build_graph(ggml_backend_vk_context * ctx, ggml_cgraph * cgr
|
|||||||
ggml_vk_add(ctx, compute_ctx, src0, src1, node);
|
ggml_vk_add(ctx, compute_ctx, src0, src1, node);
|
||||||
}
|
}
|
||||||
break;
|
break;
|
||||||
|
case GGML_OP_OUT_PROD:
|
||||||
|
ggml_vk_out_prod(ctx, compute_ctx, src0, src1, node);
|
||||||
|
break;
|
||||||
case GGML_OP_SUB:
|
case GGML_OP_SUB:
|
||||||
ggml_vk_sub(ctx, compute_ctx, src0, src1, node);
|
ggml_vk_sub(ctx, compute_ctx, src0, src1, node);
|
||||||
|
|
||||||
@@ -17655,6 +17685,10 @@ static bool ggml_backend_vk_device_supports_op(ggml_backend_dev_t dev, const ggm
|
|||||||
case GGML_OP_OPT_STEP_ADAMW:
|
case GGML_OP_OPT_STEP_ADAMW:
|
||||||
case GGML_OP_OPT_STEP_SGD:
|
case GGML_OP_OPT_STEP_SGD:
|
||||||
return ggml_is_contiguous(op->src[0]) && op->src[0]->type == GGML_TYPE_F32;
|
return ggml_is_contiguous(op->src[0]) && op->src[0]->type == GGML_TYPE_F32;
|
||||||
|
case GGML_OP_OUT_PROD:
|
||||||
|
return ggml_is_contiguous(op->src[0]) && op->src[0]->type == GGML_TYPE_F32
|
||||||
|
&& ggml_is_contiguous(op->src[1]) && op->src[1]->type == GGML_TYPE_F32
|
||||||
|
&& op->type == GGML_TYPE_F32;
|
||||||
case GGML_OP_LOG:
|
case GGML_OP_LOG:
|
||||||
case GGML_OP_TRI:
|
case GGML_OP_TRI:
|
||||||
case GGML_OP_DIAG:
|
case GGML_OP_DIAG:
|
||||||
|
|||||||
@@ -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("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"}});
|
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