Skip to content
Open
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
52 changes: 52 additions & 0 deletions ggml/src/ggml-cuda/concat.cu
Original file line number Diff line number Diff line change
Expand Up @@ -139,6 +139,41 @@ static __global__ void __launch_bounds__(CUDA_CONCAT_BLOCK_SIZE)
}
}

// Transpose a dimension-0 view while prepending a small contiguous state. The
// generic non-contiguous kernel assigns a CTA to each output row, which makes
// every warp gather with a large stride for this layout.
static __global__ void concat_dim0_transpose_u32(
const uint32_t * src0, const char * src1, uint32_t * dst,
int ne00, int rows, int cols, uint64_t src1_col_stride, int dst_row_stride) {
__shared__ uint32_t tile[32][33];

const int input_x = (int) blockIdx.x*32 + threadIdx.x; // output row / input column
const int input_y0 = (int) blockIdx.y*32 + threadIdx.y; // output column / input row

#pragma unroll
for (int j = 0; j < 32; j += 8) {
const int input_y = input_y0 + j;
if (input_x < rows && input_y < cols) {
tile[threadIdx.y + j][threadIdx.x] =
*(const uint32_t *)(src1 + (uint64_t) input_y*src1_col_stride + (uint64_t) input_x*sizeof(uint32_t));
}
}
__syncthreads();

const int output_row0 = (int) blockIdx.x*32 + threadIdx.y;
const int output_col = (int) blockIdx.y*32 + threadIdx.x;
#pragma unroll
for (int j = 0; j < 32; j += 8) {
const int output_row = output_row0 + j;
if (output_row < rows && output_col < cols) {
dst[(int64_t) output_row*dst_row_stride + ne00 + output_col] = tile[threadIdx.x][threadIdx.y + j];
}
if (blockIdx.y == 0 && threadIdx.x < ne00 && output_row < rows) {
dst[(int64_t) output_row*dst_row_stride + threadIdx.x] = src0[(int64_t) output_row*ne00 + threadIdx.x];
}
}
}

