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:
Pascal
2026-10-06 10:38:12 +03:00
committed by Georgi Gerganov
parent d0af734dfd
commit c1b3fb1094
3 changed files with 24 additions and 11 deletions
+16 -10
View File
@@ -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);
}
}
}
+4 -1
View File
@@ -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;
+4
View File
@@ -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;
}