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
265 changes: 265 additions & 0 deletions ggml/src/ggml-vulkan/ggml-vulkan.cpp
Original file line number Diff line number Diff line change
Expand Up @@ -815,6 +815,11 @@ struct vk_device_struct {
bool subgroup_clustered;
bool subgroup_vote;
bool multi_add;
// DeepSeek-V4 hyper-connection ops, per-op so a kernel can be bisected
// against the unfused graph.
bool dsv4_hc_comb;
bool dsv4_hc_pre;
bool dsv4_hc_post;
bool shader_int64;
bool buffer_device_address;
bool vulkan_memory_model;
Expand Down Expand Up @@ -1047,6 +1052,9 @@ struct vk_device_struct {
vk_pipeline pipeline_cumsum_multipass2_f32;
vk_pipeline pipeline_argmax_f32;
vk_pipeline pipeline_count_equal_i32;
vk_pipeline pipeline_dsv4_hc_comb_f32;
vk_pipeline pipeline_dsv4_hc_pre_f32;
vk_pipeline pipeline_dsv4_hc_post_f32;
std::map<vk_solve_tri_pipeline_state, vk_pipeline> pipeline_solve_tri_f32;
vk_pipeline pipeline_im2col_f32, pipeline_im2col_f32_f16;
vk_pipeline pipeline_im2col_3d_f32, pipeline_im2col_3d_f32_f16;
Expand Down Expand Up @@ -1402,6 +1410,53 @@ struct vk_op_fwht_push_constants {
float scale;
};

struct vk_op_dsv4_hc_comb_push_constants {
uint32_t n_tokens;

uint32_t nbm0; uint32_t nbm1;
uint32_t nbs0;
uint32_t nbb0;
uint32_t nbd0; uint32_t nbd1; uint32_t nbd2;

uint32_t m_offset;
uint32_t s_offset;
uint32_t b_offset;
uint32_t d_offset;

float eps;
uint32_t n_iter;
};

struct vk_op_dsv4_hc_pre_push_constants {
uint32_t n_embd;
uint32_t n_tokens;

uint32_t nbx0; uint32_t nbx1; uint32_t nbx2;
uint32_t nbw0; uint32_t nbw1;
uint32_t nbd0; uint32_t nbd1;

uint32_t x_offset;
uint32_t w_offset;
uint32_t d_offset;
};

struct vk_op_dsv4_hc_post_push_constants {
uint32_t n_embd;
uint32_t n_tokens;

uint32_t nbx0; uint32_t nbx1;
uint32_t nbr0; uint32_t nbr1; uint32_t nbr2;
uint32_t nbp0; uint32_t nbp1;
uint32_t nbc0; uint32_t nbc1; uint32_t nbc2;
uint32_t nbd0; uint32_t nbd1; uint32_t nbd2;

uint32_t x_offset;
uint32_t r_offset;
uint32_t p_offset;
uint32_t c_offset;
uint32_t d_offset;
};

struct vk_op_count_experts_push_constants {
uint32_t ne00;
uint32_t ne01;
Expand Down Expand Up @@ -2498,6 +2553,32 @@ template <> void init_pushconst_tensor_offsets(ggml_backend_vk_context * ctx, vk
GGML_UNUSED(src3);
}

template <> void init_pushconst_tensor_offsets(ggml_backend_vk_context * ctx, vk_op_dsv4_hc_comb_push_constants &p, const ggml_tensor * src0, const ggml_tensor * src1, const ggml_tensor * src2, const ggml_tensor * src3, ggml_tensor * dst) {
p.m_offset = get_misalign_bytes(ctx, src0) / ggml_type_size(src0->type);
p.s_offset = get_misalign_bytes(ctx, src1) / ggml_type_size(src1->type);
p.b_offset = get_misalign_bytes(ctx, src2) / ggml_type_size(src2->type);
p.d_offset = get_misalign_bytes(ctx, dst) / ggml_type_size(dst->type);

GGML_UNUSED(src3);
}

template <> void init_pushconst_tensor_offsets(ggml_backend_vk_context * ctx, vk_op_dsv4_hc_pre_push_constants &p, const ggml_tensor * src0, const ggml_tensor * src1, const ggml_tensor * src2, const ggml_tensor * src3, ggml_tensor * dst) {
p.x_offset = get_misalign_bytes(ctx, src0) / ggml_type_size(src0->type);
p.w_offset = get_misalign_bytes(ctx, src1) / ggml_type_size(src1->type);
p.d_offset = get_misalign_bytes(ctx, dst) / ggml_type_size(dst->type);

GGML_UNUSED(src2);
GGML_UNUSED(src3);
}

template <> void init_pushconst_tensor_offsets(ggml_backend_vk_context * ctx, vk_op_dsv4_hc_post_push_constants &p, const ggml_tensor * src0, const ggml_tensor * src1, const ggml_tensor * src2, const ggml_tensor * src3, ggml_tensor * dst) {
p.x_offset = get_misalign_bytes(ctx, src0) / ggml_type_size(src0->type);
p.r_offset = get_misalign_bytes(ctx, src1) / ggml_type_size(src1->type);
p.p_offset = get_misalign_bytes(ctx, src2) / ggml_type_size(src2->type);
p.c_offset = get_misalign_bytes(ctx, src3) / ggml_type_size(src3->type);
p.d_offset = get_misalign_bytes(ctx, dst) / ggml_type_size(dst->type);
}

struct ggml_backend_vk_buffer_context {
vk_device_ref device;
vk_buffer dev_buffer;
Expand Down Expand Up @@ -5784,6 +5865,16 @@ static void ggml_vk_load_shaders(vk_device& device, vk_pipeline requested) {

ggml_vk_create_pipeline(device, device->pipeline_count_experts, "count_experts", count_experts_len, count_experts_data, "main", 2, sizeof(vk_op_count_experts_push_constants), {1, 1, 1}, {}, 1, true);

// comb holds a token's 4x4 matrix in one 16-lane slice of a subgroup, so it
// needs at least 16 lanes, pinned to a known size.
if (device->subgroup_basic && device->subgroup_shuffle && device->subgroup_require_full_support && device->subgroup_size >= 16) {
const uint32_t tokens_per_workgroup = 4 * (device->subgroup_size / 16);
ggml_vk_create_pipeline(device, device->pipeline_dsv4_hc_comb_f32, "dsv4_hc_comb_f32", dsv4_hc_comb_f32_len, dsv4_hc_comb_f32_data, "main", 4, sizeof(vk_op_dsv4_hc_comb_push_constants), {tokens_per_workgroup, 1, 1}, { device->subgroup_size }, 1, true, true, device->subgroup_size);
}

ggml_vk_create_pipeline(device, device->pipeline_dsv4_hc_pre_f32, "dsv4_hc_pre_f32", dsv4_hc_pre_f32_len, dsv4_hc_pre_f32_data, "main", 3, sizeof(vk_op_dsv4_hc_pre_push_constants), {256, 1, 1}, { 256 }, 1);
ggml_vk_create_pipeline(device, device->pipeline_dsv4_hc_post_f32, "dsv4_hc_post_f32", dsv4_hc_post_f32_len, dsv4_hc_post_f32_data, "main", 5, sizeof(vk_op_dsv4_hc_post_push_constants), {256, 1, 1}, { 256 }, 1);

for (auto &s : device->pipeline_solve_tri_f32) {
const vk_solve_tri_pipeline_state &state = s.first;

Expand Down Expand Up @@ -6679,6 +6770,11 @@ static vk_device ggml_vk_get_device(size_t idx) {
device->properties.limits.maxPushConstantsSize >= sizeof(vk_op_multi_add_push_constants) &&
getenv("GGML_VK_DISABLE_MULTI_ADD") == nullptr;

const bool dsv4_hc_all = getenv("GGML_VK_DISABLE_DSV4_HC") == nullptr;
device->dsv4_hc_comb = dsv4_hc_all && getenv("GGML_VK_DISABLE_DSV4_HC_COMB") == nullptr;
device->dsv4_hc_pre = dsv4_hc_all && getenv("GGML_VK_DISABLE_DSV4_HC_PRE") == nullptr;
device->dsv4_hc_post = dsv4_hc_all && getenv("GGML_VK_DISABLE_DSV4_HC_POST") == nullptr;

device->shader_int64 = device_features2.features.shaderInt64;
device->buffer_device_address = vk12_features.bufferDeviceAddress;
device->vulkan_memory_model = vk12_features.vulkanMemoryModel;
Expand Down Expand Up @@ -9999,6 +10095,122 @@ static void ggml_vk_fwht(ggml_backend_vk_context * ctx, vk_context& subctx, cons
ggml_vk_dispatch_pipeline(ctx, subctx, pipeline, { src_buf, dst_buf }, pc, { workgroups_x, 1, 1 });
}

// Element stride; supports_op checks that the byte strides divide evenly.
static uint32_t ggml_vk_nb_elem(const ggml_tensor * t, int i) {
return (uint32_t)(t->nb[i] / ggml_type_size(t->type));
}

// The HC shaders address their tensors in whole f32 elements.
static bool ggml_vk_dsv4_hc_strides_ok(const ggml_tensor * op) {
auto const & strides_ok = [](const ggml_tensor * t) {
const size_t ts = ggml_type_size(t->type);
for (int i = 0; i < GGML_MAX_DIMS; ++i) {
if (t->nb[i] % ts != 0) {
return false;
}
}
return true;
};

if (!strides_ok(op)) {
return false;
}
for (uint32_t i = 0; i < GGML_MAX_SRC; ++i) {
if (op->src[i] && !strides_ok(op->src[i])) {
return false;
}
}
return true;
}

static void ggml_vk_dsv4_hc_comb(ggml_backend_vk_context * ctx, vk_context& subctx, const ggml_tensor * mixes, const ggml_tensor * scale, const ggml_tensor * base, ggml_tensor * dst) {
VK_LOG_DEBUG("ggml_vk_dsv4_hc_comb(" << mixes << ", " << scale << ", " << base << ", " << dst << ")");

vk_pipeline pipeline = ctx->device->pipeline_dsv4_hc_comb_f32;
GGML_ASSERT(pipeline != nullptr);

const uint32_t n_tokens = (uint32_t)mixes->ne[1];

ggml_pipeline_request_descriptor_sets(ctx, pipeline, 1);

const vk_subbuffer mixes_buf = ggml_vk_tensor_subbuffer(ctx, mixes, true);
const vk_subbuffer scale_buf = ggml_vk_tensor_subbuffer(ctx, scale, true);
const vk_subbuffer base_buf = ggml_vk_tensor_subbuffer(ctx, base, true);
const vk_subbuffer dst_buf = ggml_vk_tensor_subbuffer(ctx, dst, true);

vk_op_dsv4_hc_comb_push_constants pc = {
n_tokens,
ggml_vk_nb_elem(mixes, 0), ggml_vk_nb_elem(mixes, 1),
ggml_vk_nb_elem(scale, 0),
ggml_vk_nb_elem(base, 0),
ggml_vk_nb_elem(dst, 0), ggml_vk_nb_elem(dst, 1), ggml_vk_nb_elem(dst, 2),
0, 0, 0, 0,
ggml_get_op_params_f32(dst, 0),
(uint32_t)ggml_get_op_params_i32(dst, 1),
};
init_pushconst_tensor_offsets(ctx, pc, mixes, scale, base, nullptr, dst);

ggml_vk_dispatch_pipeline(ctx, subctx, pipeline, { mixes_buf, scale_buf, base_buf, dst_buf }, pc, { n_tokens, 1, 1 });
}

static void ggml_vk_dsv4_hc_pre(ggml_backend_vk_context * ctx, vk_context& subctx, const ggml_tensor * x, const ggml_tensor * weights, ggml_tensor * dst) {
VK_LOG_DEBUG("ggml_vk_dsv4_hc_pre(" << x << ", " << weights << ", " << dst << ")");

vk_pipeline pipeline = ctx->device->pipeline_dsv4_hc_pre_f32;
GGML_ASSERT(pipeline != nullptr);

const uint32_t n_embd = (uint32_t)x->ne[0];
const uint32_t n_tokens = (uint32_t)x->ne[2];

ggml_pipeline_request_descriptor_sets(ctx, pipeline, 1);

const vk_subbuffer x_buf = ggml_vk_tensor_subbuffer(ctx, x, true);
const vk_subbuffer w_buf = ggml_vk_tensor_subbuffer(ctx, weights, true);
const vk_subbuffer d_buf = ggml_vk_tensor_subbuffer(ctx, dst, true);

vk_op_dsv4_hc_pre_push_constants pc = {
n_embd, n_tokens,
ggml_vk_nb_elem(x, 0), ggml_vk_nb_elem(x, 1), ggml_vk_nb_elem(x, 2),
ggml_vk_nb_elem(weights, 0), ggml_vk_nb_elem(weights, 1),
ggml_vk_nb_elem(dst, 0), ggml_vk_nb_elem(dst, 1),
0, 0, 0,
};
init_pushconst_tensor_offsets(ctx, pc, x, weights, nullptr, nullptr, dst);

ggml_vk_dispatch_pipeline(ctx, subctx, pipeline, { x_buf, w_buf, d_buf }, pc, { n_embd, n_tokens, 1 });
}

static void ggml_vk_dsv4_hc_post(ggml_backend_vk_context * ctx, vk_context& subctx, const ggml_tensor * x, const ggml_tensor * residual, const ggml_tensor * post, const ggml_tensor * comb, ggml_tensor * dst) {
VK_LOG_DEBUG("ggml_vk_dsv4_hc_post(" << x << ", " << residual << ", " << post << ", " << comb << ", " << dst << ")");

vk_pipeline pipeline = ctx->device->pipeline_dsv4_hc_post_f32;
GGML_ASSERT(pipeline != nullptr);

const uint32_t n_embd = (uint32_t)x->ne[0];
const uint32_t n_tokens = (uint32_t)x->ne[1];

ggml_pipeline_request_descriptor_sets(ctx, pipeline, 1);

const vk_subbuffer x_buf = ggml_vk_tensor_subbuffer(ctx, x, true);
const vk_subbuffer r_buf = ggml_vk_tensor_subbuffer(ctx, residual, true);
const vk_subbuffer p_buf = ggml_vk_tensor_subbuffer(ctx, post, true);
const vk_subbuffer c_buf = ggml_vk_tensor_subbuffer(ctx, comb, true);
const vk_subbuffer d_buf = ggml_vk_tensor_subbuffer(ctx, dst, true);

vk_op_dsv4_hc_post_push_constants pc = {
n_embd, n_tokens,
ggml_vk_nb_elem(x, 0), ggml_vk_nb_elem(x, 1),
ggml_vk_nb_elem(residual, 0), ggml_vk_nb_elem(residual, 1), ggml_vk_nb_elem(residual, 2),
ggml_vk_nb_elem(post, 0), ggml_vk_nb_elem(post, 1),
ggml_vk_nb_elem(comb, 0), ggml_vk_nb_elem(comb, 1), ggml_vk_nb_elem(comb, 2),
ggml_vk_nb_elem(dst, 0), ggml_vk_nb_elem(dst, 1), ggml_vk_nb_elem(dst, 2),
0, 0, 0, 0, 0,
};
init_pushconst_tensor_offsets(ctx, pc, x, residual, post, comb, dst);

ggml_vk_dispatch_pipeline(ctx, subctx, pipeline, { x_buf, r_buf, p_buf, c_buf, d_buf }, pc, { n_embd, n_tokens, 1 });
}

static void ggml_vk_mul_mat(ggml_backend_vk_context * ctx, vk_context& subctx, const struct ggml_cgraph * cgraph, int node_idx) {
ggml_tensor * dst = cgraph->nodes[node_idx];
ggml_tensor * src0 = dst->src[0];
Expand Down Expand Up @@ -15588,6 +15800,18 @@ static bool ggml_vk_build_graph(ggml_backend_vk_context * ctx, ggml_cgraph * cgr
case GGML_OP_CUMSUM:
ggml_vk_cumsum(ctx, compute_ctx, src0, node);

break;
case GGML_OP_DSV4_HC_COMB:
ggml_vk_dsv4_hc_comb(ctx, compute_ctx, src0, src1, src2, node);

break;
case GGML_OP_DSV4_HC_PRE:
ggml_vk_dsv4_hc_pre(ctx, compute_ctx, src0, src1, node);

break;
case GGML_OP_DSV4_HC_POST:
ggml_vk_dsv4_hc_post(ctx, compute_ctx, src0, src1, src2, src3, node);

break;
case GGML_OP_MEAN:
ggml_vk_mean(ctx, compute_ctx, src0, node);
Expand Down Expand Up @@ -18399,6 +18623,40 @@ static bool ggml_backend_vk_device_supports_op(ggml_backend_dev_t dev, const ggm
}
return false;
}
case GGML_OP_DSV4_HC_COMB:
case GGML_OP_DSV4_HC_PRE:
case GGML_OP_DSV4_HC_POST:
{
const bool enabled =
op->op == GGML_OP_DSV4_HC_COMB ? device->dsv4_hc_comb :
op->op == GGML_OP_DSV4_HC_PRE ? device->dsv4_hc_pre :
device->dsv4_hc_post;
if (!enabled || op->type != GGML_TYPE_F32) {
return false;
}
for (uint32_t i = 0; i < GGML_MAX_SRC; ++i) {
if (op->src[i] && op->src[i]->type != GGML_TYPE_F32) {
return false;
}
}
if (!ggml_vk_dsv4_hc_strides_ok(op)) {
return false;
}
// hc is hardcoded to 4 in the shaders. ggml only constrains it
// to 4 for COMB, so PRE/POST have to be checked here.
if (op->op == GGML_OP_DSV4_HC_PRE && op->src[0]->ne[1] != 4) {
return false;
}
if (op->op == GGML_OP_DSV4_HC_POST && op->src[1]->ne[1] != 4) {
return false;
}
if (op->op == GGML_OP_DSV4_HC_COMB) {
return device->pipeline_dsv4_hc_comb_f32 != nullptr;
}
// pre/post launch one workgroup row per token
const uint32_t n_tokens = (uint32_t)(op->op == GGML_OP_DSV4_HC_PRE ? op->src[0]->ne[2] : op->src[0]->ne[1]);
return n_tokens <= device->properties.limits.maxComputeWorkGroupCount[1];
}
case GGML_OP_SOLVE_TRI:
{
if (op->type != GGML_TYPE_F32 || op->src[0]->type != GGML_TYPE_F32) {
Expand Down Expand Up @@ -19339,6 +19597,13 @@ static void ggml_vk_check_results_0(ggml_backend_vk_context * ctx, ggml_cgraph *
tensor_clone = ggml_sum_rows(ggml_ctx, src_clone[0]);
} else if (tensor->op == GGML_OP_CUMSUM) {
tensor_clone = ggml_cumsum(ggml_ctx, src_clone[0]);
} else if (tensor->op == GGML_OP_DSV4_HC_COMB) {
tensor_clone = ggml_dsv4_hc_comb(ggml_ctx, src_clone[0], src_clone[1], src_clone[2],
ggml_get_op_params_f32(tensor, 0), ggml_get_op_params_i32(tensor, 1));
} else if (tensor->op == GGML_OP_DSV4_HC_PRE) {
tensor_clone = ggml_dsv4_hc_pre(ggml_ctx, src_clone[0], src_clone[1]);
} else if (tensor->op == GGML_OP_DSV4_HC_POST) {
tensor_clone = ggml_dsv4_hc_post(ggml_ctx, src_clone[0], src_clone[1], src_clone[2], src_clone[3]);
} else if (tensor->op == GGML_OP_MEAN) {
tensor_clone = ggml_mean(ggml_ctx, src_clone[0]);
} else if (tensor->op == GGML_OP_ARGMAX) {
Expand Down
Loading