Skip to content

Commit bf94216

Browse files
Implemented vulkan cross_entropy_loss and cross_entropy_loss_back (ggml-org#27216)
1 parent d0132a6 commit bf94216

6 files changed

Lines changed: 299 additions & 6 deletions

File tree

‎docs/ops.md‎

Lines changed: 2 additions & 2 deletions
Original file line numberDiff line numberDiff line change
@@ -35,8 +35,8 @@ Legend:
3535
| COS | ❌ | ✅ | ✅ | ✅ | ❌ | ✅ | ❌ | ✅ | ✅ | ✅ | ❌ | ❌ |
3636
| COUNT_EQUAL | ❌ | ✅ | ✅ | ✅ | ❌ | ✅ | ❌ | ✅ | ✅ | ❌ | ❌ | ❌ |
3737
| CPY | ❌ | 🟡 | 🟡 | 🟡 | 🟡 | 🟡 | 🟡 | 🟡 | 🟡 | 🟡 | ❌ | ❌ |
38-
| CROSS_ENTROPY_LOSS | ❌ | ✅ | ✅ | ✅ | ❌ | ❌ | ❌ | ✅ | ❌ | ❌ | ❌ | ❌ |
39-
| CROSS_ENTROPY_LOSS_BACK | ❌ | ❌ | ✅ | ✅ | ❌ | ❌ | ❌ | ✅ | ❌ | ❌ | ❌ | ❌ |
38+
| CROSS_ENTROPY_LOSS | ❌ | ✅ | ✅ | ✅ | ❌ | ❌ | ❌ | ✅ | ✅ | ❌ | ❌ | ❌ |
39+
| CROSS_ENTROPY_LOSS_BACK | ❌ | ❌ | ✅ | ✅ | ❌ | ❌ | ❌ | ✅ | ✅ | ❌ | ❌ | ❌ |
4040
| CUMSUM | ❌ | ❌ | ✅ | ✅ | ✅ | ✅ | ❌ | ✅ | ✅ | ✅ | ❌ | ❌ |
4141
| DIAG | ❌ | ❌ | ✅ | ✅ | 🟡 | ✅ | ❌ | ✅ | ✅ | ✅ | ❌ | ❌ |
4242
| DIAG_MASK_INF | ❌ | ✅ | ✅ | ✅ | ❌ | ❌ | 🟡 | ✅ | ✅ | ❌ | ❌ | ❌ |

‎docs/ops/Vulkan.csv‎

Lines changed: 4 additions & 4 deletions
Original file line numberDiff line numberDiff line change
@@ -19292,10 +19292,10 @@
1929219292
"Vulkan0","FLASH_ATTN_EXT","hsk=128,hsv=128,nh=8,nr23=[4,1],kv=4096,nb=512,mask=1,sinks=0,max_bias=0.000000,logit_softcap=0.000000,prec=f32,type_K=f16,type_V=f16,permute=[0,1,2,3]","support","1","yes","Vulkan"
1929319293
"Vulkan0","FLASH_ATTN_EXT","hsk=256,hsv=256,nh=4,nr23=[6,1],kv=16384,nb=512,mask=1,sinks=0,max_bias=0.000000,logit_softcap=0.000000,prec=f32,type_K=f16,type_V=f16,permute=[0,1,2,3]","support","1","yes","Vulkan"
1929419294
"Vulkan0","FLASH_ATTN_EXT","hsk=128,hsv=128,nh=8,nr23=[4,1],kv=16384,nb=512,mask=1,sinks=0,max_bias=0.000000,logit_softcap=0.000000,prec=f32,type_K=f16,type_V=f16,permute=[0,1,2,3]","support","1","yes","Vulkan"
19295-
"Vulkan0","CROSS_ENTROPY_LOSS","type=f32,ne=[10,5,4,3]","support","0","no","Vulkan"
19296-
"Vulkan0","CROSS_ENTROPY_LOSS","type=f32,ne=[30000,1,1,1]","support","0","no","Vulkan"
19297-
"Vulkan0","CROSS_ENTROPY_LOSS_BACK","type=f32,ne=[10,5,4,3]","support","0","no","Vulkan"
19298-
"Vulkan0","CROSS_ENTROPY_LOSS_BACK","type=f32,ne=[30000,1,1,1]","support","0","no","Vulkan"
19295+
"Vulkan0","CROSS_ENTROPY_LOSS","type=f32,ne=[10,5,4,3]","support","1","yes","Vulkan"
19296+
"Vulkan0","CROSS_ENTROPY_LOSS","type=f32,ne=[30000,1,1,1]","support","1","yes","Vulkan"
19297+
"Vulkan0","CROSS_ENTROPY_LOSS_BACK","type=f32,ne=[10,5,4,3]","support","1","yes","Vulkan"
19298+
"Vulkan0","CROSS_ENTROPY_LOSS_BACK","type=f32,ne=[30000,1,1,1]","support","1","yes","Vulkan"
1929919299
"Vulkan0","OPT_STEP_ADAMW","type=f32,ne=[10,5,4,3]","support","1","yes","Vulkan"
1930019300
"Vulkan0","OPT_STEP_SGD","type=f32,ne=[10,5,4,3]","support","1","yes","Vulkan"
1930119301
"Vulkan0","GATED_DELTA_NET","type=f32,head_count=32,head_size=128,n_seq_tokens=1,n_seqs=1,v_repeat=1,permuted=0,kda=0,K=1","support","1","yes","Vulkan"

‎ggml/src/ggml-vulkan/ggml-vulkan.cpp‎

Lines changed: 138 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -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+
1394514059
static 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) {
Lines changed: 78 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -0,0 +1,78 @@
1+
#version 450
2+
3+
#include "generic_head.glsl"
4+
#include "types.glsl"
5+
6+
#extension GL_EXT_control_flow_attributes : enable
7+
8+
layout(constant_id = 0) const uint BLOCK_SIZE = 32;
9+
layout(local_size_x_id = 0, local_size_y = 1, local_size_z = 1) in;
10+
11+
layout (binding = 0) readonly buffer A {A_TYPE data_a[];};
12+
layout (binding = 1) readonly buffer B {B_TYPE data_b[];};
13+
layout (binding = 2) writeonly buffer D {D_TYPE data_d[];};
14+
15+
shared FLOAT_TYPE tmp[BLOCK_SIZE];
16+
17+
FLOAT_TYPE wg_reduce_max(FLOAT_TYPE v) {
18+
const uint tid = gl_LocalInvocationID.x;
19+
tmp[tid] = v;
20+
barrier();
21+
[[unroll]] for (uint s = BLOCK_SIZE / 2; s > 0; s >>= 1) {
22+
if (tid < s) {
23+
tmp[tid] = max(tmp[tid], tmp[tid + s]);
24+
}
25+
barrier();
26+
}
27+
v = tmp[0];
28+
barrier();
29+
return v;
30+
}
31+
32+
FLOAT_TYPE wg_reduce_sum(FLOAT_TYPE v) {
33+
const uint tid = gl_LocalInvocationID.x;
34+
tmp[tid] = v;
35+
barrier();
36+
[[unroll]] for (uint s = BLOCK_SIZE / 2; s > 0; s >>= 1) {
37+
if (tid < s) {
38+
tmp[tid] += tmp[tid + s];
39+
}
40+
barrier();
41+
}
42+
v = tmp[0];
43+
barrier();
44+
return v;
45+
}
46+
47+
void main() {
48+
const uint row = gl_WorkGroupID.z * 262144 + gl_WorkGroupID.y * 512 + gl_WorkGroupID.x;
49+
const uint tid = gl_LocalInvocationID.x;
50+
51+
if (row >= p.KY) {
52+
return;
53+
}
54+
55+
const uint off = row * p.KX;
56+
57+
FLOAT_TYPE max_logit = FLOAT_TYPE(uintBitsToFloat(0xFF800000));
58+
for (uint i = tid; i < p.KX; i += BLOCK_SIZE) {
59+
max_logit = max(max_logit, FLOAT_TYPE(data_a[off + i]));
60+
}
61+
max_logit = wg_reduce_max(max_logit);
62+
63+
FLOAT_TYPE sum_exp = FLOAT_TYPE(0.0f);
64+
for (uint i = tid; i < p.KX; i += BLOCK_SIZE) {
65+
sum_exp += exp(FLOAT_TYPE(data_a[off + i]) - max_logit);
66+
}
67+
const FLOAT_TYPE log_sum = log(wg_reduce_sum(sum_exp));
68+
69+
FLOAT_TYPE loss = FLOAT_TYPE(0.0f);
70+
for (uint i = tid; i < p.KX; i += BLOCK_SIZE) {
71+
loss += (FLOAT_TYPE(data_a[off + i]) - max_logit - log_sum) * FLOAT_TYPE(data_b[off + i]);
72+
}
73+
loss = -wg_reduce_sum(loss) / FLOAT_TYPE(p.KY);
74+
75+
if (tid == 0) {
76+
data_d[row] = D_TYPE(loss);
77+
}
78+
}
Lines changed: 75 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -0,0 +1,75 @@
1+
#version 450
2+
3+
#include "generic_head.glsl"
4+
#include "types.glsl"
5+
6+
#extension GL_EXT_control_flow_attributes : enable
7+
8+
layout(constant_id = 0) const uint BLOCK_SIZE = 32;
9+
layout(local_size_x_id = 0, local_size_y = 1, local_size_z = 1) in;
10+
11+
layout (binding = 0) readonly buffer G {A_TYPE data_g[];};
12+
layout (binding = 1) readonly buffer X {B_TYPE data_x[];};
13+
layout (binding = 2) readonly buffer Y {B_TYPE data_y[];};
14+
layout (binding = 3) writeonly buffer D {D_TYPE data_d[];};
15+
16+
shared FLOAT_TYPE tmp[BLOCK_SIZE];
17+
18+
FLOAT_TYPE wg_reduce_max(FLOAT_TYPE v) {
19+
const uint tid = gl_LocalInvocationID.x;
20+
tmp[tid] = v;
21+
barrier();
22+
[[unroll]] for (uint s = BLOCK_SIZE / 2; s > 0; s >>= 1) {
23+
if (tid < s) {
24+
tmp[tid] = max(tmp[tid], tmp[tid + s]);
25+
}
26+
barrier();
27+
}
28+
v = tmp[0];
29+
barrier();
30+
return v;
31+
}
32+
33+
FLOAT_TYPE wg_reduce_sum(FLOAT_TYPE v) {
34+
const uint tid = gl_LocalInvocationID.x;
35+
tmp[tid] = v;
36+
barrier();
37+
[[unroll]] for (uint s = BLOCK_SIZE / 2; s > 0; s >>= 1) {
38+
if (tid < s) {
39+
tmp[tid] += tmp[tid + s];
40+
}
41+
barrier();
42+
}
43+
v = tmp[0];
44+
barrier();
45+
return v;
46+
}
47+
48+
void main() {
49+
const uint row = gl_WorkGroupID.z * 262144 + gl_WorkGroupID.y * 512 + gl_WorkGroupID.x;
50+
const uint tid = gl_LocalInvocationID.x;
51+
52+
if (row >= p.KY) {
53+
return;
54+
}
55+
56+
const uint off = row * p.KX;
57+
const FLOAT_TYPE d_by_nrows = FLOAT_TYPE(data_g[0]) / FLOAT_TYPE(p.KY);
58+
59+
FLOAT_TYPE max_logit = FLOAT_TYPE(uintBitsToFloat(0xFF800000));
60+
for (uint i = tid; i < p.KX; i += BLOCK_SIZE) {
61+
max_logit = max(max_logit, FLOAT_TYPE(data_x[off + i]));
62+
}
63+
max_logit = wg_reduce_max(max_logit);
64+
65+
FLOAT_TYPE sum_exp = FLOAT_TYPE(0.0f);
66+
for (uint i = tid; i < p.KX; i += BLOCK_SIZE) {
67+
sum_exp += exp(FLOAT_TYPE(data_x[off + i]) - max_logit);
68+
}
69+
const FLOAT_TYPE inv_sum = FLOAT_TYPE(1.0f) / wg_reduce_sum(sum_exp);
70+
71+
for (uint i = tid; i < p.KX; i += BLOCK_SIZE) {
72+
const FLOAT_TYPE sm = exp(FLOAT_TYPE(data_x[off + i]) - max_logit) * inv_sum;
73+
data_d[off + i] = D_TYPE((sm - FLOAT_TYPE(data_y[off + i])) * d_by_nrows);
74+
}
75+
}

‎ggml/src/ggml-vulkan/vulkan-shaders/vulkan-shaders-gen.cpp‎

Lines changed: 2 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -1029,6 +1029,8 @@ void process_shaders() {
10291029

10301030
string_to_spv("argmax_f32", "argmax.comp", merge_maps(base_dict, {{"A_TYPE", "float"}, {"D_TYPE", "int"}}));
10311031
string_to_spv("sum_rows_f32", "sum_rows.comp", merge_maps(base_dict, {{"A_TYPE", "float"}, {"D_TYPE", "float"}}));
1032+
string_to_spv("cross_entropy_loss_f32", "cross_entropy_loss.comp", merge_maps(base_dict, {{"A_TYPE", "float"}, {"B_TYPE", "float"}, {"D_TYPE", "float"}}));
1033+
string_to_spv("cross_entropy_loss_back_f32", "cross_entropy_loss_back.comp", merge_maps(base_dict, {{"A_TYPE", "float"}, {"B_TYPE", "float"}, {"D_TYPE", "float"}}));
10321034
string_to_spv("fwht_f32", "fwht.comp", {});
10331035
string_to_spv("fwht_shmem_f32", "fwht.comp", {{"FWHT_SHMEM", "1"}});
10341036
string_to_spv("count_equal_i32", "count_equal.comp", merge_maps(base_dict, {{"A_TYPE", "int"}, {"B_TYPE", "int"}, {"D_TYPE", "int"}}));

0 commit comments

Comments
 (0)