Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
4 changes: 2 additions & 2 deletions ggml/include/ggml-rpc.h
Original file line number Diff line number Diff line change
Expand Up @@ -8,10 +8,10 @@ extern "C" {

#define RPC_PROTO_MAJOR_VERSION 7
#define RPC_PROTO_MINOR_VERSION 0
#define RPC_PROTO_PATCH_VERSION 0
#define RPC_PROTO_PATCH_VERSION 1

#ifdef __cplusplus
static_assert(GGML_OP_COUNT == 101, "GGML_OP_COUNT has changed - update RPC_PROTO_PATCH_VERSION");
static_assert(GGML_OP_COUNT == 102, "GGML_OP_COUNT has changed - update RPC_PROTO_PATCH_VERSION");
#endif

#define GGML_RPC_MAX_SERVERS 16
Expand Down
18 changes: 18 additions & 0 deletions ggml/include/ggml.h
Original file line number Diff line number Diff line change
Expand Up @@ -601,6 +601,10 @@ extern "C" {

GGML_OP_GLU,

// Project-local inference op: select the recurrent-state bank row on device.
// Appended to preserve all existing operation IDs.
GGML_OP_GATED_DELTA_NET_INDEXED,

GGML_OP_COUNT,
};

Expand Down Expand Up @@ -2659,6 +2663,20 @@ extern "C" {
struct ggml_tensor * state,
int64_t K);

// Read the initial recurrent state directly from a persistent state bank
// using one I32 row index per sequence.
GGML_API struct ggml_tensor * ggml_gated_delta_net_indexed(
struct ggml_context * ctx,
struct ggml_tensor * q,
struct ggml_tensor * k,
struct ggml_tensor * v,
struct ggml_tensor * g,
struct ggml_tensor * beta,
struct ggml_tensor * bank,
struct ggml_tensor * rows,
struct ggml_tensor * state_dependency,
int64_t K);

// DSA lightning indexer
//
// q: [n_embd_idx, n_head_idx, n_batch, ne3 ]
Expand Down
3 changes: 3 additions & 0 deletions ggml/src/ggml-cpu/ggml-cpu.c
Original file line number Diff line number Diff line change
Expand Up @@ -2090,6 +2090,7 @@ static void ggml_compute_forward(struct ggml_compute_params * params, struct ggm
ggml_compute_forward_solve_tri(params, tensor);
} break;
case GGML_OP_GATED_DELTA_NET:
case GGML_OP_GATED_DELTA_NET_INDEXED:
{
ggml_compute_forward_gated_delta_net(params, tensor);
} break;
Expand Down Expand Up @@ -2289,6 +2290,7 @@ 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_GATED_DELTA_NET_INDEXED:
case GGML_OP_DSV4_HC_COMB:
case GGML_OP_DSV4_HC_PRE:
case GGML_OP_DSV4_HC_POST:
Expand Down Expand Up @@ -3031,6 +3033,7 @@ struct ggml_cplan ggml_graph_plan(
cur = ggml_type_size(node->type)*(n_tasks + node->src[0]->ne[0]*n_tasks);
} break;
case GGML_OP_GATED_DELTA_NET:
case GGML_OP_GATED_DELTA_NET_INDEXED:
{
const int64_t S_v = node->src[2]->ne[0];
const int64_t K = ggml_get_op_params_i32(node, 0);
Expand Down
14 changes: 11 additions & 3 deletions ggml/src/ggml-cpu/ops.cpp
Original file line number Diff line number Diff line change
Expand Up @@ -10980,9 +10980,17 @@ static void ggml_compute_forward_gated_delta_net_one_chunk(
? state_work
: state_out_base + (iv3 * H + iv1) * S_v * S_v;

// copy input state into the working buffer and operate in-place
// state layout [S_v, S_v, H, n_seqs]: seq iv3 starts at iv3 * state_seq_stride.
const float * s_in = state_in_base + iv3 * state_seq_stride + iv1 * S_v * S_v;
// copy input state into the working buffer and operate in-place.
// Ordinary GDN addresses the state by sequence. The indexed variant reads
// the dynamic row selector from src[6] so graph replay can reuse the same
// graph while selecting a different persistent recurrent-state row.
int64_t state_row = iv3;
if (dst->op == GGML_OP_GATED_DELTA_NET_INDEXED) {
GGML_ASSERT(dst->src[6] && dst->src[6]->type == GGML_TYPE_I32);
state_row = ((const int32_t *) dst->src[6]->data)[iv3];
GGML_ASSERT(state_row >= 0 && state_row < src_state->ne[3]);
}
const float * s_in = state_in_base + state_row * state_seq_stride + iv1 * S_v * S_v;
memcpy(s_out, s_in, S_v * S_v * sizeof(float));

// attn output pointer for first token of this (head, seq)
Expand Down
48 changes: 39 additions & 9 deletions ggml/src/ggml-cuda/gated_delta_net.cu
Original file line number Diff line number Diff line change
@@ -1,7 +1,7 @@
#include "gated_delta_net.cuh"
#include "ggml-cuda/common.cuh"

template <int S_v, bool KDA, bool keep_rs_t>
template <int S_v, bool KDA, bool keep_rs_t, bool indexed_state_t = false>
__global__ void __launch_bounds__((ggml_cuda_get_physical_warp_size() < S_v ? ggml_cuda_get_physical_warp_size() : S_v) * 4, 2)
gated_delta_net_cuda(const float * q,
const float * k,
Expand All @@ -27,7 +27,8 @@ gated_delta_net_cuda(const float * q,
const uint3 rq3_magic,
float scale,
int64_t state_slot_stride,
int K) {
int K,
const int32_t * state_rows = nullptr) {
const uint32_t h_idx = blockIdx.x;
const uint32_t sequence = blockIdx.y;
// each warp owns one column, using warp-level primitives to reduce across rows
Expand All @@ -39,9 +40,13 @@ gated_delta_net_cuda(const float * q,

float * attn_data = dst;

// input state holds s0 only: [S_v, S_v, H, n_seqs] — seq stride is D = H * S_v * S_v.
// output state layout (per-slot D * n_seqs) — same per-(seq,head) offset as before.
const int64_t state_in_offset = sequence * H * S_v * S_v + h_idx * S_v * S_v;
// Ordinary input state is [S_v,S_v,H,n_seqs]. The indexed variant uses
// a persistent state bank and a dynamic row selector.
int64_t state_row = sequence;
if constexpr (indexed_state_t) {
state_row = state_rows[sequence];
}
const int64_t state_in_offset = state_row * H * S_v * S_v + h_idx * S_v * S_v;
const int64_t state_out_offset = (sequence * H + h_idx) * S_v * S_v;
state += state_out_offset;
curr_state += state_in_offset + col * S_v;
Expand Down Expand Up @@ -192,26 +197,26 @@ static void launch_gated_delta_net(
ggml_cuda_kernel_launch(gated_delta_net_cuda<16, KDA, keep_rs_t>, launch_params,
q_d, k_d, v_d, g_d, b_d, s_d, dst_d, state_d, H,
n_tokens, n_seqs, sq1, sq2, sq3, sv1, sv2, sv3,
sb1, sb2, sb3, neqk1_magic, rq3_magic, scale, state_slot_stride, K);
sb1, sb2, sb3, neqk1_magic, rq3_magic, scale, state_slot_stride, K, (const int32_t *) nullptr);
break;
case 32:
ggml_cuda_kernel_launch(gated_delta_net_cuda<32, KDA, keep_rs_t>, launch_params,
q_d, k_d, v_d, g_d, b_d, s_d, dst_d, state_d, H,
n_tokens, n_seqs, sq1, sq2, sq3, sv1, sv2, sv3,
sb1, sb2, sb3, neqk1_magic, rq3_magic, scale, state_slot_stride, K);
sb1, sb2, sb3, neqk1_magic, rq3_magic, scale, state_slot_stride, K, (const int32_t *) nullptr);
break;
case 64: {
ggml_cuda_kernel_launch(gated_delta_net_cuda<64, KDA, keep_rs_t>, launch_params,
q_d, k_d, v_d, g_d, b_d, s_d, dst_d, state_d, H,
n_tokens, n_seqs, sq1, sq2, sq3, sv1, sv2, sv3,
sb1, sb2, sb3, neqk1_magic, rq3_magic, scale, state_slot_stride, K);
sb1, sb2, sb3, neqk1_magic, rq3_magic, scale, state_slot_stride, K, (const int32_t *) nullptr);
break;
}
case 128: {
ggml_cuda_kernel_launch(gated_delta_net_cuda<128, KDA, keep_rs_t>, launch_params,
q_d, k_d, v_d, g_d, b_d, s_d, dst_d, state_d, H,
n_tokens, n_seqs, sq1, sq2, sq3, sv1, sv2, sv3,
sb1, sb2, sb3, neqk1_magic, rq3_magic, scale, state_slot_stride, K);
sb1, sb2, sb3, neqk1_magic, rq3_magic, scale, state_slot_stride, K, (const int32_t *) nullptr);
break;
}
default:
Expand Down Expand Up @@ -294,6 +299,31 @@ static void ggml_cuda_op_gated_delta_net_impl(
state_slot_stride = cache->slot_stride;
}

if (dst->op == GGML_OP_GATED_DELTA_NET_INDEXED) {
#if defined(GGML_USE_HIP)
GGML_ASSERT(GGML_CUDA_CC_IS_RDNA4(ggml_cuda_info().devices[ctx.device].cc));
GGML_ASSERT(S_v == 128 && H == 48 && n_seqs == 1 && K == 3 && !kda);
GGML_ASSERT(n_tokens >= 1 && n_tokens <= 3 && neqk1 == 16 && rq3 == 1);
GGML_ASSERT(dst->src[6] && dst->src[6]->type == GGML_TYPE_I32 && ggml_is_contiguous(dst->src[6]));

const int warp_size = ggml_cuda_info().devices[ctx.device].warp_size;
const int num_warps = 4;
const ggml_cuda_kernel_launch_params launch_params(
dim3(H, n_seqs, (S_v + num_warps - 1) / num_warps),
dim3(warp_size <= S_v ? warp_size : S_v, num_warps, 1),
0, stream);

ggml_cuda_kernel_launch(gated_delta_net_cuda<128, false, true, true>, launch_params,
q_d, k_d, v_d, g_d, b_d, s_d, dst_d, state_d, H,
n_tokens, n_seqs, sq1, sq2, sq3, sv1, sv2, sv3,
sb1, sb2, sb3, init_fastdiv_values(neqk1), init_fastdiv_values(rq3),
scale, state_slot_stride, K, (const int32_t *) dst->src[6]->data);
return;
#else
GGML_ABORT("indexed GDN requires the HIP RDNA4 backend");
#endif
}

if (kda) {
if (keep_rs) {
launch_gated_delta_net<true, true>(q_d, k_d, v_d, g_d, b_d, s_d, dst_d, state_d,
Expand Down
34 changes: 34 additions & 0 deletions ggml/src/ggml-cuda/ggml-cuda.cu
Original file line number Diff line number Diff line change
Expand Up @@ -2390,6 +2390,7 @@ static bool ggml_cuda_compute_forward(ggml_backend_cuda_context & ctx, struct gg
ggml_cuda_op_gated_linear_attn(ctx, dst);
break;
case GGML_OP_GATED_DELTA_NET:
case GGML_OP_GATED_DELTA_NET_INDEXED:
ggml_cuda_op_gated_delta_net(ctx, dst);
break;
case GGML_OP_DSV4_HC_COMB:
Expand Down Expand Up @@ -5492,6 +5493,18 @@ static bool ggml_backend_cuda_device_supports_op(ggml_backend_dev_t dev, const g
case GGML_OP_GATED_LINEAR_ATTN:
case GGML_OP_RWKV_WKV7:
return true;
case GGML_OP_GATED_DELTA_NET_INDEXED:
#if defined(GGML_USE_HIP)
return GGML_CUDA_CC_IS_RDNA4(ggml_cuda_info().devices[dev_ctx->device].cc)
&& op->src[2]->ne[0] == 128 && op->src[2]->ne[1] == 48 && op->src[2]->ne[3] == 1
&& op->src[2]->ne[2] >= 1 && op->src[2]->ne[2] <= 3
&& op->src[0]->ne[1] == 16 && op->src[0]->ne[3] == 1
&& op->src[3]->ne[0] == 1 && ggml_get_op_params_i32(op, 0) == 3
&& op->src[5]->type == GGML_TYPE_F32 && ggml_is_contiguous(op->src[5])
&& op->src[6] && op->src[6]->type == GGML_TYPE_I32 && ggml_is_contiguous(op->src[6]);
#else
return false;
#endif
case GGML_OP_GATED_DELTA_NET:
//TODO: enable once MUSA compiler is solved https://github.com/ggml-org/llama.cpp/pull/19504#issuecomment-4018634327
#ifdef GGML_USE_MUSA
Expand Down Expand Up @@ -5684,8 +5697,29 @@ static ggml_backend_feature * ggml_backend_cuda_get_features(ggml_backend_reg_t
GGML_UNUSED(reg);
}

static bool ggml_backend_rocm_gdn_indexed_bank_supported(ggml_backend_buffer_type_t buft) {
#if defined(GGML_USE_HIP)
if (!buft || !ggml_backend_buft_is_cuda(buft)) {
return false;
}
const auto dev = ggml_backend_buft_get_device(buft);
if (!dev) {
return false;
}
const auto * dev_ctx = (const ggml_backend_cuda_device_context *) dev->context;
return buft == ggml_backend_cuda_buffer_type(dev_ctx->device)
&& GGML_CUDA_CC_IS_RDNA4(ggml_cuda_info().devices[dev_ctx->device].cc);
#else
GGML_UNUSED(buft);
return false;
#endif
}

static void * ggml_backend_cuda_reg_get_proc_address(ggml_backend_reg_t reg, const char * name) {
GGML_UNUSED(reg);
if (strcmp(name, "ggml_backend_rocm_gdn_indexed_bank_supported") == 0) {
return (void *) ggml_backend_rocm_gdn_indexed_bank_supported;
}
if (strcmp(name, "ggml_backend_comm_init") == 0) {
return (void *)ggml_backend_cuda_comm_init;
}
Expand Down
36 changes: 34 additions & 2 deletions ggml/src/ggml.c
Original file line number Diff line number Diff line change
Expand Up @@ -1099,9 +1099,10 @@ static const char * GGML_OP_NAME[GGML_OP_COUNT] = {
"OPT_STEP_SGD",

"GLU",
"GATED_DELTA_NET_INDEXED",
};

static_assert(GGML_OP_COUNT == 101, "GGML_OP_COUNT != 101");
static_assert(GGML_OP_COUNT == 102, "GGML_OP_COUNT != 102");

static const char * GGML_OP_SYMBOL[GGML_OP_COUNT] = {
"none",
Expand Down Expand Up @@ -1214,9 +1215,10 @@ static const char * GGML_OP_SYMBOL[GGML_OP_COUNT] = {
"sgd(x)",

"glu(x)",
"gated_delta_net_indexed(q,k,v,g,b,bank,rows)",
};

static_assert(GGML_OP_COUNT == 101, "GGML_OP_COUNT != 101");
static_assert(GGML_OP_COUNT == 102, "GGML_OP_COUNT != 102");

static_assert(GGML_OP_POOL_COUNT == 2, "GGML_OP_POOL_COUNT != 2");

Expand Down Expand Up @@ -6418,6 +6420,36 @@ struct ggml_tensor * ggml_gated_delta_net(
return result;
}

// ggml_gated_delta_net_indexed

struct ggml_tensor * ggml_gated_delta_net_indexed(
struct ggml_context * ctx,
struct ggml_tensor * q,
struct ggml_tensor * k,
struct ggml_tensor * v,
struct ggml_tensor * g,
struct ggml_tensor * beta,
struct ggml_tensor * bank,
struct ggml_tensor * rows,
struct ggml_tensor * state_dependency,
int64_t K) {
GGML_ASSERT(bank->type == GGML_TYPE_F32 && ggml_is_contiguous(bank));
GGML_ASSERT(rows->type == GGML_TYPE_I32 && ggml_is_contiguous(rows));
GGML_ASSERT(ggml_is_vector(rows) && ggml_nelements(rows) == v->ne[3]);
GGML_ASSERT(bank->ne[3] >= v->ne[3]);

struct ggml_tensor * shape = ggml_view_4d(ctx, bank,
bank->ne[0], bank->ne[1], bank->ne[2], v->ne[3],
bank->nb[1], bank->nb[2], bank->nb[3], 0);

struct ggml_tensor * result = ggml_gated_delta_net(ctx, q, k, v, g, beta, shape, K);
result->op = GGML_OP_GATED_DELTA_NET_INDEXED;
result->src[5] = bank;
result->src[6] = rows;
result->src[7] = state_dependency;
return result;
}

// ggml_lightning_indexer

struct ggml_tensor * ggml_lightning_indexer(
Expand Down
14 changes: 10 additions & 4 deletions src/models/delta-net-base.cpp
Original file line number Diff line number Diff line change
Expand Up @@ -533,17 +533,20 @@ ggml_tensor * llm_build_delta_net_base::build_recurrent_attn(
ggml_tensor * g,
ggml_tensor * b,
ggml_tensor * s,
int il) {
int il,
ggml_tensor * state_rows,
ggml_tensor * state_dependency) {
const auto * mctx_cur = inp->mctx;
const auto kv_head = mctx_cur->get_head();
const uint32_t mem_size = mctx_cur->get_size();

const int64_t S_v = s->ne[0];
const int64_t H_v = s->ne[2];
const int64_t n_seqs = s->ne[3];
const int64_t n_seqs = state_rows ? v->ne[3] : s->ne[3];
const int64_t n_seq_tokens = q->ne[2];

const bool keep = cparams.n_rs_seq > 0;
GGML_ASSERT(!state_rows || keep);

if (!keep) {
auto attn_out = build_delta_net(q, k, v, g, b, s, il);
Expand All @@ -563,8 +566,11 @@ ggml_tensor * llm_build_delta_net_base::build_recurrent_attn(
const int64_t D = S_v * S_v * H_v;
const int64_t K = cparams.n_rs_seq + 1;

// state s is 4D [S_v, S_v, H_v, n_seqs]; K snapshot slots are written into the output.
ggml_tensor * gdn_out = ggml_gated_delta_net(ctx0, q, k, v, g, b, s, K);
// state s is either a packed per-sequence state or the full persistent bank.
// The indexed op selects the required row directly on the backend.
ggml_tensor * gdn_out = state_rows
? ggml_gated_delta_net_indexed(ctx0, q, k, v, g, b, s, state_rows, state_dependency, K)
: ggml_gated_delta_net(ctx0, q, k, v, g, b, s, K);
if (n_seq_tokens > 1) {
res->add_fused_node({LLM_FUSED_OP_GDN_CH, gdn_out, il});
} else {
Expand Down
4 changes: 3 additions & 1 deletion src/models/models.h
Original file line number Diff line number Diff line change
Expand Up @@ -97,7 +97,9 @@ struct llm_build_delta_net_base : public llm_graph_context {
ggml_tensor * g,
ggml_tensor * b,
ggml_tensor * s,
int il);
int il,
ggml_tensor * state_rows = nullptr,
ggml_tensor * state_dependency = nullptr);
};

struct llm_build_rwkv6_base : public llm_graph_context {
Expand Down
Loading
Loading