From 9cb62bc695a5c301cbd452c6987d9cd713cc8097 Mon Sep 17 00:00:00 2001 From: Fdy <73393700+MrFadiAi@users.noreply.github.com> Date: Wed, 30 Sep 2026 14:59:22 +0200 Subject: [PATCH] vulkan: bf16 SSM state pools + gated_delta_net rows mode - gated_delta_net.comp: USE_STATE_ROWS + STATE_BF16 variants (bf16 state read via uint16<<16, rows-indexed initial state from int32 index buffer) - 8 new pipelines (rows f32 / rows bf16state x reduce modes x kda) - scale_bf16 pipeline + GGML_OP_SCALE BF16 dispatch - dispatch/support accept src[6] index tensor with F32 or BF16 state - llama-model: opt-in env LLAMA_SSM_BF16_STATE / LLAMA_SSM_BF16_CONV allocate hybrid recurrent pools as BF16 (default unchanged) - ggml.c: ggml_gated_delta_net_rows accepts BF16 state Measured, Radeon 890M (RDNA 3.5), Bonsai-2-27B Q2_0-fork + MTP n-max 1: single-stream 10.4 -> 12.7 t/s (+22%); np16 aggregate 18.46 -> 20.30 t/s. Quality: 5/5 short-form gates + three 900-token factual generations. Co-authored-by: Hermes Agent --- ggml/src/ggml-vulkan/ggml-vulkan.cpp | 62 +++++++++++++++++-- .../vulkan-shaders/gated_delta_net.comp | 23 +++++++ .../vulkan-shaders/vulkan-shaders-gen.cpp | 5 ++ ggml/src/ggml.c | 2 +- src/llama-model.cpp | 8 ++- 5 files changed, 93 insertions(+), 7 deletions(-) diff --git a/ggml/src/ggml-vulkan/ggml-vulkan.cpp b/ggml/src/ggml-vulkan/ggml-vulkan.cpp index f4fa37d2cc92..64392233704f 100644 --- a/ggml/src/ggml-vulkan/ggml-vulkan.cpp +++ b/ggml/src/ggml-vulkan/ggml-vulkan.cpp @@ -960,6 +960,7 @@ struct vk_device_struct { vk_pipeline pipeline_concat_i8, pipeline_concat_i16, pipeline_concat_i32, pipeline_concat_i64; vk_pipeline pipeline_upscale_nearest_f32, pipeline_upscale_bilinear_f32, pipeline_upscale_bicubic_f32, pipeline_upscale_bilinear_antialias_f32; vk_pipeline pipeline_scale_f32; + vk_pipeline pipeline_scale_bf16; vk_pipeline pipeline_log[2]; vk_pipeline pipeline_tri[2]; vk_pipeline pipeline_diag[2]; @@ -1081,6 +1082,8 @@ struct vk_device_struct { vk_pipeline pipeline_gated_linear_attn_f32; // [size_idx][kda] where size_idx: 0=d16, 1=d32, 2=d64, 3=d128 vk_pipeline pipeline_gated_delta_net[4][2]; + vk_pipeline pipeline_gated_delta_net_rows[4][2]; + vk_pipeline pipeline_gated_delta_net_rows_bf16state[4][2]; vk_pipeline pipeline_ssm_scan_f32_d128; vk_pipeline pipeline_ssm_scan_f32_d256; vk_pipeline pipeline_ssm_conv_f32; @@ -5669,6 +5672,7 @@ static void ggml_vk_load_shaders(vk_device& device, vk_pipeline requested) { ggml_vk_create_pipeline(device, device->pipeline_upscale_bilinear_antialias_f32, "upscale_f32", upscale_f32_len, upscale_f32_data, "main", 2, sizeof(vk_op_upscale_push_constants), {512, 1, 1}, {GGML_SCALE_MODE_BILINEAR | GGML_SCALE_FLAG_ANTIALIAS}, 1); ggml_vk_create_pipeline(device, device->pipeline_scale_f32, "scale_f32", scale_f32_len, scale_f32_data, "main", 2, sizeof(vk_op_unary_push_constants), {512, 1, 1}, {}, 1); + ggml_vk_create_pipeline(device, device->pipeline_scale_bf16, "scale_bf16", scale_bf16_len, scale_bf16_data, "main", 2, sizeof(vk_op_unary_push_constants), {512, 1, 1}, {}, 1); ggml_vk_create_pipeline(device, device->pipeline_log[0], "log_f32", log_f32_len, log_f32_data, "main", 2, sizeof(vk_op_unary_push_constants), {512, 1, 1}, {}, 1); ggml_vk_create_pipeline(device, device->pipeline_log[1], "log_f16", log_f16_len, log_f16_data, "main", 2, sizeof(vk_op_unary_push_constants), {512, 1, 1}, {}, 1); @@ -5962,6 +5966,24 @@ static void ggml_vk_load_shaders(vk_device& device, vk_pipeline requested) { gdn_names[si][kda], gdn_len, gdn_data, "main", 7, sizeof(vk_op_gated_delta_net_push_constants), wg_denoms, {S_V, kda, device->subgroup_size, lanes_per_column}, 1, true, use_subgroup_ops, device->subgroup_size); } + size_t gdnr_len; const void * gdnr_data; + if (use_clustered_reduce) { gdnr_len = gated_delta_net_rows_f32_len; gdnr_data = (const void *)gated_delta_net_rows_f32_data; } + else if (use_subgroup_reduce) { gdnr_len = gated_delta_net_rows_f32_nocluster_len; gdnr_data = (const void *)gated_delta_net_rows_f32_nocluster_data; } + else { gdnr_len = gated_delta_net_rows_f32_shmem_len; gdnr_data = (const void *)gated_delta_net_rows_f32_shmem_data; } + for (uint32_t kda = 0; kda < 2; kda++) { + ggml_vk_create_pipeline(device, device->pipeline_gated_delta_net_rows[si][kda], + gdn_names[si][kda], gdnr_len, gdnr_data, "main", 8, sizeof(vk_op_gated_delta_net_push_constants), + wg_denoms, {S_V, kda, device->subgroup_size, lanes_per_column}, 1, true, use_subgroup_ops, device->subgroup_size); + } + size_t gdnrb_len; const void * gdnrb_data; + if (use_clustered_reduce) { gdnrb_len = gated_delta_net_rows_bf16state_f32_len; gdnrb_data = (const void *)gated_delta_net_rows_bf16state_f32_data; } + else if (use_subgroup_reduce) { gdnrb_len = gated_delta_net_rows_bf16state_f32_nocluster_len; gdnrb_data = (const void *)gated_delta_net_rows_bf16state_f32_nocluster_data; } + else { gdnrb_len = gated_delta_net_rows_bf16state_f32_shmem_len; gdnrb_data = (const void *)gated_delta_net_rows_bf16state_f32_shmem_data; } + for (uint32_t kda = 0; kda < 2; kda++) { + ggml_vk_create_pipeline(device, device->pipeline_gated_delta_net_rows_bf16state[si][kda], + gdn_names[si][kda], gdnrb_len, gdnrb_data, "main", 8, sizeof(vk_op_gated_delta_net_push_constants), + wg_denoms, {S_V, kda, device->subgroup_size, lanes_per_column}, 1, true, use_subgroup_ops, device->subgroup_size); + } } } @@ -11390,6 +11412,9 @@ static vk_pipeline ggml_vk_op_get_pipeline(ggml_backend_vk_context * ctx, const if (src0->type == GGML_TYPE_F32 && dst->type == GGML_TYPE_F32) { return ctx->device->pipeline_scale_f32; } + if (src0->type == GGML_TYPE_BF16 && dst->type == GGML_TYPE_BF16) { + return ctx->device->pipeline_scale_bf16; + } return nullptr; case GGML_OP_SQR: if (src0->type == dst->type && @@ -11792,6 +11817,12 @@ static vk_pipeline ggml_vk_op_get_pipeline(ggml_backend_vk_context * ctx, const case 128: si = 3; break; default: return nullptr; } + if (dst->src[5]->type == GGML_TYPE_BF16) { + return ctx->device->pipeline_gated_delta_net_rows_bf16state[si][kda]; + } + if (dst->src[6] != nullptr) { + return ctx->device->pipeline_gated_delta_net_rows[si][kda]; + } return ctx->device->pipeline_gated_delta_net[si][kda]; } return nullptr; @@ -12882,6 +12913,12 @@ static void ggml_vk_gated_delta_net(ggml_backend_vk_context * ctx, vk_context& s for (int i = 0; i < 6; i++) { src_buf[i] = ggml_vk_tensor_subbuffer(ctx, dst->src[i]); } + // rows mode: extra index buffer at binding 7 (dst->src[6]) + const bool gdn_rows = dst->src[6] != nullptr; + vk_subbuffer rows_buf = {}; + if (gdn_rows) { + rows_buf = ggml_vk_tensor_subbuffer(ctx, dst->src[6]); + } const uint32_t sq1 = (uint32_t)(src_q->nb[1] / sizeof(float)); const uint32_t sq2 = (uint32_t)(src_q->nb[2] / sizeof(float)); @@ -12907,9 +12944,15 @@ static void ggml_vk_gated_delta_net(ggml_backend_vk_context * ctx, vk_context& s K }; - ggml_vk_dispatch_pipeline(ctx, subctx, pipeline, - {src_buf[0], src_buf[1], src_buf[2], src_buf[3], src_buf[4], src_buf[5], dst_buf}, - pc, { H, n_seqs, S_v }); + if (gdn_rows) { + ggml_vk_dispatch_pipeline(ctx, subctx, pipeline, + {src_buf[0], src_buf[1], src_buf[2], src_buf[3], src_buf[4], src_buf[5], dst_buf, rows_buf}, + pc, { H, n_seqs, S_v }); + } else { + ggml_vk_dispatch_pipeline(ctx, subctx, pipeline, + {src_buf[0], src_buf[1], src_buf[2], src_buf[3], src_buf[4], src_buf[5], dst_buf}, + pc, { H, n_seqs, S_v }); + } } static void ggml_vk_ssm_scan(ggml_backend_vk_context * ctx, vk_context& subctx, ggml_tensor * dst) { @@ -18661,15 +18704,26 @@ static bool ggml_backend_vk_device_supports_op(ggml_backend_dev_t dev, const ggm if (ggml_get_op_params_i32(op, 1) != 0) { return false; } + // rows-indexed state read (src[6]) requires an int32 index tensor + if (op->src[6] != nullptr && op->src[6]->type != GGML_TYPE_I32) { + return false; + } + if (op->src[6] != nullptr && op->src[5]->type != GGML_TYPE_F32 && op->src[5]->type != GGML_TYPE_BF16) { + return false; + } const uint32_t S_v = op->src[2]->ne[0]; if (S_v != 16 && S_v != 32 && S_v != 64 && S_v != 128) { return false; } - for (int i = 0; i < 6; i++) { +for (int i = 0; i < 6; i++) { + if (i == 5 && op->src[5]->type == GGML_TYPE_BF16) { + continue; // bf16 state handled by the dedicated rows pipeline + } if (op->src[i] == nullptr || op->src[i]->type != GGML_TYPE_F32) { return false; } } + } return op->type == GGML_TYPE_F32; } case GGML_OP_SSM_SCAN: diff --git a/ggml/src/ggml-vulkan/vulkan-shaders/gated_delta_net.comp b/ggml/src/ggml-vulkan/vulkan-shaders/gated_delta_net.comp index 0e384330b9b9..076b61e2a57c 100644 --- a/ggml/src/ggml-vulkan/vulkan-shaders/gated_delta_net.comp +++ b/ggml/src/ggml-vulkan/vulkan-shaders/gated_delta_net.comp @@ -1,6 +1,9 @@ #version 450 #extension GL_EXT_control_flow_attributes : require +#if STATE_BF16 +#extension GL_EXT_shader_16bit_storage : require +#endif #extension GL_KHR_shader_subgroup_basic : enable #if USE_SUBGROUP_CLUSTERED #extension GL_KHR_shader_subgroup_clustered : enable @@ -39,7 +42,16 @@ layout(binding = 1) readonly buffer KBuf { FLOAT_TYPE data_k[]; }; layout(binding = 2) readonly buffer VBuf { FLOAT_TYPE data_v[]; }; layout(binding = 3) readonly buffer GBuf { FLOAT_TYPE data_g[]; }; layout(binding = 4) readonly buffer BetaBuf { FLOAT_TYPE data_beta[]; }; +#if STATE_BF16 +layout(binding = 5) readonly buffer StateBuf { uint16_t data_state_h[]; }; +// idx is in logical element units; uint16_t elements advance 2 bytes each, matching bf16 stride +float state_load(uint idx) { return uintBitsToFloat(uint(data_state_h[idx]) << 16); } +#else layout(binding = 5) readonly buffer StateBuf { FLOAT_TYPE data_state[]; }; +#endif +#if USE_STATE_ROWS +layout(binding = 7) readonly buffer StateRowsBuf { int data_state_rows[]; }; +#endif layout(binding = 6) buffer DstBuf { FLOAT_TYPE data_dst[]; }; #if !USE_SUBGROUP_ADD && !USE_SUBGROUP_CLUSTERED @@ -102,15 +114,26 @@ void main() { const uint iq3 = seq_id / rq3; const uint state_size = S_V * S_V; +#if USE_STATE_ROWS + // rows mode: state is a 2D cache view [D, n_rows]; row for this seq comes + // from the index buffer; row size (in floats) is H*state_size. + const uint row = uint(data_state_rows[seq_id]); + const uint state_in_base = row * H * state_size + head_id * state_size; +#else // input state holds s0 only [S_v, S_v, H, n_seqs]: per-seq stride is H*D. const uint state_in_base = (seq_id * H + head_id) * state_size; +#endif // output state layout per slot: same per-(seq,head) offset as the single-slot case. const uint state_out_base = (seq_id * H + head_id) * state_size; const uint state_size_per_snap = state_size * H * n_seqs; FLOAT_TYPE s_shard[ROWS_PER_LANE]; [[unroll]] for (uint r = 0; r < ROWS_PER_LANE; r++) { +#if STATE_BF16 + s_shard[r] = FLOAT_TYPE(state_load(state_in_base + col * S_V + r * LANES_PER_COLUMN + lane)); +#else s_shard[r] = FLOAT_TYPE(data_state[state_in_base + col * S_V + r * LANES_PER_COLUMN + lane]); +#endif } // snapshot slot mapping: slot 0 = most recent state, slot s = s tokens back. diff --git a/ggml/src/ggml-vulkan/vulkan-shaders/vulkan-shaders-gen.cpp b/ggml/src/ggml-vulkan/vulkan-shaders/vulkan-shaders-gen.cpp index f2239186d922..f491d5569a9e 100644 --- a/ggml/src/ggml-vulkan/vulkan-shaders/vulkan-shaders-gen.cpp +++ b/ggml/src/ggml-vulkan/vulkan-shaders/vulkan-shaders-gen.cpp @@ -902,6 +902,7 @@ void process_shaders() { string_to_spv("repeat_i16", "repeat.comp", {{"A_TYPE", "int16_t"}, {"D_TYPE", "int16_t"}}); string_to_spv("scale_f32", "scale.comp", {{"A_TYPE", "float"}, {"D_TYPE", "float"}, {"FLOAT_TYPE", "float"}}); + string_to_spv("scale_bf16", "scale.comp", {{"A_TYPE", "float16_t"}, {"D_TYPE", "float16_t"}, {"FLOAT_TYPE", "float"}}); string_to_spv("pad_f32", "pad.comp", {{"A_TYPE", "float"}, {"D_TYPE", "float"}}); string_to_spv("pad_reflect_1d_f32", "pad_reflect_1d.comp", {{"A_TYPE", "float"}, {"D_TYPE", "float"}}); @@ -1082,6 +1083,10 @@ void process_shaders() { string_to_spv("gated_delta_net_f32", "gated_delta_net.comp", merge_maps(base_dict, {{"FLOAT_TYPE", "float"}, {"USE_SUBGROUP_ADD", "1"}, {"USE_SUBGROUP_CLUSTERED", "1"}})); string_to_spv("gated_delta_net_f32_nocluster", "gated_delta_net.comp", merge_maps(base_dict, {{"FLOAT_TYPE", "float"}, {"USE_SUBGROUP_ADD", "1"}, {"USE_SUBGROUP_CLUSTERED", "0"}})); string_to_spv("gated_delta_net_f32_shmem", "gated_delta_net.comp", merge_maps(base_dict, {{"FLOAT_TYPE", "float"}, {"USE_SUBGROUP_ADD", "0"}, {"USE_SUBGROUP_CLUSTERED", "0"}})); + string_to_spv("gated_delta_net_rows_f32", "gated_delta_net.comp", merge_maps(base_dict, {{"FLOAT_TYPE", "float"}, {"USE_SUBGROUP_ADD", "1"}, {"USE_SUBGROUP_CLUSTERED", "1"}, {"USE_STATE_ROWS", "1"}})); + string_to_spv("gated_delta_net_rows_f32_nocluster", "gated_delta_net.comp", merge_maps(base_dict, {{"FLOAT_TYPE", "float"}, {"USE_SUBGROUP_ADD", "1"}, {"USE_SUBGROUP_CLUSTERED", "0"}, {"USE_STATE_ROWS", "1"}})); + string_to_spv("gated_delta_net_rows_f32_shmem", "gated_delta_net.comp", merge_maps(base_dict, {{"FLOAT_TYPE", "float"}, {"USE_SUBGROUP_ADD", "0"}, {"USE_SUBGROUP_CLUSTERED", "0"}, {"USE_STATE_ROWS", "1"}})); + string_to_spv("gated_delta_net_rows_bf16state_f32", "gated_delta_net.comp", merge_maps(base_dict, {{"FLOAT_TYPE", "float"}, {"USE_SUBGROUP_ADD", "1"}, {"USE_SUBGROUP_CLUSTERED", "1"}, {"USE_STATE_ROWS", "1"}, {"STATE_BF16", "1"}})); string_to_spv("opt_step_adamw_f32", "opt_step_adamw.comp", merge_maps(base_dict, {{"A_TYPE", "float"}})); string_to_spv("opt_step_sgd_f32", "opt_step_sgd.comp", merge_maps(base_dict, {{"A_TYPE", "float"}})); diff --git a/ggml/src/ggml.c b/ggml/src/ggml.c index fa2a9c44f5cf..5af5f2cb39a5 100644 --- a/ggml/src/ggml.c +++ b/ggml/src/ggml.c @@ -6387,7 +6387,7 @@ struct ggml_tensor * ggml_gated_delta_net_rows( GGML_ASSERT(v->type == GGML_TYPE_F32); GGML_ASSERT(g->type == GGML_TYPE_F32); GGML_ASSERT(beta->type == GGML_TYPE_F32); - GGML_ASSERT(states->type == GGML_TYPE_F32); + GGML_ASSERT(states->type == GGML_TYPE_F32 || states->type == GGML_TYPE_BF16); GGML_ASSERT(rows->type == GGML_TYPE_I32); const int64_t S_v = v->ne[0]; diff --git a/src/llama-model.cpp b/src/llama-model.cpp index c2ac6a50b316..b36a9f0cbd05 100644 --- a/src/llama-model.cpp +++ b/src/llama-model.cpp @@ -2814,6 +2814,10 @@ llama_memory_i * llama_model::create_memory(const llama_memory_params & params, /* filter_attn */ std::move(filter_attn), /* filter_recr */ std::move(filter_recr)); } else { + // bf16 SSM state pools: recurrence is bandwidth-bound on these; + // bf16 halves the traffic. Opt-in via env until kernels are audited. + const bool bf16_ssm_state = getenv("LLAMA_SSM_BF16_STATE") != nullptr; + const bool bf16_ssm_conv = getenv("LLAMA_SSM_BF16_CONV") != nullptr; res = new llama_memory_hybrid( /* model */ *this, /* attn_type_k */ params.type_k, @@ -2823,8 +2827,8 @@ llama_memory_i * llama_model::create_memory(const llama_memory_params & params, /* attn_n_pad */ 1, /* attn_n_swa */ hparams.n_swa, /* attn_swa_type */ hparams.swa_type, - /* recurrent_type_k */ GGML_TYPE_F32, - /* recurrent_type_v */ GGML_TYPE_F32, + /* recurrent_type_k */ bf16_ssm_conv ? GGML_TYPE_BF16 : GGML_TYPE_F32, + /* recurrent_type_v */ bf16_ssm_state ? GGML_TYPE_BF16 : GGML_TYPE_F32, /* recurrent_kv_size */ std::max((uint32_t) 1, cparams.n_seq_max), /* n_seq_max */ cparams.n_seq_max, /* n_rs_seq */ cparams.n_rs_seq,