diff --git a/ggml/include/ggml.h b/ggml/include/ggml.h index d6807b6dd47a..aa238f6bc56f 100644 --- a/ggml/include/ggml.h +++ b/ggml/include/ggml.h @@ -568,6 +568,10 @@ extern "C" { GGML_OP_RWKV_WKV7, GGML_OP_SOLVE_TRI, GGML_OP_GATED_DELTA_NET, + GGML_OP_LIGHTNING_INDEXER, + GGML_OP_DSV4_HC_COMB, + GGML_OP_DSV4_HC_PRE, + GGML_OP_DSV4_HC_POST, GGML_OP_UNARY, @@ -2573,6 +2577,40 @@ extern "C" { struct ggml_tensor * state, int64_t K); + GGML_API struct ggml_tensor * ggml_lightning_indexer( + struct ggml_context * ctx, + struct ggml_tensor * q, + struct ggml_tensor * k, + struct ggml_tensor * weights, + float scale_embd, + float scale_heads); + + // DeepSeek V4 hyper-connection helpers. + // hc_comb: mixes [(2 + hc)*hc, n_tokens], scale [3], base [(2 + hc)*hc] + // -> [dst_hc, src_hc, n_tokens] + // hc_pre : x [n_embd, hc, n_tokens], weights [hc, n_tokens] -> [n_embd, n_tokens] + // hc_post: x [n_embd, n_tokens], residual [n_embd, hc, n_tokens], + // post [hc, n_tokens], comb [dst_hc, src_hc, n_tokens] -> [n_embd, hc, n_tokens] + GGML_API struct ggml_tensor * ggml_dsv4_hc_comb( + struct ggml_context * ctx, + struct ggml_tensor * mixes, + struct ggml_tensor * scale, + struct ggml_tensor * base, + float eps, + int32_t n_iter); + + GGML_API struct ggml_tensor * ggml_dsv4_hc_pre( + struct ggml_context * ctx, + struct ggml_tensor * x, + struct ggml_tensor * weights); + + GGML_API struct ggml_tensor * ggml_dsv4_hc_post( + struct ggml_context * ctx, + struct ggml_tensor * x, + struct ggml_tensor * residual, + struct ggml_tensor * post, + struct ggml_tensor * comb); + // custom operators typedef void (*ggml_custom1_op_t)(struct ggml_tensor * dst , const struct ggml_tensor * a, int ith, int nth, void * userdata); diff --git a/ggml/src/ggml-backend-meta.cpp b/ggml/src/ggml-backend-meta.cpp index 0a36f099000f..e09362af1396 100644 --- a/ggml/src/ggml-backend-meta.cpp +++ b/ggml/src/ggml-backend-meta.cpp @@ -984,6 +984,11 @@ static struct ggml_backend_meta_split_state ggml_backend_meta_get_split_state( case GGML_OP_GATED_DELTA_NET: { split_state = handle_gated_delta_net(src_ss); } break; + case GGML_OP_DSV4_HC_COMB: + case GGML_OP_DSV4_HC_PRE: + case GGML_OP_DSV4_HC_POST: { + split_state = handle_generic(src_ss, /*scalar_only =*/ true); + } break; case GGML_OP_UNARY: { split_state = handle_generic(src_ss, /*scalar_only =*/ false); } break; diff --git a/ggml/src/ggml-cpu/ggml-cpu.c b/ggml/src/ggml-cpu/ggml-cpu.c index eb8341c9aecc..80e67d85c007 100644 --- a/ggml/src/ggml-cpu/ggml-cpu.c +++ b/ggml/src/ggml-cpu/ggml-cpu.c @@ -2051,6 +2051,22 @@ static void ggml_compute_forward(struct ggml_compute_params * params, struct ggm { ggml_compute_forward_gated_delta_net(params, tensor); } break; + case GGML_OP_LIGHTNING_INDEXER: + { + ggml_compute_forward_lightning_indexer(params, tensor); + } break; + case GGML_OP_DSV4_HC_COMB: + { + ggml_compute_forward_dsv4_hc_comb(params, tensor); + } break; + case GGML_OP_DSV4_HC_PRE: + { + ggml_compute_forward_dsv4_hc_pre(params, tensor); + } break; + case GGML_OP_DSV4_HC_POST: + { + ggml_compute_forward_dsv4_hc_post(params, tensor); + } break; case GGML_OP_MAP_CUSTOM1: { ggml_compute_forward_map_custom1(params, tensor); @@ -2231,6 +2247,9 @@ static int ggml_get_n_tasks(struct ggml_tensor * node, int n_threads) { case GGML_OP_COUNT_EQUAL: case GGML_OP_SOLVE_TRI: case GGML_OP_GATED_DELTA_NET: + case GGML_OP_DSV4_HC_COMB: + case GGML_OP_DSV4_HC_PRE: + case GGML_OP_DSV4_HC_POST: { n_tasks = n_threads; } break; @@ -2371,6 +2390,7 @@ static int ggml_get_n_tasks(struct ggml_tensor * node, int n_threads) { case GGML_OP_FLASH_ATTN_BACK: case GGML_OP_SSM_CONV: case GGML_OP_SSM_SCAN: + case GGML_OP_LIGHTNING_INDEXER: { n_tasks = n_threads; } break; @@ -2956,6 +2976,12 @@ struct ggml_cplan ggml_graph_plan( { GGML_ABORT("fatal error"); } + case GGML_OP_LIGHTNING_INDEXER: + { + // temp buffer for dequantizing lightning indexer keys + const int64_t ne10 = node->src[1]->ne[0]; + cur += sizeof(float)*ne10*n_tasks; + } break; default: break; } diff --git a/ggml/src/ggml-cpu/ops.cpp b/ggml/src/ggml-cpu/ops.cpp index 6724686b8ae2..02c87bb64a50 100644 --- a/ggml/src/ggml-cpu/ops.cpp +++ b/ggml/src/ggml-cpu/ops.cpp @@ -10823,6 +10823,291 @@ void ggml_compute_forward_gated_delta_net( } } + +// ggml_compute_forward_dsv4_hc_comb + +static void ggml_dsv4_hc_comb_norm_cols(float * comb, float eps) { + constexpr int64_t hc = 4; + + for (int64_t idst = 0; idst < hc; ++idst) { + float sum = eps; + for (int64_t isrc = 0; isrc < hc; ++isrc) { + sum += comb[idst + hc*isrc]; + } + + const float inv_sum = 1.0f / sum; + for (int64_t isrc = 0; isrc < hc; ++isrc) { + comb[idst + hc*isrc] *= inv_sum; + } + } +} + +static void ggml_dsv4_hc_comb_norm_rows(float * comb, float eps) { + constexpr int64_t hc = 4; + + for (int64_t isrc = 0; isrc < hc; ++isrc) { + float sum = eps; + for (int64_t idst = 0; idst < hc; ++idst) { + sum += comb[idst + hc*isrc]; + } + + const float inv_sum = 1.0f / sum; + for (int64_t idst = 0; idst < hc; ++idst) { + comb[idst + hc*isrc] *= inv_sum; + } + } +} + +static void ggml_compute_forward_dsv4_hc_comb_f32( + const ggml_compute_params * params, + ggml_tensor * dst) { + const ggml_tensor * mixes = dst->src[0]; + const ggml_tensor * scale = dst->src[1]; + const ggml_tensor * base = dst->src[2]; + + GGML_ASSERT(mixes->type == GGML_TYPE_F32); + GGML_ASSERT(scale->type == GGML_TYPE_F32); + GGML_ASSERT(base->type == GGML_TYPE_F32); + GGML_ASSERT(dst->type == GGML_TYPE_F32); + + constexpr int64_t hc = 4; + constexpr int64_t comb_offset = 2*hc; + constexpr int64_t hc_mix_dim = (2 + hc)*hc; + + const int64_t n_tokens = mixes->ne[1]; + + GGML_ASSERT(mixes->ne[0] == hc_mix_dim); + GGML_ASSERT(dst->ne[0] == hc); + GGML_ASSERT(dst->ne[1] == hc); + GGML_ASSERT(dst->ne[2] == n_tokens); + GGML_ASSERT(scale->ne[0] >= 3); + GGML_ASSERT(base->ne[0] == hc_mix_dim); + + GGML_TENSOR_LOCALS(size_t, nbm, mixes, nb); + GGML_TENSOR_LOCALS(size_t, nbs, scale, nb); + GGML_TENSOR_LOCALS(size_t, nbb, base, nb); + GGML_TENSOR_LOCALS(size_t, nbd, dst, nb); + + const float eps = ggml_get_op_params_f32(dst, 0); + const int32_t n_iter = ggml_get_op_params_i32(dst, 1); + GGML_ASSERT(n_iter > 0); + + const int ith = params->ith; + const int nth = params->nth; + + const int64_t dr = (n_tokens + nth - 1) / nth; + const int64_t it0 = dr * ith; + const int64_t it1 = MIN(it0 + dr, n_tokens); + + const float scale_comb = *(const float *) ((const char *) scale->data + 2*nbs0); + + for (int64_t it = it0; it < it1; ++it) { + float comb[hc*hc]; + + for (int64_t isrc = 0; isrc < hc; ++isrc) { + float max = -INFINITY; + for (int64_t idst = 0; idst < hc; ++idst) { + const int64_t idx = idst + hc*isrc; + const float xv = *(const float *) ((const char *) mixes->data + (comb_offset + idx)*nbm0 + it*nbm1); + const float bv = *(const float *) ((const char *) base->data + (comb_offset + idx)*nbb0); + const float v = xv * scale_comb + bv; + comb[idx] = v; + max = MAX(max, v); + } + + float sum = 0.0f; + for (int64_t idst = 0; idst < hc; ++idst) { + const int64_t idx = idst + hc*isrc; + const float v = expf(comb[idx] - max); + comb[idx] = v; + sum += v; + } + + const float inv_sum = 1.0f / sum; + for (int64_t idst = 0; idst < hc; ++idst) { + const int64_t idx = idst + hc*isrc; + comb[idx] = comb[idx] * inv_sum + eps; + } + } + + ggml_dsv4_hc_comb_norm_cols(comb, eps); + for (int32_t i = 1; i < n_iter; ++i) { + ggml_dsv4_hc_comb_norm_rows(comb, eps); + ggml_dsv4_hc_comb_norm_cols(comb, eps); + } + + for (int64_t isrc = 0; isrc < hc; ++isrc) { + for (int64_t idst = 0; idst < hc; ++idst) { + const int64_t idx = idst + hc*isrc; + *(float *) ((char *) dst->data + idst*nbd0 + isrc*nbd1 + it*nbd2) = comb[idx]; + } + } + } +} + +void ggml_compute_forward_dsv4_hc_comb( + const ggml_compute_params * params, + ggml_tensor * dst) { + const ggml_tensor * src0 = dst->src[0]; + + switch (src0->type) { + case GGML_TYPE_F32: + { + ggml_compute_forward_dsv4_hc_comb_f32(params, dst); + } break; + default: + { + GGML_ABORT("fatal error"); + } + } +} + +// ggml_compute_forward_dsv4_hc_pre + +static void ggml_compute_forward_dsv4_hc_pre_f32( + const ggml_compute_params * params, + ggml_tensor * dst) { + const ggml_tensor * x = dst->src[0]; + const ggml_tensor * weights = dst->src[1]; + + GGML_ASSERT(x->type == GGML_TYPE_F32); + GGML_ASSERT(weights->type == GGML_TYPE_F32); + GGML_ASSERT(dst->type == GGML_TYPE_F32); + + const int64_t n_embd = x->ne[0]; + const int64_t hc = x->ne[1]; + const int64_t n_tokens = x->ne[2]; + + GGML_ASSERT(dst->ne[0] == n_embd); + GGML_ASSERT(dst->ne[1] == n_tokens); + GGML_ASSERT(weights->ne[0] == hc); + GGML_ASSERT(weights->ne[1] == n_tokens); + + GGML_TENSOR_LOCALS(size_t, nbx, x, nb); + GGML_TENSOR_LOCALS(size_t, nbw, weights, nb); + GGML_TENSOR_LOCALS(size_t, nbd, dst, nb); + + const int ith = params->ith; + const int nth = params->nth; + + const int64_t nr = n_embd * n_tokens; + const int64_t dr = (nr + nth - 1) / nth; + const int64_t ir0 = dr * ith; + const int64_t ir1 = MIN(ir0 + dr, nr); + + for (int64_t ir = ir0; ir < ir1; ++ir) { + const int64_t i0 = ir % n_embd; + const int64_t it = ir / n_embd; + + float sum = 0.0f; + for (int64_t ih = 0; ih < hc; ++ih) { + const float xv = *(const float *) ((const char *) x->data + i0*nbx0 + ih*nbx1 + it*nbx2); + const float wv = *(const float *) ((const char *) weights->data + ih*nbw0 + it*nbw1); + sum += xv * wv; + } + + *(float *) ((char *) dst->data + i0*nbd0 + it*nbd1) = sum; + } +} + +void ggml_compute_forward_dsv4_hc_pre( + const ggml_compute_params * params, + ggml_tensor * dst) { + const ggml_tensor * src0 = dst->src[0]; + + switch (src0->type) { + case GGML_TYPE_F32: + { + ggml_compute_forward_dsv4_hc_pre_f32(params, dst); + } break; + default: + { + GGML_ABORT("fatal error"); + } + } +} + +// ggml_compute_forward_dsv4_hc_post + +static void ggml_compute_forward_dsv4_hc_post_f32( + const ggml_compute_params * params, + ggml_tensor * dst) { + const ggml_tensor * x = dst->src[0]; + const ggml_tensor * residual = dst->src[1]; + const ggml_tensor * post = dst->src[2]; + const ggml_tensor * comb = dst->src[3]; + + GGML_ASSERT(x->type == GGML_TYPE_F32); + GGML_ASSERT(residual->type == GGML_TYPE_F32); + GGML_ASSERT(post->type == GGML_TYPE_F32); + GGML_ASSERT(comb->type == GGML_TYPE_F32); + GGML_ASSERT(dst->type == GGML_TYPE_F32); + + const int64_t n_embd = x->ne[0]; + const int64_t n_tokens = x->ne[1]; + const int64_t hc = residual->ne[1]; + + GGML_ASSERT(dst->ne[0] == n_embd); + GGML_ASSERT(dst->ne[1] == hc); + GGML_ASSERT(dst->ne[2] == n_tokens); + GGML_ASSERT(residual->ne[0] == n_embd); + GGML_ASSERT(residual->ne[2] == n_tokens); + GGML_ASSERT(post->ne[0] == hc); + GGML_ASSERT(post->ne[1] == n_tokens); + GGML_ASSERT(comb->ne[0] == hc); + GGML_ASSERT(comb->ne[1] == hc); + GGML_ASSERT(comb->ne[2] == n_tokens); + + GGML_TENSOR_LOCALS(size_t, nbx, x, nb); + GGML_TENSOR_LOCALS(size_t, nbr, residual, nb); + GGML_TENSOR_LOCALS(size_t, nbp, post, nb); + GGML_TENSOR_LOCALS(size_t, nbc, comb, nb); + GGML_TENSOR_LOCALS(size_t, nbd, dst, nb); + + const int ith = params->ith; + const int nth = params->nth; + + const int64_t nr = n_embd * hc * n_tokens; + const int64_t dr = (nr + nth - 1) / nth; + const int64_t ir0 = dr * ith; + const int64_t ir1 = MIN(ir0 + dr, nr); + + for (int64_t ir = ir0; ir < ir1; ++ir) { + const int64_t i0 = ir % n_embd; + const int64_t idst = (ir / n_embd) % hc; + const int64_t it = ir / (n_embd * hc); + + const float xv = *(const float *) ((const char *) x->data + i0*nbx0 + it*nbx1); + const float pv = *(const float *) ((const char *) post->data + idst*nbp0 + it*nbp1); + + float sum = xv * pv; + for (int64_t isrc = 0; isrc < hc; ++isrc) { + const float rv = *(const float *) ((const char *) residual->data + i0*nbr0 + isrc*nbr1 + it*nbr2); + const float cv = *(const float *) ((const char *) comb->data + idst*nbc0 + isrc*nbc1 + it*nbc2); + sum += rv * cv; + } + + *(float *) ((char *) dst->data + i0*nbd0 + idst*nbd1 + it*nbd2) = sum; + } +} + +void ggml_compute_forward_dsv4_hc_post( + const ggml_compute_params * params, + ggml_tensor * dst) { + const ggml_tensor * src0 = dst->src[0]; + + switch (src0->type) { + case GGML_TYPE_F32: + { + ggml_compute_forward_dsv4_hc_post_f32(params, dst); + } break; + default: + { + GGML_ABORT("fatal error"); + } + } +} + // ggml_compute_forward_rwkv_wkv7 static void ggml_compute_forward_rwkv_wkv7_f32( @@ -11512,3 +11797,76 @@ void ggml_compute_forward_fwht(const ggml_compute_params * params, ggml_tensor * } } } + +// ggml_compute_forward_lightning_indexer + +void ggml_compute_forward_lightning_indexer( + const ggml_compute_params * params, + ggml_tensor * dst) { + + const ggml_tensor * src0 = dst->src[0]; // q + const ggml_tensor * src1 = dst->src[1]; // k + const ggml_tensor * src2 = dst->src[2]; // weights + + const float scale_embd = ggml_get_op_params_f32(dst, 0); + const float scale_heads = ggml_get_op_params_f32(dst, 1); + + GGML_ASSERT(dst->type == GGML_TYPE_F32); + GGML_ASSERT(src0->type == GGML_TYPE_F32); + GGML_ASSERT(src2->type == GGML_TYPE_F32); + + GGML_TENSOR_TERNARY_OP_LOCALS + + GGML_ASSERT( nb0 == sizeof(float)); + GGML_ASSERT(nb00 == sizeof(float)); + + int n_embd = src0->ne[0]; + int n_head = src0->ne[1]; + int n_batch = src0->ne[2]; + int n_stream = src0->ne[3]; + int n_kv = src1->ne[2]; + + ggml_to_float_t const k_to_float = ggml_get_type_traits(src1->type)->to_float; + GGML_ASSERT((src1->type == GGML_TYPE_F32 || k_to_float) && "lightning indexer: unsupported K-type"); + + const int nr = n_kv; + const int ith = params->ith; + const int nth = params->nth; + + // (temporary) buffer for K converted to float + float * src1_row_f32 = (float *) params->wdata + ith*(1*n_embd + CACHE_LINE_SIZE_F32); + + // rows per thread + const int dr = (nr + nth - 1)/nth; + + // row range for this thread + const int ir0 = dr*ith; + const int ir1 = MIN(ir0 + dr, nr); + + for (int i_stream = 0; i_stream < n_stream; ++i_stream) { + for (int i_batch = 0; i_batch < n_batch; ++i_batch) { + for (int i_kv = ir0; i_kv < ir1; ++i_kv) { + char * src1_row = (char *) src1->data + i_kv*nb12 + i_stream*nb13; + if (k_to_float) { + k_to_float(src1_row, src1_row_f32, n_embd); + } else { + src1_row_f32 = (float *) src1_row; + } + float * src2_row = (float *) ((char *) src2->data + i_batch*nb21 + i_stream*nb23); + float * dst_row = (float *) ((char *) dst->data + i_batch*nb1 + i_stream*nb3); + float score = 0.0f; + for (int i_head = 0; i_head < n_head; ++i_head) { + // dot product of q and k for head i_head + float qk = 0.0f; + float * src0_row = (float *) ((char *) src0->data + i_head*nb01 + i_batch*nb02 + i_stream*nb03); + ggml_vec_dot_f32(n_embd, &qk, 0, src0_row, 0, src1_row_f32, 0, 1); + qk *= scale_embd; + // ReLU and weights + score += MAX(qk, 0.0f) * src2_row[i_head]; + } + score *= scale_heads; + dst_row[i_kv] = score; + } + } + } +} diff --git a/ggml/src/ggml-cpu/ops.h b/ggml/src/ggml-cpu/ops.h index a8e18c716db7..4c1642a67603 100644 --- a/ggml/src/ggml-cpu/ops.h +++ b/ggml/src/ggml-cpu/ops.h @@ -105,6 +105,10 @@ void ggml_compute_forward_rwkv_wkv7(const struct ggml_compute_params * params, s void ggml_compute_forward_solve_tri(const struct ggml_compute_params * params, struct ggml_tensor * dst); void ggml_compute_forward_gla(const struct ggml_compute_params * params, struct ggml_tensor * dst); void ggml_compute_forward_gated_delta_net(const struct ggml_compute_params * params, struct ggml_tensor * dst); +void ggml_compute_forward_lightning_indexer(const struct ggml_compute_params * params, struct ggml_tensor * dst); +void ggml_compute_forward_dsv4_hc_comb(const struct ggml_compute_params * params, struct ggml_tensor * dst); +void ggml_compute_forward_dsv4_hc_pre(const struct ggml_compute_params * params, struct ggml_tensor * dst); +void ggml_compute_forward_dsv4_hc_post(const struct ggml_compute_params * params, struct ggml_tensor * dst); void ggml_compute_forward_map_custom1(const struct ggml_compute_params * params, struct ggml_tensor * dst); void ggml_compute_forward_map_custom2(const struct ggml_compute_params * params, struct ggml_tensor * dst); void ggml_compute_forward_map_custom3(const struct ggml_compute_params * params, struct ggml_tensor * dst); diff --git a/ggml/src/ggml-cuda/argsort.cu b/ggml/src/ggml-cuda/argsort.cu index c4f08091e79a..33a38c23e87e 100644 --- a/ggml/src/ggml-cuda/argsort.cu +++ b/ggml/src/ggml-cuda/argsort.cu @@ -28,6 +28,20 @@ static __global__ void init_offsets(int * offsets, const int ncols, const int nr #endif // STRIDED_ITERATOR_AVAILABLE #ifdef GGML_CUDA_USE_CUB + +// returns the suggested maximum number of rows to process during one argsort_f32_i32_cuda_cub() call +int argsort_f32_i32_cuda_cub_chunk_nrows(const size_t nb01, const int64_t nrows) { + // perform argsort in chunks up to approximately this size (currently 64MB) + // to avoid excessive temporary buffers memory usage + const int chunk_bytes = 1 << 26; + + // calculate how many rows will fit in one chunk (must be at least one) + const int chunk_nrows = chunk_bytes > nb01 ? chunk_bytes / nb01 : 1; + + // limit the resulting amount to total nrows + return nrows < chunk_nrows ? nrows : chunk_nrows; +} + void argsort_f32_i32_cuda_cub(ggml_cuda_pool & pool, const float * x, int * dst, @@ -254,11 +268,22 @@ void ggml_cuda_op_argsort(ggml_backend_cuda_context & ctx, ggml_tensor * dst) { const size_t shared_mem = ncols_pad * sizeof(int); const size_t max_shared_mem = ggml_cuda_info().devices[ggml_cuda_get_device()].smpb; - if (shared_mem > max_shared_mem || ncols > 1024) { - ggml_cuda_pool & pool = ctx.pool(); - argsort_f32_i32_cuda_cub(pool, src0_d, (int *) dst_d, ncols, nrows, order, stream); - } else { - argsort_f32_i32_cuda_bitonic(src0_d, (int *) dst_d, ncols, nrows, order, stream); + // early return if we can use bitonic argsort + if (shared_mem <= max_shared_mem && ncols <= 1024) { + return argsort_f32_i32_cuda_bitonic(src0_d, (int *) dst_d, ncols, nrows, order, stream); + } + + const int chunk_nrows = argsort_f32_i32_cuda_cub_chunk_nrows(src0->nb[1], nrows); + + ggml_cuda_pool & pool = ctx.pool(); + + for (int64_t i = 0; i < nrows; i += chunk_nrows) { + int iter_nrows = chunk_nrows < nrows - i ? chunk_nrows : nrows - i; + + argsort_f32_i32_cuda_cub(pool, src0_d, (int *) dst_d, ncols, iter_nrows, order, stream); + + src0_d += ncols * iter_nrows; + dst_d += ncols * iter_nrows; } #else argsort_f32_i32_cuda_bitonic(src0_d, (int *) dst_d, ncols, nrows, order, stream); diff --git a/ggml/src/ggml-cuda/argsort.cuh b/ggml/src/ggml-cuda/argsort.cuh index 22b7306f2020..3abb6448a057 100644 --- a/ggml/src/ggml-cuda/argsort.cuh +++ b/ggml/src/ggml-cuda/argsort.cuh @@ -3,6 +3,7 @@ void ggml_cuda_op_argsort(ggml_backend_cuda_context & ctx, ggml_tensor * dst); #ifdef GGML_CUDA_USE_CUB +int argsort_f32_i32_cuda_cub_chunk_nrows(const size_t nb01, const int64_t nrows); void argsort_f32_i32_cuda_cub(ggml_cuda_pool & pool, const float * x, int * dst, diff --git a/ggml/src/ggml-cuda/dsv4-hc.cu b/ggml/src/ggml-cuda/dsv4-hc.cu new file mode 100644 index 000000000000..4dc9055f519e --- /dev/null +++ b/ggml/src/ggml-cuda/dsv4-hc.cu @@ -0,0 +1,297 @@ +#include "common.cuh" +#include "dsv4-hc.cuh" + + +static __device__ void dsv4_hc_comb_norm_cols(float * comb, float eps) { + constexpr int64_t hc = 4; + + for (int64_t idst = 0; idst < hc; ++idst) { + float sum = eps; + for (int64_t isrc = 0; isrc < hc; ++isrc) { + sum += comb[idst + hc*isrc]; + } + + const float inv_sum = 1.0f / sum; + for (int64_t isrc = 0; isrc < hc; ++isrc) { + comb[idst + hc*isrc] *= inv_sum; + } + } +} + +static __device__ void dsv4_hc_comb_norm_rows(float * comb, float eps) { + constexpr int64_t hc = 4; + + for (int64_t isrc = 0; isrc < hc; ++isrc) { + float sum = eps; + for (int64_t idst = 0; idst < hc; ++idst) { + sum += comb[idst + hc*isrc]; + } + + const float inv_sum = 1.0f / sum; + for (int64_t idst = 0; idst < hc; ++idst) { + comb[idst + hc*isrc] *= inv_sum; + } + } +} + +static __global__ void dsv4_hc_comb_f32( + const float * mixes, + const float * scale, + const float * base, + float * dst, + int64_t n_tokens, + int64_t sm0, + int64_t sm1, + int64_t ss0, + int64_t sb0, + int64_t sd0, + int64_t sd1, + int64_t sd2, + float eps, + int32_t n_iter) { + constexpr int64_t hc = 4; + constexpr int64_t comb_offset = 2*hc; + + ggml_cuda_pdl_lc(); + const int64_t it = (int64_t) blockIdx.x * blockDim.x + threadIdx.x; + + if (it >= n_tokens) { + return; + } + + ggml_cuda_pdl_sync(); + + const float scale_comb = scale[2*ss0]; + float comb[hc*hc]; + + for (int64_t isrc = 0; isrc < hc; ++isrc) { + float max = -INFINITY; + for (int64_t idst = 0; idst < hc; ++idst) { + const int64_t idx = idst + hc*isrc; + const float v = mixes[(comb_offset + idx)*sm0 + it*sm1] * scale_comb + base[(comb_offset + idx)*sb0]; + comb[idx] = v; + max = fmaxf(max, v); + } + + float sum = 0.0f; + for (int64_t idst = 0; idst < hc; ++idst) { + const int64_t idx = idst + hc*isrc; + const float v = expf(comb[idx] - max); + comb[idx] = v; + sum += v; + } + + const float inv_sum = 1.0f / sum; + for (int64_t idst = 0; idst < hc; ++idst) { + const int64_t idx = idst + hc*isrc; + comb[idx] = comb[idx] * inv_sum + eps; + } + } + + dsv4_hc_comb_norm_cols(comb, eps); + for (int32_t i = 1; i < n_iter; ++i) { + dsv4_hc_comb_norm_rows(comb, eps); + dsv4_hc_comb_norm_cols(comb, eps); + } + + for (int64_t isrc = 0; isrc < hc; ++isrc) { + for (int64_t idst = 0; idst < hc; ++idst) { + const int64_t idx = idst + hc*isrc; + dst[idst*sd0 + isrc*sd1 + it*sd2] = comb[idx]; + } + } +} + +static __global__ void dsv4_hc_pre_f32( + const float * x, + const float * weights, + float * dst, + int64_t n_embd, + int64_t hc, + int64_t n_tokens, + int64_t sx0, + int64_t sx1, + int64_t sx2, + int64_t sw0, + int64_t sw1, + int64_t sd0, + int64_t sd1) { + ggml_cuda_pdl_lc(); + const int64_t ir = (int64_t) blockIdx.x * blockDim.x + threadIdx.x; + const int64_t nr = n_embd * n_tokens; + + if (ir >= nr) { + return; + } + + ggml_cuda_pdl_sync(); + + const int64_t i0 = ir % n_embd; + const int64_t it = ir / n_embd; + + float sum = __fmul_rn(x[i0*sx0 + it*sx2], weights[it*sw1]); + for (int64_t ih = 1; ih < hc; ++ih) { + const float xv = x[i0*sx0 + ih*sx1 + it*sx2]; + const float wv = weights[ih*sw0 + it*sw1]; + sum = __fadd_rn(sum, __fmul_rn(xv, wv)); + } + + dst[i0*sd0 + it*sd1] = sum; +} + +static __global__ void dsv4_hc_post_f32( + const float * x, + const float * residual, + const float * post, + const float * comb, + float * dst, + int64_t n_embd, + int64_t hc, + int64_t n_tokens, + int64_t sx0, + int64_t sx1, + int64_t sr0, + int64_t sr1, + int64_t sr2, + int64_t sp0, + int64_t sp1, + int64_t sc0, + int64_t sc1, + int64_t sc2, + int64_t sd0, + int64_t sd1, + int64_t sd2) { + ggml_cuda_pdl_lc(); + const int64_t ir = (int64_t) blockIdx.x * blockDim.x + threadIdx.x; + const int64_t nr = n_embd * hc * n_tokens; + + if (ir >= nr) { + return; + } + + ggml_cuda_pdl_sync(); + + const int64_t i0 = ir % n_embd; + const int64_t idst = (ir / n_embd) % hc; + const int64_t it = ir / (n_embd * hc); + + float sum = x[i0*sx0 + it*sx1] * post[idst*sp0 + it*sp1]; + for (int64_t isrc = 0; isrc < hc; ++isrc) { + sum += residual[i0*sr0 + isrc*sr1 + it*sr2] * comb[idst*sc0 + isrc*sc1 + it*sc2]; + } + + dst[i0*sd0 + idst*sd1 + it*sd2] = sum; +} + +void ggml_cuda_op_dsv4_hc_comb(ggml_backend_cuda_context & ctx, ggml_tensor * dst) { + const ggml_tensor * mixes = dst->src[0]; + const ggml_tensor * scale = dst->src[1]; + const ggml_tensor * base = dst->src[2]; + + GGML_ASSERT(mixes->type == GGML_TYPE_F32); + GGML_ASSERT(scale->type == GGML_TYPE_F32); + GGML_ASSERT(base->type == GGML_TYPE_F32); + GGML_ASSERT(dst->type == GGML_TYPE_F32); + + constexpr int64_t hc = 4; + constexpr int64_t hc_mix_dim = (2 + hc)*hc; + + GGML_ASSERT(mixes->ne[0] == hc_mix_dim); + GGML_ASSERT(dst->ne[0] == hc); + GGML_ASSERT(dst->ne[1] == hc); + GGML_ASSERT(dst->ne[2] == mixes->ne[1]); + GGML_ASSERT(scale->ne[0] >= 3); + GGML_ASSERT(base->ne[0] == hc_mix_dim); + + GGML_TENSOR_LOCALS(size_t, nbm, mixes, nb); + GGML_TENSOR_LOCALS(size_t, nbs, scale, nb); + GGML_TENSOR_LOCALS(size_t, nbb, base, nb); + GGML_TENSOR_LOCALS(size_t, nbd, dst, nb); + + const int64_t n_tokens = mixes->ne[1]; + const float eps = ggml_get_op_params_f32(dst, 0); + const int32_t n_iter = ggml_get_op_params_i32(dst, 1); + + const int block_size = 256; + const dim3 block_dims(block_size, 1, 1); + const dim3 grid_dims((n_tokens + block_size - 1) / block_size, 1, 1); + const ggml_cuda_kernel_launch_params launch_params = ggml_cuda_kernel_launch_params(grid_dims, block_dims, 0, ctx.stream()); + + ggml_cuda_kernel_launch(dsv4_hc_comb_f32, launch_params, + (const float *) mixes->data, (const float *) scale->data, (const float *) base->data, (float *) dst->data, + n_tokens, + nbm0 / sizeof(float), nbm1 / sizeof(float), + nbs0 / sizeof(float), + nbb0 / sizeof(float), + nbd0 / sizeof(float), nbd1 / sizeof(float), nbd2 / sizeof(float), + eps, n_iter); +} + +void ggml_cuda_op_dsv4_hc_pre(ggml_backend_cuda_context & ctx, ggml_tensor * dst) { + const ggml_tensor * x = dst->src[0]; + const ggml_tensor * weights = dst->src[1]; + + GGML_ASSERT(x->type == GGML_TYPE_F32); + GGML_ASSERT(weights->type == GGML_TYPE_F32); + GGML_ASSERT(dst->type == GGML_TYPE_F32); + + GGML_TENSOR_LOCALS(size_t, nbx, x, nb); + GGML_TENSOR_LOCALS(size_t, nbw, weights, nb); + GGML_TENSOR_LOCALS(size_t, nbd, dst, nb); + + const int64_t n_embd = x->ne[0]; + const int64_t hc = x->ne[1]; + const int64_t n_tokens = x->ne[2]; + + const int block_size = 256; + const int64_t nr = n_embd * n_tokens; + const dim3 block_dims(block_size, 1, 1); + const dim3 grid_dims((nr + block_size - 1) / block_size, 1, 1); + const ggml_cuda_kernel_launch_params launch_params = ggml_cuda_kernel_launch_params(grid_dims, block_dims, 0, ctx.stream()); + + ggml_cuda_kernel_launch(dsv4_hc_pre_f32, launch_params, + (const float *) x->data, (const float *) weights->data, (float *) dst->data, + n_embd, hc, n_tokens, + nbx0 / sizeof(float), nbx1 / sizeof(float), nbx2 / sizeof(float), + nbw0 / sizeof(float), nbw1 / sizeof(float), + nbd0 / sizeof(float), nbd1 / sizeof(float)); +} + +void ggml_cuda_op_dsv4_hc_post(ggml_backend_cuda_context & ctx, ggml_tensor * dst) { + const ggml_tensor * x = dst->src[0]; + const ggml_tensor * residual = dst->src[1]; + const ggml_tensor * post = dst->src[2]; + const ggml_tensor * comb = dst->src[3]; + + GGML_ASSERT(x->type == GGML_TYPE_F32); + GGML_ASSERT(residual->type == GGML_TYPE_F32); + GGML_ASSERT(post->type == GGML_TYPE_F32); + GGML_ASSERT(comb->type == GGML_TYPE_F32); + GGML_ASSERT(dst->type == GGML_TYPE_F32); + + GGML_TENSOR_LOCALS(size_t, nbx, x, nb); + GGML_TENSOR_LOCALS(size_t, nbr, residual, nb); + GGML_TENSOR_LOCALS(size_t, nbp, post, nb); + GGML_TENSOR_LOCALS(size_t, nbc, comb, nb); + GGML_TENSOR_LOCALS(size_t, nbd, dst, nb); + + const int64_t n_embd = x->ne[0]; + const int64_t n_tokens = x->ne[1]; + const int64_t hc = residual->ne[1]; + + const int block_size = 256; + const int64_t nr = n_embd * hc * n_tokens; + const dim3 block_dims(block_size, 1, 1); + const dim3 grid_dims((nr + block_size - 1) / block_size, 1, 1); + const ggml_cuda_kernel_launch_params launch_params = ggml_cuda_kernel_launch_params(grid_dims, block_dims, 0, ctx.stream()); + + ggml_cuda_kernel_launch(dsv4_hc_post_f32, launch_params, + (const float *) x->data, (const float *) residual->data, + (const float *) post->data, (const float *) comb->data, (float *) dst->data, + n_embd, hc, n_tokens, + nbx0 / sizeof(float), nbx1 / sizeof(float), + nbr0 / sizeof(float), nbr1 / sizeof(float), nbr2 / sizeof(float), + nbp0 / sizeof(float), nbp1 / sizeof(float), + nbc0 / sizeof(float), nbc1 / sizeof(float), nbc2 / sizeof(float), + nbd0 / sizeof(float), nbd1 / sizeof(float), nbd2 / sizeof(float)); +} diff --git a/ggml/src/ggml-cuda/dsv4-hc.cuh b/ggml/src/ggml-cuda/dsv4-hc.cuh new file mode 100644 index 000000000000..2379aaefb41b --- /dev/null +++ b/ggml/src/ggml-cuda/dsv4-hc.cuh @@ -0,0 +1,6 @@ +#include "common.cuh" +#include "ggml.h" + +void ggml_cuda_op_dsv4_hc_comb(ggml_backend_cuda_context & ctx, ggml_tensor * dst); +void ggml_cuda_op_dsv4_hc_pre(ggml_backend_cuda_context & ctx, ggml_tensor * dst); +void ggml_cuda_op_dsv4_hc_post(ggml_backend_cuda_context & ctx, ggml_tensor * dst); diff --git a/ggml/src/ggml-cuda/ggml-cuda.cu b/ggml/src/ggml-cuda/ggml-cuda.cu index f3fb32452d2c..6ae5289a31b6 100644 --- a/ggml/src/ggml-cuda/ggml-cuda.cu +++ b/ggml/src/ggml-cuda/ggml-cuda.cu @@ -58,6 +58,7 @@ #include "ggml-cuda/wkv.cuh" #include "ggml-cuda/gla.cuh" #include "ggml-cuda/gated_delta_net.cuh" +#include "ggml-cuda/dsv4-hc.cuh" #include "ggml-cuda/set.cuh" #include "ggml-cuda/set-rows.cuh" #include "ggml-cuda/pad_reflect_1d.cuh" @@ -65,6 +66,7 @@ #include "ggml-cuda/tri.cuh" #include "ggml-cuda/cumsum.cuh" #include "ggml-cuda/fill.cuh" +#include "ggml-cuda/lightning-indexer.cuh" #include "ggml.h" #include @@ -3100,6 +3102,15 @@ static bool ggml_cuda_compute_forward(ggml_backend_cuda_context & ctx, struct gg case GGML_OP_GATED_DELTA_NET: ggml_cuda_op_gated_delta_net(ctx, dst); break; + case GGML_OP_DSV4_HC_COMB: + ggml_cuda_op_dsv4_hc_comb(ctx, dst); + break; + case GGML_OP_DSV4_HC_PRE: + ggml_cuda_op_dsv4_hc_pre(ctx, dst); + break; + case GGML_OP_DSV4_HC_POST: + ggml_cuda_op_dsv4_hc_post(ctx, dst); + break; case GGML_OP_RWKV_WKV7: ggml_cuda_op_rwkv_wkv7(ctx, dst); break; @@ -3118,6 +3129,9 @@ static bool ggml_cuda_compute_forward(ggml_backend_cuda_context & ctx, struct gg case GGML_OP_FILL: ggml_cuda_op_fill(ctx, dst); break; + case GGML_OP_LIGHTNING_INDEXER: + ggml_cuda_op_lightning_indexer(ctx, dst); + break; default: return false; } @@ -5453,6 +5467,16 @@ static bool ggml_backend_cuda_device_supports_op(ggml_backend_dev_t dev, const g #else return true; #endif // GGML_USE_MUSA + case GGML_OP_DSV4_HC_COMB: + return op->src[0]->type == GGML_TYPE_F32 && op->src[1]->type == GGML_TYPE_F32 && + op->src[2]->type == GGML_TYPE_F32 && op->type == GGML_TYPE_F32; + case GGML_OP_DSV4_HC_PRE: + return op->src[0]->type == GGML_TYPE_F32 && op->src[1]->type == GGML_TYPE_F32 && + op->type == GGML_TYPE_F32; + case GGML_OP_DSV4_HC_POST: + return op->src[0]->type == GGML_TYPE_F32 && op->src[1]->type == GGML_TYPE_F32 && + op->src[2]->type == GGML_TYPE_F32 && op->src[3]->type == GGML_TYPE_F32 && + op->type == GGML_TYPE_F32; case GGML_OP_FLASH_ATTN_EXT: return ggml_cuda_flash_attn_ext_supported(dev_ctx->device, op); case GGML_OP_CROSS_ENTROPY_LOSS: @@ -5464,6 +5488,7 @@ static bool ggml_backend_cuda_device_supports_op(ggml_backend_dev_t dev, const g case GGML_OP_TRI: case GGML_OP_DIAG: case GGML_OP_SOLVE_TRI: + case GGML_OP_LIGHTNING_INDEXER: return true; default: diff --git a/ggml/src/ggml-cuda/lightning-indexer.cu b/ggml/src/ggml-cuda/lightning-indexer.cu new file mode 100644 index 000000000000..6cda36efd157 --- /dev/null +++ b/ggml/src/ggml-cuda/lightning-indexer.cu @@ -0,0 +1,507 @@ +#include "common.cuh" +#include "lightning-indexer.cuh" +#include "fattn-common.cuh" +#include "convert.cuh" + +#if !defined(GGML_USE_HIP) && !defined(GGML_USE_MUSA) +#if defined(__CUDA_ARCH__) && __CUDA_ARCH__ >= GGML_CUDA_CC_AMPERE + +typedef union { + int2 i2; + half2 h2[2]; +} half4; + +#include +namespace wmma = nvcuda::wmma; + +template +static __global__ void lightning_indexer_kernel_wmma( + const float * src0, const char * src1, const float * src2, float * dst, + const float scale_embd, const float scale_heads, + int64_t n_stream, int64_t n_batch, int64_t n_kv, + size_t nb1, size_t nb2, size_t nb3, + size_t nb01, size_t nb02, size_t nb03, + size_t nb11, size_t nb12, size_t nb13, + size_t nb21, size_t nb22, size_t nb23 + ) { + + constexpr int K_VECS_PER_BLOCK = 32; + constexpr int WARPS_PER_BLOCK = 8; + constexpr int THREADS_PER_BLOCK = WARPS_PER_BLOCK * WARP_SIZE; + constexpr int HEADS_PER_INNER_LOOP = 8; + constexpr int K_EMBD_PER_INNER_LOOP = 16; + constexpr int n_embd_padded = n_embd + 8; + + const int i_batch = blockIdx.y; + const int i_stream = blockIdx.z; + const int i_warp = threadIdx.y; + const int i_lane = threadIdx.x; + const int tid = i_warp * WARP_SIZE + i_lane; + + // each block processes K_VECS_PER_BLOCK K vectors + const int start_kv = blockIdx.x * K_VECS_PER_BLOCK; + + const char * q_base = (const char *) src0 + i_batch*nb02 + i_stream*nb03; + const float * w_base = (const float *) ((const char *) src2 + i_batch*nb21 + i_stream*nb23); + + // phase 1 - load weights and first Q tile to shared memory + + __shared__ float w_shared[n_head]; + __shared__ int2 q_shared_h[HEADS_PER_INNER_LOOP][n_embd_padded / 4]; + + if (tid < n_head) { + w_shared[tid] = w_base[tid]; + } + + // total number of half4 elements in HEADS_PER_INNER_LOOP x n_embd Q tile + constexpr int n_q_tile = HEADS_PER_INNER_LOOP * (n_embd / 4); + // number of registers needed in each thread to store Q tile in thread block + constexpr int n_q_next = (n_q_tile + THREADS_PER_BLOCK - 1) / THREADS_PER_BLOCK; + + #pragma unroll + for (int i_q = tid; i_q < n_q_tile; i_q += THREADS_PER_BLOCK) { + const int i_head = i_q / (n_embd / 4); + const int i_embd = i_q % (n_embd / 4); + const float4 q = *(const float4 *) (q_base + i_head*nb01 + i_embd*sizeof(float4)); + half4 q_packed; + q_packed.h2[0] = __float22half2_rn(make_float2(q.x, q.y)); + q_packed.h2[1] = __float22half2_rn(make_float2(q.z, q.w)); + q_shared_h[i_head][i_embd] = q_packed.i2; + } + + // phase 2 - load (and dequantize if needed) K to shared mem + + __shared__ half2 k_shared_h[K_VECS_PER_BLOCK][n_embd_padded / 4][2]; + + constexpr int n_k = K_VECS_PER_BLOCK * (n_embd / 4); + + if constexpr (type_K == GGML_TYPE_F16) { + #pragma unroll + for (int i_k = tid; i_k < n_k; i_k += THREADS_PER_BLOCK) { + const int i_k_vec = i_k / (n_embd / 4); + const int i_embd = i_k % (n_embd / 4); + const int i_kv = start_kv + i_k_vec; + if (i_kv < n_kv) { + const int2 * k_base = (const int2 *) ((const char *) src1 + i_kv*nb12 + i_stream*nb13); + *(int2*) &k_shared_h[i_k_vec][i_embd] = k_base[i_embd]; + } else { + *(int2*) &k_shared_h[i_k_vec][i_embd] = make_int2(0, 0); + } + } + } else { + constexpr dequantize_V_t dequantize_k = get_dequantize_V(); + #pragma unroll + for (int i_k = tid; i_k < n_k; i_k += THREADS_PER_BLOCK) { + const int i_k_vec = i_k / (n_embd / 4); + const int i_embd = i_k % (n_embd / 4); + const int i_kv = start_kv + i_k_vec; + if (i_kv < n_kv) { + const void * k_base = (const void *) ((const char *) src1 + i_kv*nb12 + i_stream*nb13); + dequantize_k(k_base, &k_shared_h[i_k_vec][i_embd][0], i_embd * 4); + } else { + *(int2*) &k_shared_h[i_k_vec][i_embd] = make_int2(0, 0); + } + } + } + + __syncthreads(); + + // phase 3 - calculate lightning indexer scores + + __shared__ float qk_shared[WARPS_PER_BLOCK][HEADS_PER_INNER_LOOP][K_VECS_PER_BLOCK]; + + // load K fragment + wmma::fragment frag_k; + wmma::load_matrix_sync(frag_k, (half*) &k_shared_h[0][i_warp * K_EMBD_PER_INNER_LOOP / 4], n_embd_padded); + + float score_k = 0.0f; + + for (int i_head_0 = 0; i_head_0 < n_head; i_head_0 += HEADS_PER_INNER_LOOP) { + const int i_head_next = i_head_0 + HEADS_PER_INNER_LOOP; + + // we don't use accumulator for anything, fill it with zeros + wmma::fragment frag_acc; + wmma::fill_fragment(frag_acc, 0.0f); + + // load Q fragment + wmma::fragment frag_q; + wmma::load_matrix_sync(frag_q, (half*) &q_shared_h[0][i_warp * K_EMBD_PER_INNER_LOOP / 4], n_embd_padded); + + // preload next Q tile to registers during matrix multiplication + float4 q_next[n_q_next]; + + if (i_head_next < n_head) { + #pragma unroll + for (int i_q = tid, i_q_next = 0; i_q < n_q_tile; i_q += THREADS_PER_BLOCK) { + const int i_head = i_head_next + i_q / (n_embd / 4); + const int i_embd = i_q % (n_embd / 4); + q_next[i_q_next++] = *(const float4 *) (q_base + i_head*nb01 + i_embd*sizeof(float4)); + } + } + + // perform matrix multiplication + wmma::mma_sync(frag_acc, frag_q, frag_k, frag_acc); + wmma::store_matrix_sync((float*) &qk_shared[i_warp][0][0], frag_acc, K_VECS_PER_BLOCK, wmma::mem_row_major); + + // make sure all threads finished using q_shared_h so we can store next tile + __syncthreads(); + + // write preloaded Q tile to shared memory + if (i_head_next < n_head) { + #pragma unroll + for (int i_q = tid, i_q_next = 0; i_q < n_q_tile; i_q += THREADS_PER_BLOCK) { + const int i_head = i_q / (n_embd / 4); + const int i_embd = i_q % (n_embd / 4); + half4 q_packed; + q_packed.h2[0] = __float22half2_rn(make_float2(q_next[i_q_next].x, q_next[i_q_next].y)); + q_packed.h2[1] = __float22half2_rn(make_float2(q_next[i_q_next].z, q_next[i_q_next].w)); + q_shared_h[i_head][i_embd] = q_packed.i2; + ++i_q_next; + } + } + + // accumulate QK multiplication results from all block warps + // (there are 256 threads in block and 256 matmul outputs) + // TODO it will break if WARP_SIZE is not 32 + const int h = tid / K_VECS_PER_BLOCK; + const int k = tid % K_VECS_PER_BLOCK; + const float w_val = w_shared[i_head_0 + h]; + + float sum = 0.0f; + #pragma unroll + for (int w = 0; w < WARPS_PER_BLOCK; ++w) { + sum += qk_shared[w][h][k]; + } + + // scale_embd, ReLU, weight + sum *= scale_embd; + sum = sum > 0.0f ? sum : 0.0f; + sum *= w_val; + + // wait until qk_shared[0] is no longer used + __syncthreads(); + + // reuse qk_shared[0] for storing partial results + qk_shared[0][h][k] = sum; + + // wait until all threads write their results + __syncthreads(); + + // accumulate result over heads + if (tid < K_VECS_PER_BLOCK) { + #pragma unroll + for (int i_head = 0; i_head < HEADS_PER_INNER_LOOP; ++i_head) { + score_k += qk_shared[0][i_head][tid]; + } + } + + // make sure all threads finished using qk_shared + __syncthreads(); + } + + // phase 4 - store output to VRAM + + if (tid < K_VECS_PER_BLOCK) { + const int i_kv = start_kv + tid; + if (i_kv < n_kv) { + float * dst_base = (float *) ((char *) dst + i_batch*nb1 + i_stream*nb3); + dst_base[i_kv] = score_k * scale_heads; + } + } +} + +#else // defined(__CUDA_ARCH__) && __CUDA_ARCH__ >= GGML_CUDA_CC_AMPERE + +template +static __global__ void lightning_indexer_kernel_wmma( + const float * src0, const char * src1, const float * src2, float * dst, + const float scale_embd, const float scale_heads, + int64_t n_stream, int64_t n_batch, int64_t n_kv, + size_t nb1, size_t nb2, size_t nb3, + size_t nb01, size_t nb02, size_t nb03, + size_t nb11, size_t nb12, size_t nb13, + size_t nb21, size_t nb22, size_t nb23 + ) { + GGML_UNUSED_VARS(src0, src1, src2, dst, + scale_embd, scale_heads, + n_stream, n_batch, n_kv, + nb1, nb2, nb3, + nb01, nb02, nb03, + nb11, nb12, nb13, + nb21, nb22, nb23); + NO_DEVICE_CODE; +} + +#endif // defined(__CUDA_ARCH__) && __CUDA_ARCH__ >= GGML_CUDA_CC_AMPERE +#endif // !defined(GGML_USE_HIP) && !defined(GGML_USE_MUSA) && defined(__CUDA_ARCH__) && __CUDA_ARCH__ >= GGML_CUDA_CC_AMPERE + +// TODO there is one ugly assumption used in this kernel - that WARP_SIZE is equal to 32 +// thanks to that one warp operating on float4 processes whole indexer K/Q vectors +// 32 * 4 = 128 (n_embd) + +template +static __global__ void lightning_indexer_kernel_vec( + const float * src0, const char * src1, const float * src2, float * dst, + const float scale_embd, const float scale_heads, + int64_t n_stream, int64_t n_batch, int64_t n_kv, + size_t nb1, size_t nb2, size_t nb3, + size_t nb01, size_t nb02, size_t nb03, + size_t nb11, size_t nb12, size_t nb13, + size_t nb21, size_t nb22, size_t nb23 + ) { + + constexpr int K_VECS_PER_WARP = 8; + constexpr int WARPS_PER_BLOCK = 8; + constexpr int THREADS_PER_BLOCK = WARPS_PER_BLOCK * WARP_SIZE; + + const int i_batch = blockIdx.y; + const int i_stream = blockIdx.z; + const int i_warp = threadIdx.y; + const int i_lane = threadIdx.x; + const int tid = i_warp * WARP_SIZE + i_lane; + + // each warp processes K_VECS_PER_WARP K vectors + const int start_kv_block = blockIdx.x * (WARPS_PER_BLOCK * K_VECS_PER_WARP); + const int start_kv = start_kv_block + i_warp * K_VECS_PER_WARP; + + const char * q_base = (const char *) src0 + i_batch*nb02 + i_stream*nb03; + const float * w_base = (const float *) ((const char *) src2 + i_batch*nb21 + i_stream*nb23); + + // phase 1 - load (and dequantize if needed) K to registers + + float4 k_reg_f[K_VECS_PER_WARP]; + + if constexpr (type_K == GGML_TYPE_F32) { + // direct copy of float4 + #pragma unroll + for (int k = 0; k < K_VECS_PER_WARP; ++k) { + int i_kv = start_kv + k; + if (i_kv < n_kv) { + const float4 * k_base = (const float4 *) ((const char *) src1 + i_kv*nb12 + i_stream*nb13); + k_reg_f[k] = k_base[i_lane]; + } else { + k_reg_f[k] = make_float4(0, 0, 0, 0); + } + } + } else { + // dequantize remaining types to float + constexpr dequantize_V_t dequantize_k = get_dequantize_V(); + #pragma unroll + for (int k = 0; k < K_VECS_PER_WARP; ++k) { + int i_kv = start_kv + k; + if (i_kv < n_kv) { + const void * k_base = (const void *) ((const char *) src1 + i_kv*nb12 + i_stream*nb13); + dequantize_k(k_base, &k_reg_f[k], i_lane * 4); + } else { + k_reg_f[k] = make_float4(0, 0, 0, 0); + } + } + } + + float score_k[K_VECS_PER_WARP] = { 0.0f }; + + // load weights and Q only for n_head_inner heads at once to reduce shared memory usage + constexpr int n_head_inner = n_head / 4; + + for (int i_head_0 = 0; i_head_0 < n_head; i_head_0 += n_head_inner) { + // phase 2 - load weights and Q to shared memory + + __shared__ float w_shared[n_head_inner]; + __shared__ float4 q_shared_f[n_head_inner][n_embd / 4]; + + if (tid < n_head_inner) { + w_shared[tid] = w_base[i_head_0 + tid]; + } + + constexpr int n_q = n_head_inner * (n_embd / 4); + #pragma unroll + for (int i_q = tid; i_q < n_q; i_q += THREADS_PER_BLOCK) { + const int i_head_inner = i_q / (n_embd / 4); + const int i_head = i_head_0 + i_head_inner; + const int i_embd = i_q % (n_embd / 4); + q_shared_f[i_head_inner][i_embd] = *(const float4 *) (q_base + i_head*nb01 + i_embd*sizeof(float4)); + } + + __syncthreads(); + + // phase 3 - calculate lightning indexer scores + + for (int i_head_inner = 0; i_head_inner < n_head_inner; ++i_head_inner) { + const float w_val = w_shared[i_head_inner]; + float qk[K_VECS_PER_WARP] = { 0.0f }; + + // dot product of floats + const float4 q_vec = q_shared_f[i_head_inner][i_lane]; + + #pragma unroll + for (int k = 0; k < K_VECS_PER_WARP; ++k) { + ggml_cuda_mad(qk[k], q_vec.x, k_reg_f[k].x); + ggml_cuda_mad(qk[k], q_vec.y, k_reg_f[k].y); + ggml_cuda_mad(qk[k], q_vec.z, k_reg_f[k].z); + ggml_cuda_mad(qk[k], q_vec.w, k_reg_f[k].w); + } + + #pragma unroll + for (int k = 0; k < K_VECS_PER_WARP; ++k) { + float sum = warp_reduce_sum(qk[k]); + + // scale_embd, ReLU, weight + if (i_lane == 0) { + sum *= scale_embd; + sum = (sum > 0.0f) ? sum : 0.0f; + score_k[k] += sum * w_val; + } + } + } + + __syncthreads(); + } + + // phase 4 - store outputs to shared memory + + __shared__ float dst_shared[WARPS_PER_BLOCK * K_VECS_PER_WARP]; + + if (i_lane == 0) { + #pragma unroll + for (int k = 0; k < K_VECS_PER_WARP; ++k) { + dst_shared[i_warp * K_VECS_PER_WARP + k] = score_k[k] * scale_heads; + } + } + + __syncthreads(); + + // phase 5 - write from shared memory to VRAM in coalesced manner + + if (tid < WARPS_PER_BLOCK * K_VECS_PER_WARP) { + int i_kv = start_kv_block + tid; + if (i_kv < n_kv) { + float * dst_base = (float *) ((char *) dst + i_batch*nb1 + i_stream*nb3); + dst_base[i_kv] = dst_shared[tid]; + } + } +} + +#define DECL_LIGHTNING_INDEXER_CASE(lightning_indexer_kernel, n_embd, n_head, type_K) \ + template __global__ void lightning_indexer_kernel( \ + const float * src0, const char * src1, const float * src2, float * dst, \ + const float scale_embd, const float scale_heads, \ + int64_t n_stream, int64_t n_batch, int64_t n_kv, \ + size_t nb1, size_t nb2, size_t nb3, \ + size_t nb01, size_t nb02, size_t nb03, \ + size_t nb11, size_t nb12, size_t nb13, \ + size_t nb21, size_t nb22, size_t nb23); + +#if !defined(GGML_USE_HIP) && !defined(GGML_USE_MUSA) +DECL_LIGHTNING_INDEXER_CASE(lightning_indexer_kernel_wmma, 128, 64, GGML_TYPE_F16) +DECL_LIGHTNING_INDEXER_CASE(lightning_indexer_kernel_wmma, 128, 64, GGML_TYPE_Q4_0) +DECL_LIGHTNING_INDEXER_CASE(lightning_indexer_kernel_wmma, 128, 64, GGML_TYPE_Q4_1) +DECL_LIGHTNING_INDEXER_CASE(lightning_indexer_kernel_wmma, 128, 64, GGML_TYPE_Q5_0) +DECL_LIGHTNING_INDEXER_CASE(lightning_indexer_kernel_wmma, 128, 64, GGML_TYPE_Q5_1) +DECL_LIGHTNING_INDEXER_CASE(lightning_indexer_kernel_wmma, 128, 64, GGML_TYPE_Q8_0) +#endif // !defined(GGML_USE_HIP) && !defined(GGML_USE_MUSA) + +DECL_LIGHTNING_INDEXER_CASE(lightning_indexer_kernel_vec, 128, 64, GGML_TYPE_F16) +DECL_LIGHTNING_INDEXER_CASE(lightning_indexer_kernel_vec, 128, 64, GGML_TYPE_Q4_0) +DECL_LIGHTNING_INDEXER_CASE(lightning_indexer_kernel_vec, 128, 64, GGML_TYPE_Q4_1) +DECL_LIGHTNING_INDEXER_CASE(lightning_indexer_kernel_vec, 128, 64, GGML_TYPE_Q5_0) +DECL_LIGHTNING_INDEXER_CASE(lightning_indexer_kernel_vec, 128, 64, GGML_TYPE_Q5_1) +DECL_LIGHTNING_INDEXER_CASE(lightning_indexer_kernel_vec, 128, 64, GGML_TYPE_Q8_0) +DECL_LIGHTNING_INDEXER_CASE(lightning_indexer_kernel_vec, 128, 64, GGML_TYPE_BF16) +DECL_LIGHTNING_INDEXER_CASE(lightning_indexer_kernel_vec, 128, 64, GGML_TYPE_F32) + +#define LIGHTNING_INDEXER_CASE(lightning_indexer_kernel, n_embd, n_head, K, type_K) \ + if (K->type == (type_K)) { \ + lightning_indexer_kernel<<>>( \ + src0_d, src1_d, src2_d, dst_d, scale_embd, scale_heads, \ + n_stream, n_batch, n_kv, \ + nb1, nb2, nb3, \ + nb01, nb02, nb03, \ + nb11, nb12, nb13, \ + nb21, nb22, nb23 \ + ); \ + } else + +void ggml_cuda_op_lightning_indexer(ggml_backend_cuda_context & ctx, ggml_tensor * dst) { + const ggml_tensor * src0 = dst->src[0]; + const ggml_tensor * src1 = dst->src[1]; + const ggml_tensor * src2 = dst->src[2]; + + const float scale_embd = ggml_get_op_params_f32(dst, 0); + const float scale_heads = ggml_get_op_params_f32(dst, 1); + + GGML_ASSERT(dst->type == GGML_TYPE_F32); + GGML_ASSERT(src0->type == GGML_TYPE_F32); + GGML_ASSERT(src2->type == GGML_TYPE_F32); + + GGML_TENSOR_TERNARY_OP_LOCALS + + // input tensor rows must be contiguous + GGML_ASSERT(nb00 == ggml_type_size(src0->type)); + GGML_ASSERT(nb10 == ggml_type_size(src1->type)); + GGML_ASSERT(nb20 == ggml_type_size(src2->type)); + + // dst cannot be transposed or permuted + GGML_ASSERT(nb0 == sizeof(float)); + GGML_ASSERT(nb0 <= nb1); + GGML_ASSERT(nb1 <= nb2); + GGML_ASSERT(nb2 <= nb3); + + const int n_embd = src0->ne[0]; + const int n_head = src0->ne[1]; + const int n_batch = src0->ne[2]; + const int n_stream = src0->ne[3]; + const int n_kv = src1->ne[2]; + + const float * src0_d = (const float *) src0->data; + const char * src1_d = (const char *) src1->data; + const float * src2_d = (const float *) src2->data; + float * dst_d = (float *) dst->data; + + const int device = ggml_cuda_get_device(); + const int cc = ggml_cuda_info().devices[device].cc; + + if (n_embd == 128 && n_head == 64) { +#if !defined(GGML_USE_HIP) && !defined(GGML_USE_MUSA) + if (GGML_CUDA_CC_IS_NVIDIA(cc) && ampere_mma_available(cc) && src1->type != GGML_TYPE_F32 && src1->type != GGML_TYPE_BF16) { + // use wmma kernel + constexpr int K_VECS_PER_BLOCK = 32; + constexpr int WARPS_PER_BLOCK = 8; + + dim3 block(32, WARPS_PER_BLOCK); + int num_kv_blocks = (n_kv + (K_VECS_PER_BLOCK) - 1) / (K_VECS_PER_BLOCK); + dim3 grid(num_kv_blocks, n_batch, n_stream); + + LIGHTNING_INDEXER_CASE(lightning_indexer_kernel_wmma, 128, 64, src1, GGML_TYPE_F16) + LIGHTNING_INDEXER_CASE(lightning_indexer_kernel_wmma, 128, 64, src1, GGML_TYPE_Q4_0) + LIGHTNING_INDEXER_CASE(lightning_indexer_kernel_wmma, 128, 64, src1, GGML_TYPE_Q4_1) + LIGHTNING_INDEXER_CASE(lightning_indexer_kernel_wmma, 128, 64, src1, GGML_TYPE_Q5_0) + LIGHTNING_INDEXER_CASE(lightning_indexer_kernel_wmma, 128, 64, src1, GGML_TYPE_Q5_1) + LIGHTNING_INDEXER_CASE(lightning_indexer_kernel_wmma, 128, 64, src1, GGML_TYPE_Q8_0) + GGML_ABORT("fatal error"); + } else { +#else // !defined(GGML_USE_HIP) && !defined(GGML_USE_MUSA) + { +#endif // !defined(GGML_USE_HIP) && !defined(GGML_USE_MUSA) + // use vector kernel + constexpr int K_VECS_PER_WARP = 8; + constexpr int WARPS_PER_BLOCK = 8; + constexpr int K_VECS_PER_BLOCK = K_VECS_PER_WARP * WARPS_PER_BLOCK; + + dim3 block(32, WARPS_PER_BLOCK); + int num_kv_blocks = (n_kv + (K_VECS_PER_BLOCK) - 1) / (K_VECS_PER_BLOCK); + dim3 grid(num_kv_blocks, n_batch, n_stream); + + LIGHTNING_INDEXER_CASE(lightning_indexer_kernel_vec, 128, 64, src1, GGML_TYPE_F16) + LIGHTNING_INDEXER_CASE(lightning_indexer_kernel_vec, 128, 64, src1, GGML_TYPE_Q4_0) + LIGHTNING_INDEXER_CASE(lightning_indexer_kernel_vec, 128, 64, src1, GGML_TYPE_Q4_1) + LIGHTNING_INDEXER_CASE(lightning_indexer_kernel_vec, 128, 64, src1, GGML_TYPE_Q5_0) + LIGHTNING_INDEXER_CASE(lightning_indexer_kernel_vec, 128, 64, src1, GGML_TYPE_Q5_1) + LIGHTNING_INDEXER_CASE(lightning_indexer_kernel_vec, 128, 64, src1, GGML_TYPE_Q8_0) + LIGHTNING_INDEXER_CASE(lightning_indexer_kernel_vec, 128, 64, src1, GGML_TYPE_BF16) + LIGHTNING_INDEXER_CASE(lightning_indexer_kernel_vec, 128, 64, src1, GGML_TYPE_F32) + GGML_ABORT("fatal error"); + } + } else { + GGML_ABORT("fatal error"); + } +} diff --git a/ggml/src/ggml-cuda/lightning-indexer.cuh b/ggml/src/ggml-cuda/lightning-indexer.cuh new file mode 100644 index 000000000000..31fcc7d5ae0a --- /dev/null +++ b/ggml/src/ggml-cuda/lightning-indexer.cuh @@ -0,0 +1,3 @@ +#include "common.cuh" + +void ggml_cuda_op_lightning_indexer(ggml_backend_cuda_context & ctx, ggml_tensor * dst); diff --git a/ggml/src/ggml-cuda/top-k.cu b/ggml/src/ggml-cuda/top-k.cu index db1d39e2dc71..5e708e6c5ed4 100644 --- a/ggml/src/ggml-cuda/top-k.cu +++ b/ggml/src/ggml-cuda/top-k.cu @@ -75,17 +75,26 @@ void ggml_cuda_op_top_k(ggml_backend_cuda_context & ctx, ggml_tensor * dst) { const int ncols_pad = next_power_of_2(ncols); const size_t shared_mem = ncols_pad * sizeof(int); const size_t max_shared_mem = ggml_cuda_info().devices[ggml_cuda_get_device()].smpb; + const bool use_bitonic = shared_mem <= max_shared_mem && ncols <= 1024; + const int chunk_nrows = argsort_f32_i32_cuda_cub_chunk_nrows(src0->nb[1], nrows); - ggml_cuda_pool_alloc temp_dst_alloc(pool, ncols * nrows); + ggml_cuda_pool_alloc temp_dst_alloc(pool, ncols * chunk_nrows); int * tmp_dst = temp_dst_alloc.get(); - if (shared_mem > max_shared_mem || ncols > 1024) { - argsort_f32_i32_cuda_cub(pool, src0_d, tmp_dst, ncols, nrows, GGML_SORT_ORDER_DESC, stream); - } else { - argsort_f32_i32_cuda_bitonic(src0_d, tmp_dst, ncols, nrows, GGML_SORT_ORDER_DESC, stream); + for (int64_t i = 0; i < nrows; i += chunk_nrows) { + int iter_nrows = chunk_nrows < nrows - i ? chunk_nrows : nrows - i; + + if (use_bitonic) { + argsort_f32_i32_cuda_bitonic(src0_d, tmp_dst, ncols, iter_nrows, GGML_SORT_ORDER_DESC, stream); + } else { + argsort_f32_i32_cuda_cub(pool, src0_d, tmp_dst, ncols, iter_nrows, GGML_SORT_ORDER_DESC, stream); + } + CUDA_CHECK(cudaMemcpy2DAsync(dst_d, k * sizeof(int), tmp_dst, ncols * sizeof(int), k * sizeof(int), iter_nrows, + cudaMemcpyDeviceToDevice, stream)); + + src0_d += ncols * iter_nrows; + dst_d += k * iter_nrows; } - CUDA_CHECK(cudaMemcpy2DAsync(dst_d, k * sizeof(int), tmp_dst, ncols * sizeof(int), k * sizeof(int), nrows, - cudaMemcpyDeviceToDevice, stream)); #else // GGML_CUDA_USE_CUB ggml_cuda_pool_alloc temp_dst_alloc(pool, ncols * nrows); int * tmp_dst = temp_dst_alloc.get(); diff --git a/ggml/src/ggml-metal/ggml-metal-device.cpp b/ggml/src/ggml-metal/ggml-metal-device.cpp index 0e1f1de4577d..8751d0f26e3a 100644 --- a/ggml/src/ggml-metal/ggml-metal-device.cpp +++ b/ggml/src/ggml-metal/ggml-metal-device.cpp @@ -473,6 +473,62 @@ ggml_metal_pipeline_with_params ggml_metal_library_get_pipeline_soft_max(ggml_me return res; } +ggml_metal_pipeline_with_params ggml_metal_library_get_pipeline_lightning_indexer(ggml_metal_library_t lib, const ggml_tensor * op) { + GGML_ASSERT(op->op == GGML_OP_LIGHTNING_INDEXER); + + const ggml_tensor * src0 = op->src[0]; + const ggml_tensor * src1 = op->src[1]; + + GGML_ASSERT(src0->ne[0] == 128); // n_embd + GGML_ASSERT(src0->ne[1] == 64); // n_head + + char base[256]; + char name[256]; + + snprintf(base, 256, "kernel_lightning_indexer_%s", ggml_type_name(src1->type)); + snprintf(name, 256, "%s", base); + + ggml_metal_pipeline_with_params res = ggml_metal_library_get_pipeline(lib, name); + if (!res.pipeline) { + res = ggml_metal_library_compile_pipeline(lib, base, name, nullptr); + } + + return res; +} + +ggml_metal_pipeline_with_params ggml_metal_library_get_pipeline_dsv4_hc_comb(ggml_metal_library_t lib, const ggml_tensor * op) { + GGML_ASSERT(op->op == GGML_OP_DSV4_HC_COMB); + + ggml_metal_pipeline_with_params res = ggml_metal_library_get_pipeline(lib, "kernel_dsv4_hc_comb"); + if (!res.pipeline) { + res = ggml_metal_library_compile_pipeline(lib, "kernel_dsv4_hc_comb", "kernel_dsv4_hc_comb", nullptr); + } + + return res; +} + +ggml_metal_pipeline_with_params ggml_metal_library_get_pipeline_dsv4_hc_pre(ggml_metal_library_t lib, const ggml_tensor * op) { + GGML_ASSERT(op->op == GGML_OP_DSV4_HC_PRE); + + ggml_metal_pipeline_with_params res = ggml_metal_library_get_pipeline(lib, "kernel_dsv4_hc_pre"); + if (!res.pipeline) { + res = ggml_metal_library_compile_pipeline(lib, "kernel_dsv4_hc_pre", "kernel_dsv4_hc_pre", nullptr); + } + + return res; +} + +ggml_metal_pipeline_with_params ggml_metal_library_get_pipeline_dsv4_hc_post(ggml_metal_library_t lib, const ggml_tensor * op) { + GGML_ASSERT(op->op == GGML_OP_DSV4_HC_POST); + + ggml_metal_pipeline_with_params res = ggml_metal_library_get_pipeline(lib, "kernel_dsv4_hc_post"); + if (!res.pipeline) { + res = ggml_metal_library_compile_pipeline(lib, "kernel_dsv4_hc_post", "kernel_dsv4_hc_post", nullptr); + } + + return res; +} + ggml_metal_pipeline_with_params ggml_metal_library_get_pipeline_ssm_conv(ggml_metal_library_t lib, const ggml_tensor * op) { GGML_ASSERT(op->src[0]->type == GGML_TYPE_F32); GGML_ASSERT(op->src[1]->type == GGML_TYPE_F32); diff --git a/ggml/src/ggml-metal/ggml-metal-device.h b/ggml/src/ggml-metal/ggml-metal-device.h index d465f31c083b..dfe6fb241239 100644 --- a/ggml/src/ggml-metal/ggml-metal-device.h +++ b/ggml/src/ggml-metal/ggml-metal-device.h @@ -124,6 +124,10 @@ struct ggml_metal_pipeline_with_params ggml_metal_library_get_pipeline_cumsum_bl struct ggml_metal_pipeline_with_params ggml_metal_library_get_pipeline_cumsum_add (ggml_metal_library_t lib, const struct ggml_tensor * op); struct ggml_metal_pipeline_with_params ggml_metal_library_get_pipeline_tri (ggml_metal_library_t lib, const struct ggml_tensor * op); struct ggml_metal_pipeline_with_params ggml_metal_library_get_pipeline_soft_max (ggml_metal_library_t lib, const struct ggml_tensor * op); +struct ggml_metal_pipeline_with_params ggml_metal_library_get_pipeline_lightning_indexer (ggml_metal_library_t lib, const struct ggml_tensor * op); +struct ggml_metal_pipeline_with_params ggml_metal_library_get_pipeline_dsv4_hc_comb (ggml_metal_library_t lib, const struct ggml_tensor * op); +struct ggml_metal_pipeline_with_params ggml_metal_library_get_pipeline_dsv4_hc_pre (ggml_metal_library_t lib, const struct ggml_tensor * op); +struct ggml_metal_pipeline_with_params ggml_metal_library_get_pipeline_dsv4_hc_post (ggml_metal_library_t lib, const struct ggml_tensor * op); struct ggml_metal_pipeline_with_params ggml_metal_library_get_pipeline_ssm_conv (ggml_metal_library_t lib, const struct ggml_tensor * op); struct ggml_metal_pipeline_with_params ggml_metal_library_get_pipeline_ssm_conv_batched (ggml_metal_library_t lib, const struct ggml_tensor * op, int ssm_conv_bs); struct ggml_metal_pipeline_with_params ggml_metal_library_get_pipeline_ssm_scan (ggml_metal_library_t lib, const struct ggml_tensor * op); diff --git a/ggml/src/ggml-metal/ggml-metal-device.m b/ggml/src/ggml-metal/ggml-metal-device.m index a7cbc60ebe41..ddf6156e7c01 100644 --- a/ggml/src/ggml-metal/ggml-metal-device.m +++ b/ggml/src/ggml-metal/ggml-metal-device.m @@ -1255,6 +1255,59 @@ bool ggml_metal_device_supports_op(ggml_metal_device_t dev, const struct ggml_te return false; } return has_simdgroup_mm; // TODO: over-restricted for vec-kernels + case GGML_OP_LIGHTNING_INDEXER: + { + // DeepSeek V4 lightning indexer: n_embd=128, n_head=64 + const int64_t n_embd = op->src[0]->ne[0]; + const int64_t n_head = op->src[0]->ne[1]; + + if (n_embd != 128 || n_head != 64) { + return false; + } + + if (op->src[0]->type != GGML_TYPE_F32) { + return false; + } + if (op->src[2]->type != GGML_TYPE_F32) { + return false; + } + + switch (op->src[1]->type) { + case GGML_TYPE_F32: + case GGML_TYPE_F16: + case GGML_TYPE_Q8_0: + case GGML_TYPE_Q4_0: + case GGML_TYPE_Q4_1: + case GGML_TYPE_Q5_0: + case GGML_TYPE_Q5_1: + break; + case GGML_TYPE_BF16: + if (!has_bfloat) { + return false; + } + break; + default: + return false; + } + + return true; + } + case GGML_OP_DSV4_HC_COMB: + if (op->src[0]->type != GGML_TYPE_F32 || op->src[1]->type != GGML_TYPE_F32 || op->src[2]->type != GGML_TYPE_F32) { + return false; + } + return true; + case GGML_OP_DSV4_HC_PRE: + if (op->src[0]->type != GGML_TYPE_F32 || op->src[1]->type != GGML_TYPE_F32) { + return false; + } + return true; + case GGML_OP_DSV4_HC_POST: + if (op->src[0]->type != GGML_TYPE_F32 || op->src[1]->type != GGML_TYPE_F32 || + op->src[2]->type != GGML_TYPE_F32 || op->src[3]->type != GGML_TYPE_F32) { + return false; + } + return true; case GGML_OP_SSM_CONV: case GGML_OP_SSM_SCAN: return has_simdgroup_reduction; diff --git a/ggml/src/ggml-metal/ggml-metal-impl.h b/ggml/src/ggml-metal/ggml-metal-impl.h index ff74cafb5b79..11c5c4f3c7cc 100644 --- a/ggml/src/ggml-metal/ggml-metal-impl.h +++ b/ggml/src/ggml-metal/ggml-metal-impl.h @@ -101,6 +101,10 @@ #define FC_SUM_ROWS 1400 #define FC_UPSCALE 1500 #define FC_GATED_DELTA_NET 1600 +#define FC_LIGHTNING_INDEXER 1700 +#define FC_DSV4_HC_COMB 1800 +#define FC_DSV4_HC_PRE 1900 +#define FC_DSV4_HC_POST 2000 // op-specific constants #define OP_FLASH_ATTN_EXT_NQPSG 8 @@ -1172,4 +1176,67 @@ typedef struct { int64_t np; } ggml_metal_kargs_opt_step_sgd; +typedef struct { + int32_t n_kv; + int32_t n_head; + uint64_t nb1; + uint64_t nb2; + uint64_t nb3; + uint64_t nb01; + uint64_t nb02; + uint64_t nb03; + uint64_t nb11; + uint64_t nb12; + uint64_t nb13; + uint64_t nb21; + uint64_t nb22; + uint64_t nb23; + float scale_embd; + float scale_heads; +} ggml_metal_kargs_lightning_indexer; + +typedef struct { + uint64_t nd0; + uint64_t nd1; + uint64_t nd2; + uint64_t nm0; + uint64_t nm1; + uint64_t ns0; + uint64_t nb0; + int32_t n_tokens; + int32_t n_iter; + float eps; + int32_t pad; +} ggml_metal_kargs_dsv4_hc_comb; + +typedef struct { + uint64_t nbx0; + uint64_t nbx1; + uint64_t nbx2; + uint64_t nbw0; + uint64_t nbw1; + uint64_t nbd0; + uint64_t nbd1; + int32_t n_embd; + int32_t n_tokens; +} ggml_metal_kargs_dsv4_hc_pre; + +typedef struct { + uint64_t nbx0; + uint64_t nbx1; + uint64_t nbr0; + uint64_t nbr1; + uint64_t nbr2; + uint64_t nbp0; + uint64_t nbp1; + uint64_t nbc0; + uint64_t nbc1; + uint64_t nbc2; + uint64_t nbd0; + uint64_t nbd1; + uint64_t nbd2; + int32_t n_embd; + int32_t n_tokens; +} ggml_metal_kargs_dsv4_hc_post; + #endif // GGML_METAL_IMPL diff --git a/ggml/src/ggml-metal/ggml-metal-ops.cpp b/ggml/src/ggml-metal/ggml-metal-ops.cpp index 18656b346f21..41413e5be5f1 100644 --- a/ggml/src/ggml-metal/ggml-metal-ops.cpp +++ b/ggml/src/ggml-metal/ggml-metal-ops.cpp @@ -316,6 +316,22 @@ static int ggml_metal_op_encode_impl(ggml_metal_op_t ctx, int idx) { { n_fuse = ggml_metal_op_cumsum(ctx, idx); } break; + case GGML_OP_LIGHTNING_INDEXER: + { + n_fuse = ggml_metal_op_lightning_indexer(ctx, idx); + } break; + case GGML_OP_DSV4_HC_COMB: + { + n_fuse = ggml_metal_op_dsv4_hc_comb(ctx, idx); + } break; + case GGML_OP_DSV4_HC_PRE: + { + n_fuse = ggml_metal_op_dsv4_hc_pre(ctx, idx); + } break; + case GGML_OP_DSV4_HC_POST: + { + n_fuse = ggml_metal_op_dsv4_hc_post(ctx, idx); + } break; case GGML_OP_SOFT_MAX: { n_fuse = ggml_metal_op_soft_max(ctx, idx); @@ -1289,6 +1305,260 @@ int ggml_metal_op_diag(ggml_metal_op_t ctx, int idx) { return 1; } +int ggml_metal_op_dsv4_hc_comb(ggml_metal_op_t ctx, int idx) { + ggml_tensor * dst = ctx->node(idx); + + ggml_metal_library_t lib = ctx->lib; + ggml_metal_encoder_t enc = ctx->enc; + + GGML_ASSERT(dst->op == GGML_OP_DSV4_HC_COMB); + + const ggml_tensor * mixes = dst->src[0]; + const ggml_tensor * scale = dst->src[1]; + const ggml_tensor * base = dst->src[2]; + + GGML_ASSERT(mixes->type == GGML_TYPE_F32); + GGML_ASSERT(scale->type == GGML_TYPE_F32); + GGML_ASSERT(base->type == GGML_TYPE_F32); + GGML_ASSERT(dst->type == GGML_TYPE_F32); + + GGML_TENSOR_LOCALS(size_t, nbm, mixes, nb); + GGML_TENSOR_LOCALS(size_t, nbs, scale, nb); + GGML_TENSOR_LOCALS(size_t, nbb, base, nb); + GGML_TENSOR_LOCALS(size_t, nbd, dst, nb); + + const int32_t n_tokens = (int32_t) mixes->ne[1]; + const float eps = ggml_get_op_params_f32(dst, 0); + const int32_t n_iter = ggml_get_op_params_i32(dst, 1); + + ggml_metal_kargs_dsv4_hc_comb args = { + /*.nd0 =*/ nbd0, + /*.nd1 =*/ nbd1, + /*.nd2 =*/ nbd2, + /*.nm0 =*/ nbm0, + /*.nm1 =*/ nbm1, + /*.ns0 =*/ nbs0, + /*.nb0 =*/ nbb0, + /*.n_tokens=*/ n_tokens, + /*.n_iter =*/ n_iter, + /*.eps =*/ eps, + /*.pad =*/ 0, + }; + + auto pipeline = ggml_metal_library_get_pipeline_dsv4_hc_comb(lib, dst); + + const int block_size = 256; + const int grid_size = (n_tokens + block_size - 1) / block_size; + + ggml_metal_encoder_set_pipeline(enc, pipeline); + ggml_metal_encoder_set_bytes(enc, &args, sizeof(args), 0); + ggml_metal_encoder_set_buffer(enc, ggml_metal_get_buffer_id(mixes), 1); + ggml_metal_encoder_set_buffer(enc, ggml_metal_get_buffer_id(scale), 2); + ggml_metal_encoder_set_buffer(enc, ggml_metal_get_buffer_id(base), 3); + ggml_metal_encoder_set_buffer(enc, ggml_metal_get_buffer_id(dst), 4); + + ggml_metal_encoder_dispatch_threadgroups(enc, grid_size, 1, 1, block_size, 1, 1); + + return 1; +} + +int ggml_metal_op_dsv4_hc_pre(ggml_metal_op_t ctx, int idx) { + ggml_tensor * dst = ctx->node(idx); + + ggml_metal_library_t lib = ctx->lib; + ggml_metal_encoder_t enc = ctx->enc; + + GGML_ASSERT(dst->op == GGML_OP_DSV4_HC_PRE); + + const ggml_tensor * x = dst->src[0]; + const ggml_tensor * weights = dst->src[1]; + + GGML_ASSERT(x->type == GGML_TYPE_F32); + GGML_ASSERT(weights->type == GGML_TYPE_F32); + GGML_ASSERT(dst->type == GGML_TYPE_F32); + + GGML_TENSOR_LOCALS(size_t, nbx, x, nb); + GGML_TENSOR_LOCALS(size_t, nbw, weights, nb); + GGML_TENSOR_LOCALS(size_t, nbd, dst, nb); + + const int32_t n_embd = (int32_t) x->ne[0]; + const int32_t n_tokens = (int32_t) x->ne[2]; + + ggml_metal_kargs_dsv4_hc_pre args = { + /*.nbx0 =*/ nbx0, + /*.nbx1 =*/ nbx1, + /*.nbx2 =*/ nbx2, + /*.nbw0 =*/ nbw0, + /*.nbw1 =*/ nbw1, + /*.nbd0 =*/ nbd0, + /*.nbd1 =*/ nbd1, + /*.n_embd =*/ n_embd, + /*.n_tokens=*/ n_tokens, + }; + + auto pipeline = ggml_metal_library_get_pipeline_dsv4_hc_pre(lib, dst); + + const int block_size = 256; + const int nr = n_embd * n_tokens; + const int grid_size = (nr + block_size - 1) / block_size; + + ggml_metal_encoder_set_pipeline(enc, pipeline); + ggml_metal_encoder_set_bytes(enc, &args, sizeof(args), 0); + ggml_metal_encoder_set_buffer(enc, ggml_metal_get_buffer_id(x), 1); + ggml_metal_encoder_set_buffer(enc, ggml_metal_get_buffer_id(weights), 2); + ggml_metal_encoder_set_buffer(enc, ggml_metal_get_buffer_id(dst), 3); + + ggml_metal_encoder_dispatch_threadgroups(enc, grid_size, 1, 1, block_size, 1, 1); + + return 1; +} + +int ggml_metal_op_dsv4_hc_post(ggml_metal_op_t ctx, int idx) { + ggml_tensor * dst = ctx->node(idx); + + ggml_metal_library_t lib = ctx->lib; + ggml_metal_encoder_t enc = ctx->enc; + + GGML_ASSERT(dst->op == GGML_OP_DSV4_HC_POST); + + const ggml_tensor * x = dst->src[0]; + const ggml_tensor * residual = dst->src[1]; + const ggml_tensor * post = dst->src[2]; + const ggml_tensor * comb = dst->src[3]; + + GGML_ASSERT(x->type == GGML_TYPE_F32); + GGML_ASSERT(residual->type == GGML_TYPE_F32); + GGML_ASSERT(post->type == GGML_TYPE_F32); + GGML_ASSERT(comb->type == GGML_TYPE_F32); + GGML_ASSERT(dst->type == GGML_TYPE_F32); + + GGML_TENSOR_LOCALS(size_t, nbx, x, nb); + GGML_TENSOR_LOCALS(size_t, nbr, residual, nb); + GGML_TENSOR_LOCALS(size_t, nbp, post, nb); + GGML_TENSOR_LOCALS(size_t, nbc, comb, nb); + GGML_TENSOR_LOCALS(size_t, nbd, dst, nb); + + const int32_t n_embd = (int32_t) x->ne[0]; + const int32_t n_tokens = (int32_t) x->ne[1]; + + ggml_metal_kargs_dsv4_hc_post args = { + /*.nbx0 =*/ nbx0, + /*.nbx1 =*/ nbx1, + /*.nbr0 =*/ nbr0, + /*.nbr1 =*/ nbr1, + /*.nbr2 =*/ nbr2, + /*.nbp0 =*/ nbp0, + /*.nbp1 =*/ nbp1, + /*.nbc0 =*/ nbc0, + /*.nbc1 =*/ nbc1, + /*.nbc2 =*/ nbc2, + /*.nbd0 =*/ nbd0, + /*.nbd1 =*/ nbd1, + /*.nbd2 =*/ nbd2, + /*.n_embd =*/ n_embd, + /*.n_tokens=*/ n_tokens, + }; + + auto pipeline = ggml_metal_library_get_pipeline_dsv4_hc_post(lib, dst); + + constexpr int hc = 4; + const int block_size = 256; + const int nr = n_embd * hc * n_tokens; + const int grid_size = (nr + block_size - 1) / block_size; + + ggml_metal_encoder_set_pipeline(enc, pipeline); + ggml_metal_encoder_set_bytes(enc, &args, sizeof(args), 0); + ggml_metal_encoder_set_buffer(enc, ggml_metal_get_buffer_id(x), 1); + ggml_metal_encoder_set_buffer(enc, ggml_metal_get_buffer_id(residual), 2); + ggml_metal_encoder_set_buffer(enc, ggml_metal_get_buffer_id(post), 3); + ggml_metal_encoder_set_buffer(enc, ggml_metal_get_buffer_id(comb), 4); + ggml_metal_encoder_set_buffer(enc, ggml_metal_get_buffer_id(dst), 5); + + ggml_metal_encoder_dispatch_threadgroups(enc, grid_size, 1, 1, block_size, 1, 1); + + return 1; +} + +int ggml_metal_op_lightning_indexer(ggml_metal_op_t ctx, int idx) { + ggml_tensor * dst = ctx->node(idx); + + ggml_metal_library_t lib = ctx->lib; + ggml_metal_encoder_t enc = ctx->enc; + + GGML_ASSERT(dst->op == GGML_OP_LIGHTNING_INDEXER); + + const ggml_tensor * src0 = dst->src[0]; + const ggml_tensor * src1 = dst->src[1]; + const ggml_tensor * src2 = dst->src[2]; + + GGML_ASSERT(src0->type == GGML_TYPE_F32); + GGML_ASSERT(src2->type == GGML_TYPE_F32); + GGML_ASSERT(dst->type == GGML_TYPE_F32); + + GGML_TENSOR_TERNARY_OP_LOCALS + + // input tensor rows must be contiguous + GGML_ASSERT(nb00 == ggml_type_size(src0->type)); + GGML_ASSERT(nb10 == ggml_type_size(src1->type)); + GGML_ASSERT(nb20 == ggml_type_size(src2->type)); + + // dst cannot be transposed or permuted + GGML_ASSERT(nb0 == sizeof(float)); + GGML_ASSERT(nb0 <= nb1); + GGML_ASSERT(nb1 <= nb2); + GGML_ASSERT(nb2 <= nb3); + + const int n_embd = (int) src0->ne[0]; + const int n_head = (int) src0->ne[1]; + const int n_batch = (int) src0->ne[2]; + const int n_stream = (int) src0->ne[3]; + const int n_kv = (int) src1->ne[2]; + + const float scale_embd = ggml_get_op_params_f32(dst, 0); + const float scale_heads = ggml_get_op_params_f32(dst, 1); + + GGML_ASSERT(n_embd == 128); + GGML_ASSERT(n_head == 64); + + ggml_metal_kargs_lightning_indexer args = { + /*.n_kv =*/ n_kv, + /*.n_head =*/ n_head, + /*.nb1 =*/ nb1, + /*.nb2 =*/ nb2, + /*.nb3 =*/ nb3, + /*.nb01 =*/ nb01, + /*.nb02 =*/ nb02, + /*.nb03 =*/ nb03, + /*.nb11 =*/ nb11, + /*.nb12 =*/ nb12, + /*.nb13 =*/ nb13, + /*.nb21 =*/ nb21, + /*.nb22 =*/ nb22, + /*.nb23 =*/ nb23, + /*.scale_embd =*/ scale_embd, + /*.scale_heads=*/ scale_heads, + }; + + auto pipeline = ggml_metal_library_get_pipeline_lightning_indexer(lib, dst); + + constexpr int K_VECS_PER_SG = 8; + constexpr int N_SG_PER_TG = 8; + constexpr int K_VECS_PER_TG = K_VECS_PER_SG * N_SG_PER_TG; + + int num_kv_blocks = (n_kv + K_VECS_PER_TG - 1) / K_VECS_PER_TG; + + ggml_metal_encoder_set_pipeline(enc, pipeline); + ggml_metal_encoder_set_bytes(enc, &args, sizeof(args), 0); + ggml_metal_encoder_set_buffer(enc, ggml_metal_get_buffer_id(src0), 1); + ggml_metal_encoder_set_buffer(enc, ggml_metal_get_buffer_id(src1), 2); + ggml_metal_encoder_set_buffer(enc, ggml_metal_get_buffer_id(src2), 3); + ggml_metal_encoder_set_buffer(enc, ggml_metal_get_buffer_id(dst), 4); + + ggml_metal_encoder_dispatch_threadgroups(enc, num_kv_blocks, n_batch, n_stream, 32, N_SG_PER_TG, 1); + + return 1; +} + int ggml_metal_op_soft_max(ggml_metal_op_t ctx, int idx) { ggml_tensor * op = ctx->node(idx); diff --git a/ggml/src/ggml-metal/ggml-metal-ops.h b/ggml/src/ggml-metal/ggml-metal-ops.h index 36c61071b4fa..c69be0057a64 100644 --- a/ggml/src/ggml-metal/ggml-metal-ops.h +++ b/ggml/src/ggml-metal/ggml-metal-ops.h @@ -54,6 +54,10 @@ int ggml_metal_op_cumsum (ggml_metal_op_t ctx, int idx); int ggml_metal_op_get_rows (ggml_metal_op_t ctx, int idx); int ggml_metal_op_set_rows (ggml_metal_op_t ctx, int idx); int ggml_metal_op_diag (ggml_metal_op_t ctx, int idx); +int ggml_metal_op_lightning_indexer (ggml_metal_op_t ctx, int idx); +int ggml_metal_op_dsv4_hc_comb (ggml_metal_op_t ctx, int idx); +int ggml_metal_op_dsv4_hc_pre (ggml_metal_op_t ctx, int idx); +int ggml_metal_op_dsv4_hc_post (ggml_metal_op_t ctx, int idx); int ggml_metal_op_soft_max (ggml_metal_op_t ctx, int idx); int ggml_metal_op_ssm_conv (ggml_metal_op_t ctx, int idx); int ggml_metal_op_ssm_scan (ggml_metal_op_t ctx, int idx); diff --git a/ggml/src/ggml-metal/ggml-metal.metal b/ggml/src/ggml-metal/ggml-metal.metal index 25e78e100898..f2752a5fd988 100644 --- a/ggml/src/ggml-metal/ggml-metal.metal +++ b/ggml/src/ggml-metal/ggml-metal.metal @@ -10752,3 +10752,436 @@ kernel void kernel_count_equal( typedef decltype(kernel_count_equal) kernel_count_equal_t; template [[host_name("kernel_count_equal_i32")]] kernel kernel_count_equal_t kernel_count_equal; + +// Lightning indexer kernel +// Follows the CUDA vec kernel approach: +// - Each threadgroup processes K_VECS_PER_TG K vectors +// - Each SIMD group processes K_VECS_PER_SG K vectors +// - Each thread (lane) handles one float4 (4 elements) of the 128-element vectors +// - n_embd=128, n_head=64 are hardcoded (matching DeepSeek V4) + +constexpr constant int LI_N_EMBD = 128; +constexpr constant int LI_N_HEAD = 64; +constexpr constant int LI_N_EMBD_4 = LI_N_EMBD / 4; // 32 float4s per vector +constexpr constant int LI_K_VECS_PER_SG = 8; +constexpr constant int LI_N_SG_PER_TG = 8; +constexpr constant int LI_K_VECS_PER_TG = LI_K_VECS_PER_SG * LI_N_SG_PER_TG; // 64 +constexpr constant int LI_N_HEAD_INNER = LI_N_HEAD / 4; // 16 + +// shared compute logic, after K has been loaded to float4 registers +// threadgroup memory pointers are passed in from the kernel entry point +void kernel_lightning_indexer_compute( + constant ggml_metal_kargs_lightning_indexer & args, + device const float * q_base, + device const float * w_base, + device float * dst, + thread float4 * k_reg, + threadgroup float * w_shared, + threadgroup float4 * q_shared, + threadgroup float * dst_shared, + int start_kv_block, + int i_batch, int i_stream, + ushort tiisg, ushort sgitg) { + + float score_k[LI_K_VECS_PER_SG] = { 0.0f }; + + for (int i_head_0 = 0; i_head_0 < LI_N_HEAD; i_head_0 += LI_N_HEAD_INNER) { + const int tid_tg = (int) (tiisg + sgitg * N_SIMDWIDTH); + if (tid_tg < LI_N_HEAD_INNER) { + w_shared[tid_tg] = w_base[i_head_0 + tid_tg]; + } + + const int n_q = LI_N_HEAD_INNER * LI_N_EMBD_4; + const int n_tg = LI_N_SG_PER_TG * N_SIMDWIDTH; + + for (int i_q = tid_tg; i_q < n_q; i_q += n_tg) { + const int i_head_inner = i_q / LI_N_EMBD_4; + const int i_head = i_head_0 + i_head_inner; + const int i_embd = i_q % LI_N_EMBD_4; + q_shared[i_head_inner * LI_N_EMBD_4 + i_embd] = + *(device const float4 *) ((device const char *) q_base + i_head*args.nb01 + i_embd*sizeof(float4)); + } + + threadgroup_barrier(mem_flags::mem_threadgroup); + + for (int i_head_inner = 0; i_head_inner < LI_N_HEAD_INNER; ++i_head_inner) { + const float w_val = w_shared[i_head_inner]; + float qk[LI_K_VECS_PER_SG] = { 0.0f }; + + const float4 q_vec = q_shared[i_head_inner * LI_N_EMBD_4 + tiisg]; + + for (int k = 0; k < LI_K_VECS_PER_SG; ++k) { + qk[k] += q_vec.x * k_reg[k].x; + qk[k] += q_vec.y * k_reg[k].y; + qk[k] += q_vec.z * k_reg[k].z; + qk[k] += q_vec.w * k_reg[k].w; + } + + for (int k = 0; k < LI_K_VECS_PER_SG; ++k) { + float sum = simd_sum(qk[k]); + + if (tiisg == 0) { + sum *= args.scale_embd; + sum = (sum > 0.0f) ? sum : 0.0f; + score_k[k] += sum * w_val; + } + } + } + + threadgroup_barrier(mem_flags::mem_threadgroup); + } + + if (tiisg == 0) { + for (int k = 0; k < LI_K_VECS_PER_SG; ++k) { + dst_shared[sgitg * LI_K_VECS_PER_SG + k] = score_k[k] * args.scale_heads; + } + } + + threadgroup_barrier(mem_flags::mem_threadgroup); + + const int tid_tg = (int) (tiisg + sgitg * N_SIMDWIDTH); + if (tid_tg < LI_K_VECS_PER_TG) { + int i_kv = start_kv_block + tid_tg; + if (i_kv < args.n_kv) { + device float * dst_base = (device float *) ((device char *) dst + i_batch*args.nb1 + i_stream*args.nb3); + dst_base[i_kv] = dst_shared[tid_tg]; + } + } +} + +// kernel entry point for F32 K type +kernel void kernel_lightning_indexer_f32( + constant ggml_metal_kargs_lightning_indexer & args, + device const float * src0, + device const char * src1, + device const float * src2, + device float * dst, + uint3 tgpig[[threadgroup_position_in_grid]], + ushort tiisg[[thread_index_in_simdgroup]], + ushort sgitg[[simdgroup_index_in_threadgroup]]) { + + threadgroup float w_shared[LI_N_HEAD_INNER]; + threadgroup float4 q_shared[LI_N_HEAD_INNER * LI_N_EMBD_4]; + threadgroup float dst_shared[LI_K_VECS_PER_TG]; + + const int i_batch = (int) tgpig.y; + const int i_stream = (int) tgpig.z; + const int start_kv_block = (int) tgpig.x * LI_K_VECS_PER_TG; + const int start_kv = start_kv_block + (int) sgitg * LI_K_VECS_PER_SG; + + device const float * q_base = (device const float *) ((device const char *) src0 + i_batch*args.nb02 + i_stream*args.nb03); + device const float * w_base = (device const float *) ((device const char *) src2 + i_batch*args.nb21 + i_stream*args.nb23); + + float4 k_reg[LI_K_VECS_PER_SG]; + for (int k = 0; k < LI_K_VECS_PER_SG; ++k) { + int i_kv = start_kv + k; + if (i_kv < args.n_kv) { + device const float4 * k_base = (device const float4 *) ((device const char *) src1 + i_kv*args.nb12 + i_stream*args.nb13); + k_reg[k] = k_base[tiisg]; + } else { + k_reg[k] = float4(0); + } + } + kernel_lightning_indexer_compute(args, q_base, w_base, dst, k_reg, + w_shared, q_shared, dst_shared, start_kv_block, i_batch, i_stream, tiisg, sgitg); +} + +// kernel entry point for F16 K type +kernel void kernel_lightning_indexer_f16( + constant ggml_metal_kargs_lightning_indexer & args, + device const float * src0, + device const char * src1, + device const float * src2, + device float * dst, + uint3 tgpig[[threadgroup_position_in_grid]], + ushort tiisg[[thread_index_in_simdgroup]], + ushort sgitg[[simdgroup_index_in_threadgroup]]) { + + threadgroup float w_shared[LI_N_HEAD_INNER]; + threadgroup float4 q_shared[LI_N_HEAD_INNER * LI_N_EMBD_4]; + threadgroup float dst_shared[LI_K_VECS_PER_TG]; + + const int i_batch = (int) tgpig.y; + const int i_stream = (int) tgpig.z; + const int start_kv_block = (int) tgpig.x * LI_K_VECS_PER_TG; + const int start_kv = start_kv_block + (int) sgitg * LI_K_VECS_PER_SG; + + device const float * q_base = (device const float *) ((device const char *) src0 + i_batch*args.nb02 + i_stream*args.nb03); + device const float * w_base = (device const float *) ((device const char *) src2 + i_batch*args.nb21 + i_stream*args.nb23); + + float4 k_reg[LI_K_VECS_PER_SG]; + for (int k = 0; k < LI_K_VECS_PER_SG; ++k) { + int i_kv = start_kv + k; + if (i_kv < args.n_kv) { + device const half4 * k_base = (device const half4 *) ((device const char *) src1 + i_kv*args.nb12 + i_stream*args.nb13); + k_reg[k] = float4(k_base[tiisg]); + } else { + k_reg[k] = float4(0); + } + } + kernel_lightning_indexer_compute(args, q_base, w_base, dst, k_reg, + w_shared, q_shared, dst_shared, start_kv_block, i_batch, i_stream, tiisg, sgitg); +} + +#if defined(GGML_METAL_HAS_BF16) +// kernel entry point for BF16 K type +kernel void kernel_lightning_indexer_bf16( + constant ggml_metal_kargs_lightning_indexer & args, + device const float * src0, + device const char * src1, + device const float * src2, + device float * dst, + uint3 tgpig[[threadgroup_position_in_grid]], + ushort tiisg[[thread_index_in_simdgroup]], + ushort sgitg[[simdgroup_index_in_threadgroup]]) { + + threadgroup float w_shared[LI_N_HEAD_INNER]; + threadgroup float4 q_shared[LI_N_HEAD_INNER * LI_N_EMBD_4]; + threadgroup float dst_shared[LI_K_VECS_PER_TG]; + + const int i_batch = (int) tgpig.y; + const int i_stream = (int) tgpig.z; + const int start_kv_block = (int) tgpig.x * LI_K_VECS_PER_TG; + const int start_kv = start_kv_block + (int) sgitg * LI_K_VECS_PER_SG; + + device const float * q_base = (device const float *) ((device const char *) src0 + i_batch*args.nb02 + i_stream*args.nb03); + device const float * w_base = (device const float *) ((device const char *) src2 + i_batch*args.nb21 + i_stream*args.nb23); + + float4 k_reg[LI_K_VECS_PER_SG]; + for (int k = 0; k < LI_K_VECS_PER_SG; ++k) { + int i_kv = start_kv + k; + if (i_kv < args.n_kv) { + device const bfloat4 * k_base = (device const bfloat4 *) ((device const char *) src1 + i_kv*args.nb12 + i_stream*args.nb13); + k_reg[k] = float4(k_base[tiisg]); + } else { + k_reg[k] = float4(0); + } + } + kernel_lightning_indexer_compute(args, q_base, w_base, dst, k_reg, + w_shared, q_shared, dst_shared, start_kv_block, i_batch, i_stream, tiisg, sgitg); +} +#endif + +// quantized type kernels: template with function pointer for dequantize + +template +kernel void kernel_lightning_indexer_quantized( + constant ggml_metal_kargs_lightning_indexer & args, + device const float * src0, + device const char * src1, + device const float * src2, + device float * dst, + uint3 tgpig[[threadgroup_position_in_grid]], + ushort tiisg[[thread_index_in_simdgroup]], + ushort sgitg[[simdgroup_index_in_threadgroup]]) { + + threadgroup float w_shared[LI_N_HEAD_INNER]; + threadgroup float4 q_shared[LI_N_HEAD_INNER * LI_N_EMBD_4]; + threadgroup float dst_shared[LI_K_VECS_PER_TG]; + + const int i_batch = (int) tgpig.y; + const int i_stream = (int) tgpig.z; + const int start_kv_block = (int) tgpig.x * LI_K_VECS_PER_TG; + const int start_kv = start_kv_block + (int) sgitg * LI_K_VECS_PER_SG; + + device const float * q_base = (device const float *) ((device const char *) src0 + i_batch*args.nb02 + i_stream*args.nb03); + device const float * w_base = (device const float *) ((device const char *) src2 + i_batch*args.nb21 + i_stream*args.nb23); + + // dequantize K to float4 registers + // n_embd=128, block_size=32 -> 4 blocks per K vector, 8 positions per block + constexpr int positions_per_block = 32 / 4; // 8 + const int il = (int) tiisg % positions_per_block; + const int block_idx = (int) tiisg / positions_per_block; + + float4 k_reg[LI_K_VECS_PER_SG]; + for (int k = 0; k < LI_K_VECS_PER_SG; ++k) { + int i_kv = start_kv + k; + if (i_kv < args.n_kv) { + device const block_t * k_block = (device const block_t *) ((device const char *) src1 + i_kv*args.nb12 + i_stream*args.nb13); + deq_t4(k_block + block_idx, (short) il, k_reg[k]); + } else { + k_reg[k] = float4(0); + } + } + + kernel_lightning_indexer_compute(args, q_base, w_base, dst, k_reg, + w_shared, q_shared, dst_shared, start_kv_block, i_batch, i_stream, tiisg, sgitg); +} + +typedef decltype(kernel_lightning_indexer_quantized) kernel_lightning_indexer_quantized_t; + +template [[host_name("kernel_lightning_indexer_q4_0")]] kernel kernel_lightning_indexer_quantized_t kernel_lightning_indexer_quantized; +template [[host_name("kernel_lightning_indexer_q4_1")]] kernel kernel_lightning_indexer_quantized_t kernel_lightning_indexer_quantized; +template [[host_name("kernel_lightning_indexer_q5_0")]] kernel kernel_lightning_indexer_quantized_t kernel_lightning_indexer_quantized; +template [[host_name("kernel_lightning_indexer_q5_1")]] kernel kernel_lightning_indexer_quantized_t kernel_lightning_indexer_quantized; +template [[host_name("kernel_lightning_indexer_q8_0")]] kernel kernel_lightning_indexer_quantized_t kernel_lightning_indexer_quantized; + +// DeepSeek V4 hierarchical connection kernels + +// HC_PRE: weighted sum over hc slices +// x[n_embd, hc, n_tokens] * weights[hc, n_tokens] -> dst[n_embd, n_tokens] +kernel void kernel_dsv4_hc_pre( + constant ggml_metal_kargs_dsv4_hc_pre & args, + device const float * x, + device const float * weights, + device float * dst, + uint3 tgpig[[threadgroup_position_in_grid]], + uint tiitg[[thread_index_in_threadgroup]]) { + + const int ir = (int) tgpig.x * 256 + (int) tiitg; + const int nr = args.n_embd * args.n_tokens; + + if (ir >= nr) { + return; + } + + const int i0 = ir % args.n_embd; + const int it = ir / args.n_embd; + + constexpr int hc = 4; + + device const float * w_row = (device const float *) ((device const char *) weights + it*args.nbw1); + + float sum = *(device const float *) ((device const char *) x + i0*args.nbx0 + it*args.nbx2) * w_row[0]; + for (int ih = 1; ih < hc; ++ih) { + sum += *(device const float *) ((device const char *) x + i0*args.nbx0 + ih*args.nbx1 + it*args.nbx2) * + w_row[ih*args.nbw0/sizeof(float)]; + } + + *(device float *) ((device char *) dst + i0*args.nbd0 + it*args.nbd1) = sum; +} + +// HC_POST: residual blend +// dst[i_embd, idst, it] = x[i_embd, it] * post[idst, it] +// + sum_{isrc} residual[i_embd, isrc, it] * comb[idst, isrc, it] +kernel void kernel_dsv4_hc_post( + constant ggml_metal_kargs_dsv4_hc_post & args, + device const float * x, + device const float * residual, + device const float * post, + device const float * comb, + device float * dst, + uint3 tgpig[[threadgroup_position_in_grid]], + uint tiitg[[thread_index_in_threadgroup]]) { + + const int ir = (int) tgpig.x * 256 + (int) tiitg; + + constexpr int hc = 4; + const int nr = args.n_embd * hc * args.n_tokens; + + if (ir >= nr) { + return; + } + + const int i0 = ir % args.n_embd; + const int idst = (ir / args.n_embd) % hc; + const int it = ir / (args.n_embd * hc); + + const float xv = *(device const float *) ((device const char *) x + i0*args.nbx0 + it*args.nbx1); + const float pv = *(device const float *) ((device const char *) post + idst*args.nbp0 + it*args.nbp1); + + float sum = xv * pv; + for (int isrc = 0; isrc < hc; ++isrc) { + const float rv = *(device const float *) ((device const char *) residual + i0*args.nbr0 + isrc*args.nbr1 + it*args.nbr2); + const float cv = *(device const float *) ((device const char *) comb + idst*args.nbc0 + isrc*args.nbc1 + it*args.nbc2); + sum += rv * cv; + } + + *(device float *) ((device char *) dst + i0*args.nbd0 + idst*args.nbd1 + it*args.nbd2) = sum; +} + +// HC_COMB: Sinkhorn normalization of combination matrix +// mixes[24, n_tokens], scale[3], base[24] -> comb[4, 4, n_tokens] +kernel void kernel_dsv4_hc_comb( + constant ggml_metal_kargs_dsv4_hc_comb & args, + device const float * mixes, + device const float * scale, + device const float * base, + device float * dst, + uint3 tgpig[[threadgroup_position_in_grid]], + uint tiitg[[thread_index_in_threadgroup]]) { + + const int it = (int) tgpig.x * 256 + (int) tiitg; + + if (it >= args.n_tokens) { + return; + } + + constexpr int hc = 4; + constexpr int comb_offset = 2 * hc; // 8 + + device const float * mixes_row = (device const float *) ((device const char *) mixes + it*args.nm1); + device const float * base_data = (device const float *) base; + + const float scale_comb = scale[2*args.ns0/sizeof(float)]; + + float comb[hc * hc]; + + // row softmax with scale + base affine + for (int isrc = 0; isrc < hc; ++isrc) { + float max = -INFINITY; + for (int idst = 0; idst < hc; ++idst) { + const int idx = idst + hc*isrc; + const int mix_idx = comb_offset + idx; + const float v = mixes_row[mix_idx*args.nm0/sizeof(float)] * scale_comb + base_data[mix_idx*args.nb0/sizeof(float)]; + comb[idx] = v; + max = fmax(max, v); + } + + float sum = 0.0f; + for (int idst = 0; idst < hc; ++idst) { + const int idx = idst + hc*isrc; + const float v = exp(comb[idx] - max); + comb[idx] = v; + sum += v; + } + + const float inv_sum = 1.0f / sum; + for (int idst = 0; idst < hc; ++idst) { + const int idx = idst + hc*isrc; + comb[idx] = comb[idx] * inv_sum + args.eps; + } + } + + // Sinkhorn iterations: normalize columns, then rows alternately + auto norm_cols = [&]() { + for (int idst = 0; idst < hc; ++idst) { + float sum = args.eps; + for (int isrc = 0; isrc < hc; ++isrc) { + sum += comb[idst + hc*isrc]; + } + const float inv_sum = 1.0f / sum; + for (int isrc = 0; isrc < hc; ++isrc) { + comb[idst + hc*isrc] *= inv_sum; + } + } + }; + + auto norm_rows = [&]() { + for (int isrc = 0; isrc < hc; ++isrc) { + float sum = args.eps; + for (int idst = 0; idst < hc; ++idst) { + sum += comb[idst + hc*isrc]; + } + const float inv_sum = 1.0f / sum; + for (int idst = 0; idst < hc; ++idst) { + comb[idst + hc*isrc] *= inv_sum; + } + } + }; + + norm_cols(); + for (int i = 1; i < args.n_iter; ++i) { + norm_rows(); + norm_cols(); + } + + // store output + device float * dst_row = (device float *) ((device char *) dst + it*args.nd2); + for (int isrc = 0; isrc < hc; ++isrc) { + for (int idst = 0; idst < hc; ++idst) { + const int idx = idst + hc*isrc; + *(device float *) ((device char *) dst + idst*args.nd0 + isrc*args.nd1 + it*args.nd2) = comb[idx]; + } + } +} diff --git a/ggml/src/ggml.c b/ggml/src/ggml.c index 0f682fd1856c..c0a03c3b530a 100644 --- a/ggml/src/ggml.c +++ b/ggml/src/ggml.c @@ -1061,6 +1061,10 @@ static const char * GGML_OP_NAME[GGML_OP_COUNT] = { "RWKV_WKV7", "SOLVE_TRI", "GATED_DELTA_NET", + "LIGHTNING_INDEXER", + "DSV4_HC_COMB", + "DSV4_HC_PRE", + "DSV4_HC_POST", "UNARY", @@ -1078,7 +1082,7 @@ static const char * GGML_OP_NAME[GGML_OP_COUNT] = { "GLU", }; -static_assert(GGML_OP_COUNT == 97, "GGML_OP_COUNT != 97"); +static_assert(GGML_OP_COUNT == 101, "GGML_OP_COUNT != 101"); static const char * GGML_OP_SYMBOL[GGML_OP_COUNT] = { "none", @@ -1172,6 +1176,10 @@ static const char * GGML_OP_SYMBOL[GGML_OP_COUNT] = { "rwkv_wkv7(r, w, k, v, a, b, s)", "A X = B, A triangular, solve X", "gated_delta_net(q, k, v, g, beta, s)", + "lightning_indexer(q, k, weights, scale_embd, scale_heads)", + "dsv4_hc_comb(mixes, scale, base)", + "dsv4_hc_pre(x, weights)", + "dsv4_hc_post(x, residual, post, comb)", "unary(x)", @@ -1189,7 +1197,7 @@ static const char * GGML_OP_SYMBOL[GGML_OP_COUNT] = { "glu(x)", }; -static_assert(GGML_OP_COUNT == 97, "GGML_OP_COUNT != 97"); +static_assert(GGML_OP_COUNT == 101, "GGML_OP_COUNT != 101"); static_assert(GGML_OP_POOL_COUNT == 2, "GGML_OP_POOL_COUNT != 2"); @@ -6268,6 +6276,166 @@ struct ggml_tensor * ggml_gated_delta_net( return result; } +// ggml_lightning_indexer + +struct ggml_tensor * ggml_lightning_indexer( + struct ggml_context * ctx, + struct ggml_tensor * q, + struct ggml_tensor * k, + struct ggml_tensor * weights, + float scale_embd, + float scale_heads) { + + GGML_ASSERT(q->type == GGML_TYPE_F32); + GGML_ASSERT(weights->type == GGML_TYPE_F32); + GGML_ASSERT(q->ne[0] == k->ne[0]); + GGML_ASSERT(q->ne[1] == weights->ne[0]); + GGML_ASSERT(k->ne[1] == 1); + GGML_ASSERT(q->ne[2] == weights->ne[1]); + GGML_ASSERT(weights->ne[2] == 1); + GGML_ASSERT(q->ne[3] == k->ne[3]); + GGML_ASSERT(k->ne[3] == weights->ne[3]); + + int64_t ne[4] = { k->ne[2], q->ne[2], 1, q->ne[3] }; + struct ggml_tensor * result = ggml_new_tensor(ctx, GGML_TYPE_F32, 4, ne); + + ggml_set_op_params_f32(result, 0, scale_embd); + ggml_set_op_params_f32(result, 1, scale_heads); + + result->op = GGML_OP_LIGHTNING_INDEXER; + result->src[0] = q; + result->src[1] = k; + result->src[2] = weights; + + return result; +} + +// ggml_dsv4_hc_comb + +struct ggml_tensor * ggml_dsv4_hc_comb( + struct ggml_context * ctx, + struct ggml_tensor * mixes, + struct ggml_tensor * scale, + struct ggml_tensor * base, + float eps, + int32_t n_iter) { + GGML_ASSERT(mixes->type == GGML_TYPE_F32); + GGML_ASSERT(scale->type == GGML_TYPE_F32); + GGML_ASSERT(base->type == GGML_TYPE_F32); + GGML_ASSERT(n_iter > 0); + + const int64_t hc_mix_dim = mixes->ne[0]; + const int64_t n_tokens = mixes->ne[1]; + + int64_t hc = 0; + for (int64_t i = 1; i*i + 2*i <= hc_mix_dim; ++i) { + if ((2 + i)*i == hc_mix_dim) { + hc = i; + break; + } + } + + GGML_ASSERT(hc > 0); + GGML_ASSERT(hc == 4); + GGML_ASSERT(mixes->ne[2] == 1); + GGML_ASSERT(mixes->ne[3] == 1); + GGML_ASSERT(scale->ne[0] >= 3); + GGML_ASSERT(scale->ne[1] == 1); + GGML_ASSERT(scale->ne[2] == 1); + GGML_ASSERT(scale->ne[3] == 1); + GGML_ASSERT(base->ne[0] == hc_mix_dim); + GGML_ASSERT(base->ne[1] == 1); + GGML_ASSERT(base->ne[2] == 1); + GGML_ASSERT(base->ne[3] == 1); + + struct ggml_tensor * result = ggml_new_tensor_3d(ctx, GGML_TYPE_F32, hc, hc, n_tokens); + + ggml_set_op_params_f32(result, 0, eps); + ggml_set_op_params_i32(result, 1, n_iter); + + result->op = GGML_OP_DSV4_HC_COMB; + result->src[0] = mixes; + result->src[1] = scale; + result->src[2] = base; + + return result; +} + +// ggml_dsv4_hc_pre + +struct ggml_tensor * ggml_dsv4_hc_pre( + struct ggml_context * ctx, + struct ggml_tensor * x, + struct ggml_tensor * weights) { + GGML_ASSERT(x->type == GGML_TYPE_F32); + GGML_ASSERT(weights->type == GGML_TYPE_F32); + + const int64_t n_embd = x->ne[0]; + const int64_t hc = x->ne[1]; + const int64_t n_tokens = x->ne[2]; + + GGML_ASSERT(hc > 0); + GGML_ASSERT(x->ne[3] == 1); + GGML_ASSERT(weights->ne[0] == hc); + GGML_ASSERT(weights->ne[1] == n_tokens); + GGML_ASSERT(weights->ne[2] == 1); + GGML_ASSERT(weights->ne[3] == 1); + + struct ggml_tensor * result = ggml_new_tensor_2d(ctx, GGML_TYPE_F32, n_embd, n_tokens); + + result->op = GGML_OP_DSV4_HC_PRE; + result->src[0] = x; + result->src[1] = weights; + + return result; +} + +// ggml_dsv4_hc_post + +struct ggml_tensor * ggml_dsv4_hc_post( + struct ggml_context * ctx, + struct ggml_tensor * x, + struct ggml_tensor * residual, + struct ggml_tensor * post, + struct ggml_tensor * comb) { + GGML_ASSERT(x->type == GGML_TYPE_F32); + GGML_ASSERT(residual->type == GGML_TYPE_F32); + GGML_ASSERT(post->type == GGML_TYPE_F32); + GGML_ASSERT(comb->type == GGML_TYPE_F32); + + const int64_t n_embd = x->ne[0]; + const int64_t n_tokens = x->ne[1]; + const int64_t hc = residual->ne[1]; + + GGML_ASSERT(hc > 0); + GGML_ASSERT(x->ne[2] == 1); + GGML_ASSERT(x->ne[3] == 1); + + GGML_ASSERT(residual->ne[0] == n_embd); + GGML_ASSERT(residual->ne[2] == n_tokens); + GGML_ASSERT(residual->ne[3] == 1); + + GGML_ASSERT(post->ne[0] == hc); + GGML_ASSERT(post->ne[1] == n_tokens); + GGML_ASSERT(post->ne[2] == 1); + GGML_ASSERT(post->ne[3] == 1); + + GGML_ASSERT(comb->ne[0] == hc); + GGML_ASSERT(comb->ne[1] == hc); + GGML_ASSERT(comb->ne[2] == n_tokens); + GGML_ASSERT(comb->ne[3] == 1); + + struct ggml_tensor * result = ggml_new_tensor_3d(ctx, GGML_TYPE_F32, n_embd, hc, n_tokens); + + result->op = GGML_OP_DSV4_HC_POST; + result->src[0] = x; + result->src[1] = residual; + result->src[2] = post; + result->src[3] = comb; + + return result; +} + //////////////////////////////////////////////////////////////////////////////// struct ggml_hash_set ggml_hash_set_new(size_t size) { diff --git a/src/models/deepseek4.cpp b/src/models/deepseek4.cpp index 759654228e36..4d98ff58b58d 100644 --- a/src/models/deepseek4.cpp +++ b/src/models/deepseek4.cpp @@ -4,6 +4,8 @@ #include #include +#include +#include #include #include @@ -15,6 +17,41 @@ static float dsv4_rope_attn_factor(float freq_scale, float ext_factor) { return 1.0f / (1.0f + 0.1f*logf(1.0f/freq_scale)); } +static bool dsv4_env_enabled(const char * name, bool default_enabled) { + const char * value = getenv(name); + if (!value || !*value) { + return default_enabled; + } + + return strcmp(value, "0") != 0 && + strcmp(value, "false") != 0 && + strcmp(value, "FALSE") != 0 && + strcmp(value, "off") != 0 && + strcmp(value, "OFF") != 0 && + strcmp(value, "no") != 0 && + strcmp(value, "NO") != 0; +} + +static bool dsv4_fuse_hc_pre() { + static const bool enabled = dsv4_env_enabled("LLAMA_DSV4_FUSE_HC_PRE", true); + return enabled; +} + +static bool dsv4_fuse_hc_post() { + static const bool enabled = dsv4_env_enabled("LLAMA_DSV4_FUSE_HC_POST", true); + return enabled; +} + +static bool dsv4_fuse_hc_comb() { + static const bool enabled = dsv4_env_enabled("LLAMA_DSV4_FUSE_HC_COMB", true); + return enabled; +} + +static bool dsv4_fuse_hc_head() { + static const bool enabled = dsv4_env_enabled("LLAMA_DSV4_FUSE_HC_HEAD", false); + return enabled; +} + void llama_model_deepseek4::load_arch_hparams(llama_model_loader & ml) { ml.get_key(LLM_KV_ATTENTION_LAYERNORM_RMS_EPS, hparams.f_norm_rms_eps); ml.get_key(LLM_KV_ATTENTION_Q_LORA_RANK, hparams.n_lora_q); @@ -225,20 +262,30 @@ static ggml_tensor * dsv4_hc_affine( ggml_tensor * llama_model_deepseek4::graph::build_hc_weighted_sum( ggml_tensor * x, - ggml_tensor * weights) const { + ggml_tensor * weights, + bool fuse) const { + GGML_ASSERT(x->ne[0] == n_embd); + GGML_ASSERT(x->ne[1] == hparams.dsv4_hc_mult); + const int64_t hc = hparams.dsv4_hc_mult; const int64_t nt = x->ne[2]; - ggml_tensor * acc = nullptr; - for (int64_t ih = 0; ih < hc; ++ih) { - ggml_tensor * xh = ggml_view_2d(ctx0, x, n_embd, nt, x->nb[2], ih*x->nb[1]); - ggml_tensor * wh = ggml_view_2d(ctx0, weights, 1, nt, weights->nb[1], ih*weights->nb[0]); + auto build_ref = [&]() { + ggml_tensor * acc = nullptr; + for (int64_t ih = 0; ih < hc; ++ih) { + ggml_tensor * xh = ggml_view_2d(ctx0, x, n_embd, nt, x->nb[2], ih*x->nb[1]); + ggml_tensor * wh = ggml_view_2d(ctx0, weights, 1, nt, weights->nb[1], ih*weights->nb[0]); + ggml_tensor * cur = ggml_mul(ctx0, xh, wh); + acc = acc ? ggml_add(ctx0, acc, cur) : cur; + } + return acc; + }; - ggml_tensor * cur = ggml_mul(ctx0, xh, wh); - acc = acc ? ggml_add(ctx0, acc, cur) : cur; + if (fuse) { + return ggml_dsv4_hc_pre(ctx0, x, weights); } - return acc; + return build_ref(); } ggml_tensor * llama_model_deepseek4::graph::build_hc_sinkhorn( @@ -301,11 +348,9 @@ ggml_tensor * llama_model_deepseek4::graph::build_hc_pre( ggml_tensor * scale_pre = dsv4_view_1d(ctx0, hc_scale, 1, 0); ggml_tensor * scale_post = dsv4_view_1d(ctx0, hc_scale, 1, 1); - ggml_tensor * scale_comb = dsv4_view_1d(ctx0, hc_scale, 1, 2); ggml_tensor * base_pre = dsv4_view_1d(ctx0, hc_base, hc, 0); ggml_tensor * base_post = dsv4_view_1d(ctx0, hc_base, hc, hc); - ggml_tensor * base_comb = dsv4_view_1d(ctx0, hc_base, hc*hc, 2*hc); ggml_tensor * pre = dsv4_view_2d(ctx0, mixes, hc, nt, 0); pre = dsv4_hc_affine(ctx0, pre, scale_pre, base_pre); @@ -319,13 +364,21 @@ ggml_tensor * llama_model_deepseek4::graph::build_hc_pre( *post = ggml_scale(ctx0, *post, 2.0f); cb(*post, "hc_post", il); - *comb = dsv4_view_2d(ctx0, mixes, hc*hc, nt, 2*hc); - *comb = dsv4_hc_affine(ctx0, *comb, scale_comb, base_comb); - *comb = ggml_reshape_3d(ctx0, *comb, hc, hc, nt); - *comb = build_hc_sinkhorn(*comb, il); + if (dsv4_fuse_hc_comb()) { + *comb = ggml_dsv4_hc_comb(ctx0, mixes, hc_scale, hc_base, hparams.dsv4_hc_eps, + (int32_t) hparams.dsv4_hc_sinkhorn_iters); + } else { + ggml_tensor * scale_comb = dsv4_view_1d(ctx0, hc_scale, 1, 2); + ggml_tensor * base_comb = dsv4_view_1d(ctx0, hc_base, hc*hc, 2*hc); + + *comb = dsv4_view_2d(ctx0, mixes, hc*hc, nt, 2*hc); + *comb = dsv4_hc_affine(ctx0, *comb, scale_comb, base_comb); + *comb = ggml_reshape_3d(ctx0, *comb, hc, hc, nt); + *comb = build_hc_sinkhorn(*comb, il); + } cb(*comb, "hc_comb", il); - return build_hc_weighted_sum(x, pre); + return build_hc_weighted_sum(x, pre, dsv4_fuse_hc_pre()); } ggml_tensor * llama_model_deepseek4::graph::build_hc_post( @@ -336,6 +389,13 @@ ggml_tensor * llama_model_deepseek4::graph::build_hc_post( int il) const { GGML_UNUSED(il); + GGML_ASSERT(x->ne[0] == n_embd); + GGML_ASSERT(residual->ne[1] == hparams.dsv4_hc_mult); + + if (dsv4_fuse_hc_post()) { + return ggml_dsv4_hc_post(ctx0, x, residual, post, comb); + } + const int64_t hc = hparams.dsv4_hc_mult; const int64_t nt = x->ne[1]; @@ -346,7 +406,8 @@ ggml_tensor * llama_model_deepseek4::graph::build_hc_post( for (int64_t src = 0; src < hc; ++src) { ggml_tensor * res_src = ggml_view_2d(ctx0, residual, n_embd, nt, residual->nb[2], src*residual->nb[1]); - ggml_tensor * comb_src_dst = ggml_view_2d(ctx0, comb, 1, nt, comb->nb[2], dst*comb->nb[0] + src*comb->nb[1]); + ggml_tensor * comb_src_dst = ggml_view_2d(ctx0, comb, 1, nt, comb->nb[2], + dst*comb->nb[0] + src*comb->nb[1]); cur = ggml_add(ctx0, cur, ggml_mul(ctx0, res_src, comb_src_dst)); } @@ -376,7 +437,7 @@ ggml_tensor * llama_model_deepseek4::graph::build_hc_head( pre = ggml_scale_bias(ctx0, pre, 1.0f, hparams.dsv4_hc_eps); cb(pre, "hc_head_pre", -1); - return build_hc_weighted_sum(x, pre); + return build_hc_weighted_sum(x, pre, dsv4_fuse_hc_head()); } ggml_tensor * llama_model_deepseek4::graph::build_hca_compressed_kv_from_state( @@ -561,7 +622,6 @@ ggml_tensor * llama_model_deepseek4::graph::build_lid_top_k( cb(indexer_q, "lid_q_rot", il); ggml_tensor * indexer_weights = build_lora_mm(layer.indexer_proj, cur); - indexer_weights = ggml_scale(ctx0, indexer_weights, 1.0f/sqrtf(float(n_embd_indexer_head*n_indexer_head))); cb(indexer_weights, "lid_weights", il); ggml_tensor * indexer_k = inp_dsv4->mctx->get_lid()->get_k(ctx0, il); @@ -582,6 +642,10 @@ ggml_tensor * llama_model_deepseek4::graph::build_lid_top_k( indexer_weights->ne[0], indexer_weights->ne[1]/n_stream, indexer_weights->ne[2], n_stream, indexer_weights->nb[1], indexer_weights->nb[2]/n_stream, indexer_weights->nb[3]/n_stream, 0); +#if 1 + ggml_tensor * indexer_score = ggml_lightning_indexer(ctx0, indexer_q, indexer_k, indexer_weights, 1.0f / sqrtf(float(n_embd_indexer_head)), 1.0f / sqrtf(float(n_indexer_head))); + cb(indexer_score, "indexer_score", il); +#else indexer_q = ggml_permute(ctx0, indexer_q, 0, 2, 1, 3); cb(indexer_q, "lid_q", il); indexer_k = ggml_permute(ctx0, indexer_k, 0, 2, 1, 3); @@ -593,12 +657,15 @@ ggml_tensor * llama_model_deepseek4::graph::build_lid_top_k( indexer_kq = ggml_cont(ctx0, ggml_permute(ctx0, indexer_kq, 2, 1, 0, 3)); cb(indexer_kq, "lid_kq", il); + indexer_weights = ggml_scale(ctx0, indexer_weights, 1.0f/sqrtf(float(n_embd_indexer_head*n_indexer_head))); + cb(indexer_weights, "lid_weights", il); + ggml_tensor * indexer_score = ggml_relu(ctx0, indexer_kq); indexer_score = ggml_mul(ctx0, indexer_score, indexer_weights); indexer_score = ggml_sum_rows(ctx0, indexer_score); indexer_score = ggml_cont(ctx0, ggml_permute(ctx0, indexer_score, 2, 1, 0, 3)); cb(indexer_score, "lid_score", il); - +#endif indexer_score = ggml_add(ctx0, indexer_score, inp_lid.kq_mask); cb(indexer_score, "lid_score_masked", il); diff --git a/src/models/models.h b/src/models/models.h index 7a52e7bc1ab7..20e87bc40891 100644 --- a/src/models/models.h +++ b/src/models/models.h @@ -1189,7 +1189,8 @@ struct llama_model_deepseek4 : public llama_model_base { ggml_tensor * build_hc_weighted_sum( ggml_tensor * x, - ggml_tensor * weights) const; + ggml_tensor * weights, + bool fuse) const; ggml_tensor * build_hc_sinkhorn( ggml_tensor * comb, diff --git a/tests/test-backend-ops.cpp b/tests/test-backend-ops.cpp index 15b50209c850..f8aa4287e38e 100644 --- a/tests/test-backend-ops.cpp +++ b/tests/test-backend-ops.cpp @@ -1138,6 +1138,12 @@ struct test_case { virtual ggml_tensor * build_graph(ggml_context * ctx) = 0; + virtual bool compare_with_reference_graph() { return false; } + + virtual ggml_tensor * build_reference_graph(ggml_context * ctx) { + return build_graph(ctx); + } + virtual double max_nmse_err() { return 1e-7; } @@ -1308,10 +1314,151 @@ struct test_case { } } + test_status_t eval_reference_graph(ggml_backend_t backend1, + ggml_backend_t backend2, + const char * op_names_filter, + printer * output_printer) { + mode = MODE_TEST; + + ggml_init_params params = { + /* .mem_size = */ ggml_tensor_overhead()*256 + ggml_graph_overhead(), + /* .mem_base = */ NULL, + /* .no_alloc = */ true, + }; + ggml_context * ctx_eval = ggml_init(params); + ggml_context * ctx_ref = ggml_init(params); + GGML_ASSERT(ctx_eval); + GGML_ASSERT(ctx_ref); + + gf = ggml_new_graph(ctx_eval); + ggml_cgraph * gf_ref = ggml_new_graph(ctx_ref); + + ggml_tensor * out_eval = build_graph(ctx_eval); + current_op_name = op_desc(out_eval); + check_for_f16_tensor(ctx_eval); + + if (!matches_filter(out_eval, op_names_filter)) { + ggml_free(ctx_eval); + ggml_free(ctx_ref); + return test_status_t::SKIPPED; + } + + ggml_tensor * out_ref = build_reference_graph(ctx_ref); + + bool supported = true; + for (ggml_tensor * t = ggml_get_first_tensor(ctx_eval); t != NULL; t = ggml_get_next_tensor(ctx_eval, t)) { + if (!ggml_backend_supports_op(backend1, t)) { + supported = false; + break; + } + } + for (ggml_tensor * t = ggml_get_first_tensor(ctx_ref); supported && t != NULL; t = ggml_get_next_tensor(ctx_ref, t)) { + if (!ggml_backend_supports_op(backend2, t)) { + supported = false; + break; + } + } + + if (!supported) { + test_result result(ggml_backend_name(backend1), current_op_name, vars(), "test", + false, false, "not supported"); + print_test_result_locked(output_printer, result); + ggml_free(ctx_eval); + ggml_free(ctx_ref); + return test_status_t::NOT_SUPPORTED; + } + + ggml_backend_buffer_t buf_eval = ggml_backend_alloc_ctx_tensors(ctx_eval, backend1); + ggml_backend_buffer_t buf_ref = ggml_backend_alloc_ctx_tensors(ctx_ref, backend2); + + if (buf_eval == NULL || buf_ref == NULL) { + printf("failed to allocate tensors [%s] ", ggml_backend_name(backend1)); + if (buf_eval) { + ggml_backend_buffer_free(buf_eval); + } + if (buf_ref) { + ggml_backend_buffer_free(buf_ref); + } + ggml_free(ctx_eval); + ggml_free(ctx_ref); + return test_status_t::FAIL; + } + + ggml_build_forward_expand(gf, out_eval); + ggml_build_forward_expand(gf_ref, out_ref); + + initialize_tensors(ctx_eval); + initialize_tensors(ctx_ref); + + bool ok = true; + ggml_status status = ggml_backend_graph_compute(backend1, gf); + if (status != GGML_STATUS_SUCCESS) { + fprintf(stderr, "%s: ggml_backend_graph_compute failed. status=%s \n", __func__, ggml_status_to_string(status)); + ok = false; + } + status = ggml_backend_graph_compute(backend2, gf_ref); + if (status != GGML_STATUS_SUCCESS) { + fprintf(stderr, "%s: ggml_backend_graph_compute failed. status=%s \n", __func__, ggml_status_to_string(status)); + ok = false; + } + + if (ok) { + const char * bn1 = ggml_backend_name(backend1); + const char * bn2 = ggml_backend_name(backend2); + std::vector f1 = tensor_to_float(out_eval); + std::vector f2 = tensor_to_float(out_ref); + + GGML_ASSERT(f1.size() == f2.size()); + for (size_t i = 0; i < f1.size(); i++) { + if (std::isnan(f1[i]) || std::isnan(f2[i])) { + printf("[%s] NaN at index %zu (%s=%f %s=%f) ", current_op_name.c_str(), i, bn1, f1[i], bn2, f2[i]); + ok = false; + break; + } + if (isinf_or_max(f1[i]) || isinf_or_max(f2[i])) { + if (isinf_or_max(f1[i]) && isinf_or_max(f2[i])) { + if (std::signbit(f1[i]) != std::signbit(f2[i])) { + printf("[%s] inf sign mismatch: %s=%f %s=%f ", current_op_name.c_str(), bn1, f1[i], bn2, f2[i]); + ok = false; + break; + } + } else { + printf("[%s] inf mismatch: %s=%f %s=%f ", current_op_name.c_str(), bn1, f1[i], bn2, f2[i]); + ok = false; + break; + } + } + } + + if (ok) { + const double error = err(f1.data(), f2.data(), f1.size()); + if (error > max_err(backend1)) { + printf("[%s] ERR = %.9f > %.9f ", current_op_name.c_str(), error, max_err(backend1)); + ok = false; + } + } + } + + ggml_backend_buffer_free(buf_eval); + ggml_backend_buffer_free(buf_ref); + ggml_free(ctx_eval); + ggml_free(ctx_ref); + + test_result result(ggml_backend_name(backend1), current_op_name, vars(), "test", supported, ok, + ok ? "" : "test failed"); + print_test_result_locked(output_printer, result); + + return ok ? test_status_t::OK : test_status_t::FAIL; + } + test_status_t eval(ggml_backend_t backend1, ggml_backend_t backend2, const char * op_names_filter, printer * output_printer) { + if (compare_with_reference_graph()) { + return eval_reference_graph(backend1, backend2, op_names_filter, output_printer); + } + mode = MODE_TEST; ggml_init_params params = { @@ -3709,6 +3856,395 @@ struct test_snake_fuse : public test_case { } }; + +struct test_dsv4_hc : public test_case { + static constexpr int64_t hc = 4; + + ggml_tensor * out = nullptr; + + bool compare_with_reference_graph() override { return true; } + + double err(const float * a, const float * b, size_t n) override { + double max_abs = 0.0; + for (size_t i = 0; i < n; ++i) { + max_abs = std::max(max_abs, fabsf(a[i] - b[i])); + } + return max_abs; + } + + double max_err() override { + return 1e-5; + } + + double max_err(ggml_backend_t backend) override { + GGML_UNUSED(backend); + return max_err(); + } + + static ggml_tensor * view_1d(ggml_context * ctx, ggml_tensor * t, int64_t ne0, int64_t i0) { + return ggml_view_1d(ctx, t, ne0, i0*ggml_element_size(t)); + } + + static ggml_tensor * view_2d(ggml_context * ctx, ggml_tensor * t, int64_t ne0, int64_t ne1, int64_t i0) { + return ggml_view_2d(ctx, t, ne0, ne1, t->nb[1], i0*ggml_element_size(t)); + } + + static ggml_tensor * hc_affine(ggml_context * ctx, ggml_tensor * x, ggml_tensor * scale, ggml_tensor * base) { + x = ggml_mul(ctx, x, scale); + x = ggml_add(ctx, x, base); + return x; + } + + static ggml_tensor * hc_sinkhorn_ref(ggml_context * ctx, ggml_tensor * comb, float eps, int32_t n_iter) { + comb = ggml_soft_max(ctx, comb); + + ggml_tensor * eps_t = ::ggml_new_tensor_1d(ctx, GGML_TYPE_F32, 1); + eps_t = ggml_fill(ctx, eps_t, eps); + + comb = ggml_add(ctx, comb, eps_t); + + auto norm_cols = [&]() { + ggml_tensor * comb_src_dst = ggml_cont(ctx, ggml_permute(ctx, comb, 1, 0, 2, 3)); + ggml_tensor * col_sum = ggml_sum_rows(ctx, comb_src_dst); + col_sum = ggml_add(ctx, col_sum, eps_t); + col_sum = ggml_permute(ctx, col_sum, 1, 0, 2, 3); + comb = ggml_div(ctx, comb, col_sum); + }; + + auto norm_rows = [&]() { + ggml_tensor * row_sum = ggml_sum_rows(ctx, comb); + row_sum = ggml_add(ctx, row_sum, eps_t); + comb = ggml_div(ctx, comb, row_sum); + }; + + norm_cols(); + for (int32_t i = 1; i < n_iter; ++i) { + norm_rows(); + norm_cols(); + } + + return comb; + } + + static uint32_t tensor_seed(const ggml_tensor * t) { + uint32_t seed = 2166136261u; + for (const char * p = ggml_get_name(t); *p; ++p) { + seed ^= (uint8_t) *p; + seed *= 16777619u; + } + for (int i = 0; i < GGML_MAX_DIMS; ++i) { + seed ^= (uint32_t) t->ne[i]; + seed *= 16777619u; + } + return seed; + } + + static bool tensor_range(const std::string & name, float & lo, float & hi) { + if (name.rfind("sent_", 0) == 0) { + lo = -1.0f; hi = 1.0f; return true; + } + if (name == "mixes") { + lo = -2.0f; hi = 2.0f; return true; + } + if (name == "scale") { + lo = -0.5f; hi = 0.5f; return true; + } + if (name == "base") { + lo = -0.25f; hi = 0.25f; return true; + } + if (name == "weights" || name == "comb") { + lo = 0.0f; hi = 1.0f; return true; + } + if (name == "post") { + lo = 0.0f; hi = 2.0f; return true; + } + if (name == "x" || name == "residual") { + lo = -1.0f; hi = 1.0f; return true; + } + return false; + } + + void initialize_tensors(ggml_context * ctx) override { + for (ggml_tensor * t = ggml_get_first_tensor(ctx); t != nullptr; t = ggml_get_next_tensor(ctx, t)) { + const std::string name = ggml_get_name(t); + float lo; + float hi; + if (!tensor_range(name, lo, hi)) { + continue; + } + + GGML_ASSERT(t->type == GGML_TYPE_F32); + std::mt19937 rng(tensor_seed(t)); + std::uniform_real_distribution dist(lo, hi); + std::vector data(ggml_nelements(t)); + for (float & v : data) { + v = dist(rng); + } + ggml_backend_tensor_set(t, data.data(), 0, data.size()*sizeof(float)); + } + } +}; + +struct test_dsv4_hc_comb : public test_dsv4_hc { + const int64_t n_tokens; + const int32_t n_iter; + const float eps; + + std::string op_desc(ggml_tensor * t) override { + GGML_UNUSED(t); + return "DSV4_HC_COMB"; + } + + std::string vars() override { + return VARS_TO_STR3(n_tokens, n_iter, eps); + } + + test_dsv4_hc_comb(int64_t n_tokens = 17, int32_t n_iter = 4, float eps = 1e-6f) + : n_tokens(n_tokens), n_iter(n_iter), eps(eps) {} + + ggml_tensor * build_graph(ggml_context * ctx) override { + ggml_tensor * mixes = ggml_new_tensor_2d(ctx, GGML_TYPE_F32, (2 + hc)*hc, n_tokens); + ggml_set_name(mixes, "mixes"); + + ggml_tensor * scale = ggml_new_tensor_1d(ctx, GGML_TYPE_F32, 3); + ggml_set_name(scale, "scale"); + + ggml_tensor * base = ggml_new_tensor_1d(ctx, GGML_TYPE_F32, (2 + hc)*hc); + ggml_set_name(base, "base"); + + out = ggml_dsv4_hc_comb(ctx, mixes, scale, base, eps, n_iter); + ggml_set_name(out, "out"); + return out; + } + + ggml_tensor * build_reference_graph(ggml_context * ctx) override { + ggml_tensor * mixes = ggml_new_tensor_2d(ctx, GGML_TYPE_F32, (2 + hc)*hc, n_tokens); + ggml_set_name(mixes, "mixes"); + + ggml_tensor * scale = ggml_new_tensor_1d(ctx, GGML_TYPE_F32, 3); + ggml_set_name(scale, "scale"); + + ggml_tensor * base = ggml_new_tensor_1d(ctx, GGML_TYPE_F32, (2 + hc)*hc); + ggml_set_name(base, "base"); + + ggml_tensor * scale_comb = view_1d(ctx, scale, 1, 2); + ggml_tensor * base_comb = view_1d(ctx, base, hc*hc, 2*hc); + out = view_2d(ctx, mixes, hc*hc, n_tokens, 2*hc); + out = hc_affine(ctx, out, scale_comb, base_comb); + out = ggml_reshape_3d(ctx, out, hc, hc, n_tokens); + out = hc_sinkhorn_ref(ctx, out, eps, n_iter); + ggml_set_name(out, "out"); + return out; + } +}; + +struct test_dsv4_hc_pre : public test_dsv4_hc { + const int64_t n_embd; + const int64_t n_tokens; + const bool model_graph; + + std::string op_desc(ggml_tensor * t) override { + GGML_UNUSED(t); + return "DSV4_HC_PRE"; + } + + std::string vars() override { + return VARS_TO_STR3(n_embd, n_tokens, model_graph); + } + + test_dsv4_hc_pre(int64_t n_embd = 31, int64_t n_tokens = 17, bool model_graph = false) + : n_embd(n_embd), n_tokens(n_tokens), model_graph(model_graph) {} + + ggml_tensor * build_x(ggml_context * ctx) { + if (!model_graph) { + ggml_tensor * x = ggml_new_tensor_3d(ctx, GGML_TYPE_F32, n_embd, hc, n_tokens); + ggml_set_name(x, "x"); + return x; + } + + ggml_tensor * sent = ggml_new_tensor_2d(ctx, GGML_TYPE_F32, n_embd, n_tokens); + ggml_set_name(sent, "sent_x"); + + ggml_tensor * x = ggml_reshape_3d(ctx, sent, n_embd, 1, n_tokens); + x = ggml_repeat_4d(ctx, x, n_embd, hc, n_tokens, 1); + return x; + } + + ggml_tensor * build_weights(ggml_context * ctx) { + if (!model_graph) { + ggml_tensor * weights = ggml_new_tensor_2d(ctx, GGML_TYPE_F32, hc, n_tokens); + ggml_set_name(weights, "weights"); + return weights; + } + + ggml_tensor * mixes = ggml_new_tensor_2d(ctx, GGML_TYPE_F32, (2 + hc)*hc, n_tokens); + ggml_set_name(mixes, "mixes"); + + ggml_tensor * scale = ggml_new_tensor_1d(ctx, GGML_TYPE_F32, 3); + ggml_set_name(scale, "scale"); + + ggml_tensor * base = ggml_new_tensor_1d(ctx, GGML_TYPE_F32, (2 + hc)*hc); + ggml_set_name(base, "base"); + + ggml_tensor * scale_pre = view_1d(ctx, scale, 1, 0); + ggml_tensor * base_pre = view_1d(ctx, base, hc, 0); + + ggml_tensor * weights = view_2d(ctx, mixes, hc, n_tokens, 0); + weights = hc_affine(ctx, weights, scale_pre, base_pre); + weights = ggml_sigmoid(ctx, weights); + weights = ggml_scale_bias(ctx, weights, 1.0f, 1e-6f); + return weights; + } + + ggml_tensor * build_graph(ggml_context * ctx) override { + ggml_tensor * x = build_x(ctx); + ggml_tensor * weights = build_weights(ctx); + + out = ggml_dsv4_hc_pre(ctx, x, weights); + ggml_set_name(out, "out"); + return out; + } + + ggml_tensor * build_reference_graph(ggml_context * ctx) override { + ggml_tensor * x = build_x(ctx); + ggml_tensor * weights = build_weights(ctx); + + out = nullptr; + for (int64_t ih = 0; ih < hc; ++ih) { + ggml_tensor * xh = ggml_view_2d(ctx, x, n_embd, n_tokens, x->nb[2], ih*x->nb[1]); + ggml_tensor * wh = ggml_view_2d(ctx, weights, 1, n_tokens, weights->nb[1], ih*weights->nb[0]); + ggml_tensor * cur = ggml_mul(ctx, xh, wh); + out = out ? ggml_add(ctx, out, cur) : cur; + } + + ggml_set_name(out, "out"); + return out; + } +}; + +struct test_dsv4_hc_pre_ref : public test_dsv4_hc_pre { + test_dsv4_hc_pre_ref(int64_t n_embd = 31, int64_t n_tokens = 17, bool model_graph = false) + : test_dsv4_hc_pre(n_embd, n_tokens, model_graph) {} + + std::string op_desc(ggml_tensor * t) override { + GGML_UNUSED(t); + return "DSV4_HC_PRE_REF"; + } + + ggml_tensor * build_graph(ggml_context * ctx) override { + return build_reference_graph(ctx); + } +}; + +struct test_dsv4_hc_post : public test_dsv4_hc { + const int64_t n_embd; + const int64_t n_tokens; + + std::string op_desc(ggml_tensor * t) override { + GGML_UNUSED(t); + return "DSV4_HC_POST"; + } + + std::string vars() override { + return VARS_TO_STR2(n_embd, n_tokens); + } + + test_dsv4_hc_post(int64_t n_embd = 31, int64_t n_tokens = 17) + : n_embd(n_embd), n_tokens(n_tokens) {} + + ggml_tensor * build_graph(ggml_context * ctx) override { + ggml_tensor * x = ggml_new_tensor_2d(ctx, GGML_TYPE_F32, n_embd, n_tokens); + ggml_set_name(x, "x"); + + ggml_tensor * residual = ggml_new_tensor_3d(ctx, GGML_TYPE_F32, n_embd, hc, n_tokens); + ggml_set_name(residual, "residual"); + + ggml_tensor * post = ggml_new_tensor_2d(ctx, GGML_TYPE_F32, hc, n_tokens); + ggml_set_name(post, "post"); + + ggml_tensor * comb = ggml_new_tensor_3d(ctx, GGML_TYPE_F32, hc, hc, n_tokens); + ggml_set_name(comb, "comb"); + + out = ggml_dsv4_hc_post(ctx, x, residual, post, comb); + ggml_set_name(out, "out"); + return out; + } + + ggml_tensor * build_reference_graph(ggml_context * ctx) override { + ggml_tensor * x = ggml_new_tensor_2d(ctx, GGML_TYPE_F32, n_embd, n_tokens); + ggml_set_name(x, "x"); + + ggml_tensor * residual = ggml_new_tensor_3d(ctx, GGML_TYPE_F32, n_embd, hc, n_tokens); + ggml_set_name(residual, "residual"); + + ggml_tensor * post = ggml_new_tensor_2d(ctx, GGML_TYPE_F32, hc, n_tokens); + ggml_set_name(post, "post"); + + ggml_tensor * comb = ggml_new_tensor_3d(ctx, GGML_TYPE_F32, hc, hc, n_tokens); + ggml_set_name(comb, "comb"); + + out = nullptr; + for (int64_t dst = 0; dst < hc; ++dst) { + ggml_tensor * post_dst = ggml_view_2d(ctx, post, 1, n_tokens, post->nb[1], dst*post->nb[0]); + ggml_tensor * cur = ggml_mul(ctx, x, post_dst); + + for (int64_t src = 0; src < hc; ++src) { + ggml_tensor * res_src = ggml_view_2d(ctx, residual, n_embd, n_tokens, residual->nb[2], src*residual->nb[1]); + ggml_tensor * comb_src_dst = ggml_view_2d(ctx, comb, 1, n_tokens, comb->nb[2], + dst*comb->nb[0] + src*comb->nb[1]); + cur = ggml_add(ctx, cur, ggml_mul(ctx, res_src, comb_src_dst)); + } + + cur = ggml_reshape_3d(ctx, cur, n_embd, 1, n_tokens); + out = out ? ggml_concat(ctx, out, cur, 1) : cur; + } + + ggml_set_name(out, "out"); + return out; + } +}; + +struct test_lightning_indexer : public test_case { + const int64_t n_batch; + const int64_t n_kv; + const int64_t n_stream; + + std::string vars() override { + return VARS_TO_STR3(n_batch, n_kv, n_stream); + } + + double err(const float * a, const float * b, size_t n) override { + double max_abs = 0.0; + for (size_t i = 0; i < n; ++i) { + max_abs = std::max(max_abs, fabsf(a[i] - b[i])); + } + return max_abs; + } + + double max_err() override { + return 1e-4; + } + + test_lightning_indexer(int64_t n_batch = 3, int64_t n_kv = 65, int64_t n_stream = 2) + : n_batch(n_batch), n_kv(n_kv), n_stream(n_stream) {} + + ggml_tensor * build_graph(ggml_context * ctx) override { + ggml_tensor * q = ggml_new_tensor_4d(ctx, GGML_TYPE_F32, 128, 64, n_batch, n_stream); + ggml_set_name(q, "q"); + + ggml_tensor * k = ggml_new_tensor_4d(ctx, GGML_TYPE_F32, 128, 1, n_kv, n_stream); + ggml_set_name(k, "k"); + + ggml_tensor * weights = ggml_new_tensor_4d(ctx, GGML_TYPE_F32, 64, n_batch, 1, n_stream); + ggml_set_name(weights, "weights"); + + ggml_tensor * out = ggml_lightning_indexer(ctx, q, k, weights, 1.0f/sqrtf(128.0f), 1.0f/sqrtf(64.0f)); + ggml_set_name(out, "out"); + return out; + } +}; + + // GGML_OP_SSM_CONV struct test_ssm_conv : public test_case { const ggml_type type; @@ -7710,6 +8246,22 @@ static std::vector> make_test_cases_eval() { test_cases.emplace_back(new test_snake_fuse(type, { 64, 32, 2, 3})); // ne[2] > 1 and ne[3] > 1 } + test_cases.emplace_back(new test_dsv4_hc_comb(1, 1)); + test_cases.emplace_back(new test_dsv4_hc_comb(17, 4)); + test_cases.emplace_back(new test_dsv4_hc_comb(257, 8)); + + test_cases.emplace_back(new test_dsv4_hc_pre(1, 1)); + test_cases.emplace_back(new test_dsv4_hc_pre(31, 17)); + test_cases.emplace_back(new test_dsv4_hc_pre(128, 257)); + test_cases.emplace_back(new test_dsv4_hc_pre(4096, 21, true)); + test_cases.emplace_back(new test_dsv4_hc_pre_ref(4096, 21, true)); + + test_cases.emplace_back(new test_dsv4_hc_post(1, 1)); + test_cases.emplace_back(new test_dsv4_hc_post(31, 17)); + test_cases.emplace_back(new test_dsv4_hc_post(128, 257)); + + test_cases.emplace_back(new test_lightning_indexer()); + // glu ops for (ggml_type type : {GGML_TYPE_F16, GGML_TYPE_F32}) { for (int v : {0, 1}) {