Skip to content

Commit 0489875

Browse files
illsilinassistant-librarian[bot]
authored andcommitted
[rocm-libraries] ROCm/rocm-libraries#12203 (commit 3d0a0e8)
[CK] enable new gfx1250-strict target. ## Motivation These changes will enable building CK for the new gfx1250-strict target by limiting access to some of the unsupported builtins. I have verified that the code builds successfully for the gfx1250-strict with a rocm/compiler installation that supports it. JIRA ID : AICK-2239 ## Technical Details <!-- Explain the changes along with any relevant GitHub links. --> ## Test Plan <!-- Explain any relevant testing done to verify this PR. --> ## Test Result <!-- Briefly summarize test outcomes. --> ## Submission Checklist - [ ] Look over the contributing guidelines at https://github.com/ROCm/TheRock/blob/main/GOVERNANCE.md#pull-requests.
1 parent 398d985 commit 0489875

40 files changed

Lines changed: 161 additions & 125 deletions

File tree

CMakeLists.txt

Lines changed: 1 addition & 1 deletion
Original file line numberDiff line numberDiff line change
@@ -196,7 +196,7 @@ function(_ck_gpu_target_string_to_id TARGET_STR OUT_VAR)
196196
set(${OUT_VAR} "0x1201" PARENT_SCOPE)
197197
elseif(_tgt MATCHES "^gfx12-generic$")
198198
set(${OUT_VAR} "0x12FF" PARENT_SCOPE)
199-
elseif(_tgt STREQUAL "gfx1250")
199+
elseif(_tgt STREQUAL "gfx1250" OR _tgt STREQUAL "gfx1250-strict")
200200
set(${OUT_VAR} "0x1250" PARENT_SCOPE)
201201
else()
202202
message(WARNING "_ck_gpu_target_string_to_id: unknown GPU target '${TARGET_STR}', skipping")

example/ck_tile/42_mx_gemm/run_mx_gemm.inc