template <typename T>
static void concat_cuda(const ggml_tensor * src0, const ggml_tensor * src1, ggml_tensor * dst, int dim, cudaStream_t stream) {
if (dim != 3 && ggml_is_contiguous_to_3(src0) && ggml_is_contiguous_to_3(src1)) {
Expand All @@ -164,6 +199,23 @@ static void concat_cuda(const ggml_tensor * src0, const ggml_tensor * src1, ggml
GGML_ASSERT(!ggml_is_quantized(src0->type));

dim3 grid_dim(dst->ne[1], dst->ne[2], dst->ne[3]);
if constexpr (sizeof(T) == sizeof(uint32_t)) {
const bool transpose_dim0 = ggml_cuda_info().devices[ggml_cuda_get_device()].cc == GGML_CUDA_CC_DGX_SPARK &&
dim == 0 && src0->ne[2] == 1 && src0->ne[3] == 1 && src1->ne[2] == 1 && src1->ne[3] == 1 &&
dst->ne[2] == 1 && dst->ne[3] == 1 && src0->ne[0] <= 8 &&
src0->nb[0] == sizeof(uint32_t) && src0->nb[1] == (uint64_t) src0->ne[0]*sizeof(uint32_t) &&
src1->nb[1] == sizeof(uint32_t) && dst->nb[0] == sizeof(uint32_t) &&
dst->nb[1] == (uint64_t) dst->ne[0]*sizeof(uint32_t) &&
src0->ne[1] == src1->ne[1] && src0->ne[1] == dst->ne[1] &&
src0->ne[0] + src1->ne[0] == dst->ne[0];
if (transpose_dim0) {
const dim3 grid((dst->ne[1] + 31)/32, (src1->ne[0] + 31)/32, 1);
concat_dim0_transpose_u32<<<grid, dim3(32, 8, 1), 0, stream>>>(
(const uint32_t *) src0->data, (const char *) src1->data, (uint32_t *) dst->data,
src0->ne[0], dst->ne[1], src1->ne[0], src1->nb[0], dst->nb[1]/sizeof(uint32_t));
return;
}
}
auto launch_kernel = [&](auto dim) {
concat_non_cont<T, dim><<<grid_dim, CUDA_CONCAT_BLOCK_SIZE, 0, stream>>>(
(const char *) src0->data, (const char *) src1->data, (char *) dst->data,
Expand Down
145 changes: 97 additions & 48 deletions ggml/src/ggml-cuda/gated_delta_net.cu
Original file line number Diff line number Diff line change
@@ -1,7 +1,14 @@
#include "gated_delta_net.cuh"
#include "ggml-cuda/common.cuh"

template <int S_v, bool KDA, bool keep_rs_t>
static __global__ void gdn_precompute_exp(const float * g, float * g_exp, int64_t n) {
for (int64_t i = (int64_t) blockIdx.x*blockDim.x + threadIdx.x; i < n;
i += (int64_t) blockDim.x*gridDim.x) {
g_exp[i] = expf(g[i]);
}
}

template <int S_v, bool KDA, bool keep_rs_t, bool G_PRECOMPUTED>
__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 Down Expand Up @@ -30,9 +37,14 @@ gated_delta_net_cuda(const float * q,
int K) {
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
// Each warp owns one or more columns, using warp-level primitives to reduce across rows.
const int lane = threadIdx.x;
const int col = blockIdx.z * blockDim.y + threadIdx.y;
#if defined(__CUDA_ARCH__) && __CUDA_ARCH__ == GGML_CUDA_CC_DGX_SPARK
constexpr int cols_per_warp = S_v == 128 && !KDA ? 4 : 1;
#else
constexpr int cols_per_warp = 1;
#endif
const int col = (blockIdx.z * blockDim.y + threadIdx.y) * cols_per_warp;

const uint32_t iq1 = fastmodulo(h_idx, neqk1_magic);
const uint32_t iq3 = fastdiv(sequence, rq3_magic);
Expand All @@ -44,20 +56,23 @@ gated_delta_net_cuda(const float * q,
const int64_t state_in_offset = 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 + col * S_v;
curr_state += state_in_offset;
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];
float s_shard[cols_per_warp][rows_per_lane];
// state is stored transposed: M[col][i] = S[i][col], row col is contiguous

ggml_cuda_pdl_sync();
#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 c = 0; c < cols_per_warp; ++c) {
#pragma unroll
for (int r = 0; r < rows_per_lane; r++) {
const int i = r * warp_size + lane;
s_shard[c][r] = curr_state[(col + c) * S_v + i];
}
}

for (int t = 0; t < n_tokens; t++) {
Expand All @@ -82,40 +97,40 @@ gated_delta_net_cuda(const float * q,
}

if constexpr (!KDA) {
const float g_val = expf(*g_t);
const float g_val = G_PRECOMPUTED ? *g_t : expf(*g_t);

// kv[col] = (S^T @ k)[col] = sum_i S[i][col] * k[i]
float kv_shard = 0.0f;
// Each warp owns one or more columns and reuses the common q/k registers.
#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<warp_size>(kv_shard);
for (int c = 0; c < cols_per_warp; ++c) {
float kv_shard = 0.0f;
#pragma unroll
for (int r = 0; r < rows_per_lane; r++) {
kv_shard += s_shard[c][r] * k_reg[r];
}
float kv_col = warp_reduce_sum<warp_size>(kv_shard);

// delta[col] = (v[col] - g * kv[col]) * beta
float delta_col = (v_t[col] - g_val * kv_col) * beta_val;
float delta_col = (v_t[col + c] - g_val * kv_col) * beta_val;

// fused: S[i][col] = g * S[i][col] + k[i] * delta[col]
// attn[col] = (S^T @ q)[col] = sum_i S[i][col] * q[i]
float attn_partial = 0.0f;
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];
}
for (int r = 0; r < rows_per_lane; r++) {
s_shard[c][r] = g_val * s_shard[c][r] + k_reg[r] * delta_col;
attn_partial += s_shard[c][r] * q_reg[r];
}

float attn_col = warp_reduce_sum<warp_size>(attn_partial);
float attn_col = warp_reduce_sum<warp_size>(attn_partial);

if (lane == 0) {
attn_data[col] = attn_col * scale;
if (lane == 0) {
attn_data[col + c] = attn_col * scale;
}
}
} else {
// kv[col] = sum_i g[i] * S[i][col] * k[i]
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];
kv_shard += expf(g_t[i]) * s_shard[0][r] * k_reg[r];
}

float kv_col = warp_reduce_sum<warp_size>(kv_shard);
Expand All @@ -129,8 +144,8 @@ gated_delta_net_cuda(const float * q,
#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];
s_shard[0][r] = expf(g_t[i]) * s_shard[0][r] + k_reg[r] * delta_col;
attn_partial += s_shard[0][r] * q_reg[r];
}

float attn_col = warp_reduce_sum<warp_size>(attn_partial);
Expand All @@ -149,24 +164,31 @@ gated_delta_net_cuda(const float * q,
if (target_slot >= 0 && target_slot < K) {
float * curr_state = state + target_slot * state_slot_stride;
#pragma unroll
for (int r = 0; r < rows_per_lane; r++) {
const int i = r * warp_size + lane;
curr_state[col * S_v + i] = s_shard[r];
for (int c = 0; c < cols_per_warp; ++c) {
#pragma unroll
for (int r = 0; r < rows_per_lane; r++) {
const int i = r * warp_size + lane;
curr_state[(col + c) * S_v + i] = s_shard[c][r];
}
}
}
}

}

if constexpr (!keep_rs_t) {
#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];
for (int c = 0; c < cols_per_warp; ++c) {
#pragma unroll
for (int r = 0; r < rows_per_lane; r++) {
const int i = r * warp_size + lane;
state[(col + c) * S_v + i] = s_shard[c][r];
}
}
}
}

template <bool KDA, bool keep_rs_t>
template <bool KDA, bool keep_rs_t, bool G_PRECOMPUTED = false>
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,
Expand All @@ -179,8 +201,10 @@ static void launch_gated_delta_net(
float scale, int64_t state_slot_stride, int K, 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;
const int num_warps = 4;
dim3 grid_dims(H, n_seqs, (S_v + num_warps - 1) / num_warps);
const int cols_per_warp = cc == GGML_CUDA_CC_DGX_SPARK && S_v == 128 && !KDA ? 4 : 1;
dim3 grid_dims(H, n_seqs, (S_v + num_warps * cols_per_warp - 1) / (num_warps * cols_per_warp));
dim3 block_dims(warp_size <= S_v ? warp_size : S_v, num_warps, 1);

const uint3 neqk1_magic = init_fastdiv_values(neqk1);
Expand All @@ -189,26 +213,26 @@ static void launch_gated_delta_net(
const ggml_cuda_kernel_launch_params launch_params = ggml_cuda_kernel_launch_params(grid_dims, block_dims, 0, stream);
switch (S_v) {
case 16:
ggml_cuda_kernel_launch(gated_delta_net_cuda<16, KDA, keep_rs_t>, launch_params,
ggml_cuda_kernel_launch(gated_delta_net_cuda<16, KDA, keep_rs_t, G_PRECOMPUTED>, 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);
break;
case 32:
ggml_cuda_kernel_launch(gated_delta_net_cuda<32, KDA, keep_rs_t>, launch_params,
ggml_cuda_kernel_launch(gated_delta_net_cuda<32, KDA, keep_rs_t, G_PRECOMPUTED>, 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);
break;
case 64: {
ggml_cuda_kernel_launch(gated_delta_net_cuda<64, KDA, keep_rs_t>, launch_params,
ggml_cuda_kernel_launch(gated_delta_net_cuda<64, KDA, keep_rs_t, G_PRECOMPUTED>, 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);
break;
}
case 128: {
ggml_cuda_kernel_launch(gated_delta_net_cuda<128, KDA, keep_rs_t>, launch_params,
ggml_cuda_kernel_launch(gated_delta_net_cuda<128, KDA, keep_rs_t, G_PRECOMPUTED>, 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);
Expand Down Expand Up @@ -294,6 +318,19 @@ static void ggml_cuda_op_gated_delta_net_impl(
state_slot_stride = cache->slot_stride;
}

ggml_cuda_pool_alloc<float> g_exp_alloc(ctx.pool());
bool g_precomputed = false;
if (!kda && S_v == 128 && n_tokens >= 32 &&
ggml_cuda_info().devices[ggml_cuda_get_device()].cc == GGML_CUDA_CC_DGX_SPARK) {
const int64_t n_g = ggml_nelements(src_g);
g_exp_alloc.alloc(n_g);
const int block = 256;
const int grid = std::min<int64_t>((n_g + block - 1)/block, 4096);
gdn_precompute_exp<<<grid, block, 0, stream>>>(g_d, g_exp_alloc.ptr, n_g);
g_d = g_exp_alloc.ptr;
g_precomputed = true;
}

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 All @@ -306,13 +343,25 @@ static void ggml_cuda_op_gated_delta_net_impl(
}
} else {
if (keep_rs) {
launch_gated_delta_net<false, true>(q_d, k_d, v_d, g_d, b_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);
if (g_precomputed) {
launch_gated_delta_net<false, true, true>(q_d, k_d, v_d, g_d, b_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);
} else {
launch_gated_delta_net<false, true>(q_d, k_d, v_d, g_d, b_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);
}
} else {
launch_gated_delta_net<false, false>(q_d, k_d, v_d, g_d, b_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);
if (g_precomputed) {
launch_gated_delta_net<false, false, true>(q_d, k_d, v_d, g_d, b_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);
} else {
launch_gated_delta_net<false, false>(q_d, k_d, v_d, g_d, b_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);
}
}
}
}
Expand Down
Loading