@@ -1042,6 +1042,8 @@ struct vk_device_struct {
10421042 vk_pipeline pipeline_argsort_large_f32[num_argsort_pipelines];
10431043 vk_pipeline pipeline_topk_f32[num_topk_pipelines];
10441044 vk_pipeline pipeline_sum_rows_f32;
1045+ vk_pipeline pipeline_cross_entropy_loss_f32, pipeline_cross_entropy_loss_f32_wg512;
1046+ vk_pipeline pipeline_cross_entropy_loss_back_f32, pipeline_cross_entropy_loss_back_f32_wg512;
10451047 vk_pipeline pipeline_fwht_f32[4];
10461048 vk_pipeline pipeline_cumsum_f32;
10471049 vk_pipeline pipeline_cumsum_small_f32;
@@ -5758,6 +5760,10 @@ static void ggml_vk_load_shaders(vk_device& device, vk_pipeline requested) {
57585760 ggml_vk_create_pipeline(device, device->pipeline_argmax_f32, "argmax_f32", argmax_f32_len, argmax_f32_data, "main", 2, sizeof(vk_op_push_constants), {1, 1, 1}, { device->subgroup_size }, 1);
57595761
57605762 ggml_vk_create_pipeline(device, device->pipeline_sum_rows_f32, "sum_rows_f32", sum_rows_f32_len, sum_rows_f32_data, "main", 2, sizeof(vk_op_sum_rows_push_constants), {1, 1, 1}, { device->subgroup_size }, 1);
5763+ ggml_vk_create_pipeline(device, device->pipeline_cross_entropy_loss_f32, "cross_entropy_loss_f32", cross_entropy_loss_f32_len, cross_entropy_loss_f32_data, "main", 3, sizeof(vk_op_push_constants), {1, 1, 1}, { device->subgroup_size }, 1);
5764+ ggml_vk_create_pipeline(device, device->pipeline_cross_entropy_loss_f32_wg512, "cross_entropy_loss_f32_wg512", cross_entropy_loss_f32_len, cross_entropy_loss_f32_data, "main", 3, sizeof(vk_op_push_constants), {1, 1, 1}, { 512 }, 1);
5765+ ggml_vk_create_pipeline(device, device->pipeline_cross_entropy_loss_back_f32, "cross_entropy_loss_back_f32", cross_entropy_loss_back_f32_len, cross_entropy_loss_back_f32_data, "main", 4, sizeof(vk_op_push_constants), {1, 1, 1}, { device->subgroup_size }, 1);
5766+ ggml_vk_create_pipeline(device, device->pipeline_cross_entropy_loss_back_f32_wg512, "cross_entropy_loss_back_f32_wg512", cross_entropy_loss_back_f32_len, cross_entropy_loss_back_f32_data, "main", 4, sizeof(vk_op_push_constants), {1, 1, 1}, { 512 }, 1);
57615767 // Intel Windows driver in range [32.0.101.8509, 32.0.101.8860) will crash when using fwht kernels so we gate that here
57625768 const bool can_use_fwht = device->driver_id != vk::DriverId::eIntelProprietaryWindows ||
57635769 !ggml_vk_intel_windows_driver_in_range(device->properties.driverVersion, 101, 8509, 101, 8860);
@@ -11577,6 +11583,17 @@ static vk_pipeline ggml_vk_op_get_pipeline(ggml_backend_vk_context * ctx, const
1157711583 return ctx->device->pipeline_sum_rows_f32;
1157811584 }
1157911585 return nullptr;
11586+ case GGML_OP_CROSS_ENTROPY_LOSS:
11587+ if (src0->type == GGML_TYPE_F32 && src1->type == GGML_TYPE_F32 && dst->type == GGML_TYPE_F32) {
11588+ return src0->ne[0] > 1024 ? ctx->device->pipeline_cross_entropy_loss_f32_wg512 : ctx->device->pipeline_cross_entropy_loss_f32;
11589+ }
11590+ return nullptr;
11591+ case GGML_OP_CROSS_ENTROPY_LOSS_BACK:
11592+ // src0 is the scalar grad; src1 is logits
11593+ if (src0->type == GGML_TYPE_F32 && src1->type == GGML_TYPE_F32 && src2 && src2->type == GGML_TYPE_F32 && dst->type == GGML_TYPE_F32) {
11594+ return src1->ne[0] > 1024 ? ctx->device->pipeline_cross_entropy_loss_back_f32_wg512 : ctx->device->pipeline_cross_entropy_loss_back_f32;
11595+ }
11596+ return nullptr;
1158011597 case GGML_OP_CUMSUM:
1158111598 if (src0->type == GGML_TYPE_F32 && dst->type == GGML_TYPE_F32) {
1158211599 if (src0->ne[0] <= 512) {
@@ -13942,6 +13959,103 @@ static void ggml_vk_cumsum(ggml_backend_vk_context * ctx, vk_context& subctx, co
1394213959 ctx->prealloc_split_k_need_sync = true;
1394313960}
1394413961
13962+ static std::array<uint32_t, 3> ggml_vk_nrows_elements(uint32_t nr) {
13963+ if (nr > 262144) {
13964+ return { 512, 512, CEIL_DIV(nr, 262144) };
13965+ }
13966+ if (nr > 512) {
13967+ return { 512, CEIL_DIV(nr, 512), 1 };
13968+ }
13969+ return { nr, 1, 1 };
13970+ }
13971+
13972+ static void ggml_vk_cross_entropy_loss(ggml_backend_vk_context * ctx, vk_context& subctx, ggml_tensor * dst) {
13973+ const ggml_tensor * src0 = dst->src[0];
13974+ const ggml_tensor * src1 = dst->src[1];
13975+
13976+ GGML_ASSERT(src0->type == GGML_TYPE_F32);
13977+ GGML_ASSERT(src1->type == GGML_TYPE_F32);
13978+ GGML_ASSERT(dst->type == GGML_TYPE_F32);
13979+ GGML_ASSERT(ggml_is_contiguous(src0));
13980+ GGML_ASSERT(ggml_is_contiguous(src1));
13981+ GGML_ASSERT(ggml_is_contiguous(dst));
13982+ GGML_ASSERT(ggml_are_same_shape(src0, src1));
13983+ GGML_ASSERT(ggml_is_scalar(dst));
13984+
13985+ const uint32_t nclasses = (uint32_t)src0->ne[0];
13986+ const uint32_t nrows = (uint32_t)ggml_nrows(src0);
13987+
13988+ vk_pipeline pipeline = ggml_vk_op_get_pipeline(ctx, src0, src1, nullptr, dst, GGML_OP_CROSS_ENTROPY_LOSS);
13989+ GGML_ASSERT(pipeline != nullptr);
13990+
13991+ ggml_pipeline_request_descriptor_sets(ctx, pipeline, 1);
13992+ ggml_pipeline_request_descriptor_sets(ctx, ctx->device->pipeline_sum_rows_f32, 1);
13993+
13994+ vk_subbuffer src0_buf = ggml_vk_tensor_subbuffer(ctx, src0);
13995+ vk_subbuffer src1_buf = ggml_vk_tensor_subbuffer(ctx, src1);
13996+ vk_subbuffer dst_buf = ggml_vk_tensor_subbuffer(ctx, dst, true);
13997+
13998+ const vk_op_push_constants pc = { nclasses, nrows, 0.0f, 0.0f, 0.0f, 0.0f };
13999+
14000+ const size_t tmp_size = (size_t)nrows * sizeof(float);
14001+ if (ctx->prealloc_size_x < tmp_size) {
14002+ ctx->prealloc_size_x = tmp_size;
14003+ ggml_vk_preallocate_buffers(ctx, subctx);
14004+ }
14005+ if (ctx->prealloc_x_need_sync) {
14006+ ggml_vk_sync_buffers(ctx, subctx);
14007+ }
14008+
14009+ vk_subbuffer tmp_buf = { ctx->prealloc_x, 0, tmp_size };
14010+ ggml_vk_dispatch_pipeline(ctx, subctx, pipeline, { src0_buf, src1_buf, tmp_buf }, pc, ggml_vk_nrows_elements(nrows));
14011+ ggml_vk_sync_buffers(ctx, subctx);
14012+
14013+ vk_op_sum_rows_push_constants sp = {};
14014+ sp.n_cols = nrows;
14015+ sp.ne01 = 1;
14016+ sp.ne02 = 1;
14017+ sp.weight = 1.0f;
14018+ init_pushconst_fastdiv(sp);
14019+ sp.misalign_offsets = get_misalign_bytes(ctx, dst) / ggml_type_size(dst->type);
14020+
14021+ ggml_vk_dispatch_pipeline(ctx, subctx, ctx->device->pipeline_sum_rows_f32, { tmp_buf, dst_buf }, sp, { 1, 1, 1 });
14022+ ctx->prealloc_x_need_sync = true;
14023+ }
14024+
14025+ static void ggml_vk_cross_entropy_loss_back(ggml_backend_vk_context * ctx, vk_context& subctx, ggml_tensor * dst) {
14026+ const ggml_tensor * grad = dst->src[0];
14027+ const ggml_tensor * logits = dst->src[1];
14028+ const ggml_tensor * labels = dst->src[2];
14029+
14030+ GGML_ASSERT(grad->type == GGML_TYPE_F32);
14031+ GGML_ASSERT(logits->type == GGML_TYPE_F32);
14032+ GGML_ASSERT(labels->type == GGML_TYPE_F32);
14033+ GGML_ASSERT(dst->type == GGML_TYPE_F32);
14034+ GGML_ASSERT(ggml_is_scalar(grad));
14035+ GGML_ASSERT(ggml_is_contiguous(grad));
14036+ GGML_ASSERT(ggml_is_contiguous(logits));
14037+ GGML_ASSERT(ggml_is_contiguous(labels));
14038+ GGML_ASSERT(ggml_is_contiguous(dst));
14039+ GGML_ASSERT(ggml_are_same_shape(logits, labels));
14040+ GGML_ASSERT(ggml_are_same_shape(logits, dst));
14041+
14042+ const uint32_t nclasses = (uint32_t)logits->ne[0];
14043+ const uint32_t nrows = (uint32_t)ggml_nrows(logits);
14044+
14045+ vk_pipeline pipeline = ggml_vk_op_get_pipeline(ctx, grad, logits, labels, dst, GGML_OP_CROSS_ENTROPY_LOSS_BACK);
14046+ GGML_ASSERT(pipeline != nullptr);
14047+
14048+ ggml_pipeline_request_descriptor_sets(ctx, pipeline, 1);
14049+
14050+ vk_subbuffer grad_buf = ggml_vk_tensor_subbuffer(ctx, grad);
14051+ vk_subbuffer logits_buf = ggml_vk_tensor_subbuffer(ctx, logits);
14052+ vk_subbuffer labels_buf = ggml_vk_tensor_subbuffer(ctx, labels);
14053+ vk_subbuffer dst_buf = ggml_vk_tensor_subbuffer(ctx, dst);
14054+
14055+ const vk_op_push_constants pc = { nclasses, nrows, 0.0f, 0.0f, 0.0f, 0.0f };
14056+ ggml_vk_dispatch_pipeline(ctx, subctx, pipeline, { grad_buf, logits_buf, labels_buf, dst_buf }, pc, ggml_vk_nrows_elements(nrows));
14057+ }
14058+
1394514059static void ggml_vk_argmax(ggml_backend_vk_context * ctx, vk_context& subctx, const ggml_tensor * src0, ggml_tensor * dst) {
1394614060 ggml_vk_op_f32<vk_op_push_constants>(ctx, subctx, src0, nullptr, nullptr, nullptr, dst, GGML_OP_ARGMAX, { (uint32_t)src0->ne[0], (uint32_t)src0->ne[1], 0.0f, 0.0f, 0.0f, 0.0f });
1394714061}
@@ -15687,6 +15801,14 @@ static bool ggml_vk_build_graph(ggml_backend_vk_context * ctx, ggml_cgraph * cgr
1568715801 case GGML_OP_ARGMAX:
1568815802 ggml_vk_argmax(ctx, compute_ctx, src0, node);
1568915803
15804+ break;
15805+ case GGML_OP_CROSS_ENTROPY_LOSS:
15806+ ggml_vk_cross_entropy_loss(ctx, compute_ctx, node);
15807+
15808+ break;
15809+ case GGML_OP_CROSS_ENTROPY_LOSS_BACK:
15810+ ggml_vk_cross_entropy_loss_back(ctx, compute_ctx, node);
15811+
1569015812 break;
1569115813 case GGML_OP_COUNT_EQUAL:
1569215814 ggml_vk_count_equal(ctx, compute_ctx, src0, src1, node);
@@ -18511,6 +18633,18 @@ static bool ggml_backend_vk_device_supports_op(ggml_backend_dev_t dev, const ggm
1851118633 }
1851218634 case GGML_OP_ARGMAX:
1851318635 return ggml_is_contiguous(op->src[0]) && op->src[0]->type == GGML_TYPE_F32;
18636+ case GGML_OP_CROSS_ENTROPY_LOSS:
18637+ return ggml_is_contiguous(op->src[0]) && op->src[0]->type == GGML_TYPE_F32
18638+ && ggml_is_contiguous(op->src[1]) && op->src[1]->type == GGML_TYPE_F32
18639+ && ggml_are_same_shape(op->src[0], op->src[1])
18640+ && ggml_is_contiguous(op) && ggml_is_scalar(op) && op->type == GGML_TYPE_F32;
18641+ case GGML_OP_CROSS_ENTROPY_LOSS_BACK:
18642+ return ggml_is_contiguous(op->src[0]) && op->src[0]->type == GGML_TYPE_F32 && ggml_is_scalar(op->src[0])
18643+ && ggml_is_contiguous(op->src[1]) && op->src[1]->type == GGML_TYPE_F32
18644+ && ggml_is_contiguous(op->src[2]) && op->src[2]->type == GGML_TYPE_F32
18645+ && ggml_are_same_shape(op->src[1], op->src[2])
18646+ && ggml_are_same_shape(op->src[1], op)
18647+ && ggml_is_contiguous(op) && op->type == GGML_TYPE_F32;
1851418648 case GGML_OP_COUNT_EQUAL:
1851518649 return ggml_is_contiguous(op->src[0]) && op->src[0]->type == GGML_TYPE_I32
1851618650 && ggml_is_contiguous(op->src[1]) && op->src[1]->type == GGML_TYPE_I32;
@@ -19437,6 +19571,10 @@ static void ggml_vk_check_results_0(ggml_backend_vk_context * ctx, ggml_cgraph *
1943719571 tensor_clone = ggml_mean(ggml_ctx, src_clone[0]);
1943819572 } else if (tensor->op == GGML_OP_ARGMAX) {
1943919573 tensor_clone = ggml_argmax(ggml_ctx, src_clone[0]);
19574+ } else if (tensor->op == GGML_OP_CROSS_ENTROPY_LOSS) {
19575+ tensor_clone = ggml_cross_entropy_loss(ggml_ctx, src_clone[0], src_clone[1]);
19576+ } else if (tensor->op == GGML_OP_CROSS_ENTROPY_LOSS_BACK) {
19577+ tensor_clone = ggml_cross_entropy_loss_back(ggml_ctx, src_clone[0], src_clone[1], src_clone[2]);
1944019578 } else if (tensor->op == GGML_OP_COUNT_EQUAL) {
1944119579 tensor_clone = ggml_count_equal(ggml_ctx, src_clone[0], src_clone[1]);
1944219580 } else if (tensor->op == GGML_OP_SOLVE_TRI) {
0 commit comments