Skip to content

Commit 32408c8

Browse files
authored
moe fp8 blockscale use nt (#3524)
* nt on fp8 blockscale * some improve and tests needs to be fixed * update * fix format * revert useless change * revert any change in amd_buffer_coherence
1 parent 4216d43 commit 32408c8

3 files changed

Lines changed: 63 additions & 34 deletions

File tree

example/65_gemm_multiply_multiply/moe_gemm1_xdl_fp8_blockscale_splitk.cpp

Lines changed: 15 additions & 11 deletions
Original file line numberDiff line numberDiff line change
@@ -119,7 +119,7 @@ static constexpr ck::index_t ActOP = 0; // 0: gelu_and_mul, 1: silu_an
119119
static constexpr bool MulRoutedWeight = false; // splitk gemm1 does not do routedWeight.
120120

121121
#if 1
122-
static constexpr ck::index_t MPerBlock = 32;
122+
static constexpr ck::index_t MPerBlock = 64;
123123
static constexpr ck::index_t NPerBlock = 128;
124124
static constexpr ck::index_t MNPerXDL = 16;
125125
static constexpr ck::index_t MXDLPerWave = MPerBlock / (MNPerXDL * 1);
@@ -156,7 +156,8 @@ using DeviceOpInstance = ck::tensor_operation::device::DeviceMoeGemmBlockScale
156156
// MXdlPerWave| NXdlPerWave| _MBlock_MWaveMPerXdl| ScalarPerVector|
157157
// PerShuffle| PerShuffle| _NBlock_NWaveNPerXdl| _NWaveNPerXdl|
158158
CShuffleMXDLPerWave, CShuffleNXDLPerWave, S<1, 32, 1, 8>, S<EVec, D0Vec, D1Vec, 1>,
159-
ck::BlockGemmPipelineScheduler::Intrawave, ck::BlockGemmPipelineVersion::v1, ActOP, Nswizzle, IsInputGemm, IsSplitK, MulRoutedWeight, int32_t, A0DataType>;
159+
ck::BlockGemmPipelineScheduler::Intrawave, ck::BlockGemmPipelineVersion::v1, ActOP, Nswizzle, IsInputGemm, IsSplitK, MulRoutedWeight,
160+
int32_t, A0DataType, A0DataType, A0DataType, A0DataType, true>;
160161
#else
161162
162163
static constexpr ck::index_t MPerBlock = 64; using DeviceOpInstance = ck::tensor_operation::device::DeviceMoeGemmBlockScale<
@@ -171,7 +172,8 @@ static constexpr ck::index_t MPerBlock = 64; using DeviceOpInstance = ck::tensor
171172
S<8, 32, 1>, S<1, 0, 2>, S<1, 0, 2>, 2, 16, 16, 0,
172173
S<8, 32, 1>, S<1, 0, 2>, S<1, 0, 2>, 2, 16, 16, 0,
173174
4, 2, S<1, 32, 1, 8>, S<2, 1, 1, 1>,
174-
ck::BlockGemmPipelineScheduler::Intrawave, ck::BlockGemmPipelineVersion::v3, ActOP, Nswizzle, IsInputGemm, IsSplitK, MulRoutedWeight, int32_t, A0DataType>;
175+
ck::BlockGemmPipelineScheduler::Intrawave, ck::BlockGemmPipelineVersion::v3, ActOP, Nswizzle, IsInputGemm, IsSplitK, MulRoutedWeight,
176+
int32_t, A0DataType, A0DataType, A0DataType, A0DataType, false>;
175177
#endif
176178
// clang-format on
177179

@@ -182,12 +184,14 @@ int main(int argc, char* argv[])
182184
bool time_kernel = true;
183185
#if 1
184186
// GEMM shape
185-
ck::index_t N = 4096;
186-
ck::index_t K = 6144;
187+
ck::index_t N = 1536;
188+
ck::index_t K = 4096;
189+
// ck::index_t N = 4096;
190+
// ck::index_t K = 6144;
187191
// ck::index_t N = 128;
188192
// ck::index_t K = 512;
189-
ck::index_t experts = 8;
190-
ck::index_t topk = 2;
193+
ck::index_t experts = 16;
194+
ck::index_t topk = 8;
191195
// ck::index_t sorted_tile_num = 515;
192196
// ck::index_t valid_tile_num = 512;
193197
// ck::index_t tokens = 208;
@@ -196,9 +200,9 @@ int main(int argc, char* argv[])
196200
// ck::index_t sorted_tile_num = 259;
197201
// ck::index_t valid_tile_num = 256;
198202
// ck::index_t tokens = 4096;
199-
ck::index_t sorted_tile_num = 2;
200-
ck::index_t valid_tile_num = 2;
201-
ck::index_t tokens = 32;
203+
ck::index_t sorted_tile_num = 16;
204+
ck::index_t valid_tile_num = 16;
205+
ck::index_t tokens = 4;
202206
#else
203207
// deepseek
204208
ck::index_t N = 2048;
@@ -209,7 +213,7 @@ int main(int argc, char* argv[])
209213
ck::index_t sorted_tile_num = 261;
210214
ck::index_t valid_tile_num = 256;
211215
#endif
212-
ck::index_t KBatch = 6;
216+
ck::index_t KBatch = 1;
213217
if(argc == 1)
214218
{
215219
// use default case

include/ck/tensor_operation/gpu/device/impl/device_moe_gemm_blockscale.hpp

Lines changed: 4 additions & 2 deletions
Original file line numberDiff line numberDiff line change
@@ -80,7 +80,8 @@ template <typename ALayout,
8080
typename ComputeTypeA = CDataType,
8181
typename ComputeTypeB = ComputeTypeA,
8282
typename LDSTypeA = ComputeTypeA,
83-
typename LDSTypeB = ComputeTypeB>
83+
typename LDSTypeB = ComputeTypeB,
84+
bool NonTemporalLoadB = false>
8485
struct DeviceMoeGemmBlockScale
8586
: public DeviceGemmMultipleD_BlockScale_BPreshuffle<ALayout,
8687
BLayout,
@@ -163,7 +164,8 @@ struct DeviceMoeGemmBlockScale
163164
ComputeTypeA,
164165
ComputeTypeB,
165166
LDSTypeA,
166-
LDSTypeB>;
167+
LDSTypeB,
168+
NonTemporalLoadB>;
167169
using GridwiseGemm64 = GridwiseGemmBase<math::max(NXdlPerWave64, 1)>;
168170
using GridwiseGemm32 = GridwiseGemmBase<NXdlPerWave32>;
169171

include/ck/tensor_operation/gpu/grid/gridwise_moe_gemm_blockscale.hpp

Lines changed: 44 additions & 21 deletions
Original file line numberDiff line numberDiff line change
@@ -173,7 +173,8 @@ template <typename ALayout,
173173
typename ComputeTypeA = CDataType,
174174
typename ComputeTypeB = ComputeTypeA,
175175
typename LDSTypeA = ADataType,
176-
typename LDSTypeB = BDataType>
176+
typename LDSTypeB = BDataType,
177+
bool NonTemporalLoadB = false>
177178
struct GridwiseMoeGemmBlockScale
178179
{
179180
using AScaleType = float;
@@ -1202,6 +1203,13 @@ struct GridwiseMoeGemmBlockScale
12021203
BElementwiseOperation b_element_op,
12031204
CElementwiseOperation c_element_op)
12041205
{
1206+
#if defined(__gfx942__) || defined(__gfx950__)
1207+
constexpr auto b_coherence_flag = NonTemporalLoadB
1208+
? AmdBufferCoherenceEnum::WAVE_NT1
1209+
: AmdBufferCoherenceEnum::DefaultCoherence;
1210+
#else
1211+
constexpr auto b_coherence_flag = AmdBufferCoherenceEnum::DefaultCoherence;
1212+
#endif
12051213
ignore = b_element_op;
12061214
index_t BN0Shuffled = CalculateBN0Shuffled(problem.N * (IsInputGemm && IsSplitK ? 2 : 1));
12071215
index_t BK0Shuffled = CalculateBK0Shuffled(problem.K);
@@ -1300,15 +1308,16 @@ struct GridwiseMoeGemmBlockScale
13001308

13011309
const auto a_grid_buf = make_dynamic_buffer<AddressSpaceEnum::Global>(
13021310
p_a_grid, a_grid_desc_ak0_m_ak1.GetElementSpaceSize());
1303-
const auto b_grid_buf = make_dynamic_buffer<AddressSpaceEnum::Global>(
1311+
const auto b_grid_buf = make_dynamic_buffer<AddressSpaceEnum::Global, b_coherence_flag>(
13041312
p_b_grid + expert_id * static_cast<long_index_t>(expert_stride) / BPackedSize,
13051313
b_grid_desc_bpreshuffled.GetElementSpaceSize());
13061314

13071315
const auto a_scale_grid_buf = make_dynamic_buffer<AddressSpaceEnum::Global>(
13081316
p_a_scale_grid, a_scale_grid_desc_am_ak.GetElementSpaceSize());
1309-
const auto b_scale_grid_buf = make_dynamic_buffer<AddressSpaceEnum::Global>(
1310-
p_b_scale_grid + expert_id * expert_scale_stride,
1311-
b_scale_grid_desc_bn_ak.GetElementSpaceSize());
1317+
const auto b_scale_grid_buf =
1318+
make_dynamic_buffer<AddressSpaceEnum::Global, b_coherence_flag>(
1319+
p_b_scale_grid + expert_id * expert_scale_stride,
1320+
b_scale_grid_desc_bn_ak.GetElementSpaceSize());
13121321

13131322
// A matrix in LDS memory, dst of blockwise copy
13141323
constexpr auto a_block_desc_ak0_m_ak1 = GetABlockDescriptor_AK0PerBlock_MPerBlock_AK1();
@@ -1465,9 +1474,11 @@ struct GridwiseMoeGemmBlockScale
14651474
if constexpr(IsInputGemm && !IsSplitK)
14661475
{
14671476
const BDataType* p_b_grid_up = p_b_grid + expert_stride / 2 / BPackedSize;
1468-
const auto b_grid_buf_up = make_dynamic_buffer<AddressSpaceEnum::Global>(
1469-
p_b_grid_up + expert_id * static_cast<long_index_t>(expert_stride) / BPackedSize,
1470-
b_grid_desc_bpreshuffled.GetElementSpaceSize());
1477+
const auto b_grid_buf_up =
1478+
make_dynamic_buffer<AddressSpaceEnum::Global, b_coherence_flag>(
1479+
p_b_grid_up +
1480+
expert_id * static_cast<long_index_t>(expert_stride) / BPackedSize,
1481+
b_grid_desc_bpreshuffled.GetElementSpaceSize());
14711482
auto b_blockwise_copy_up = ThreadwiseTensorSliceTransfer_v2<
14721483
BDataType,
14731484
BDataType,
@@ -1485,9 +1496,10 @@ struct GridwiseMoeGemmBlockScale
14851496
KPack / KGroup * (get_thread_local_1d_id() % WarpSize)));
14861497
const BScaleType* p_b_scale_grid_up =
14871498
p_b_scale_grid + expert_scale_stride / 2 / BPackedSize;
1488-
const auto b_scale_grid_buf_up = make_dynamic_buffer<AddressSpaceEnum::Global>(
1489-
p_b_scale_grid_up + expert_id * expert_scale_stride,
1490-
b_scale_grid_desc_bn_ak.GetElementSpaceSize());
1499+
const auto b_scale_grid_buf_up =
1500+
make_dynamic_buffer<AddressSpaceEnum::Global, b_coherence_flag>(
1501+
p_b_scale_grid_up + expert_id * expert_scale_stride,
1502+
b_scale_grid_desc_bn_ak.GetElementSpaceSize());
14911503
auto b_scale_thread_copy_up =
14921504
ThreadwiseTensorSliceTransfer_v2<BScaleType,
14931505
BScaleType,
@@ -1958,6 +1970,13 @@ struct GridwiseMoeGemmBlockScale
19581970
BElementwiseOperation b_element_op,
19591971
CElementwiseOperation c_element_op)
19601972
{
1973+
#if defined(__gfx942__) || defined(__gfx950__)
1974+
constexpr auto b_coherence_flag = NonTemporalLoadB
1975+
? AmdBufferCoherenceEnum::WAVE_NT1
1976+
: AmdBufferCoherenceEnum::DefaultCoherence;
1977+
#else
1978+
constexpr auto b_coherence_flag = AmdBufferCoherenceEnum::DefaultCoherence;
1979+
#endif
19611980
ignore = b_element_op;
19621981
index_t BN0Shuffled = CalculateBN0Shuffled(problem.N);
19631982
index_t BK0Shuffled = CalculateBK0Shuffled(problem.K);
@@ -2054,15 +2073,16 @@ struct GridwiseMoeGemmBlockScale
20542073

20552074
const auto a_grid_buf = make_dynamic_buffer<AddressSpaceEnum::Global>(
20562075
p_a_grid, a_grid_desc_ak0_m_ak1.GetElementSpaceSize());
2057-
const auto b_grid_buf = make_dynamic_buffer<AddressSpaceEnum::Global>(
2076+
const auto b_grid_buf = make_dynamic_buffer<AddressSpaceEnum::Global, b_coherence_flag>(
20582077
p_b_grid + expert_id * static_cast<long_index_t>(expert_stride) / BPackedSize,
20592078
b_grid_desc_bpreshuffled.GetElementSpaceSize());
20602079

20612080
const auto a_scale_grid_buf = make_dynamic_buffer<AddressSpaceEnum::Global>(
20622081
p_a_scale_grid, a_scale_grid_desc_am_ak.GetElementSpaceSize());
2063-
const auto b_scale_grid_buf = make_dynamic_buffer<AddressSpaceEnum::Global>(
2064-
p_b_scale_grid + expert_id * expert_scale_stride,
2065-
b_scale_grid_desc_bn_ak.GetElementSpaceSize());
2082+
const auto b_scale_grid_buf =
2083+
make_dynamic_buffer<AddressSpaceEnum::Global, b_coherence_flag>(
2084+
p_b_scale_grid + expert_id * expert_scale_stride,
2085+
b_scale_grid_desc_bn_ak.GetElementSpaceSize());
20662086

20672087
// A matrix in LDS memory, dst of blockwise copy
20682088
constexpr auto a_block_desc_ak0_m_ak1 = GetABlockDescriptor_AK0PerBlock_MPerBlock_AK1();
@@ -2227,9 +2247,11 @@ struct GridwiseMoeGemmBlockScale
22272247
if constexpr(IsInputGemm && !IsSplitK)
22282248
{
22292249
const BDataType* p_b_grid_up = p_b_grid + expert_stride / 2 / BPackedSize;
2230-
const auto b_grid_buf_up = make_dynamic_buffer<AddressSpaceEnum::Global>(
2231-
p_b_grid_up + expert_id * static_cast<long_index_t>(expert_stride) / BPackedSize,
2232-
b_grid_desc_bpreshuffled.GetElementSpaceSize());
2250+
const auto b_grid_buf_up =
2251+
make_dynamic_buffer<AddressSpaceEnum::Global, b_coherence_flag>(
2252+
p_b_grid_up +
2253+
expert_id * static_cast<long_index_t>(expert_stride) / BPackedSize,
2254+
b_grid_desc_bpreshuffled.GetElementSpaceSize());
22332255
auto b_blockwise_copy_up = ThreadwiseTensorSliceTransfer_v2<
22342256
BDataType,
22352257
BDataType,
@@ -2247,9 +2269,10 @@ struct GridwiseMoeGemmBlockScale
22472269
KPack / KGroup * (get_thread_local_1d_id() % WarpSize)));
22482270
const BScaleType* p_b_scale_grid_up =
22492271
p_b_scale_grid + expert_scale_stride / 2 / BPackedSize;
2250-
const auto b_scale_grid_buf_up = make_dynamic_buffer<AddressSpaceEnum::Global>(
2251-
p_b_scale_grid_up + expert_id * expert_scale_stride / BPackedSize,
2252-
b_scale_grid_desc_bn_ak.GetElementSpaceSize());
2272+
const auto b_scale_grid_buf_up =
2273+
make_dynamic_buffer<AddressSpaceEnum::Global, b_coherence_flag>(
2274+
p_b_scale_grid_up + expert_id * expert_scale_stride / BPackedSize,
2275+
b_scale_grid_desc_bn_ak.GetElementSpaceSize());
22532276
auto b_scale_thread_copy_up =
22542277
ThreadwiseTensorSliceTransfer_v2<BScaleType,
22552278
BScaleType,

0 commit comments

Comments
 (0)