Skip to content

Commit 569640d

Browse files
authored
Revert "Implement device grouped gemm fixed nk multi abd for rdna4 (#3619)" (#3705)
This reverts commit 301eb5c.
1 parent 8cbd09c commit 569640d

24 files changed

Lines changed: 120 additions & 3517 deletions

client_example/31_grouped_gemm_bf16Aint8B/grouped_gemm_bias_fastgelu_xdl_bf16_i8.cpp

Lines changed: 0 additions & 2 deletions
Original file line numberDiff line numberDiff line change
@@ -15,8 +15,6 @@
1515

1616
#include "ck/library/tensor_operation_instance/gpu/grouped_gemm_multi_abd_fixed_nk.hpp"
1717

18-
#include "ck/host_utility/hip_check_error.hpp"
19-
2018
using ::ck::hip_check_error;
2119

2220
template <ck::index_t... Is>

example/59_grouped_gemm_multi_ABD/CMakeLists.txt

Lines changed: 0 additions & 8 deletions
Original file line numberDiff line numberDiff line change
@@ -8,11 +8,3 @@ add_example_dependencies(example_grouped_gemm_xdl_multi_abd example_grouped_gemm
88

99
add_example_executable(example_grouped_gemm_multi_abd_xdl_fixed_nk_bias_bf16_i8 grouped_gemm_multi_abd_xdl_fixed_nk_bias_bf16_i8.cpp)
1010
add_example_dependencies(example_grouped_gemm_xdl_multi_abd example_grouped_gemm_multi_abd_xdl_fixed_nk_bias_bf16_i8)
11-
12-
add_custom_target(example_grouped_gemm_wmma_multi_abd)
13-
14-
add_example_executable(example_grouped_gemm_multi_abd_wmma_fixed_nk_bias_fp16 grouped_gemm_multi_abd_wmma_fixed_nk_bias_fp16.cpp)
15-
add_example_dependencies(example_grouped_gemm_wmma_multi_abd example_grouped_gemm_multi_abd_wmma_fixed_nk_bias_fp16)
16-
17-
add_example_executable(example_grouped_gemm_multi_abd_wmma_fixed_nk_bias_bf16_i8 grouped_gemm_multi_abd_wmma_fixed_nk_bias_bf16_i8.cpp)
18-
add_example_dependencies(example_grouped_gemm_wmma_multi_abd example_grouped_gemm_multi_abd_wmma_fixed_nk_bias_bf16_i8)

example/59_grouped_gemm_multi_ABD/grouped_gemm_multi_abd_wmma_fixed_nk_bias_bf16_i8.cpp

Lines changed: 0 additions & 400 deletions
This file was deleted.

example/59_grouped_gemm_multi_ABD/grouped_gemm_multi_abd_wmma_fixed_nk_bias_fp16.cpp

Lines changed: 0 additions & 396 deletions
This file was deleted.

example/59_grouped_gemm_multi_ABD/grouped_gemm_multi_abd_xdl_fixed_nk_bias_bf16_i8.cpp

Lines changed: 17 additions & 10 deletions
Original file line numberDiff line numberDiff line change
@@ -20,8 +20,6 @@
2020
#include "ck/library/utility/literals.hpp"
2121
#include "ck/library/reference_tensor_operation/cpu/reference_gemm.hpp"
2222

23-
#include "ck/host_utility/hip_check_error.hpp"
24-
2523
using ::ck::DeviceMem;
2624
using ::ck::hip_check_error;
2725
using ::ck::HostTensorDescriptor;
@@ -222,8 +220,8 @@ bool run_grouped_gemm(const ProblemSize& problem_size, const ExecutionConfig& co
222220

223221
for(int i = 0; i < group_count; i++)
224222
{
225-
a0_tensors_device.emplace_back(std::make_unique<DeviceMem>(
226-
sizeof(A0DataType) * problem_size.Ms[i] * problem_size.Ks[i]));
223+
a0_tensors_device.emplace_back(
224+
std::make_unique<DeviceMem>(sizeof(A0DataType) * sum_of_m * problem_size.Ks[i]));
227225

228226
b0_tensors_device.emplace_back(std::make_unique<DeviceMem>(
229227
sizeof(B0DataType) * problem_size.Ns[i] * problem_size.Ks[i]));
@@ -234,12 +232,21 @@ bool run_grouped_gemm(const ProblemSize& problem_size, const ExecutionConfig& co
234232
d0_tensors_device.emplace_back(
235233
std::make_unique<DeviceMem>(sizeof(D0DataType) * problem_size.Ns[i]));
236234

237-
c_tensors_device.emplace_back(std::make_unique<DeviceMem>(
238-
sizeof(EDataType) * problem_size.Ms[i] * problem_size.Ns[i]));
235+
c_tensors_device.emplace_back(
236+
std::make_unique<DeviceMem>(sizeof(EDataType) * sum_of_m * problem_size.Ns[i]));
237+
238+
a0_tensors_device[i]->ToDevice(a0_tensors[i].mData.data(),
239+
a0_tensors[i].mDesc.GetElementSpaceSize() *
240+
sizeof(A0DataType));
241+
242+
b0_tensors_device[i]->ToDevice(b0_tensors[i].mData.data(),
243+
b0_tensors[i].mDesc.GetElementSpaceSize() *
244+
sizeof(B0DataType));
245+
246+
b1_tensors_device[i]->ToDevice(b1_tensors[i].mData.data(),
247+
b1_tensors[i].mDesc.GetElementSpaceSize() *
248+
sizeof(B1DataType));
239249

240-
a0_tensors_device[i]->ToDevice(a0_tensors[i].mData.data());
241-
b0_tensors_device[i]->ToDevice(b0_tensors[i].mData.data());
242-
b1_tensors_device[i]->ToDevice(b1_tensors[i].mData.data());
243250
d0_tensors_device[i]->ToDevice(d0_tensors[i].mData.data());
244251
c_tensors_device[i]->SetZero();
245252

@@ -391,7 +398,7 @@ int main(int argc, char* argv[])
391398
{
392399
printf("arg1: verification (0=no, 1=yes)\n");
393400
printf("arg2: initialization (0=no init, 1=integer value, 2=decimal value)\n");
394-
printf("arg3: time kernel (0=no, 1=yes)\n");
401+
printf("arg3: time kernel (0=n0, 1=yes)\n");
395402
printf("arg4: k_batch (>0)\n");
396403
exit(0);
397404
}

example/59_grouped_gemm_multi_ABD/grouped_gemm_multi_abd_xdl_fixed_nk_bias_fp16.cpp

Lines changed: 19 additions & 14 deletions
Original file line numberDiff line numberDiff line change
@@ -20,8 +20,6 @@
2020
#include "ck/library/utility/literals.hpp"
2121
#include "ck/library/reference_tensor_operation/cpu/reference_gemm.hpp"
2222

23-
#include "ck/host_utility/hip_check_error.hpp"
24-
2523
using ::ck::DeviceMem;
2624
using ::ck::hip_check_error;
2725
using ::ck::HostTensorDescriptor;
@@ -49,9 +47,9 @@ using B0DataType = F16;
4947
using BsDataType = ck::Tuple<B0DataType>;
5048
using AccDataType = F32;
5149
using CShuffleDataType = F32;
52-
using D0DataType = F16;
50+
using D0DataType = F32;
5351
using DsDataType = ck::Tuple<D0DataType>;
54-
using EDataType = F16;
52+
using EDataType = F32;
5553

5654
using A0Layout = Row;
5755
using A1Layout = Row;
@@ -212,24 +210,31 @@ bool run_grouped_gemm(const ProblemSize& problem_size, const ExecutionConfig& co
212210

213211
for(int i = 0; i < group_count; i++)
214212
{
215-
a0_tensors_device.emplace_back(std::make_unique<DeviceMem>(
216-
sizeof(A0DataType) * problem_size.Ms[i] * problem_size.Ks[i]));
213+
a0_tensors_device.emplace_back(
214+
std::make_unique<DeviceMem>(sizeof(A0DataType) * sum_of_m * problem_size.Ks[i]));
217215

218-
a1_tensors_device.emplace_back(std::make_unique<DeviceMem>(
219-
sizeof(A1DataType) * problem_size.Ms[i] * problem_size.Ks[i]));
216+
a1_tensors_device.emplace_back(
217+
std::make_unique<DeviceMem>(sizeof(A1DataType) * sum_of_m * problem_size.Ks[i]));
220218

221219
b_tensors_device.emplace_back(std::make_unique<DeviceMem>(
222220
sizeof(B0DataType) * problem_size.Ns[i] * problem_size.Ks[i]));
223221

224222
d0_tensors_device.emplace_back(
225223
std::make_unique<DeviceMem>(sizeof(D0DataType) * problem_size.Ns[i]));
226224

227-
c_tensors_device.emplace_back(std::make_unique<DeviceMem>(
228-
sizeof(EDataType) * problem_size.Ms[i] * problem_size.Ns[i]));
225+
c_tensors_device.emplace_back(
226+
std::make_unique<DeviceMem>(sizeof(EDataType) * sum_of_m * problem_size.Ns[i]));
227+
228+
a0_tensors_device[i]->ToDevice(a0_tensors[i].mData.data(),
229+
a0_tensors[i].mDesc.GetElementSpaceSize() *
230+
sizeof(A0DataType));
229231

230-
a0_tensors_device[i]->ToDevice(a0_tensors[i].mData.data());
231-
a1_tensors_device[i]->ToDevice(a1_tensors[i].mData.data());
232-
b_tensors_device[i]->ToDevice(b_tensors[i].mData.data());
232+
a1_tensors_device[i]->ToDevice(a1_tensors[i].mData.data(),
233+
a1_tensors[i].mDesc.GetElementSpaceSize() *
234+
sizeof(A1DataType));
235+
b_tensors_device[i]->ToDevice(b_tensors[i].mData.data(),
236+
b_tensors[i].mDesc.GetElementSpaceSize() *
237+
sizeof(B0DataType));
233238
d0_tensors_device[i]->ToDevice(d0_tensors[i].mData.data());
234239
c_tensors_device[i]->SetZero();
235240

@@ -389,7 +394,7 @@ int main(int argc, char* argv[])
389394
{
390395
printf("arg1: verification (0=no, 1=yes)\n");
391396
printf("arg2: initialization (0=no init, 1=integer value, 2=decimal value)\n");
392-
printf("arg3: time kernel (0=no, 1=yes)\n");
397+
printf("arg3: time kernel (0=n0, 1=yes)\n");
393398
printf("arg4: k_batch (>0)\n");
394399
exit(0);
395400
}

0 commit comments

Comments
 (0)