Lines changed: 5 additions & 6 deletions
Original file line numberDiff line numberDiff line change
@@ -188,12 +188,11 @@ int run_mx_gemm_with_layouts(int argc, char* argv[], ALayout, BLayout, CLayout)
188188
{
189189
ck_tile::preShuffleScaleBufferPermuteN_gfx950<GemmConfig::N_Warp,
190190
GemmConfig::N_Tile,
191-
XdlMNThread>(
192-
scale_b_aligned.mData.data(),
193-
scale_b_shuffled.mData.data(),
194-
N_scale_aligned,
195-
scale_k_size,
196-
true);
191+
XdlMNThread>(scale_b_aligned.mData.data(),
192+
scale_b_shuffled.mData.data(),
193+
N_scale_aligned,
194+
scale_k_size,
195+
true);
197196
}
198197
else
199198
{

include/ck/ck.hpp

Lines changed: 7 additions & 9 deletions
Original file line numberDiff line numberDiff line change
@@ -51,13 +51,12 @@
5151
#endif
5252

5353
// define general macros for various architectures
54-
#if defined(__gfx908__) || defined(__gfx90a__) || defined(__gfx942__) || defined(__gfx950__) || \
55-
defined(__gfx9_4_generic__)
56-
#define __gfx9__
57-
#endif
5854
#if defined(__gfx942__) || defined(__gfx950__) || defined(__gfx9_4_generic__)
5955
#define __gfx94__
6056
#endif
57+
#if defined(__gfx908__) || defined(__gfx90a__) || defined(__gfx94__)
58+
#define __gfx9__
59+
#endif
6160
#if defined(__gfx1010__) || defined(__gfx1011__) || defined(__gfx1012__) || \
6261
defined(__gfx1013__) || defined(__gfx10_1_generic__)
6362
#define __gfx101__
@@ -72,16 +71,15 @@
7271
defined(__gfx1152__) || defined(__gfx1153__) || defined(__gfx11_generic__)
7372
#define __gfx11__
7473
#endif
75-
#if defined(__gfx1200__) || defined(__gfx1201__) || defined(__gfx12_generic__) || \
76-
defined(__gfx1250__)
77-
#define __gfx12__
78-
#endif
7974
#if defined(__gfx1200__) || defined(__gfx1201__) || defined(__gfx12_generic__)
8075
#define __gfx120__
8176
#endif
82-
#if defined(__gfx1250__)
77+
#if defined(__gfx1250__) || defined(__gfx1250_strict__)
8378
#define __gfx125__
8479
#endif
80+
#if defined(__gfx120__) || defined(__gfx125__)
81+
#define __gfx12__
82+
#endif
8583
// buffer resource
8684
#ifndef __HIP_DEVICE_COMPILE__ // for host code
8785
#define CK_BUFFER_RESOURCE_3RD_DWORD -1

include/ck/utility/amd_wmma.hpp

Lines changed: 3 additions & 3 deletions
Original file line numberDiff line numberDiff line change
@@ -20,7 +20,7 @@ namespace ck {
2020
#define __gfx120__
2121
#endif
2222

23-
#if defined(__gfx1250__)
23+
#if defined(__gfx1250__) || defined(__gfx1250_strict__)
2424
#define __gfx125__
2525
#endif
2626

@@ -1426,7 +1426,7 @@ struct intrin_wmma_scale_f32_32x16x128_f4<32, 16, ScaleOpselB, ScaleTypeA, Scale
14261426
is_same_v<ScaleTypeB, e5m3x4_scale_t> ||
14271427
is_same_v<ScaleTypeB, e4m3x4_scale_t>,
14281428
"ScaleTypeB must be e8m0x4_bexp_t, e5m3x4_scale_t, or e4m3x4_scale_t");
1429-
#if defined(__gfx125__)
1429+
#if defined(__gfx1250__)
14301430
int32x16_t arg_a = bit_cast<int32x16_t>(reg_a);
14311431
int32x8_t arg_b = bit_cast<int32x8_t>(reg_b);
14321432
reg_c.template AsType<float16_t>()(Number<0>{}) =
@@ -1480,7 +1480,7 @@ struct intrin_wmma_scale16_f32_32x16x128_f4<32, 16, ScaleOpselB, ScaleTypeA, Sca
14801480
is_same_v<ScaleTypeB, e5m3x8_scale_t> ||
14811481
is_same_v<ScaleTypeB, e4m3x8_scale_t>,
14821482
"ScaleTypeB must be e8m0x8_bexp_t, e5m3x8_scale_t, or e4m3x8_scale_t");
1483-
#if defined(__gfx125__)
1483+
#if defined(__gfx1250__)
14841484
int32x16_t arg_a = bit_cast<int32x16_t>(reg_a);
14851485
int32x8_t arg_b = bit_cast<int32x8_t>(reg_b);
14861486
reg_c.template AsType<float16_t>()(Number<0>{}) =

include/ck/utility/e4m3.hpp

