@@ -32,14 +32,26 @@ inline bool use_gfx12_attention_schedule() {
3232 hipDeviceProp_t properties{};
3333 if (hipGetDevice (&device) == hipSuccess &&
3434 hipGetDeviceProperties (&properties, device) == hipSuccess &&
35- std::strncmp (properties.gcnArchName , " gfx12" , 5 ) == 0 ) {
35+ (std::strcmp (properties.gcnArchName , " gfx1170" ) == 0 ||
36+ std::strncmp (properties.gcnArchName , " gfx12" , 5 ) == 0 )) {
3637 return true ;
3738 }
3839 return false ;
3940 }();
4041 return selected;
4142}
4243
44+ inline bool use_gfx1170_attention_schedule () {
45+ static const bool selected = [] {
46+ int device = 0 ;
47+ hipDeviceProp_t properties{};
48+ return hipGetDevice (&device) == hipSuccess &&
49+ hipGetDeviceProperties (&properties, device) == hipSuccess &&
50+ std::strcmp (properties.gcnArchName , " gfx1170" ) == 0 ;
51+ }();
52+ return selected;
53+ }
54+
4355template <int HD , int BLOCK_M , int BLOCK_N >
4456struct Bf16AttentionShape {
4557 static_assert (HD % 16 == 0 );
@@ -50,7 +62,7 @@ struct Bf16AttentionShape {
5062 static constexpr int kDimTiles = HD / 16 ;
5163 static constexpr int kKeyTiles = BLOCK_N / 16 ;
5264 static constexpr int kKStride = HD + 4 ;
53- static constexpr int kVStride = BLOCK_N + 4 ;
65+ static constexpr int kVStride = BLOCK_N + ( BLOCK_M <= 64 ? 2 : 4 ) ;
5466};
5567
5668template <typename T>
@@ -90,7 +102,7 @@ __forceinline__ __device__ void store_bf16_acc(
90102}
91103
92104template <int HD , int BLOCK_M , int BLOCK_N , bool DIRECT_GLOBAL ,
93- bool NO_KV_TAIL = false >
105+ bool NO_KV_TAIL = false , bool STAGE_V = false >
94106__global__ __launch_bounds__ (BLOCK_M * 2 )
95107void bf16_attention_kernel(
96108 const __bf16* __restrict__ q, const __bf16* __restrict__ k,
@@ -107,7 +119,9 @@ void bf16_attention_kernel(
107119 constexpr int KT = Shape::kKeyTiles ;
108120
109121 __shared__ __align__ (16 ) __bf16 staged_k[
110- DIRECT_GLOBAL ? 1 : 2 * BLOCK_N * Shape::kKStride ];
122+ DIRECT_GLOBAL ? 1 : (STAGE_V ? 1 : 2 ) * BLOCK_N * Shape::kKStride ];
123+ __shared__ __align__ (16 ) __bf16 staged_v[
124+ STAGE_V ? HD * Shape::kVStride : 1 ];
111125
112126 const int tid = threadIdx.x ;
113127 const int lane = tid & 31 ;
@@ -158,7 +172,7 @@ void bf16_attention_kernel(
158172 constexpr int kKStageItems = DIRECT_GLOBAL
159173 ? 1
160174 : (kPacksPerTile + Shape::kThreads - 1 ) / Shape::kThreads ;
161- if constexpr (!DIRECT_GLOBAL ) {
175+ if constexpr (!DIRECT_GLOBAL && ! STAGE_V ) {
162176#pragma unroll
163177 for (int item = 0 ; item < kKStageItems ; ++item) {
164178 const int item_pack = tid + item * Shape::kThreads ;
@@ -185,8 +199,46 @@ void bf16_attention_kernel(
185199 const bool has_next_tile = tile + 1 < key_tiles;
186200 const int next_buffer = (tile + 1 ) & 1 ;
187201
202+ if constexpr (STAGE_V ) {
203+ if (tile != 0 ) __syncthreads ();
204+ #pragma unroll
205+ for (int item = 0 ; item < kKStageItems ; ++item) {
206+ const int item_pack = tid + item * Shape::kThreads ;
207+ if (item_pack >= kPacksPerTile ) continue ;
208+ const int item_key = item_pack / kPacksPerRow ;
209+ const int item_d =
210+ (item_pack - item_key * kPacksPerRow ) * kPackElems ;
211+ const int key_index = key0 + item_key;
212+ v8bf k_value{};
213+ if constexpr (NO_KV_TAIL ) {
214+ k_value = *reinterpret_cast <const v8bf*>(
215+ k_head + static_cast <int64_t >(key_index) * k_stride_n + item_d);
216+ } else if (key_index < kv_len) {
217+ k_value = *reinterpret_cast <const v8bf*>(
218+ k_head + static_cast <int64_t >(key_index) * k_stride_n + item_d);
219+ }
220+ store_bf16x8_padded (
221+ staged_k + item_key * Shape::kKStride + item_d, k_value);
222+
223+ v8bf v_value{};
224+ if constexpr (NO_KV_TAIL ) {
225+ v_value = *reinterpret_cast <const v8bf*>(
226+ v_head + static_cast <int64_t >(key_index) * v_stride_n + item_d);
227+ } else if (key_index < kv_len) {
228+ v_value = *reinterpret_cast <const v8bf*>(
229+ v_head + static_cast <int64_t >(key_index) * v_stride_n + item_d);
230+ }
231+ #pragma unroll
232+ for (int e = 0 ; e < kPackElems ; ++e) {
233+ staged_v[(item_d + e) * Shape::kVStride + item_key] =
234+ v_value[e];
235+ }
236+ }
237+ __syncthreads ();
238+ }
239+
188240 const __bf16* const current_k = staged_k +
189- (tile & 1 ) * BLOCK_N * Shape::kKStride ;
241+ (STAGE_V ? 0 : ( tile & 1 ) * BLOCK_N * Shape::kKStride ) ;
190242
191243 MmaBf16::Acc score[KT ];
192244#pragma unroll
@@ -217,7 +269,7 @@ void bf16_attention_kernel(
217269 // Prefetch the next K tile after score MMA so global latency hides
218270 // behind softmax prep instead of sitting in front of the MMA nest.
219271 v8bf next_k[kKStageItems ]{};
220- if constexpr (!DIRECT_GLOBAL ) {
272+ if constexpr (!DIRECT_GLOBAL && ! STAGE_V ) {
221273#pragma unroll
222274 for (int item = 0 ; item < kKStageItems ; ++item) {
223275 const int item_pack = tid + item * Shape::kThreads ;
@@ -275,10 +327,15 @@ void bf16_attention_kernel(
275327#pragma unroll
276328 for (int dt = 0 ; dt < DT ; ++dt) {
277329 MmaBf16::Frag v_fragment{};
278- if constexpr (!DIRECT_GLOBAL ) {
330+ if constexpr (STAGE_V ) {
331+ v_fragment = load_frag_16bit<MmaBf16>(
332+ staged_v + (dt * 16 + frag_row (lane)) * Shape::kVStride +
333+ kt * 16 ,
334+ lane);
335+ } else if constexpr (!DIRECT_GLOBAL ) {
279336 const int tile_key = key0 + kt * 16 ;
280337 if (tile_key + 16 <= kv_len) {
281- #if defined(__HIP_DEVICE_COMPILE__) && defined(COMFY_MMA_GFX12)
338+ #if defined(__HIP_DEVICE_COMPILE__) && defined(COMFY_MMA_GFX12) && defined(__GFX12__)
282339 const int source_key = ((lane >> 4 ) << 3 ) + (lane & 7 );
283340 const int source_half = (lane >> 3 ) & 1 ;
284341 const v8bf* const source = reinterpret_cast <const v8bf*>(
@@ -325,7 +382,7 @@ void bf16_attention_kernel(
325382 }
326383 row_sum += tile_sum;
327384
328- if constexpr (!DIRECT_GLOBAL ) {
385+ if constexpr (!DIRECT_GLOBAL && ! STAGE_V ) {
329386#pragma unroll
330387 for (int item = 0 ; item < kKStageItems ; ++item) {
331388 const int item_pack = tid + item * Shape::kThreads ;
@@ -356,7 +413,7 @@ void bf16_attention_kernel(
356413}
357414
358415template <int HD , int BLOCK_M , int BLOCK_N , bool DIRECT_GLOBAL ,
359- bool NO_KV_TAIL = false >
416+ bool NO_KV_TAIL = false , bool STAGE_V = false >
360417void launch_bf16_attention_shape (
361418 const void * q, const void * k, const void * v, void * output,
362419 int batch, int q_heads, int kv_heads, int q_len, int kv_len,
@@ -366,7 +423,7 @@ void launch_bf16_attention_shape(
366423 int64_t o_stride_b, int64_t o_stride_h, int64_t o_stride_n,
367424 float sm_scale, hipStream_t stream) {
368425 const dim3 grid ((q_len + BLOCK_M - 1 ) / BLOCK_M , q_heads, batch);
369- bf16_attention_kernel<HD , BLOCK_M , BLOCK_N , DIRECT_GLOBAL , NO_KV_TAIL >
426+ bf16_attention_kernel<HD , BLOCK_M , BLOCK_N , DIRECT_GLOBAL , NO_KV_TAIL , STAGE_V >
370427 <<<grid, Bf16AttentionShape<HD , BLOCK_M , BLOCK_N >::kThreads , 0 , stream>>>(
371428 static_cast <const __bf16*>(q), static_cast <const __bf16*>(k),
372429 static_cast <const __bf16*>(v), static_cast <__bf16*>(output),
@@ -393,14 +450,30 @@ void dispatch_bf16_attention(
393450 k_stride_n, v_stride_b, v_stride_h, v_stride_n, o_stride_b,
394451 o_stride_h, o_stride_n, sm_scale, stream);
395452 } else if (q_len <= 256 ) {
396- launch_bf16_attention_shape<HD , 32 , 32 , true >(
397- q, k, v, output, batch, q_heads, kv_heads, q_len, kv_len,
398- q_stride_b, q_stride_h, q_stride_n, k_stride_b, k_stride_h,
399- k_stride_n, v_stride_b, v_stride_h, v_stride_n, o_stride_b,
400- o_stride_h, o_stride_n, sm_scale, stream);
453+ if (use_gfx1170_attention_schedule () && kv_len >= 128 &&
454+ kv_len % 32 == 0 ) {
455+ launch_bf16_attention_shape<HD , 64 , 32 , false , true , true >(
456+ q, k, v, output, batch, q_heads, kv_heads, q_len, kv_len,
457+ q_stride_b, q_stride_h, q_stride_n, k_stride_b, k_stride_h,
458+ k_stride_n, v_stride_b, v_stride_h, v_stride_n, o_stride_b,
459+ o_stride_h, o_stride_n, sm_scale, stream);
460+ } else {
461+ launch_bf16_attention_shape<HD , 32 , 32 , true >(
462+ q, k, v, output, batch, q_heads, kv_heads, q_len, kv_len,
463+ q_stride_b, q_stride_h, q_stride_n, k_stride_b, k_stride_h,
464+ k_stride_n, v_stride_b, v_stride_h, v_stride_n, o_stride_b,
465+ o_stride_h, o_stride_n, sm_scale, stream);
466+ }
401467 } else {
402- if (use_gfx12_attention_schedule () && kv_len >= 1024 &&
468+ if (use_gfx1170_attention_schedule () && kv_len >= 1024 &&
403469 kv_len % 32 == 0 ) {
470+ launch_bf16_attention_shape<HD , 512 , 32 , false , true , true >(
471+ q, k, v, output, batch, q_heads, kv_heads, q_len, kv_len,
472+ q_stride_b, q_stride_h, q_stride_n, k_stride_b, k_stride_h,
473+ k_stride_n, v_stride_b, v_stride_h, v_stride_n, o_stride_b,
474+ o_stride_h, o_stride_n, sm_scale, stream);
475+ } else if (use_gfx12_attention_schedule () && kv_len >= 1024 &&
476+ kv_len % 32 == 0 ) {
404477 launch_bf16_attention_shape<HD , 128 , 32 , false , true >(
405478 q, k, v, output, batch, q_heads, kv_heads, q_len, kv_len,
406479 q_stride_b, q_stride_h, q_stride_n, k_stride_b, k_stride_h,
0 commit comments