}
}
+template <int rows_per_block>
+static __global__ void __launch_bounds__(WARP_SIZE * rows_per_block, 2)
+rwkv_wkv7_f32_t1_warp_row(const int T, const int C, const int H, const float * r, const float * w, const float * k, const float * v, const float * a, const float * b, const float * s, float * dst) {
+ constexpr int head_size = CUDA_WKV_BLOCK_SIZE;
+ constexpr int half_head = head_size / 2;
+
+ const int lane = threadIdx.x;
+ const int row = blockIdx.y * rows_per_block + threadIdx.y;
+ const int bid = blockIdx.x;
+
+ const int batch_i = bid / H;
+ const int head_i = bid % H;
+ const int state_size = C * head_size;
+ const int head_off = head_i * head_size;
+ const int t = batch_i * C + head_off + row;
+
+ __shared__ float _r[head_size], _w[head_size], _k[head_size], _a[head_size], _b[head_size];
+
+ if (threadIdx.y == 0) {
+ _r[lane] = r[batch_i * C + head_off + lane];
+ _w[lane] = w[batch_i * C + head_off + lane];
+ _k[lane] = k[batch_i * C + head_off + lane];
+ _a[lane] = a[batch_i * C + head_off + lane];
+ _b[lane] = b[batch_i * C + head_off + lane];
+
+ _r[lane + half_head] = r[batch_i * C + head_off + lane + half_head];
+ _w[lane + half_head] = w[batch_i * C + head_off + lane + half_head];
+ _k[lane + half_head] = k[batch_i * C + head_off + lane + half_head];
+ _a[lane + half_head] = a[batch_i * C + head_off + lane + half_head];
+ _b[lane + half_head] = b[batch_i * C + head_off + lane + half_head];
+ }
+ __syncthreads();
+
+ const int64_t state_base = batch_i * state_size + head_i * head_size * head_size + row * head_size;
+ const float s0 = s[state_base + lane];
+ const float s1 = s[state_base + lane + half_head];
+ const float sa = warp_reduce_sum(_a[lane] * s0 + _a[lane + half_head] * s1);
+
+ const float vt = v[t];
+ const float st0 = s0 * _w[lane] + _k[lane] * vt + sa * _b[lane];
+ const float st1 = s1 * _w[lane + half_head] + _k[lane + half_head] * vt + sa * _b[lane + half_head];
+ const float y = warp_reduce_sum(st0 * _r[lane] + st1 * _r[lane + half_head]);
+
+ dst[T * C + state_base + lane] = st0;
+ dst[T * C + state_base + lane + half_head] = st1;
+
+ if (lane == 0) {
+ dst[t] = y;
+ }
+}
+
void ggml_cuda_op_rwkv_wkv6(ggml_backend_cuda_context & ctx, ggml_tensor * dst) {
const float * k_d = (const float *)dst->src[0]->data;
const float * v_d = (const float *)dst->src[1]->data;
GGML_ASSERT(C % H == 0);
GGML_ASSERT(C / H == CUDA_WKV_BLOCK_SIZE || C / H == CUDA_WKV_BLOCK_SIZE * 2);
- if (C / H == CUDA_WKV_BLOCK_SIZE) {
+ if (T / B == 1 && C / H == CUDA_WKV_BLOCK_SIZE) {
+ constexpr int rows_per_block = 4;
+ rwkv_wkv7_f32_t1_warp_row<rows_per_block><<<dim3(B * H, CUDA_WKV_BLOCK_SIZE / rows_per_block), dim3(WARP_SIZE, rows_per_block), 0, stream>>>(T, C, H, r_d, w_d, k_d, v_d, a_d, b_d, s_d, dst_d);
+ } else if (C / H == CUDA_WKV_BLOCK_SIZE) {
rwkv_wkv7_f32<CUDA_WKV_BLOCK_SIZE><<<B * H, C / H, 0, stream>>>(B, T, C, H, r_d, w_d, k_d, v_d, a_d, b_d, s_d, dst_d);
} else {
rwkv_wkv7_f32<CUDA_WKV_BLOCK_SIZE * 2><<<B * H, C / H, 0, stream>>>(B, T, C, H, r_d, w_d, k_d, v_d, a_d, b_d, s_d, dst_d);
test_cases.emplace_back(new test_rwkv_wkv6(GGML_TYPE_F32, 32, 64, 128, 4));
test_cases.emplace_back(new test_rwkv_wkv7(GGML_TYPE_F32, 32, 64, 1, 1));
+ test_cases.emplace_back(new test_rwkv_wkv7(GGML_TYPE_F32, 32, 64, 1, 4));
test_cases.emplace_back(new test_rwkv_wkv7(GGML_TYPE_F32, 32, 64, 32, 1));
test_cases.emplace_back(new test_rwkv_wkv7(GGML_TYPE_F32, 32, 64, 32, 4));
test_cases.emplace_back(new test_rwkv_wkv7(GGML_TYPE_F32, 32, 64, 128, 4));