Skip to content

Commit 8ddd970

Browse files
[OP] support 192 head_dim (#8014)
1 parent f2f7120 commit 8ddd970

9 files changed

Lines changed: 265 additions & 97 deletions

File tree

custom_ops/gpu_ops/append_attn/append_attention_func.cuh

Lines changed: 4 additions & 5 deletions
Original file line numberDiff line numberDiff line change
@@ -163,7 +163,6 @@ __device__ __forceinline__ void load_q_global_smem_multi_warps(
163163

164164
template <uint32_t group_size,
165165
uint32_t num_frags_x,
166-
uint32_t num_frags_y,
167166
uint32_t HEAD_DIM,
168167
typename T>
169168
__device__ __forceinline__ void load_q_global_smem(
@@ -175,6 +174,7 @@ __device__ __forceinline__ void load_q_global_smem(
175174
const uint32_t qo_h_stride) {
176175
constexpr uint32_t num_vecs_per_head = HEAD_DIM / num_elems_per_128b<T>();
177176

177+
static_assert(HEAD_DIM % 64 == 0, "");
178178
const uint32_t tx = threadIdx.x, ty = threadIdx.y;
179179

180180
uint32_t q_smem_offset_w = // [NUM_WARP_Q, num_frags_x, 16, head_dim]
@@ -193,7 +193,7 @@ __device__ __forceinline__ void load_q_global_smem(
193193
const T* q_ptr =
194194
q_ptr_base + n_offset * qo_n_stride + h_offset * qo_h_stride;
195195
#pragma unroll
196-
for (uint32_t fyo = 0; fyo < num_frags_y / 4; ++fyo) {
196+
for (uint32_t fyo = 0; fyo < HEAD_DIM / 64; ++fyo) {
197197
q_smem->load_128b_async<SharedMemFillMode::kNoFill>(
198198
q_smem_offset_w, q_ptr, n_offset < qo_upper_bound);
199199
q_smem_offset_w =
@@ -202,7 +202,7 @@ __device__ __forceinline__ void load_q_global_smem(
202202
}
203203
q_smem_offset_w =
204204
q_smem->advance_offset_by_row<4, num_vecs_per_head>(q_smem_offset_w) -
205-
2 * num_frags_y; // num_frags_y / 4 * 8
205+
HEAD_DIM / 8;
206206
}
207207
}
208208
}
@@ -228,15 +228,14 @@ __device__ __forceinline__ void q_smem_inplace_multiply_sm_scale_multi_warps(
228228
}
229229
}
230230

231-
template <uint32_t num_frags_x, uint32_t num_frags_y, typename T>
231+
template <uint32_t num_frags_x, uint32_t head_dim, typename T>
232232
__device__ __forceinline__ void q_smem_inplace_multiply_sm_scale(
233233
smem_t* q_smem, // [num_frags_x * 16, num_frags_y * 16]
234234
const float sm_scale) {
235235
constexpr int vec_size = 16 / sizeof(T);
236236
using LoadT = AlignedVector<T, vec_size>;
237237
LoadT tmp_vec;
238238
const uint32_t tx = threadIdx.x, ty = threadIdx.y;
239-
constexpr uint32_t head_dim = num_frags_y * 16;
240239
constexpr uint32_t num_vecs_per_head = head_dim / num_elems_per_128b<T>();
241240

242241
#pragma unroll

custom_ops/gpu_ops/append_attn/encoder_write_cache_with_rope_impl.cuh

Lines changed: 4 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -1493,6 +1493,8 @@ __global__ void append_write_cache_kv_c8_qkv_dynamic(
14931493
const int max_blocks_per_seq,
14941494
const int num_heads,
14951495
const int kv_num_heads) {
1496+
if constexpr (HEAD_DIM == 192) return;
1497+
14961498
constexpr uint32_t num_vecs_per_head = HEAD_DIM / num_elems_per_128b<T>();
14971499
constexpr uint32_t pad_len = BLOCK_SIZE;
14981500
const uint32_t btid = blockIdx.x, kv_head_idx = blockIdx.z;
@@ -1919,6 +1921,8 @@ __global__ void append_write_cache_kv_c4_qkv(
19191921
const int max_blocks_per_seq,
19201922
const int num_heads,
19211923
const int kv_num_heads) {
1924+
if constexpr (HEAD_DIM == 192) return;
1925+
19221926
constexpr uint32_t num_vecs_per_head = HEAD_DIM / num_elems_per_128b<T>();
19231927
constexpr uint32_t pad_len = BLOCK_SIZE;
19241928
const uint32_t btid = blockIdx.x, kv_head_idx = blockIdx.z;

0 commit comments

Comments
 (0)