// SPDX-License-Identifier: MIT // Copyright (c) 2023-2026 The ggml authors // // Substantial portions derived from ggml / llama.cpp // (https://github.com/ggml-org/llama.cpp), MIT licensed. // The original copyright notice and permission notice are retained. // See libs/ai/NOTICE and, where present, LICENSE in this directory. // #pragma once // Copied from llama.cpp ggml/src/ggml-cuda/gated_delta_net.cu:3-201 // (kernel + launch_gated_delta_net). Host ggml_cuda_info trimmed to // warp_size=32. No chunked kernel — they have none (cu:159). #include "common.cuh" template __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, const float * v, const float * g, const float * beta, const float * curr_state, float * dst, int64_t H, int64_t n_tokens, int64_t n_seqs, int64_t sq1, int64_t sq2, int64_t sq3, int64_t sv1, int64_t sv2, int64_t sv3, int64_t sb1, int64_t sb2, int64_t sb3, const uint3 neqk1_magic, const uint3 rq3_magic, float scale, int64_t state_ckpt_stride) { const uint32_t h_idx = blockIdx.x; const uint32_t sequence = blockIdx.y; const int lane = threadIdx.x; const int col = blockIdx.z * blockDim.y + threadIdx.y; const uint32_t iq1 = fastmodulo(h_idx, neqk1_magic); const uint32_t iq3 = fastdiv(sequence, rq3_magic); const int64_t attn_score_elems = S_v * H * n_tokens * n_seqs; float * attn_data = dst; float * state = dst + attn_score_elems; const int64_t state_offset = (sequence * H + h_idx) * S_v * S_v; state += state_offset; curr_state += state_offset + col * S_v; attn_data += (sequence * n_tokens * H + h_idx) * S_v; constexpr int warp_size = ggml_cuda_get_physical_warp_size() < S_v ? ggml_cuda_get_physical_warp_size() : S_v; static_assert(S_v % warp_size == 0, "S_v must be a multiple of warp_size"); constexpr int rows_per_lane = (S_v + warp_size - 1) / warp_size; float s_shard[rows_per_lane]; #pragma unroll for (int r = 0; r < rows_per_lane; r++) { const int i = r * warp_size + lane; s_shard[r] = curr_state[i]; } for (int t = 0; t < n_tokens; t++) { const float * q_t = q + iq3 * sq3 + t * sq2 + iq1 * sq1; const float * k_t = k + iq3 * sq3 + t * sq2 + iq1 * sq1; const float * v_t = v + sequence * sv3 + t * sv2 + h_idx * sv1; const int64_t gb_offset = sequence * sb3 + t * sb2 + h_idx * sb1; const float * beta_t = beta + gb_offset; const float * g_t = g + gb_offset * (KDA ? S_v : 1); const float beta_val = *beta_t; float k_reg[rows_per_lane]; float q_reg[rows_per_lane]; #pragma unroll for (int r = 0; r < rows_per_lane; r++) { const int i = r * warp_size + lane; k_reg[r] = k_t[i]; q_reg[r] = q_t[i]; } if constexpr (!KDA) { const float g_val = expf(*g_t); float kv_shard = 0.0f; #pragma unroll for (int r = 0; r < rows_per_lane; r++) { kv_shard += s_shard[r] * k_reg[r]; } float kv_col = warp_reduce_sum(kv_shard); float delta_col = (v_t[col] - g_val * kv_col) * beta_val; float attn_partial = 0.0f; #pragma unroll for (int r = 0; r < rows_per_lane; r++) { s_shard[r] = g_val * s_shard[r] + k_reg[r] * delta_col; attn_partial += s_shard[r] * q_reg[r]; } float attn_col = warp_reduce_sum(attn_partial); if (lane == 0) { attn_data[col] = attn_col * scale; } } else { float kv_shard = 0.0f; #pragma unroll for (int r = 0; r < rows_per_lane; r++) { const int i = r * warp_size + lane; kv_shard += expf(g_t[i]) * s_shard[r] * k_reg[r]; } float kv_col = warp_reduce_sum(kv_shard); float delta_col = (v_t[col] - kv_col) * beta_val; float attn_partial = 0.0f; #pragma unroll for (int r = 0; r < rows_per_lane; r++) { const int i = r * warp_size + lane; s_shard[r] = expf(g_t[i]) * s_shard[r] + k_reg[r] * delta_col; attn_partial += s_shard[r] * q_reg[r]; } float attn_col = warp_reduce_sum(attn_partial); if (lane == 0) { attn_data[col] = attn_col * scale; } } // Speculative verification needs the state after EVERY token, not // just the last one, so a rejected draft can be undone by resuming // from an earlier row. Emitting them from this one call keeps the // checkpointed graph at the node count of an ordinary decode. if (state_ckpt_stride != 0) { #pragma unroll for (int r = 0; r < rows_per_lane; r++) { const int i = r * warp_size + lane; state[t * state_ckpt_stride + col * S_v + i] = s_shard[r]; } } attn_data += S_v * H; } if (state_ckpt_stride == 0) { #pragma unroll for (int r = 0; r < rows_per_lane; r++) { const int i = r * warp_size + lane; state[col * S_v + i] = s_shard[r]; } } } template static void launch_gated_delta_net( const float * q_d, const float * k_d, const float * v_d, const float * g_d, const float * b_d, const float * s_d, float * dst_d, int64_t S_v, int64_t H, int64_t n_tokens, int64_t n_seqs, int64_t sq1, int64_t sq2, int64_t sq3, 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_ckpt_stride, cudaStream_t stream) { const int warp_size = 32; const int num_warps = 4; dim3 grid_dims((unsigned) H, (unsigned) n_seqs, (unsigned) ((S_v + num_warps - 1) / num_warps)); dim3 block_dims(warp_size <= S_v ? warp_size : (int) S_v, num_warps, 1); const uint3 neqk1_magic = init_fastdiv_values((uint32_t) neqk1); const uint3 rq3_magic = init_fastdiv_values((uint32_t) rq3); switch (S_v) { case 16: gated_delta_net_cuda<16, KDA><<>>( q_d, k_d, v_d, g_d, b_d, s_d, dst_d, H, n_tokens, n_seqs, sq1, sq2, sq3, sv1, sv2, sv3, sb1, sb2, sb3, neqk1_magic, rq3_magic, scale, state_ckpt_stride); break; case 32: gated_delta_net_cuda<32, KDA><<>>( q_d, k_d, v_d, g_d, b_d, s_d, dst_d, H, n_tokens, n_seqs, sq1, sq2, sq3, sv1, sv2, sv3, sb1, sb2, sb3, neqk1_magic, rq3_magic, scale, state_ckpt_stride); break; case 64: gated_delta_net_cuda<64, KDA><<>>( q_d, k_d, v_d, g_d, b_d, s_d, dst_d, H, n_tokens, n_seqs, sq1, sq2, sq3, sv1, sv2, sv3, sb1, sb2, sb3, neqk1_magic, rq3_magic, scale, state_ckpt_stride); break; case 128: gated_delta_net_cuda<128, KDA><<>>( q_d, k_d, v_d, g_d, b_d, s_d, dst_d, H, n_tokens, n_seqs, sq1, sq2, sq3, sv1, sv2, sv3, sb1, sb2, sb3, neqk1_magic, rq3_magic, scale, state_ckpt_stride); break; default: GGML_ABORT("unsupported GDN S_v"); break; } }