if (op->src[1]->type != op->src[2]->type) {
return false;
}
+ switch (op->src[1]->type) {
+ case GGML_TYPE_F32:
+ case GGML_TYPE_F16:
+ case GGML_TYPE_Q8_0:
+ case GGML_TYPE_Q4_0:
+ case GGML_TYPE_Q4_1:
+ case GGML_TYPE_Q5_0:
+ case GGML_TYPE_Q5_1:
+ break;
+ case GGML_TYPE_BF16:
+ if (!has_bfloat) {
+ return false;
+ }
+ break;
+ default:
+ return false;
+ }
return has_simdgroup_mm; // TODO: over-restricted for vec-kernels
case GGML_OP_SSM_CONV:
case GGML_OP_SSM_SCAN: