Skip to content

Commit c44aacf

Browse files
author
Gennadij Yatskov
committed
[lib] Reduce templating in RadixSortGPU
1 parent 5163cd1 commit c44aacf

7 files changed

Lines changed: 383 additions & 435 deletions

File tree

examples/basic_sort/basic_sort.cpp

Lines changed: 1 addition & 1 deletion
Original file line numberDiff line numberDiff line change
@@ -29,7 +29,7 @@ bool sortAndVerify(ComputeState& compute, uint32_t numElements)
2929
// ------------------------------------------------------------------
3030
// 2. Sort on the GPU with a single call
3131
// ------------------------------------------------------------------
32-
RadixSortGPU<DataType> sorter;
32+
RadixSortGPU sorter;
3333
sorter.setLogStream(&std::cout);
3434

3535
std::vector<DataType> result;

examples/visualize/visualize.cpp

Lines changed: 7 additions & 7 deletions
Original file line numberDiff line numberDiff line change
@@ -1046,8 +1046,8 @@ bool sortDataZeroCopy(
10461046
constexpr auto numPasses = Params::_NUM_PASSES;
10471047

10481048
RandomDistributed<DataType> dataset(numElements);
1049-
RadixSortGPU<DataType> sorter;
1050-
[[maybe_unused]] const uint32_t nr = RadixSortGPU<DataType>::Resize(numElements);
1049+
RadixSortGPU sorter;
1050+
[[maybe_unused]] const uint32_t nr = RadixSortGPU::Resize(numElements);
10511051
assert(nr == numRounded);
10521052

10531053
// Write random data directly into the mapped Vulkan unsorted buffer.
@@ -1073,16 +1073,16 @@ bool sortDataZeroCopy(
10731073
{dstSorted, numRounded},
10741074
};
10751075

1076-
auto status = sorter.initialize(
1076+
auto status = sorter.initialize<DataType>(
10771077
compute.device(), compute.m_CLContext,
10781078
numElements, spans);
10791079
if (status != OperationStatus::OK) return false;
10801080

10811081
auto& q = compute.m_CLCommandQueue;
10821082
if (numRounded != numElements)
1083-
sorter.padGPUData(q, sizeof(DataType) * numElements);
1083+
sorter.padGPUData<DataType>(q, sizeof(DataType) * numElements);
10841084

1085-
status = sorter.uploadData(q);
1085+
status = sorter.uploadData<DataType>(q);
10861086
if (status != OperationStatus::OK) return false;
10871087

10881088
// Run pass-by-pass, capturing all intermediate buffer states.
@@ -1092,7 +1092,7 @@ bool sortDataZeroCopy(
10921092
sorter.Reorder(q, pass);
10931093

10941094
// Download all buffers to scratch
1095-
status = sorter.downloadKeys(q);
1095+
status = sorter.downloadKeys<DataType>(q);
10961096
if (status != OperationStatus::OK) return false;
10971097
status = sorter.downloadIntermediate(q);
10981098
if (status != OperationStatus::OK) return false;
@@ -1165,7 +1165,7 @@ int main()
11651165
createCommandPool(app);
11661166

11671167
// Determine padded size and create mapped Vulkan storage buffers.
1168-
const uint32_t numRounded = RadixSortGPU<uint32_t>{}.Resize(NUM_ELEMENTS);
1168+
const uint32_t numRounded = RadixSortGPU::Resize(NUM_ELEMENTS);
11691169
createMappedDataBuffers(app, numRounded);
11701170
auto* unsortedPtr = static_cast<uint32_t*>(app.mappedUnsorted);
11711171

src/HostData.h

Lines changed: 22 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -5,6 +5,7 @@
55
#include <vector>
66
#include <memory>
77
#include <span>
8+
#include <functional>
89

910
#include <cstdint>
1011

@@ -45,6 +46,27 @@ using HostSpans = HostBuffers<
4546
std::span<uint32_t>
4647
>;
4748

49+
struct HostSpansProxy {
50+
std::function<void*()> hKeysData;
51+
std::function<void*()> hHistogramsData;
52+
std::function<void*()> hGlobsumData;
53+
std::function<void*()> hPermutData;
54+
std::function<void*()> hOutputPermutData;
55+
std::function<void*()> hResultFromGPUData;
56+
57+
template <typename DataType>
58+
static HostSpansProxy FromHostSpans(HostSpans<DataType>& hostSpans) {
59+
return HostSpansProxy {
60+
[hostSpans]{ return static_cast<DataType*>(hostSpans.m_hKeys.data()); },
61+
[hostSpans]{ return static_cast<uint32_t*>(hostSpans.m_hHistograms.data()); },
62+
[hostSpans]{ return static_cast<uint32_t*>(hostSpans.m_hGlobsum.data()); },
63+
[hostSpans]{ return static_cast<uint32_t*>(hostSpans.h_Permut.data()); },
64+
[hostSpans]{ return static_cast<uint32_t*>(hostSpans.h_OutputPermut.data()); },
65+
[hostSpans]{ return static_cast<DataType*>(hostSpans.m_hResultFromGPU.data()); }
66+
};
67+
}
68+
};
69+
4870
/// @note Only used for tests
4971
template <typename T>
5072
struct HostDataWithReference

0 commit comments

Comments
 (0)