ggml : add GGML_OP_LIGHTNING_INDEXER that implements DeepSeek V3.2/V4 lightning indexer (llama/24231)
* ggml : add GGML_OP_LIGHTNING_INDEXER that implements DeepSeek V3.2/V4 lightning indexer * ggml : remove scale parameters from lightning indexer OP, add f16 mask parameter * tests : add GGML_OP_LIGHTNING_INDEXER tests * ggml : bump RPC version * chore : check if lightning indexer input tensors are not transposed * tests : count flops instead of bandwidth in lightning indexer test * chore : add missing const * chore : whitespace * ggml : renamed variables in CPU lightning indexer implementation * ggml : fix lightning indexer mask broadcasting * tests : tests for lightning indexer mask broadcasting * chore : whitespace * llama : use GGML_OP_LIGHTNING_INDEXER in DeepSeek V3.2 and DeepSeek V4 models --------- Co-authored-by: Stanisław Szymczyk <sszymczy@gmail.com>
This commit is contained in:
parent
4bc66b579f
commit
19c96ad196
|
|
@ -8,10 +8,10 @@ extern "C" {
|
||||||
|
|
||||||
#define RPC_PROTO_MAJOR_VERSION 4
|
#define RPC_PROTO_MAJOR_VERSION 4
|
||||||
#define RPC_PROTO_MINOR_VERSION 0
|
#define RPC_PROTO_MINOR_VERSION 0
|
||||||
#define RPC_PROTO_PATCH_VERSION 1
|
#define RPC_PROTO_PATCH_VERSION 2
|
||||||
|
|
||||||
#ifdef __cplusplus
|
#ifdef __cplusplus
|
||||||
static_assert(GGML_OP_COUNT == 97, "GGML_OP_COUNT has changed - update RPC_PROTO_PATCH_VERSION");
|
static_assert(GGML_OP_COUNT == 98, "GGML_OP_COUNT has changed - update RPC_PROTO_PATCH_VERSION");
|
||||||
#endif
|
#endif
|
||||||
|
|
||||||
#define GGML_RPC_MAX_SERVERS 16
|
#define GGML_RPC_MAX_SERVERS 16
|
||||||
|
|
|
||||||
|
|
@ -570,6 +570,7 @@ extern "C" {
|
||||||
GGML_OP_RWKV_WKV7,
|
GGML_OP_RWKV_WKV7,
|
||||||
GGML_OP_SOLVE_TRI,
|
GGML_OP_SOLVE_TRI,
|
||||||
GGML_OP_GATED_DELTA_NET,
|
GGML_OP_GATED_DELTA_NET,
|
||||||
|
GGML_OP_LIGHTNING_INDEXER,
|
||||||
|
|
||||||
GGML_OP_UNARY,
|
GGML_OP_UNARY,
|
||||||
|
|
||||||
|
|
@ -2575,6 +2576,24 @@ extern "C" {
|
||||||
struct ggml_tensor * state,
|
struct ggml_tensor * state,
|
||||||
int64_t K);
|
int64_t K);
|
||||||
|
|
||||||
|
// DSA lightning indexer
|
||||||
|
//
|
||||||
|
// q: [n_embd_idx, n_head_idx, n_batch, ne3 ]
|
||||||
|
// k: [n_embd_idx, 1, n_kv, ne3 ]
|
||||||
|
// weights: [n_head_idx, n_batch, 1, ne3 ] !! prescaled !!
|
||||||
|
// mask: [n_kv, n_batch, 1, ne33] !! f16 !!
|
||||||
|
// res: [n_kv, n_batch, 1, ne3 ]
|
||||||
|
//
|
||||||
|
// broadcast:
|
||||||
|
// ne3 % ne33 == 0
|
||||||
|
//
|
||||||
|
GGML_API struct ggml_tensor * ggml_lightning_indexer(
|
||||||
|
struct ggml_context * ctx,
|
||||||
|
struct ggml_tensor * q,
|
||||||
|
struct ggml_tensor * k,
|
||||||
|
struct ggml_tensor * weights,
|
||||||
|
struct ggml_tensor * mask);
|
||||||
|
|
||||||
// custom operators
|
// custom operators
|
||||||
|
|
||||||
typedef void (*ggml_custom1_op_t)(struct ggml_tensor * dst , const struct ggml_tensor * a, int ith, int nth, void * userdata);
|
typedef void (*ggml_custom1_op_t)(struct ggml_tensor * dst , const struct ggml_tensor * a, int ith, int nth, void * userdata);
|
||||||
|
|
|
||||||
|
|
@ -2060,6 +2060,10 @@ static void ggml_compute_forward(struct ggml_compute_params * params, struct ggm
|
||||||
{
|
{
|
||||||
ggml_compute_forward_gated_delta_net(params, tensor);
|
ggml_compute_forward_gated_delta_net(params, tensor);
|
||||||
} break;
|
} break;
|
||||||
|
case GGML_OP_LIGHTNING_INDEXER:
|
||||||
|
{
|
||||||
|
ggml_compute_forward_lightning_indexer(params, tensor);
|
||||||
|
} break;
|
||||||
case GGML_OP_MAP_CUSTOM1:
|
case GGML_OP_MAP_CUSTOM1:
|
||||||
{
|
{
|
||||||
ggml_compute_forward_map_custom1(params, tensor);
|
ggml_compute_forward_map_custom1(params, tensor);
|
||||||
|
|
@ -2380,6 +2384,7 @@ static int ggml_get_n_tasks(struct ggml_tensor * node, int n_threads) {
|
||||||
case GGML_OP_FLASH_ATTN_BACK:
|
case GGML_OP_FLASH_ATTN_BACK:
|
||||||
case GGML_OP_SSM_CONV:
|
case GGML_OP_SSM_CONV:
|
||||||
case GGML_OP_SSM_SCAN:
|
case GGML_OP_SSM_SCAN:
|
||||||
|
case GGML_OP_LIGHTNING_INDEXER:
|
||||||
{
|
{
|
||||||
n_tasks = n_threads;
|
n_tasks = n_threads;
|
||||||
} break;
|
} break;
|
||||||
|
|
@ -2965,6 +2970,12 @@ struct ggml_cplan ggml_graph_plan(
|
||||||
{
|
{
|
||||||
GGML_ABORT("fatal error");
|
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:
|
default:
|
||||||
break;
|
break;
|
||||||
}
|
}
|
||||||
|
|
|
||||||
|
|
@ -11568,3 +11568,87 @@ 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 * q = dst->src[0];
|
||||||
|
const ggml_tensor * k = dst->src[1];
|
||||||
|
const ggml_tensor * w = dst->src[2]; // weights
|
||||||
|
const ggml_tensor * m = dst->src[3]; // mask
|
||||||
|
|
||||||
|
GGML_ASSERT(dst->type == GGML_TYPE_F32);
|
||||||
|
GGML_ASSERT( q->type == GGML_TYPE_F32);
|
||||||
|
GGML_ASSERT( w->type == GGML_TYPE_F32);
|
||||||
|
GGML_ASSERT( m->type == GGML_TYPE_F16);
|
||||||
|
|
||||||
|
GGML_TENSOR_LOCALS(int64_t, neq, q, ne)
|
||||||
|
GGML_TENSOR_LOCALS(size_t, nbq, q, nb)
|
||||||
|
GGML_TENSOR_LOCALS(int64_t, nek, k, ne)
|
||||||
|
GGML_TENSOR_LOCALS(size_t, nbk, k, nb)
|
||||||
|
GGML_TENSOR_LOCALS(int64_t, new, w, ne)
|
||||||
|
GGML_TENSOR_LOCALS(size_t, nbw, w, nb)
|
||||||
|
GGML_TENSOR_LOCALS(int64_t, nem, m, ne)
|
||||||
|
GGML_TENSOR_LOCALS(size_t, nbm, m, nb)
|
||||||
|
GGML_TENSOR_LOCALS(int64_t, ne, dst, ne)
|
||||||
|
GGML_TENSOR_LOCALS(size_t, nb, dst, nb)
|
||||||
|
|
||||||
|
GGML_ASSERT( nb0 == ggml_type_size(dst->type));
|
||||||
|
GGML_ASSERT(nbq0 == ggml_type_size( q->type));
|
||||||
|
GGML_ASSERT(nbk0 == ggml_type_size( k->type));
|
||||||
|
GGML_ASSERT(nbw0 == ggml_type_size( w->type));
|
||||||
|
GGML_ASSERT(nbm0 == ggml_type_size( m->type));
|
||||||
|
|
||||||
|
const int n_embd = q->ne[0];
|
||||||
|
const int n_head = q->ne[1];
|
||||||
|
const int n_tokens = q->ne[2];
|
||||||
|
const int n_stream = q->ne[3];
|
||||||
|
const int n_kv = k->ne[2];
|
||||||
|
|
||||||
|
ggml_to_float_t const k_to_float = ggml_get_type_traits(k->type)->to_float;
|
||||||
|
GGML_ASSERT((k->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 * k_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 s = 0; s < n_stream; ++s) {
|
||||||
|
for (int t = 0; t < n_tokens; ++t) {
|
||||||
|
const float * w_row = (float *) ((char *) w->data + t*nbw1 + s*nbw3);
|
||||||
|
const ggml_fp16_t * m_row = (ggml_fp16_t *) ((char *) m->data + t*nbm1 + (s%nem3)*nbm3);
|
||||||
|
float * dst_row = (float *) ((char *) dst->data + t*nb1 + s*nb3 );
|
||||||
|
for (int ik = ir0; ik < ir1; ++ik) {
|
||||||
|
char * k_row = (char *) k->data + ik*nbk2 + s*nbk3;
|
||||||
|
if (k_to_float) {
|
||||||
|
k_to_float(k_row, k_row_f32, n_embd);
|
||||||
|
} else {
|
||||||
|
k_row_f32 = (float *) k_row;
|
||||||
|
}
|
||||||
|
float score = 0.0f;
|
||||||
|
for (int h = 0; h < n_head; ++h) {
|
||||||
|
// dot product of q and k for head h
|
||||||
|
float qk = 0.0f;
|
||||||
|
const float * q_row = (float *) ((char *) q->data + h*nbq1 + t*nbq2 + s*nbq3);
|
||||||
|
ggml_vec_dot_f32(n_embd, &qk, 0, q_row, 0, k_row_f32, 0, 1);
|
||||||
|
// ReLU and weights (prescaled)
|
||||||
|
score += MAX(qk, 0.0f) * w_row[h];
|
||||||
|
}
|
||||||
|
// apply mask
|
||||||
|
dst_row[ik] = score + GGML_CPU_FP16_TO_FP32(m_row[ik]);
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
|
||||||
|
|
@ -105,6 +105,7 @@ 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_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_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_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_map_custom1(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_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);
|
void ggml_compute_forward_map_custom3(const struct ggml_compute_params * params, struct ggml_tensor * dst);
|
||||||
|
|
|
||||||
|
|
@ -1079,6 +1079,7 @@ static const char * GGML_OP_NAME[GGML_OP_COUNT] = {
|
||||||
"RWKV_WKV7",
|
"RWKV_WKV7",
|
||||||
"SOLVE_TRI",
|
"SOLVE_TRI",
|
||||||
"GATED_DELTA_NET",
|
"GATED_DELTA_NET",
|
||||||
|
"LIGHTNING_INDEXER",
|
||||||
|
|
||||||
"UNARY",
|
"UNARY",
|
||||||
|
|
||||||
|
|
@ -1096,7 +1097,7 @@ static const char * GGML_OP_NAME[GGML_OP_COUNT] = {
|
||||||
"GLU",
|
"GLU",
|
||||||
};
|
};
|
||||||
|
|
||||||
static_assert(GGML_OP_COUNT == 97, "GGML_OP_COUNT != 97");
|
static_assert(GGML_OP_COUNT == 98, "GGML_OP_COUNT != 98");
|
||||||
|
|
||||||
static const char * GGML_OP_SYMBOL[GGML_OP_COUNT] = {
|
static const char * GGML_OP_SYMBOL[GGML_OP_COUNT] = {
|
||||||
"none",
|
"none",
|
||||||
|
|
@ -1190,6 +1191,7 @@ static const char * GGML_OP_SYMBOL[GGML_OP_COUNT] = {
|
||||||
"rwkv_wkv7(r, w, k, v, a, b, s)",
|
"rwkv_wkv7(r, w, k, v, a, b, s)",
|
||||||
"A X = B, A triangular, solve X",
|
"A X = B, A triangular, solve X",
|
||||||
"gated_delta_net(q, k, v, g, beta, s)",
|
"gated_delta_net(q, k, v, g, beta, s)",
|
||||||
|
"lightning_indexer(q, k, weights, mask)",
|
||||||
|
|
||||||
"unary(x)",
|
"unary(x)",
|
||||||
|
|
||||||
|
|
@ -1207,7 +1209,7 @@ static const char * GGML_OP_SYMBOL[GGML_OP_COUNT] = {
|
||||||
"glu(x)",
|
"glu(x)",
|
||||||
};
|
};
|
||||||
|
|
||||||
static_assert(GGML_OP_COUNT == 97, "GGML_OP_COUNT != 97");
|
static_assert(GGML_OP_COUNT == 98, "GGML_OP_COUNT != 98");
|
||||||
|
|
||||||
static_assert(GGML_OP_POOL_COUNT == 2, "GGML_OP_POOL_COUNT != 2");
|
static_assert(GGML_OP_POOL_COUNT == 2, "GGML_OP_POOL_COUNT != 2");
|
||||||
|
|
||||||
|
|
@ -6287,6 +6289,42 @@ struct ggml_tensor * ggml_gated_delta_net(
|
||||||
return result;
|
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,
|
||||||
|
struct ggml_tensor * mask) {
|
||||||
|
|
||||||
|
GGML_ASSERT( q->type == GGML_TYPE_F32);
|
||||||
|
GGML_ASSERT( weights->type == GGML_TYPE_F32);
|
||||||
|
GGML_ASSERT( mask->type == GGML_TYPE_F16);
|
||||||
|
GGML_ASSERT( q->ne[0] == k->ne[0]);
|
||||||
|
GGML_ASSERT( mask->ne[0] == k->ne[2]);
|
||||||
|
GGML_ASSERT( q->ne[1] == weights->ne[0]);
|
||||||
|
GGML_ASSERT( k->ne[1] == 1);
|
||||||
|
GGML_ASSERT( mask->ne[1] == q->ne[2]);
|
||||||
|
GGML_ASSERT( q->ne[2] == weights->ne[1]);
|
||||||
|
GGML_ASSERT(weights->ne[2] == 1);
|
||||||
|
GGML_ASSERT( mask->ne[2] == 1);
|
||||||
|
GGML_ASSERT( q->ne[3] == k->ne[3]);
|
||||||
|
GGML_ASSERT( k->ne[3] == weights->ne[3]);
|
||||||
|
GGML_ASSERT(weights->ne[3] % mask->ne[3] == 0);
|
||||||
|
|
||||||
|
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);
|
||||||
|
|
||||||
|
result->op = GGML_OP_LIGHTNING_INDEXER;
|
||||||
|
result->src[0] = q;
|
||||||
|
result->src[1] = k;
|
||||||
|
result->src[2] = weights;
|
||||||
|
result->src[3] = mask;
|
||||||
|
|
||||||
|
return result;
|
||||||
|
}
|
||||||
|
|
||||||
////////////////////////////////////////////////////////////////////////////////
|
////////////////////////////////////////////////////////////////////////////////
|
||||||
|
|
||||||
struct ggml_hash_set ggml_hash_set_new(size_t size) {
|
struct ggml_hash_set ggml_hash_set_new(size_t size) {
|
||||||
|
|
|
||||||
Loading…
Reference in New Issue