Skip to content
Merged
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
142 changes: 118 additions & 24 deletions ggml/src/ggml-vulkan/ggml-vulkan.cpp
Original file line number Diff line number Diff line change
Expand Up @@ -1159,6 +1159,10 @@ struct vk_device_struct {

std::map<std::pair<uint32_t, uint32_t>, vk_pipeline> pipeline_fa_mask_opt;

vk_pipeline pipeline_fa_sparse_compact;
vk_pipeline pipeline_fa_sparse_compact_subgroup;
bool fa_sparse_compact_use_subgroups;

vk_pipeline pipeline_flash_attn_split_k_reduce;
vk_pipeline pipeline_count_experts;

Expand Down Expand Up @@ -2184,6 +2188,16 @@ struct vk_op_flash_attn_mask_opt_push_constants {
uint32_t nbd3;
};

struct vk_op_flash_attn_sparse_compact_push_constants {
uint32_t KV;
uint32_t nem1;
uint32_t nem2;
uint32_t nbm1;
uint32_t nbm2;
uint32_t nbm3;
uint32_t n_kv_max;
};

// Allow pre-recording command buffers
struct vk_staging_memcpy {
vk_staging_memcpy(void * _dst, const void * _src, size_t _n) : dst(_dst), src(_src), n(_n) {}
Expand Down Expand Up @@ -4103,14 +4117,15 @@ static vk_fa_tuning_params get_fa_tuning_params(const vk_device& device, uint32_
}

static vk_fa_pipeline_state get_fa_pipeline_state(const vk_device& device, const vk_fa_tuning_params& params, uint32_t hsk, uint32_t hsv, bool aligned, bool f32acc,
bool use_mask, bool use_mask_opt, bool use_logit_softcap, ggml_type k_type, ggml_type v_type) {
bool use_mask, bool use_mask_opt, bool use_logit_softcap, bool use_sparse, ggml_type k_type, ggml_type v_type) {
const bool old_amd_windows = device->vendor_id == VK_VENDOR_ID_AMD && device->driver_id == vk::DriverId::eAmdProprietary &&
(device->architecture == AMD_GCN || device->architecture == AMD_RDNA1 || device->architecture == AMD_RDNA2);

uint32_t flags = (use_mask_opt ? 1 : 0) |
(use_mask ? 2 : 0) |
(use_logit_softcap ? 4 : 0) |
(old_amd_windows ? 8 : 0);
(old_amd_windows ? 8 : 0) |
(use_sparse ? 16 : 0);

const uint32_t subgroup_size = params.disable_subgroups ? 0 : params.subgroup_size;

Expand Down Expand Up @@ -4730,7 +4745,7 @@ static void ggml_vk_load_shaders(vk_device& device, vk_pipeline requested) {
}
name = aligned ? "flash_attn_f32_f16_aligned" : "flash_attn_f32_f16";
}
ggml_vk_create_pipeline(device, fa.second, name, spv_size, spv_data, "main", 7,
ggml_vk_create_pipeline(device, fa.second, name, spv_size, spv_data, "main", 8,
sizeof(vk_flash_attn_push_constants), {Br, 1, 1},
get_fa_spec_constants(fa.first), aligned ? Bc : 1, true,
!fa_ds, !fa_ds ? fa_sgs : 0);
Expand Down Expand Up @@ -4766,7 +4781,7 @@ static void ggml_vk_load_shaders(vk_device& device, vk_pipeline requested) {
else { spv_data = flash_attn_f32_f16_f16acc_cm1_data; spv_size = flash_attn_f32_f16_f16acc_cm1_len; }
name = aligned ? "flash_attn_f32_f16_aligned_cm1" : "flash_attn_f32_f16_cm1";
}
ggml_vk_create_pipeline(device, fa.second, name, spv_size, spv_data, "main", 7,
ggml_vk_create_pipeline(device, fa.second, name, spv_size, spv_data, "main", 8,
sizeof(vk_flash_attn_push_constants), {Br, 1, 1},
get_fa_spec_constants(fa.first), aligned ? Bc : 1, true,
!fa_ds, !fa_ds ? fa_sgs : 0);
Expand Down Expand Up @@ -4803,7 +4818,7 @@ static void ggml_vk_load_shaders(vk_device& device, vk_pipeline requested) {
if (f32acc) { spv_data = flash_attn_f32_f16_cm2_data; spv_size = flash_attn_f32_f16_cm2_len; name = "flash_attn_f32_f16_f32acc_cm2"; }
else { spv_data = flash_attn_f32_f16_f16acc_cm2_data; spv_size = flash_attn_f32_f16_f16acc_cm2_len; name = "flash_attn_f32_f16_f16acc_cm2"; }
}
ggml_vk_create_pipeline(device, fa.second, name, spv_size, spv_data, "main", 7,
ggml_vk_create_pipeline(device, fa.second, name, spv_size, spv_data, "main", 8,
sizeof(vk_flash_attn_push_constants), {Br, 1, 1},
get_fa_spec_constants(fa.first), aligned ? Bc : 1, true, false, 0);
}
Expand Down Expand Up @@ -5767,6 +5782,22 @@ static void ggml_vk_load_shaders(vk_device& device, vk_pipeline requested) {
ggml_vk_create_pipeline(device, it.second, "fa_mask_opt", fa_mask_opt_len, fa_mask_opt_data, "main", 2, sizeof(vk_op_flash_attn_mask_opt_push_constants), {1, 1, 1}, {128, 128 / device->subgroup_size, BrBc.first, BrBc.second}, 1, true, true, device->subgroup_size);
}

{
// Large workgroup so the per-row KV scan parallelizes; capped to device limits.
const uint32_t compact_max = std::min({1024u, device->properties.limits.maxComputeWorkGroupInvocations, device->properties.limits.maxComputeWorkGroupSize[0]});

// Fast ballot prefix-sum path when the device supports full subgroups; otherwise
// a shared-memory prefix-sum fallback. Both emit a deterministic ascending list.
device->fa_sparse_compact_use_subgroups = device->subgroup_ballot && device->subgroup_require_full_support;
if (device->fa_sparse_compact_use_subgroups) {
const uint32_t compact_wg = std::max(device->subgroup_size, (compact_max / device->subgroup_size) * device->subgroup_size);
const uint32_t compact_num_sg = compact_wg / device->subgroup_size;
ggml_vk_create_pipeline(device, device->pipeline_fa_sparse_compact_subgroup, "fa_sparse_compact_subgroup", fa_sparse_compact_subgroup_len, fa_sparse_compact_subgroup_data, "main", 2, sizeof(vk_op_flash_attn_sparse_compact_push_constants), {1, 1, 1}, {compact_wg, compact_num_sg}, 1, true, true, device->subgroup_size);
} else {
ggml_vk_create_pipeline(device, device->pipeline_fa_sparse_compact, "fa_sparse_compact", fa_sparse_compact_len, fa_sparse_compact_data, "main", 2, sizeof(vk_op_flash_attn_sparse_compact_push_constants), {1, 1, 1}, {compact_max}, 1, true);
}
}

if (device->subgroup_clustered && device->subgroup_require_full_support) {
ggml_vk_create_pipeline(device, device->pipeline_quantize_q8_1_x4, "quantize_q8_1_x4", quantize_q8_1_x4_subgroup_len, quantize_q8_1_x4_subgroup_data, "main", 2, sizeof(vk_quantize_q8_1_push_constants), {32 * device->subgroup_size / 8, 1, 1}, { device->subgroup_size }, 1, true, true);
} else {
Expand Down Expand Up @@ -11260,6 +11291,30 @@ static void ggml_vk_flash_attn(ggml_backend_vk_context * ctx, vk_context& subctx

tuning_params = get_fa_tuning_params(ctx->device, HSK, HSV, N, KV, k_type_eff, v_type_eff, f32acc);

float scale = 1.0f;
float max_bias = 0.0f;
float logit_softcap = 0.0f;

memcpy(&scale, (const float *) dst->op_params + 0, sizeof(float));
memcpy(&max_bias, (const float *) dst->op_params + 1, sizeof(float));
memcpy(&logit_softcap, (const float *) dst->op_params + 2, sizeof(float));

if (logit_softcap != 0) {
scale /= logit_softcap;
}

// Sparse mask hint (op_params[4]): compact the <= n_kv_max finite positions and gather only those.
const int32_t n_kv_max = mask ? ggml_get_op_params_i32(dst, 4) : 0;
static const bool disable_sparse = getenv("GGML_VK_FA_SPARSE_DISABLE") != nullptr;
// cm2 dense is fast, so it needs a larger reduction to win.
const int64_t min_ratio = tuning_params.path == FA_COOPMAT2 ? 4 : 2;
const bool use_sparse = !disable_sparse && n_kv_max > 0 && mask &&
max_bias == 0.0f && logit_softcap == 0.0f &&
k_type_eff == GGML_TYPE_F16 && v_type_eff == GGML_TYPE_F16 &&
nem0 == KV &&
(int64_t)KV >= std::max<int64_t>(4096, min_ratio * (int64_t)n_kv_max) &&
(gqa_ratio > 1 || (tuning_params.path == FA_SCALAR && N == 1));

const uint32_t q_stride = (uint32_t)(nbq1 / ggml_type_size(q->type));
uint32_t k_stride = (uint32_t)(nbk1 / ggml_type_size(k->type));
uint32_t v_stride = (uint32_t)(nbv1 / ggml_type_size(v->type));
Expand All @@ -11282,7 +11337,6 @@ static void ggml_vk_flash_attn(ggml_backend_vk_context * ctx, vk_context& subctx
nbv2_eff = (uint32_t)((uint64_t)HSV * KV * sizeof(ggml_fp16_t));
nbv3_eff = (uint32_t)((uint64_t)HSV * KV * nev2 * sizeof(ggml_fp16_t));
}

const uint32_t alignment = tuning_params.block_cols;
bool aligned = (KV % alignment) == 0 &&
// the "aligned" shader variant will forcibly align strides, for performance
Expand All @@ -11293,23 +11347,11 @@ static void ggml_vk_flash_attn(ggml_backend_vk_context * ctx, vk_context& subctx
aligned = false;
}

float scale = 1.0f;
float max_bias = 0.0f;
float logit_softcap = 0.0f;

memcpy(&scale, (const float *) dst->op_params + 0, sizeof(float));
memcpy(&max_bias, (const float *) dst->op_params + 1, sizeof(float));
memcpy(&logit_softcap, (const float *) dst->op_params + 2, sizeof(float));

if (logit_softcap != 0) {
scale /= logit_softcap;
}

// Only use mask opt when the mask is fairly large. This hasn't been tuned extensively.
bool use_mask_opt = mask && nem1 >= 32 && nem0 * nem1 > 32768 && nem0 >= tuning_params.block_cols * 16
bool use_mask_opt = mask && !use_sparse && nem1 >= 32 && nem0 * nem1 > 32768 && nem0 >= tuning_params.block_cols * 16
&& (ctx->device->architecture != vk_device_architecture::AMD_GCN || HSK > 256 || HSV > 256);
vk_fa_pipeline_state fa_pipeline_state = get_fa_pipeline_state(ctx->device, tuning_params, HSK, HSV, aligned, f32acc,
mask != nullptr, use_mask_opt, logit_softcap != 0, k_type_eff, v_type_eff);
mask != nullptr, use_mask_opt, logit_softcap != 0, use_sparse, k_type_eff, v_type_eff);

vk_pipeline pipeline = nullptr;

Expand Down Expand Up @@ -11344,7 +11386,19 @@ static void ggml_vk_flash_attn(ggml_backend_vk_context * ctx, vk_context& subctx
const uint32_t Tr = CEIL_DIV(N, Br);

// Try to use split_k when KV is large enough to be worth the overhead.
if (gqa_ratio > 1 && workgroups_x <= Br) {
// Sparse: split_kv carries n_kv_max, split_k partitions its blocks for occupancy.
if (use_sparse) {
split_kv = (uint32_t)n_kv_max;
const uint32_t total_blocks = CEIL_DIV((uint32_t)n_kv_max, Bc);
const uint32_t base_wgs = (gqa_ratio > 1 ? workgroups_x : Tr) * workgroups_y * workgroups_z;
if (base_wgs < shader_core_count * 2) {
split_k = shader_core_count * 2 / base_wgs;
}
split_k = std::max(1u, std::min(split_k, total_blocks));
// Match the shader's per-split block count so no split is empty.
const uint32_t per_blocks = CEIL_DIV(total_blocks, split_k);
split_k = CEIL_DIV(total_blocks, per_blocks);
} else if (gqa_ratio > 1 && workgroups_x <= Br) {
split_k = shader_core_count * 2 / (workgroups_x * workgroups_y * workgroups_z);
} else if (gqa_ratio <= 1) {
uint32_t total_wgs_no_split = Tr * workgroups_y * workgroups_z;
Expand All @@ -11353,7 +11407,7 @@ static void ggml_vk_flash_attn(ggml_backend_vk_context * ctx, vk_context& subctx
}
}

if (split_k > 1) {
if (!use_sparse && split_k > 1) {
// Try to evenly split KV into split_k chunks, but it needs to be a multiple
// of "align", so recompute split_k based on that.
split_kv = ROUNDUP_POW2(std::max(1u, KV / split_k), alignment);
Expand Down Expand Up @@ -11400,6 +11454,24 @@ static void ggml_vk_flash_attn(ggml_backend_vk_context * ctx, vk_context& subctx
}
}

// Sparse index scratch reuses prealloc_y (mutually exclusive with mask opt).
const uint64_t sparse_idx_size = use_sparse
? sizeof(int32_t) * (uint64_t)n_kv_max * nem1 * nem2 * nem3
: 0;
vk_pipeline sparse_compact_pipeline = ctx->device->fa_sparse_compact_use_subgroups
? ctx->device->pipeline_fa_sparse_compact_subgroup
: ctx->device->pipeline_fa_sparse_compact;
if (use_sparse) {
ggml_pipeline_request_descriptor_sets(ctx, sparse_compact_pipeline, 1);
if (ctx->prealloc_size_y < sparse_idx_size) {
ctx->prealloc_size_y = sparse_idx_size;
ggml_vk_preallocate_buffers(ctx, subctx);
}
if (ctx->prealloc_y_need_sync) {
ggml_vk_sync_buffers(ctx, subctx);
}
}

const uint32_t n_head_kv = neq2;
const uint32_t n_head_log2 = 1u << (uint32_t) floorf(log2f((float) n_head_kv));
const float m0 = powf(2.0f, -(max_bias ) / n_head_log2);
Expand All @@ -11412,6 +11484,7 @@ static void ggml_vk_flash_attn(ggml_backend_vk_context * ctx, vk_context& subctx
vk_subbuffer mask_buf = mask ? ggml_vk_tensor_subbuffer(ctx, mask) : q_buf;
vk_subbuffer sinks_buf = sinks ? ggml_vk_tensor_subbuffer(ctx, sinks) : q_buf;
vk_subbuffer mask_opt_buf = use_mask_opt ? ggml_vk_subbuffer(ctx, ctx->prealloc_y, 0) : q_buf;
vk_subbuffer sparse_buf = use_sparse ? ggml_vk_subbuffer(ctx, ctx->prealloc_y, 0) : q_buf;

if (use_dequant_kv) {
const uint64_t fp = sizeof(ggml_fp16_t);
Expand Down Expand Up @@ -11463,6 +11536,24 @@ static void ggml_vk_flash_attn(ggml_backend_vk_context * ctx, vk_context& subctx
ggml_vk_sync_buffers(ctx, subctx);
}

if (use_sparse)
{
const vk_op_flash_attn_sparse_compact_push_constants sc_pc = {
KV,
nem1,
nem2,
(uint32_t)(mask->nb[1] / sizeof(ggml_fp16_t)),
(uint32_t)(mask->nb[2] / sizeof(ggml_fp16_t)),
(uint32_t)(mask->nb[3] / sizeof(ggml_fp16_t)),
(uint32_t)n_kv_max,
};

ggml_vk_dispatch_pipeline(ctx, subctx, sparse_compact_pipeline,
{ mask_buf, sparse_buf }, sc_pc,
{ nem1, nem2, nem3 });
ggml_vk_sync_buffers(ctx, subctx);
}

const vk_flash_attn_push_constants pc = { N, KV,
(uint32_t)ne1, (uint32_t)ne2, (uint32_t)ne3,
(uint32_t)neq2, (uint32_t)neq3,
Expand Down Expand Up @@ -11495,7 +11586,7 @@ static void ggml_vk_flash_attn(ggml_backend_vk_context * ctx, vk_context& subctx

vk_subbuffer split_k_buf = ggml_vk_subbuffer(ctx, ctx->prealloc_split_k, 0);
ggml_vk_dispatch_pipeline(ctx, subctx, pipeline,
{q_buf, k_buf, v_buf, mask_buf, sinks_buf, split_k_buf, mask_opt_buf},
{q_buf, k_buf, v_buf, mask_buf, sinks_buf, split_k_buf, mask_opt_buf, sparse_buf},
pc, { dispatch_x, workgroups_y, workgroups_z });

ggml_vk_sync_buffers(ctx, subctx);
Expand All @@ -11510,13 +11601,16 @@ static void ggml_vk_flash_attn(ggml_backend_vk_context * ctx, vk_context& subctx
workgroups_x *= pipeline->wg_denoms[0];
}
ggml_vk_dispatch_pipeline(ctx, subctx, pipeline,
{q_buf, k_buf, v_buf, mask_buf, sinks_buf, dst_buf, mask_opt_buf},
{q_buf, k_buf, v_buf, mask_buf, sinks_buf, dst_buf, mask_opt_buf, sparse_buf},
pc, { workgroups_x, workgroups_y, workgroups_z });
}

if (use_dequant_kv) {
ctx->prealloc_x_need_sync = true;
}
if (use_mask_opt || use_sparse) {
ctx->prealloc_y_need_sync = true;
}
}

static vk_conv_shapes ggml_vk_conv_select_shape(ggml_backend_vk_context * ctx, uint32_t K, uint32_t NPQ) {
Expand Down
Loading
Loading