Skip to content
Closed
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
32 changes: 32 additions & 0 deletions ggml/src/ggml-cuda/common.cuh
Original file line number Diff line number Diff line change
Expand Up @@ -1453,6 +1453,35 @@ struct ggml_cuda_stream_context {
}
};

// Fused recurrent-state gather for GATED_DELTA_NET. build_rs materialises GET_ROWS(cache, s_copy)
// into a temp per layer that only the GDN kernel reads; when the graph evaluator can prove that, it
// skips the GET_ROWS and records the gather here so the kernel reads cache row ids[seq] directly.
struct ggml_cuda_gated_delta_net_gather {
const float * base = nullptr; // cache rows [row_stride floats each]
const int32_t * ids = nullptr; // per-seq row index
int64_t row_stride = 0; // in floats
};

// Owned by the backend context that evaluates the graph: registrations are keyed by node pointer,
// so they are only meaningful for the evaluation that made them. Reset at the start of every
// graph evaluation/capture; never shared between contexts or threads.
struct ggml_cuda_gdn_gather_context {
std::unordered_map<const ggml_tensor *, ggml_cuda_gated_delta_net_gather> gathers;

void reset() {
gathers.clear();
}

void set(const ggml_tensor * gdn, const ggml_cuda_gated_delta_net_gather & gather) {
gathers[gdn] = gather;
}

const ggml_cuda_gated_delta_net_gather * find(const ggml_tensor * gdn) const {
const auto it = gathers.find(gdn);
return it == gathers.end() ? nullptr : &it->second;
}
};

