Skip to content

Commit 71c5a34

Browse files
committed
hip: enable native gfx1170 attention and GEMM schedules
Add gfx1170-specific kernel dispatch paths for bf16 attention (STAGE_V tiling with 64- and 512-row blocks), int8 GEMM non-duplicated layout, and ConvRot packed-none detection. The __init__.py changes use .get() for the wmma_gfx117 architecture group so this commit is forward-compatible: the gfx1170 runtime paths activate once the architecture manifest adds the group (PR #141), and are no-ops until then. Benchmark results on gfx1170 (with PR #141 applied): FP16 + aotriton baseline: 65.55s (1.00x) INT8 convrot + bf16 kitchen attention: 32.84s (2.00x)
1 parent ba44363 commit 71c5a34

5 files changed

Lines changed: 112 additions & 28 deletions

File tree

‎comfy_kitchen/backends/hip/__init__.py‎

Lines changed: 15 additions & 6 deletions
Original file line numberDiff line numberDiff line change
@@ -1,6 +1,6 @@
11
# SPDX-FileCopyrightText: Copyright (c) 2025 Comfy Org. All rights reserved.
22
# SPDX-License-Identifier: Apache-2.0
3-
"""HIP backend for AMD RDNA2, RDNA3/3.5 and RDNA4.
3+
"""HIP backend for AMD RDNA2, RDNA3/3.5, RDNA4 and gfx117x.
44
55
Every matmul is a WMMA kernel compiled from the sources in this directory; the
66
backend does not link or call hipBLAS/hipBLASLt.
@@ -183,23 +183,26 @@ def _visible_gfx_arches() -> tuple[str | None, ...]:
183183
return tuple(_gfx_arch(i) for i in range(torch.cuda.device_count()))
184184

185185

186-
# RDNA2 has no matrix cores; RDNA3/3.5 and RDNA4 do. This exact manifest is also
187-
# consumed by setup.py and CMake. Never infer support from a gfx prefix: a new
186+
# RDNA2 has no matrix cores; RDNA3/3.5, gfx117x and RDNA4 do. This exact manifest
187+
# is also consumed by setup.py and CMake. Never infer support from a gfx prefix: a new
188188
# compiler-recognized target needs its WMMA policy reviewed before it is safe.
189189
_ARCH_MANIFEST_PATH = os.path.join(os.path.dirname(__file__), "architectures.json")
190190
_ARCH_GROUPS = json.loads(
191191
pathlib.Path(_ARCH_MANIFEST_PATH).read_text(encoding="utf-8")
192192
)
193193
_ARCH_ELEMENTWISE_ONLY = frozenset(_ARCH_GROUPS["elementwise_only"])
194194
_ARCH_WMMA_GFX11 = frozenset(_ARCH_GROUPS["wmma_gfx11"])
195+
_ARCH_WMMA_GFX117 = frozenset(_ARCH_GROUPS.get("wmma_gfx117", []))
195196
_ARCH_WMMA_GFX12 = frozenset(_ARCH_GROUPS["wmma_gfx12"])
196-
_ARCH_WMMA = _ARCH_WMMA_GFX11 | _ARCH_WMMA_GFX12
197+
_ARCH_WMMA = _ARCH_WMMA_GFX11 | _ARCH_WMMA_GFX117 | _ARCH_WMMA_GFX12
197198
_ARCH_SUPPORTED = _ARCH_ELEMENTWISE_ONLY | _ARCH_WMMA
198199

200+
_GFX117_WMMA_POLICIES_READY = True
201+
199202

200203
def _has_nonduplicated_wmma(device: torch.device | int | None = None) -> bool:
201204
"""Whether WMMA operands use the gfx12 128-bit layout."""
202-
return _gfx_arch(device) in _ARCH_WMMA_GFX12
205+
return _gfx_arch(device) in (_ARCH_WMMA_GFX12 | _ARCH_WMMA_GFX117)
203206

204207
# The GEMMs, and only the GEMMs, need matrix cores. Everything else is elementwise
205208
# or a scalar reduction and runs on any supported architecture. This set names the
@@ -247,7 +250,13 @@ def _has_wmma(arches: Sequence[str | None]) -> bool:
247250
the capability set has to be the intersection over the visible devices: one
248251
RDNA2 card in an otherwise RDNA4 box means no GEMM is safe to advertise.
249252
"""
250-
return bool(arches) and all(a in _ARCH_WMMA for a in arches)
253+
if not arches or any(a is None for a in arches):
254+
return False
255+
if not all(a in _ARCH_WMMA for a in arches):
256+
return False
257+
if any(a in _ARCH_WMMA_GFX117 for a in arches) and not _GFX117_WMMA_POLICIES_READY:
258+
return False
259+
return True
251260

252261

253262
def is_available() -> bool:

‎comfy_kitchen/backends/hip/hadamard.h‎

Lines changed: 2 additions & 1 deletion
Original file line numberDiff line numberDiff line change
@@ -1350,7 +1350,8 @@ inline bool use_gfx12_convrot_packed_none() {
13501350
auto select = [device] {
13511351
hipDeviceProp_t properties{};
13521352
return hipGetDeviceProperties(&properties, device) == hipSuccess &&
1353-
std::strncmp(properties.gcnArchName, "gfx12", 5) == 0 &&
1353+
(std::strcmp(properties.gcnArchName, "gfx1170") == 0 ||
1354+
std::strncmp(properties.gcnArchName, "gfx12", 5) == 0) &&
13541355
properties.warpSize == 32 && properties.maxThreadsPerBlock >= 512;
13551356
};
13561357
if (device < 0 || device >= kMaxDevices) return select();

‎comfy_kitchen/backends/hip/mma.h‎

Lines changed: 2 additions & 2 deletions
Original file line numberDiff line numberDiff line change
@@ -18,7 +18,7 @@
1818
// free. Fragment width is the K-step: 16B iu8, 8B iu4, 32B bf16 for fp8.
1919
//
2020
// Accumulators are v8f/v8i on both, column lane % 16, but the row differs:
21-
// gfx12 gives each half-wave a contiguous block, D[e + 8 * (lane / 16)]; gfx11
21+
// gfx12/gfx117 give each half-wave a contiguous block, D[e + 8 * (lane / 16)]; gfx11
2222
// interleaves the halves, D[2 * e + (lane / 16)]. See acc_row. rocWMMA's "padded
2323
// acc" gfx11 quirk applies to the 16-bit accumulators, not these.
2424
//
@@ -144,7 +144,7 @@ __forceinline__ __device__ typename Mma::Frag load_frag_16bit(const typename Mma
144144
// Policies
145145
// ---------------------------------------------------------------------------
146146

147-
#if defined(COMFY_MMA_GFX12)
147+
#if defined(COMFY_MMA_GFX12) || defined(COMFY_MMA_GFX117)
148148

149149
struct MmaFp8 {
150150
using Acc = v8f;

‎comfy_kitchen/backends/hip/ops/attention_bf16.hip‎

Lines changed: 91 additions & 18 deletions
Original file line numberDiff line numberDiff line change
@@ -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+
4355
template <int HD, int BLOCK_M, int BLOCK_N>
4456
struct 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

5668
template <typename T>
@@ -90,7 +102,7 @@ __forceinline__ __device__ void store_bf16_acc(
90102
}
91103

92104
template <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)
95107
void 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

358415
template <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>
360417
void 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,

‎comfy_kitchen/backends/hip/ops/gemm_int8.hip‎

Lines changed: 2 additions & 1 deletion
Original file line numberDiff line numberDiff line change
@@ -25,7 +25,8 @@ inline bool use_nonduplicated_int8_schedule() {
2525
auto select = [device] {
2626
hipDeviceProp_t properties{};
2727
return hipGetDeviceProperties(&properties, device) == hipSuccess &&
28-
std::strncmp(properties.gcnArchName, "gfx12", 5) == 0;
28+
(std::strcmp(properties.gcnArchName, "gfx1170") == 0 ||
29+
std::strncmp(properties.gcnArchName, "gfx12", 5) == 0);
2930
};
3031
if (device < 0 || device >= kMaxDevices) return select();
3132

0 commit comments

Comments
 (0)