diff --git a/ggml/src/ggml-cuda/common.cuh b/ggml/src/ggml-cuda/common.cuh index ef929d3d7842..c7e31e1fe9b5 100644 --- a/ggml/src/ggml-cuda/common.cuh +++ b/ggml/src/ggml-cuda/common.cuh @@ -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 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; @@ -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(); @@ -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); diff --git a/ggml/src/ggml-cuda/gated_delta_net.cu b/ggml/src/ggml-cuda/gated_delta_net.cu index 5cf6968a6e1b..a9e3904455df 100644 --- a/ggml/src/ggml-cuda/gated_delta_net.cu +++ b/ggml/src/ggml-cuda/gated_delta_net.cu @@ -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. @@ -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; @@ -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; @@ -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: @@ -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)); @@ -355,7 +369,7 @@ static void ggml_cuda_op_gated_delta_net_impl( #define GDN_LAUNCH(KDA_, KEEP_, RAW_, PRE_) \ launch_gated_delta_net(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); } diff --git a/ggml/src/ggml-cuda/gated_delta_net.cuh b/ggml/src/ggml-cuda/gated_delta_net.cuh index f9bf43706789..bde78c42daf6 100644 --- a/ggml/src/ggml-cuda/gated_delta_net.cuh +++ b/ggml/src/ggml-cuda/gated_delta_net.cuh @@ -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) diff --git a/ggml/src/ggml-cuda/ggml-cuda.cu b/ggml/src/ggml-cuda/ggml-cuda.cu index 7536a5d2a094..5f7f3553423f 100644 --- a/ggml/src/ggml-cuda/ggml-cuda.cu +++ b/ggml/src/ggml-cuda/ggml-cuda.cu @@ -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) { + 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( @@ -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) { @@ -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.