Lines changed: 2 additions & 2 deletions
Original file line numberDiff line numberDiff line change
@@ -35,7 +35,7 @@ struct e4m3_scale_t
3535
}
3636
__host__ __device__ explicit e4m3_scale_t(float scale)
3737
{
38-
#if defined(__gfx1250__)
38+
#if defined(__gfx125__)
3939
union
4040
{
4141
float fval;
@@ -59,7 +59,7 @@ struct e4m3_scale_t
5959

6060
__host__ __device__ explicit operator float() const
6161
{
62-
#if defined(__gfx1250__)
62+
#if defined(__gfx125__)
6363
union
6464
{
6565
unsigned int i32val;

include/ck/utility/e5m3.hpp

Lines changed: 2 additions & 2 deletions
Original file line numberDiff line numberDiff line change
@@ -35,7 +35,7 @@ struct e5m3_scale_t
3535
}
3636
__host__ __device__ explicit e5m3_scale_t(float scale)
3737
{
38-
#if defined(__gfx1250__)
38+
#if defined(__gfx125__)
3939
union
4040
{
4141
float fval;
@@ -59,7 +59,7 @@ struct e5m3_scale_t
5959

6060
__host__ __device__ explicit operator float() const
6161
{
62-
#if defined(__gfx1250__)
62+
#if defined(__gfx125__)
6363
union
6464
{
6565
unsigned int i32val;

include/ck/utility/mxf4_utils.hpp

Lines changed: 2 additions & 2 deletions
Original file line numberDiff line numberDiff line change
@@ -8,7 +8,7 @@
88
#include "ck/utility/mxfp_utils.hpp"
99
#include "dtype_vector.hpp"
1010

11-
#if CK_MX_ARCH_950 || CK_MX_ARCH_125
11+
#if CK_MX_ARCH_950 || CK_MX_ARCH_1250
1212
#define CK_MX_FP4_CVT_FAST_PATH 1
1313
#else
1414
#define CK_MX_FP4_CVT_FAST_PATH 0
@@ -287,7 +287,7 @@ static inline __device__ f4x8_t cast_to_f4_scaled(T x, float scale)
287287
return ret.vf4;
288288
}
289289

290-
#elif CK_MX_ARCH_125
290+
#elif CK_MX_ARCH_1250
291291
// from f4
292292
template <typename T>
293293
static inline __device__ enable_if_t<scalar_type<T>::vector_size == 1, T>

include/ck/utility/mxf6_utils.hpp

Lines changed: 2 additions & 2 deletions
Original file line numberDiff line numberDiff line change
@@ -7,7 +7,7 @@
77
#include "ck/utility/numeric_limits.hpp"
88
#include "ck/utility/mxfp_utils.hpp"
99

10-
#if CK_MX_ARCH_950 || CK_MX_ARCH_125
10+
#if CK_MX_ARCH_950 || CK_MX_ARCH_1250
1111
#define CK_MX_FP6_CVT_FAST_PATH 1
1212
#else
1313
#define CK_MX_FP6_CVT_FAST_PATH 0
@@ -721,7 +721,7 @@ inline __device__ T_F6 cast_to_f6_scaled(T x, float scale)
721721
}
722722
}
723723

724-
#elif CK_MX_ARCH_125
724+
#elif CK_MX_ARCH_1250
725725
// from f6
726726
template <typename T, typename T_F6>
727727
inline __device__ enable_if_t<scalar_type<T>::vector_size == 1, T> cast_from_f6_scaled(T_F6 x,

include/ck/utility/mxf8_utils.hpp

Lines changed: 3 additions & 3 deletions
Original file line numberDiff line numberDiff line change
@@ -4,7 +4,7 @@
44
#include "ck/utility/numeric_limits.hpp"
55
#include "ck/utility/mxfp_utils.hpp"
66

7-
#if CK_MX_ARCH_950 || CK_MX_ARCH_125
7+
#if CK_MX_ARCH_950 || CK_MX_ARCH_1250
88
#define CK_MX_FP8_CVT_FAST_PATH 1
99
#else
1010
#define CK_MX_FP8_CVT_FAST_PATH 0
@@ -660,7 +660,7 @@ static __device__ fp8x8_storage_t cast_to_f8_from_bf16_scaled(bhalf8_t v,
660660
return ret.v8f8x1;
661661
}
662662

663-
#elif CK_MX_ARCH_125
663+
#elif CK_MX_ARCH_1250
664664

665665
// fp8 -> float 8
666666
template <ck_fp8_interpretation_t interpret, typename Ts, int Opsel>
@@ -1011,7 +1011,7 @@ static __device__ bhalf2_t cast_to_bf16_from_f8_scaled(float scale, fp8x2_storag
10111011
out.v8x1 = cast_to_bf16_from_f8_scaled<interpret>(scale, v8);
10121012
return out.v2x4[0];
10131013
}
1014-
#endif // CK_MX_ARCH_125
1014+
#endif // CK_MX_ARCH_1250
10151015

10161016
#endif // CK_MX_FP8_CVT_FAST_PATH
10171017

include/ck/utility/mxfp_utils.hpp

Lines changed: 6 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -16,6 +16,12 @@
1616
#define CK_MX_ARCH_125 0
1717
#endif
1818

19+
#if defined(__gfx1250__) && !defined(__gfx1250_strict__) && __HIP_DEVICE_COMPILE__
20+
#define CK_MX_ARCH_1250 1
21+
#else
22+
#define CK_MX_ARCH_1250 0
23+
#endif
24+
1925
#ifdef CK_CODE_GEN_RTC
2026
#define UINT_MAX 4294967295
2127
#endif

0 commit comments

Comments
 (0)