@@ -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 >
177178struct 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