struct ggml_backend_cuda_context {
int device;
std::string name;
Expand Down Expand Up @@ -1523,6 +1552,7 @@ struct ggml_backend_cuda_context {
}

ggml_cuda_stream_context concurrent_stream_context;
ggml_cuda_gdn_gather_context gdn_gather_context;

~ggml_backend_cuda_context();

Expand All @@ -1538,6 +1568,8 @@ struct ggml_backend_cuda_context {

ggml_cuda_stream_context & stream_context() { return concurrent_stream_context; }

ggml_cuda_gdn_gather_context & gdn_gathers() { return gdn_gather_context; }

cublasHandle_t cublas_handle() {
if (cublas_handles[device][curr_stream_no] == nullptr) {
ggml_cuda_set_device(device);
Expand Down
30 changes: 22 additions & 8 deletions ggml/src/ggml-cuda/gated_delta_net.cu
Original file line number Diff line number Diff line change
Expand Up @@ -39,7 +39,9 @@ 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 * s_ids,
int64_t s_row_stride) {
const uint32_t h_idx = blockIdx.x;
const uint32_t sequence = blockIdx.y;
// Each warp owns one or more columns, using warp-level primitives to reduce across rows.
Expand All @@ -58,7 +60,9 @@ gated_delta_net_cuda(const float * q,

// 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;
// fused gather: read this sequence's live state straight out of the cache row s_ids[sequence]
const int64_t state_in_offset = (s_ids ? (int64_t) s_ids[sequence] * s_row_stride : sequence * 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;
Expand Down Expand Up @@ -212,7 +216,8 @@ static void launch_gated_delta_net(
int64_t sv1, int64_t sv2, int64_t sv3,
int64_t sb1, int64_t sb2, int64_t sb3,
int64_t neqk1, int64_t rq3,
float scale, int64_t state_slot_stride, int K, cudaStream_t stream) {
float scale, int64_t state_slot_stride, int K,
const int32_t * s_ids, int64_t s_row_stride, cudaStream_t stream) {
//TODO: Add chunked kernel for even faster pre-fill
const int warp_size = ggml_cuda_info().devices[ggml_cuda_get_device()].warp_size;
const int cc = ggml_cuda_info().devices[ggml_cuda_get_device()].cc;
Expand All @@ -230,26 +235,26 @@ static void launch_gated_delta_net(
ggml_cuda_kernel_launch(gated_delta_net_cuda<16, KDA, keep_rs_t, RAW, G_PRECOMPUTED>, launch_params,
q_d, k_d, v_d, g_d, b_d, rb_d, ra_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, s_ids, s_row_stride);
break;
case 32:
ggml_cuda_kernel_launch(gated_delta_net_cuda<32, KDA, keep_rs_t, RAW, G_PRECOMPUTED>, launch_params,
q_d, k_d, v_d, g_d, b_d, rb_d, ra_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, s_ids, s_row_stride);
break;
case 64: {
ggml_cuda_kernel_launch(gated_delta_net_cuda<64, KDA, keep_rs_t, RAW, G_PRECOMPUTED>, launch_params,
q_d, k_d, v_d, g_d, b_d, rb_d, ra_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, s_ids, s_row_stride);
break;
}
case 128: {
ggml_cuda_kernel_launch(gated_delta_net_cuda<128, KDA, keep_rs_t, RAW, G_PRECOMPUTED>, launch_params,
q_d, k_d, v_d, g_d, b_d, rb_d, ra_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, s_ids, s_row_stride);
break;
}
default:
Expand Down Expand Up @@ -296,6 +301,15 @@ static void ggml_cuda_op_gated_delta_net_impl(
const float * s_d = (const float *) src_state->data;
float * dst_d = (float *) dst->data;

// fused state gather, registered for this node by this context's graph evaluator (ggml_cuda_try_gdn_gather_skip)
const int32_t * s_ids = nullptr;
int64_t s_row_stride = 0;
if (const ggml_cuda_gated_delta_net_gather * gather = ctx.gdn_gathers().find(dst)) {
s_d = gather->base;
s_ids = gather->ids;
s_row_stride = gather->row_stride;
}

GGML_ASSERT(ggml_is_contiguous_rows(src_q));
GGML_ASSERT(ggml_is_contiguous_rows(src_k));
GGML_ASSERT(ggml_is_contiguous_rows(src_v));
Expand Down Expand Up @@ -355,7 +369,7 @@ static void ggml_cuda_op_gated_delta_net_impl(
#define GDN_LAUNCH(KDA_, KEEP_, RAW_, PRE_) \
launch_gated_delta_net<KDA_, KEEP_, RAW_, PRE_>(q_d, k_d, v_d, g_d, b_d, rb_d, ra_d, s_d, dst_d, state_d, \
S_v, H, n_tokens, n_seqs, sq1, sq2, sq3, sv1, sv2, sv3, \
sb1, sb2, sb3, neqk1, rq3, scale, state_slot_stride, K, stream)
sb1, sb2, sb3, neqk1, rq3, scale, state_slot_stride, K, s_ids, s_row_stride, stream)

if (kda) {
if (keep_rs) { GDN_LAUNCH(true, true, false, false); } else { GDN_LAUNCH(true, false, false, false); }
Expand Down
2 changes: 2 additions & 0 deletions ggml/src/ggml-cuda/gated_delta_net.cuh
Original file line number Diff line number Diff line change
Expand Up @@ -7,6 +7,8 @@ struct ggml_cuda_gated_delta_net_fused_cache {
int64_t slot_stride; // between rollback slots (0 when K==1)
};

// The fused recurrent-state gather (ggml_cuda_gated_delta_net_gather, common.cuh) is looked up in
// ctx.gdn_gathers() by node pointer; the graph evaluator registers it when it skips the GET_ROWS.
void ggml_cuda_op_gated_delta_net(ggml_backend_cuda_context & ctx, ggml_tensor * dst);

// same op, but writes the snapshot(s) into the cache instead of dst (see ggml_cuda_try_gdn_cache_fusion)
Expand Down
63 changes: 63 additions & 0 deletions ggml/src/ggml-cuda/ggml-cuda.cu
Original file line number Diff line number Diff line change
Expand Up @@ -2747,6 +2747,61 @@ static bool ggml_cuda_should_fuse_rms_norm_mul_rope(const ggml_tensor * rms_norm
return true;
}

// GET_ROWS -> [RESHAPE] -> GATED_DELTA_NET src[5]. Skip the GET_ROWS launch when that temp has one consumer and let the kernel index the cache row. The allocator still reserved the unused temp. Single-sequence only. GGML_CUDA_GDN_GATHER_FUSION=0 disables. Registry is per evaluating context.
static bool ggml_cuda_try_gdn_gather_skip(ggml_backend_cuda_context & ctx, const ggml_cgraph * cgraph, int node_idx) {

Copy link
Copy Markdown
Author

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

Fair ask. test-backend-ops already covers GDN numerics on the kernel, but it does not emit this GET_ROWS+RESHAPE+GDN graph, so it never hits the matcher or CUDA-graph replay of the skipped launch. I am not stuffing a graph-capture case into this PR; I will add one as a follow-up that runs with the fusion on and off. Disable with GGML_CUDA_GDN_GATHER_FUSION=0 in the meantime.

static const bool disabled = getenv("GGML_CUDA_GDN_GATHER_FUSION") != nullptr &&
atoi(getenv("GGML_CUDA_GDN_GATHER_FUSION")) == 0;
if (disabled) {
return false;
}
const ggml_tensor * gr = cgraph->nodes[node_idx];
if (gr->op != GGML_OP_GET_ROWS || gr->type != GGML_TYPE_F32 || (gr->flags & GGML_TENSOR_FLAG_OUTPUT) ||
!ggml_is_contiguous(gr)) {
return false;
}
const ggml_tensor * cache = gr->src[0];
const ggml_tensor * ids = gr->src[1];
if (cache->type != GGML_TYPE_F32 || ids->type != GGML_TYPE_I32 || cache->data == nullptr || ids->data == nullptr ||
cache->nb[0] != sizeof(float) || cache->nb[1] % sizeof(float) != 0 || !ggml_is_contiguous(ids) ||
ids->ne[0] != 1 || ids->ne[1] != 1 || ids->ne[2] != 1 || ids->ne[3] != 1 ||
gr->ne[1] != 1 || gr->ne[2] != 1 || gr->ne[3] != 1 || gr->ne[0] != cache->ne[0]) {
return false;
}
if (ggml_node_get_use_count(cgraph, node_idx) != 1) {
return false;
}
const ggml_tensor * cur = gr;
for (int j = node_idx + 1; j < cgraph->n_nodes; ++j) {
const ggml_tensor * n = cgraph->nodes[j];
if (n->op == GGML_OP_GATED_DELTA_NET && n->src[5] == cur) {
const ggml_tensor * v = n->src[2];
const int64_t D = v->ne[0] * v->ne[0] * v->ne[1];
if (gr->ne[0] != D || v->ne[3] != 1 || ggml_nelements(cur) != D) {
return false;
}
ggml_cuda_gated_delta_net_gather gather;
gather.base = (const float *) cache->data;
gather.ids = (const int32_t *) ids->data;
gather.row_stride = (int64_t) (cache->nb[1] / sizeof(float));
ctx.gdn_gathers().set(n, gather);
return true;
}
if (n->op == GGML_OP_RESHAPE && n->src[0] == cur) {
if (ggml_node_get_use_count(cgraph, j) != 1) {
return false;
}
cur = n;
continue;
}
for (int s = 0; s < GGML_MAX_SRC; ++s) {
if (n->src[s] == cur || (n->view_src != nullptr && n->view_src == gr)) {
return false;
}
}
}
return false;
}

// match gated_delta_net + the strided cpy that scatters its state snapshots into the cache
// (slot i -> rollback group i, slot 0 newest), so the kernel can write them and skip the cpy.
static int ggml_cuda_try_gdn_cache_fusion(
Expand Down Expand Up @@ -4281,6 +4336,8 @@ static void ggml_cuda_graph_evaluate_and_capture(ggml_backend_cuda_context * cud
stream_ctx.concurrent_events.clear();
}

cuda_ctx->gdn_gathers().reset();

for (int i = 0; i < cgraph->n_nodes; i++) {
ggml_tensor * node = cgraph->nodes[i];
if (is_concurrent_event_active) {
Expand Down Expand Up @@ -4323,6 +4380,12 @@ static void ggml_cuda_graph_evaluate_and_capture(ggml_backend_cuda_context * cud
continue;
}

// skip GET_ROWS launch; GDN indexes the cache. The gather temp stays allocated.
if (node->op == GGML_OP_GET_ROWS && !is_concurrent_event_active &&
ggml_cuda_try_gdn_gather_skip(*cuda_ctx, cgraph, i)) {
continue;
}

// The normalized pre-attention residual is consumed only by a
// group of low-bit projections. Preserve residual + one scale per
// row and let their shared Q8 quantizer apply the norm weight.
Expand Down