mirror of
https://github.com/ggml-org/whisper.cpp.git
synced 2026-10-11 08:45:40 +02:00
cpu: accept BF16 in src1 of mul_mat (llama/28937)
* cpu: accept BF16 in src1 of mul_mat ggml_conv_1d_dw builds its im2col in F32 when the kernel is BF16, then calls ggml_mul_mat(im2col, kernel), which puts F32 in src0 and BF16 in src1. The CPU backend refused that combination, so it was reported as unsupported on every backend and never compared against anything. Widen BF16 into the F32 work buffer, next to the existing packing of F32 into vec_dot_type. This is the arithmetic the Metal mat vec kernel already uses, both operands promoted to float and accumulated in float, so the two agree exactly rather than approximately. Cover it with a conv_1d_dw test over F32, F16 and BF16 kernels, plus three mul_mat cases with BF16 in src1. * vulkan: reject BF16 in src1 of mul_mat unless src0 is BF16 supports_op only checked the src1 type for non contiguous tensors, so a contiguous BF16 src1 was accepted and the pipeline lookup asserted. The only BF16 src1 path is the BF16 x BF16 multiply, every other src0 type now reports the op as unsupported and the scheduler keeps it on the CPU. The BF16 kernel case of the conv_1d_dw test needs the f32 x bf16 mat vec variants of the Metal backend, which land separately.
This commit is contained in:
@@ -1334,9 +1334,11 @@ UseGgmlGemm1:;
|
||||
const size_t nbw3 = nbw2*ne12;
|
||||
|
||||
assert(params->wsize >= ne13*nbw3);
|
||||
GGML_ASSERT(src1->type == GGML_TYPE_F32 || src1->type == GGML_TYPE_F16);
|
||||
// the F16 path below writes plain floats into wdata, so it needs an F32 vec_dot_type
|
||||
GGML_ASSERT(src1->type == GGML_TYPE_F32 || vec_dot_type == GGML_TYPE_F32);
|
||||
// src1 is either packed from F32 into vec_dot_type, or widened from F16 or BF16 into the F32 work buffer
|
||||
const bool widen = src1->type != GGML_TYPE_F32;
|
||||
|
||||
GGML_ASSERT(!widen || vec_dot_type == GGML_TYPE_F32);
|
||||
GGML_ASSERT(!widen || src1->type == GGML_TYPE_F16 || src1->type == GGML_TYPE_BF16);
|
||||
|
||||
#if 0
|
||||
for (int64_t i13 = 0; i13 < ne13; ++i13) {
|
||||
@@ -1349,20 +1351,24 @@ UseGgmlGemm1:;
|
||||
}
|
||||
}
|
||||
#else
|
||||
const int64_t bs = ggml_blck_size(vec_dot_type);
|
||||
|
||||
for (int64_t i13 = 0; i13 < ne13; ++i13) {
|
||||
for (int64_t i12 = 0; i12 < ne12; ++i12) {
|
||||
for (int64_t i11 = 0; i11 < ne11; ++i11) {
|
||||
size_t bs = ggml_blck_size(vec_dot_type);
|
||||
int64_t ne10_block_start = (ith * ne10/bs) / nth;
|
||||
int64_t ne10_block_end = ((ith + 1) * ne10/bs) / nth;
|
||||
const char * src1_block = (const char *) src1->data + i13*nb13 + i12*nb12 + i11*nb11 + ne10_block_start*bs*nb10;
|
||||
char * dst_block = wdata + i13*nbw3 + i12*nbw2 + i11*nbw1 + ne10_block_start*nbw0;
|
||||
const int64_t n_block = (ne10_block_end - ne10_block_start) * bs;
|
||||
const void * src1_row = (const char *) src1->data + i13*nb13 + i12*nb12 + i11*nb11 + ne10_block_start*bs*nb10;
|
||||
void * wdata_row = wdata + i13*nbw3 + i12*nbw2 + i11*nbw1 + ne10_block_start*nbw0;
|
||||
|
||||
if (src1->type == GGML_TYPE_F32) {
|
||||
from_float((const float *) src1_block, dst_block, n_block);
|
||||
const int64_t ne10_block_size = (ne10_block_end - ne10_block_start) * bs;
|
||||
|
||||
if (src1->type == GGML_TYPE_F16) {
|
||||
ggml_cpu_fp16_to_fp32((const ggml_fp16_t *) src1_row, (float *) wdata_row, ne10_block_size);
|
||||
} else if (src1->type == GGML_TYPE_BF16) {
|
||||
ggml_cpu_bf16_to_fp32((const ggml_bf16_t *) src1_row, (float *) wdata_row, ne10_block_size);
|
||||
} else {
|
||||
ggml_cpu_fp16_to_fp32((const ggml_fp16_t *) src1_block, (float *) dst_block, n_block);
|
||||
from_float((const float *) src1_row, wdata_row, ne10_block_size);
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
@@ -455,7 +455,10 @@ static bool ggml_backend_cpu_device_supports_op(ggml_backend_dev_t dev, const st
|
||||
src0->type == GGML_TYPE_F32 && op->type == GGML_TYPE_F32) {
|
||||
return src1->type == GGML_TYPE_F32 || src1->type == GGML_TYPE_F16;
|
||||
}
|
||||
return src1->type == GGML_TYPE_F32 || src1->type == ggml_get_type_traits_cpu(src0->type)->vec_dot_type;
|
||||
// BF16 in src1 is widened into the F32 work buffer
|
||||
return src1->type == GGML_TYPE_F32 ||
|
||||
src1->type == ggml_get_type_traits_cpu(src0->type)->vec_dot_type ||
|
||||
(src1->type == GGML_TYPE_BF16 && ggml_get_type_traits_cpu(src0->type)->vec_dot_type == GGML_TYPE_F32);
|
||||
case GGML_OP_SOFT_MAX_BACK: {
|
||||
if (op->src[0]->type != GGML_TYPE_F32 || op->src[1]->type != GGML_TYPE_F32) {
|
||||
return false;
|
||||
|
||||
@@ -15441,6 +15441,10 @@ static bool ggml_backend_vk_device_supports_op(ggml_backend_dev_t dev, const ggm
|
||||
// So don't support this combination for now.
|
||||
return false;
|
||||
}
|
||||
if (op->src[1]->type == GGML_TYPE_BF16 && op->src[0]->type != GGML_TYPE_BF16) {
|
||||
// BF16 in src1 is only served by the BF16 x BF16 pipelines
|
||||
return false;
|
||||
}
|
||||
|
||||
return true;
|
||||
}
|
||||
|
||||
Reference in New Issue
Block a user