From fe6946e984f4420b4572e6d012110382317678eb Mon Sep 17 00:00:00 2001 From: =?UTF-8?q?=E9=AB=98=E5=BA=86=E4=B8=B0?= Date: Tue, 15 Sep 2026 15:43:29 +0800 Subject: [PATCH 1/2] ggml-metal: add bf16<->f16 copy kernels CPY, SET, DUP and CONT accepted f32 conversions and same-type copies on Metal, but not the two cross combinations of bf16 and f16, so a graph copying between them could not be scheduled on the backend. Instantiate kernel_cpy_f16_bf16 and kernel_cpy_bf16_f16 (contiguous and strided) under GGML_METAL_HAS_BF16, and accept the F16<->BF16 pairs in ggml_metal_device_supports_op. --- external/ggml/src/ggml-metal/ggml-metal-device.m | 4 ++-- external/ggml/src/ggml-metal/ggml-metal.metal | 4 ++++ 2 files changed, 6 insertions(+), 2 deletions(-) diff --git a/external/ggml/src/ggml-metal/ggml-metal-device.m b/external/ggml/src/ggml-metal/ggml-metal-device.m index dca96bc5c..b4ac18298 100644 --- a/external/ggml/src/ggml-metal/ggml-metal-device.m +++ b/external/ggml/src/ggml-metal/ggml-metal-device.m @@ -1279,14 +1279,14 @@ bool ggml_metal_device_supports_op(ggml_metal_device_t dev, const struct ggml_te switch (op->type) { case GGML_TYPE_F32: case GGML_TYPE_F16: - return true; + case GGML_TYPE_BF16: return true; default: return false; } case GGML_TYPE_BF16: switch (op->type) { case GGML_TYPE_F32: - case GGML_TYPE_BF16: + case GGML_TYPE_F16: case GGML_TYPE_BF16: return true; default: return false; diff --git a/external/ggml/src/ggml-metal/ggml-metal.metal b/external/ggml/src/ggml-metal/ggml-metal.metal index b71b83c68..a2f5cf8ec 100644 --- a/external/ggml/src/ggml-metal/ggml-metal.metal +++ b/external/ggml/src/ggml-metal/ggml-metal.metal @@ -7932,6 +7932,8 @@ template [[host_name("kernel_cpy_contig_f16_f16")]] kernel kernel_cpy_contig_t k #if defined(GGML_METAL_HAS_BF16) template [[host_name("kernel_cpy_contig_bf16_f32")]] kernel kernel_cpy_contig_t kernel_cpy_contig_t_t; template [[host_name("kernel_cpy_contig_bf16_bf16")]] kernel kernel_cpy_contig_t kernel_cpy_contig_t_t; +template [[host_name("kernel_cpy_contig_f16_bf16")]] kernel kernel_cpy_contig_t kernel_cpy_contig_t_t; +template [[host_name("kernel_cpy_contig_bf16_f16")]] kernel kernel_cpy_contig_t kernel_cpy_contig_t_t; #endif template @@ -8064,6 +8066,8 @@ template [[host_name("kernel_cpy_f16_f16")]] kernel kernel_cpy_t kernel_cpy_t_ #if defined(GGML_METAL_HAS_BF16) template [[host_name("kernel_cpy_bf16_f32")]] kernel kernel_cpy_t kernel_cpy_t_t; template [[host_name("kernel_cpy_bf16_bf16")]] kernel kernel_cpy_t kernel_cpy_t_t; +template [[host_name("kernel_cpy_f16_bf16")]] kernel kernel_cpy_t kernel_cpy_t_t; +template [[host_name("kernel_cpy_bf16_f16")]] kernel kernel_cpy_t kernel_cpy_t_t; #endif template Date: Tue, 15 Sep 2026 15:43:29 +0800 Subject: [PATCH 2/2] breeze: enable bf16 activation rounding on Metal BreezeTTS rounds activations to bf16 on CUDA/HIP/Vulkan to match the reference implementation. Enable the same policy on Metal and use a bf16 KV cache there as well. Unlike CUDA/HIP/Vulkan, Metal has no fused round-to-bf16 unary op, so fused_round stays disabled and the cast is a separate graph node. --- src/models/breeze_tts/generator.cpp | 15 ++++++++------- 1 file changed, 8 insertions(+), 7 deletions(-) diff --git a/src/models/breeze_tts/generator.cpp b/src/models/breeze_tts/generator.cpp index 4304940a5..4699a1920 100644 --- a/src/models/breeze_tts/generator.cpp +++ b/src/models/breeze_tts/generator.cpp @@ -55,14 +55,13 @@ struct GgmlContextDeleter { modules::QwenDecoderActivationCastPolicy breeze_bf16_activation_policy(core::BackendType backend_type) { modules::QwenDecoderActivationCastPolicy policy; if (backend_type != core::BackendType::Cuda && backend_type != core::BackendType::Hip && - backend_type != core::BackendType::Vulkan) { + backend_type != core::BackendType::Vulkan && backend_type != core::BackendType::Metal) { return policy; } policy.enabled = true; policy.type = GGML_TYPE_BF16; // CUDA/HIP/Vulkan implement the fused round-to-bf16 unary op. - policy.fused_round = backend_type == core::BackendType::Cuda || backend_type == core::BackendType::Hip || - backend_type == core::BackendType::Vulkan; + policy.fused_round = backend_type != core::BackendType::Metal; policy.after_input_norm = true; policy.after_qkv_projection = true; policy.after_qk_norm = true; @@ -114,12 +113,13 @@ modules::QwenCausalDecodeRuntimeConfig backbone_config( out.decoder.stack.runtime.static_cache.update_mode = modules::QwenDecoderStaticCacheUpdateMode::DirectSetRows; out.decoder.stack.runtime.static_cache.set_rows_mode = modules::QwenDecoderStaticCacheSetRowsMode::BackendViewOptimized; if (backend_type == core::BackendType::Cuda || backend_type == core::BackendType::Hip || - backend_type == core::BackendType::Vulkan) { + backend_type == core::BackendType::Vulkan || backend_type == core::BackendType::Metal) { // BF16 KV cache matches the reference implementation, but flash // attention only accelerates bf16 cache with native bf16 MMA // (sm_80+); on older parts it is ~3x slower, so only HIP uses it. out.decoder.static_cache_type = - backend_type == core::BackendType::Hip ? GGML_TYPE_BF16 : GGML_TYPE_F16; + (backend_type == core::BackendType::Hip || backend_type == core::BackendType::Metal) + ? GGML_TYPE_BF16 : GGML_TYPE_F16; out.decoder.stack.activation_cast = breeze_bf16_activation_policy(backend_type); } out.decoder.logits_size = config.lm_head_size; @@ -168,11 +168,12 @@ modules::QwenCausalDecodeRuntimeConfig depth_config( out.decoder.stack.runtime.static_cache.update_mode = modules::QwenDecoderStaticCacheUpdateMode::DirectSetRows; out.decoder.stack.runtime.static_cache.set_rows_mode = modules::QwenDecoderStaticCacheSetRowsMode::BackendViewOptimized; if (backend_type == core::BackendType::Cuda || backend_type == core::BackendType::Hip || - backend_type == core::BackendType::Vulkan) { + backend_type == core::BackendType::Vulkan || backend_type == core::BackendType::Metal) { // See backbone_config: only HIP uses a bf16 KV cache; CUDA and Vulkan // keep F16. out.decoder.static_cache_type = - backend_type == core::BackendType::Hip ? GGML_TYPE_BF16 : GGML_TYPE_F16; + (backend_type == core::BackendType::Hip || backend_type == core::BackendType::Metal) + ? GGML_TYPE_BF16 : GGML_TYPE_F16; out.decoder.stack.activation_cast = breeze_bf16_activation_policy(backend_type); } out.decoder.logits_mode = modules::QwenCausalDecoderLogitsMode::LastStep;