Skip to content
Open
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
62 changes: 58 additions & 4 deletions ggml/src/ggml-vulkan/ggml-vulkan.cpp
Original file line number Diff line number Diff line change
Expand Up @@ -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];
Expand Down Expand Up @@ -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;
Expand Down Expand Up @@ -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);
Expand Down Expand Up @@ -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);
}
}
}

Expand Down Expand Up @@ -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 &&
Expand Down Expand Up @@ -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;
Expand Down Expand Up @@ -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));
Expand All @@ -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) {
Expand Down Expand Up @@ -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:
Expand Down
23 changes: 23 additions & 0 deletions ggml/src/ggml-vulkan/vulkan-shaders/gated_delta_net.comp
Original file line number Diff line number Diff line change
@@ -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
Expand Down Expand Up @@ -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
Expand Down Expand Up @@ -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.
Expand Down
5 changes: 5 additions & 0 deletions ggml/src/ggml-vulkan/vulkan-shaders/vulkan-shaders-gen.cpp
Original file line number Diff line number Diff line change
Expand Up @@ -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"}});
Expand Down Expand Up @@ -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"}}));
Expand Down
2 changes: 1 addition & 1 deletion ggml/src/ggml.c
Original file line number Diff line number Diff line change
Expand Up @@ -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];
Expand Down
8 changes: 6 additions & 2 deletions src/llama-model.cpp
Original file line number Diff line number Diff line change
Expand Up @@ -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,
Expand All @@ -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,
Expand Down