@@ -163,7 +163,6 @@ __device__ __forceinline__ void load_q_global_smem_multi_warps(
163163
164164template <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
0 commit comments