diff --git a/ggml/src/ggml-cpu/ggml-cpu.c b/ggml/src/ggml-cpu/ggml-cpu.c index 24c47569c..8620c6c7a 100644 --- a/ggml/src/ggml-cpu/ggml-cpu.c +++ b/ggml/src/ggml-cpu/ggml-cpu.c @@ -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); } } } diff --git a/ggml/src/ggml-cpu/ggml-cpu.cpp b/ggml/src/ggml-cpu/ggml-cpu.cpp index 1df0f2bb9..81ff5d79f 100644 --- a/ggml/src/ggml-cpu/ggml-cpu.cpp +++ b/ggml/src/ggml-cpu/ggml-cpu.cpp @@ -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; diff --git a/ggml/src/ggml-vulkan/ggml-vulkan.cpp b/ggml/src/ggml-vulkan/ggml-vulkan.cpp index e23d5b433..cbd3697ef 100644 --- a/ggml/src/ggml-vulkan/ggml-vulkan.cpp +++ b/ggml/src/ggml-vulkan/ggml-vulkan.cpp @@ -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; }