diff --git a/ggml/src/ggml-cuda.cu b/ggml/src/ggml-cuda.cu index e66f3de3..a812cc82 100644 --- a/ggml/src/ggml-cuda.cu +++ b/ggml/src/ggml-cuda.cu @@ -8,6 +8,7 @@ #include "ggml-cuda.h" #include "ggml.h" #include "ggml-backend-impl.h" +#include "ggml-impl.h" #include "ggml-cuda/common.cuh" #include "ggml-cuda/acc.cuh" @@ -4130,9 +4131,23 @@ static bool ggml_cuda_compute_forward(ggml_backend_cuda_context & ctx, struct gg case GGML_OP_SOLVE_TRI: ggml_cuda_op_solve_tri(ctx, dst); break; - case GGML_OP_DELTA_NET: - ggml_cuda_op_delta_net(ctx, dst); - break; + case GGML_OP_DELTA_NET: { + const int j = fusion ? ggml_delta_net_find_state_cpy(cgraph, i) : -1; + if (j >= 0) { + ggml_tensor fused = *dst; + fused.src[7] = cgraph->nodes[j]->src[1]; + ggml_cuda_op_delta_net(ctx, &fused); +#ifdef USE_CUDA_GRAPH + // claim the entry of the copy that is not going to be launched + if (ctx.cur_graph && ctx.cur_graph->use_cpy_indirection) { + ctx.cur_graph->graph_cpynode_index++; + } +#endif + i = j; + } else { + ggml_cuda_op_delta_net(ctx, dst); + } + } break; case GGML_OP_SINKHORN: ggml_cuda_op_sinkhorn(ctx, dst); break; diff --git a/ggml/src/ggml-cuda/delta-net.cu b/ggml/src/ggml-cuda/delta-net.cu index dddbf446..a7d509c1 100644 --- a/ggml/src/ggml-cuda/delta-net.cu +++ b/ggml/src/ggml-cuda/delta-net.cu @@ -34,15 +34,15 @@ __global__ void delta_net_recurrent_f32( const float * __restrict__ v, // [HEAD_DIM, n_tokens, n_heads, n_seqs] const float * __restrict__ g, // [n_tokens, 1, n_heads, n_seqs] const float * __restrict__ beta_in, // [1, n_tokens, n_heads, n_seqs] - const float * __restrict__ state_in, // [HEAD_DIM, HEAD_DIM*n_heads, 1, n_seqs] - float * __restrict__ dst, // output + new_state(s) concatenated + const float * state_in, // [HEAD_DIM, HEAD_DIM*n_heads, 1, n_seqs], aliases state_out when fused + float * __restrict__ dst, // output + float * state_out, // new state float * __restrict__ saved_states, const int64_t n_heads, const int64_t gqa_ratio, const int repeat_type, const int64_t n_tokens, const int64_t n_seqs, - const int64_t output_offset, // offset where state starts in output size_t vnb1, size_t vnb2, size_t vnb3) { constexpr int warps_per_head = HEAD_DIM/WARP_SIZE; const int batch_idx = blockIdx.x / (warps_per_head*n_heads); @@ -85,7 +85,7 @@ __global__ void delta_net_recurrent_f32( // For [dim, head, token, batch]: index = dim + head*S_v + token*S_v*H_v + batch*S_v*H_v*n_tokens float * out_base = dst + batch_idx * (HEAD_DIM * n_heads * n_tokens) + head_idx * HEAD_DIM; const int64_t out_token_stride = HEAD_DIM * n_heads; // stride between tokens - float * state_dst = dst + output_offset + batch_idx * state_batch_stride + state_head_offset; + float * state_dst = state_out + batch_idx * state_batch_stride + state_head_offset; // Shared memory for current token's Q, K, V (normalized), and intermediate results extern __shared__ float smem[]; @@ -190,6 +190,7 @@ static void delta_net_f32_cuda( const float * beta, const float * state_in, float * dst, + float * state_out, // where the new state goes: src7 when fused, otherwise dst's tail float * saved_states, const int64_t head_dim, const int64_t n_tokens, @@ -204,8 +205,6 @@ static void delta_net_f32_cuda( GGML_UNUSED(device_id); GGML_UNUSED(cc); - const int64_t output_offset = head_dim * n_tokens * n_heads * n_seqs; - if (head_dim != 64 && head_dim != 128) { GGML_ABORT("Unsupported delta net head size"); } @@ -218,19 +217,19 @@ static void delta_net_f32_cuda( constexpr int threads_per_block = 256; if (head_dim == 64) { delta_net_recurrent_f32<64, threads_per_block><<>>( - q, k, v, g, beta, state_in, dst, saved_states, n_heads, gqa_ratio, repeat_type, n_tokens, n_seqs, output_offset, vnb1, vnb2, vnb3); + q, k, v, g, beta, state_in, dst, state_out, saved_states, n_heads, gqa_ratio, repeat_type, n_tokens, n_seqs, vnb1, vnb2, vnb3); } else { delta_net_recurrent_f32<128, threads_per_block><<>>( - q, k, v, g, beta, state_in, dst, saved_states, n_heads, gqa_ratio, repeat_type, n_tokens, n_seqs, output_offset, vnb1, vnb2, vnb3); + q, k, v, g, beta, state_in, dst, state_out, saved_states, n_heads, gqa_ratio, repeat_type, n_tokens, n_seqs, vnb1, vnb2, vnb3); } } else { constexpr int threads_per_block = 128; if (head_dim == 64) { delta_net_recurrent_f32<64, threads_per_block><<>>( - q, k, v, g, beta, state_in, dst, saved_states, n_heads, gqa_ratio, repeat_type, n_tokens, n_seqs, output_offset, vnb1, vnb2, vnb3); + q, k, v, g, beta, state_in, dst, state_out, saved_states, n_heads, gqa_ratio, repeat_type, n_tokens, n_seqs, vnb1, vnb2, vnb3); } else { delta_net_recurrent_f32<128, threads_per_block><<>>( - q, k, v, g, beta, state_in, dst, saved_states, n_heads, gqa_ratio, repeat_type, n_tokens, n_seqs, output_offset, vnb1, vnb2, vnb3); + q, k, v, g, beta, state_in, dst, state_out, saved_states, n_heads, gqa_ratio, repeat_type, n_tokens, n_seqs, vnb1, vnb2, vnb3); } } @@ -246,6 +245,7 @@ void ggml_cuda_op_delta_net(ggml_backend_cuda_context & ctx, ggml_tensor * dst) const ggml_tensor * src4 = dst->src[4]; // beta const ggml_tensor * src5 = dst->src[5]; // state const ggml_tensor * src6 = dst->src[6]; // when not null, state for token 0...n_token-1 + const ggml_tensor * src7 = dst->src[7]; // when not null, the slot the fused state goes to GGML_ASSERT(src0->type == GGML_TYPE_F32); GGML_ASSERT(dst->type == GGML_TYPE_F32); @@ -279,6 +279,14 @@ void ggml_cuda_op_delta_net(ggml_backend_cuda_context & ctx, ggml_tensor * dst) GGML_ASSERT(src6->type == GGML_TYPE_F32); GGML_ASSERT(src6->ne[0] >= (n_tokens - 1)*state_size); } + if (src7) { + // the copy was fused away, so the state goes into the op's own input - the alias is + // safe for the reason spelled out at ggml_delta_net_find_state_cpy() + GGML_ASSERT(src7->data == src5->data); + GGML_ASSERT(src7->type == GGML_TYPE_F32); + GGML_ASSERT(ggml_is_contiguous(src7)); + GGML_ASSERT(ggml_nelements(src7) == state_size); + } const int64_t expected_size = output_size + state_size; GGML_ASSERT(ggml_nelements(dst) == expected_size); @@ -297,6 +305,7 @@ void ggml_cuda_op_delta_net(ggml_backend_cuda_context & ctx, ggml_tensor * dst) (const float *)src4->data, (const float *)src5->data, (float *)dst->data, + src7 ? (float *)src7->data : (float *)dst->data + output_size, src6 ? (float *)src6->data : nullptr, head_dim, n_tokens, n_heads, gqa_ratio, repeat_type, n_seqs, src2->nb[1]/sizeof(float), src2->nb[2]/sizeof(float), src2->nb[3]/sizeof(float), diff --git a/ggml/src/ggml-impl.h b/ggml/src/ggml-impl.h index 62f07e1e..32b05c1b 100644 --- a/ggml/src/ggml-impl.h +++ b/ggml/src/ggml-impl.h @@ -766,6 +766,10 @@ static size_t ggml_hash_find_or_insert(struct ggml_hash_set * hash_set, struct g GGML_ABORT("fatal error"); } +// Finds the copy the delta-net fusion can skip: the CPY that writes this node's new recurrent +// state back into the slot it was read from. Its index in cgraph, or -1 if there is none. +int ggml_delta_net_find_state_cpy(const struct ggml_cgraph * cgraph, int i); + static int32_t ggml_get_op_params_i32(const struct ggml_tensor * tensor, uint32_t i) { assert(i < GGML_MAX_OP_PARAMS / sizeof(int32_t)); return ((const int32_t *)(tensor->op_params))[i]; diff --git a/ggml/src/ggml.c b/ggml/src/ggml.c index 0e1b8527..a78ac0ff 100644 --- a/ggml/src/ggml.c +++ b/ggml/src/ggml.c @@ -23872,12 +23872,14 @@ static void ggml_compute_forward_delta_net_f32( const struct ggml_tensor * src4 = dst->src[4]; const struct ggml_tensor * src5 = dst->src[5]; const struct ggml_tensor * src6 = dst->src[6]; + const struct ggml_tensor * src7 = dst->src[7]; const int64_t head_dim = src0->ne[0]; const int64_t n_tokens = src0->ne[1]; const int64_t n_heads = src2->ne[2]; const int64_t n_seqs = src0->ne[3]; GGML_ASSERT(src2->ne[2] % src0->ne[2] == 0); + GGML_ASSERT(src2->ne[0] == head_dim); const int gqa_ratio = src2->ne[2]/src0->ne[2]; const int64_t output_size = head_dim * n_tokens * n_heads * n_seqs; @@ -23896,7 +23898,18 @@ static void ggml_compute_forward_delta_net_f32( int repeat_type = dst->op_params[0]; const int64_t state_step_stride = head_dim * head_dim * n_heads * n_seqs; - float * state_working = out_data + output_size; + // src7 is the slot the fused-away copy would have written + float * state_working = src7 ? (float *) src7->data : out_data + output_size; + + if (src7) { + // src7 signals that the copy was fused away, so the op writes the state into its own + // input - safe because each thread seeds its heads from state_in before writing them + // back, and threads own disjoint heads + GGML_ASSERT(src7->data == src5->data); + GGML_ASSERT(src7->type == GGML_TYPE_F32); + GGML_ASSERT(ggml_is_contiguous(src7)); + GGML_ASSERT(ggml_nelements(src7) == state_step_stride); + } if (src6) { GGML_ASSERT(src6->type == GGML_TYPE_F32); @@ -24017,6 +24030,59 @@ static void ggml_compute_forward_delta_net( } } +// inspired by ggml_cuda_try_gdn_cache_fusion in https://github.com/ggml-org/llama.cpp/pull/23940 +// the pattern this matches is built by delta_net::build_qkv() in src/llama-delta-net.cpp +// - the scan stops at the first real node, so a copy further down is missed +// - nothing checks that the copy is the tail's only reader +int ggml_delta_net_find_state_cpy(const struct ggml_cgraph * cgraph, int i) { + const struct ggml_tensor * dn = cgraph->nodes[i]; + assert(dn->op == GGML_OP_DELTA_NET); + + if (dn->type != GGML_TYPE_F32 || !ggml_is_contiguous(dn) || (dn->flags & GGML_TENSOR_FLAG_OUTPUT)) { + return -1; + } + + const struct ggml_tensor * q = dn->src[0]; + const struct ggml_tensor * v = dn->src[2]; + + // the sizes ggml_delta_net() allocated the result from - the total ties them to this node + const size_t tail_off = ggml_row_size(GGML_TYPE_F32, v->ne[0] * v->ne[2] * q->ne[1] * q->ne[3]); + const size_t state_size = ggml_row_size(GGML_TYPE_F32, v->ne[0] * v->ne[0] * v->ne[2] * q->ne[3]); + + if (tail_off + state_size != ggml_nbytes(dn)) { + return -1; + } + + for (int j = i + 1; j < cgraph->n_nodes; j++) { + const struct ggml_tensor * cpy = cgraph->nodes[j]; + if (ggml_is_noop(cpy)) { + continue; + } + if (cpy->op != GGML_OP_CPY || (cpy->flags & GGML_TENSOR_FLAG_OUTPUT)) { + return -1; + } + + // the tail reaches the copy through reshapes, so match on the view root and its offset + const struct ggml_tensor * src = cpy->src[0]; + const struct ggml_tensor * dst = cpy->src[1]; + + if (!ggml_is_noop(src) || src->view_src != dn || src->view_offs != tail_off || + !ggml_is_contiguous(src) || ggml_nbytes(src) != state_size || + (src->flags & GGML_TENSOR_FLAG_OUTPUT)) { + return -1; + } + // the destination has to be the slot the state was read from: the fused kernel takes it + // as a plain pointer, which on CUDA is only covered by the graph check on src[5] + if (dst->type != GGML_TYPE_F32 || !dst->data || !ggml_is_contiguous(dst) || + ggml_nbytes(dst) != state_size || dst->data != dn->src[5]->data) { + return -1; + } + return j; + } + + return -1; +} + // ggml_compute_forward_sinkhorn static void ggml_compute_forward_sinkhorn_f32( @@ -26555,7 +26621,15 @@ static int ggml_compute_forward(struct ggml_compute_params * params, struct ggml } break; case GGML_OP_DELTA_NET: { - ggml_compute_forward_delta_net(params, tensor); + const int j = fusion ? ggml_delta_net_find_state_cpy(cgraph, i) : -1; + if (j >= 0) { + struct ggml_tensor fused = *tensor; + fused.src[7] = cgraph->nodes[j]->src[1]; + ggml_compute_forward_delta_net(params, &fused); + i = j; + } else { + ggml_compute_forward_delta_net(params, tensor); + } } break; case GGML_OP_SINKHORN: { diff --git a/src/llama-delta-net.cpp b/src/llama-delta-net.cpp index fde00a55..7a622de9 100644 --- a/src/llama-delta-net.cpp +++ b/src/llama-delta-net.cpp @@ -386,9 +386,19 @@ ggml_tensor * delta_net::build_qkv(ggml_context * ctx0, ggml_tensor * state_stor cb(new_conv_states_cont, "new_conv_states_cont", il); ggml_tensor * new_conv_flat = ggml_reshape_2d(ctx0, new_conv_states_cont, conv_state_dim, 1); ggml_tensor * new_ssm_flat = ggml_reshape_2d(ctx0, new_state, ssm_state_dim, 1); - auto state_cpy = ggml_concat_inplace(ctx0, new_conv_flat, new_ssm_flat, state_dst, 0); - cb(state_cpy, "state_cpy", il); - ggml_build_forward_expand(gf, state_cpy); + ggml_tensor * conv_state_dst = ggml_view_2d(ctx0, state_dst, conv_state_dim, 1, state_row_size, 0); + ggml_tensor * ssm_dst = ggml_view_2d(ctx0, state_dst, ssm_state_dim, 1, state_row_size, + conv_state_dim * ggml_element_size(state_dst)); + + // expand this one first: ggml_delta_net_find_state_cpy() matches this copy only while nothing + // but view no-ops stands between it and the op, and the backends then write the slot directly + auto ssm_cpy = ggml_cpy(ctx0, new_ssm_flat, ssm_dst); + cb(ssm_cpy, "ssm_state_cpy", il); + ggml_build_forward_expand(gf, ssm_cpy); + + auto conv_cpy = ggml_cpy(ctx0, new_conv_flat, conv_state_dst); + cb(conv_cpy, "conv_state_cpy", il); + ggml_build_forward_expand(gf, conv_cpy); return output; }