如何为运行时记录长度的不透明数据类型高效使用std::sort
问题
我希望为运行时确定记录长度的不透明数据类型使用std::sort。
已尝试方案
- 间接排序记录指针或索引:缓存局部性极差,比较操作需要访问数据集中随机分布的行,性能表现糟糕。
- 使用C语言的
qsort函数:该函数支持运行时指定记录长度,性能比前一种方法好,但存在两个明显缺点:- 比较回调无法像lambda那样被内联,性能有损失;
- 接口易用性差,往往需要额外创建结构体来保存比较所需的上下文。
已知可行思路提示
曾看到过一种思路:实现带有特化swap的POD块伪引用迭代器,但没有找到具体的实现细节。
期望实现效果
我希望能直接调用std::sort并传入lambda作为比较函数,对运行时确定记录长度的不透明数据数组排序,示例代码如下:
#include <iostream> #include <random> #include <algorithm> #include <vector> using std::cout, std::endl; int main(int argc, const char *argv[]) { // 100条记录,每条50字节 int nrecords = 100, nbytes = 50; std::vector<uint8_t> data(nrecords * nbytes); // 生成测试用随机数据 std::uniform_int_distribution<uint8_t> d; std::mt19937_64 rng; std::generate(data.begin(), data.end(), [&]() { return d(rng); }); int key_index = 20; auto cmp = [&](const uint8_t* a, const uint8_t* b) { int aval = *(reinterpret_cast<const int*>(a + key_index)); int bval = *(reinterpret_cast<const int*>(b + key_index)); return aval < bval; }; // 如何在这里调用std::sort并支持运行时记录长度? // std::sort(data.begin(), ...) return 0; }
迭代器实现尝试及问题
我尝试自己编写迭代器,但编译失败,代码如下:
#include <algorithm> #include <random> #include <vector> struct MemoryIterator { using iterator_category = std::forward_iterator_tag; using difference_type = std::ptrdiff_t; using value_type = uint8_t *; using pointer = value_type*; using reference = value_type&; MemoryIterator(value_type ptr, size_t element_size) : ptr_(ptr) , element_size_(element_size) { } reference operator*() { return ptr_; } MemoryIterator& operator++() { ptr_ += element_size_; return *this; } MemoryIterator& operator--() { ptr_ -= element_size_; return *this; } MemoryIterator operator++(int) { auto tmp = *this; ++(*this); return tmp; } MemoryIterator operator--(int) { auto tmp = *this; --(*this); return tmp; } MemoryIterator& operator+=(size_t n) { ptr_ += n * element_size_; return *this; } MemoryIterator operator+(size_t n) { auto r = *this; r += n; return r; } MemoryIterator& operator-=(size_t n) { ptr_ -= n * element_size_; return *this; } MemoryIterator operator-(size_t n) { auto r = *this; r -= n; return r; } friend bool operator==(const MemoryIterator& a, const MemoryIterator& b) { return a.ptr_ == b.ptr_; } friend difference_type operator-(const MemoryIterator& a, const MemoryIterator& b) { return a.ptr_ - b.ptr_; } friend void swap(MemoryIterator a, MemoryIterator b) { } private: value_type ptr_; size_t element_size_; }; int main(int argc, const char *argv[]) { int nrecords = 100, nbytes = 50; std::vector<uint8_t> data(nrecords * nbytes); std::uniform_int_distribution<uint8_t> d; std::mt19937_64 rng; std::generate(data.begin(), data.end(), [&]() { return d(rng); }); int key_index = 20; auto cmp = [&](uint8_t *a, uint8_t *b) { int aval = *(int*)(a + key_index), bval = *(int*)(b + key_index); return aval < bval; }; MemoryIterator begin(data.data(), nbytes); MemoryIterator end(data.data() + data.size(), nbytes); std::sort(begin, end, cmp); return 0; }
编译错误信息
/opt/local/libexec/llvm-16/bin/../include/c++/v1/__algorithm/sort.h:533:21: error: invalid operands to binary expression ('MemoryIterator' and 'MemoryIterator') if (__i >= __j) ~~~ ^ ~~~ /opt/local/libexec/llvm-16/bin/../include/c++/v1/__algorithm/sort.h:639:8: note: in instantiation of function template specialization 'std::__introsort<std::_ClassicAlgPolicy, (lambda at sort0.cpp:92:16) &, MemoryIterator>' requested here std::__introsort<_AlgPolicy, _Compare>(__first, __last, __comp, __depth_limit); ^ /opt/local/libexec/llvm-16/bin/../include/c++/v1/__algorithm/sort.h:699:10: note: in instantiation of function template specialization 'std::__sort<(lambda at sort0.cpp:92:16) &, MemoryIterator>' requested here std::__sort<_WrappedComp>(std::__unwrap_iter(__first), std::__unwrap_iter(__last), _...
问题根源在于:std::sort要求迭代器满足随机访问迭代器的概念,而当前的MemoryIterator不仅把迭代器类别声明为std::forward_iterator_tag,还缺少>=、>、<、<=这些关系运算符的实现;同时迭代器的value_type和reference定义不符合要求,swap函数也没有实现实际的内存块交换逻辑。
解决方案:正确实现随机访问迭代器
要让std::sort正常工作,我们需要实现一个符合随机访问迭代器要求的迭代器,同时正确实现swap来交换整个记录块。以下是修正后的代码:
#include <algorithm> #include <random> #include <vector> #include <cstring> #include <cassert> // 伪引用类型,用于表示单个记录的"引用" struct RecordRef { uint8_t* ptr; size_t size; RecordRef(uint8_t* p, size_t s) : ptr(p), size(s) {} // 转换为指针,适配lambda的参数类型 operator uint8_t*() const { return ptr; } }; struct MemoryIterator { using iterator_category = std::random_access_iterator_tag; // 必须是随机访问迭代器 using difference_type = std::ptrdiff_t; using value_type = RecordRef; // 值类型为伪引用 using pointer = RecordRef*; using reference = RecordRef; // 引用类型直接用RecordRef(不需要真正的左值引用) MemoryIterator(uint8_t* ptr, size_t element_size) : ptr_(ptr) , element_size_(element_size) { } reference operator*() const { return RecordRef(ptr_, element_size_); } MemoryIterator& operator++() { ptr_ += element_size_; return *this; } MemoryIterator operator++(int) { auto tmp = *this; ++(*this); return tmp; } MemoryIterator& operator--() { ptr_ -= element_size_; return *this; } MemoryIterator operator--(int) { auto tmp = *this; --(*this); return tmp; } MemoryIterator& operator+=(difference_type n) { ptr_ += n * element_size_; return *this; } MemoryIterator operator+(difference_type n) const { auto r = *this; r += n; return r; } MemoryIterator& operator-=(difference_type n) { ptr_ -= n * element_size_; return *this; } MemoryIterator operator-(difference_type n) const { auto r = *this; r -= n; return r; } difference_type operator-(const MemoryIterator& other) const { return (ptr_ - other.ptr_) / element_size_; } bool operator==(const MemoryIterator& other) const { return ptr_ == other.ptr_; } bool operator!=(const MemoryIterator& other) const { return !(*this == other); } bool operator<(const MemoryIterator& other) const { return ptr_ < other.ptr_; } bool operator>(const MemoryIterator& other) const { return other < *this; } bool operator<=(const MemoryIterator& other) const { return !(*this > other); } bool operator>=(const MemoryIterator& other) const { return !(*this < other); } private: uint8_t* ptr_; size_t element_size_; }; // 特化swap,实现整个记录块的交换 namespace std { template<> void swap(RecordRef a, RecordRef b) { if (a.ptr == b.ptr) return; std::vector<uint8_t> tmp(a.size); std::memcpy(tmp.data(), a.ptr, a.size); std::memcpy(a.ptr, b.ptr, a.size); std::memcpy(b.ptr, tmp.data(), a.size); } } int main(int argc, const char *argv[]) { int nrecords = 100, nbytes = 50; std::vector<uint8_t> data(nrecords * nbytes); std::uniform_int_distribution<uint8_t> d; std::mt19937_64 rng; std::generate(data.begin(), data.end(), [&]() { return d(rng); }); int key_index = 20; auto cmp = [&](const uint8_t* a, const uint8_t* b) { int aval = *(reinterpret_cast<const int*>(a + key_index)); int bval = *(reinterpret_cast<const int*>(b + key_index)); return aval < bval; }; MemoryIterator begin(data.data(), nbytes); MemoryIterator end(data.data() + data.size(), nbytes); std::sort(begin, end, cmp); // 验证排序结果(可选) for (int i = 0; i < nrecords - 1; ++i) { const uint8_t* curr = &data[i * nbytes]; const uint8_t* next = &data[(i+1)*nbytes]; int curr_val = *(reinterpret_cast<const int*>(curr + key_index)); int next_val = *(reinterpret_cast<const int*>(next + key_index)); assert(curr_val <= next_val); } return 0; }
关键修正点说明
- 迭代器类别:将
iterator_category改为std::random_access_iterator_tag,满足std::sort的迭代器要求。 - 伪引用类型
RecordRef:封装单个记录的指针和大小,既可以转换为uint8_t*适配lambda参数,又能让swap函数知道要交换的内存块大小。 - 完整的关系运算符:实现了
<、>、<=、>=等所有随机访问迭代器需要的关系运算符,解决编译错误。 - 正确的
swap实现:特化std::swap来交换整个记录块的内存,确保排序过程中能正确交换元素。 - 迭代器差值计算:修正
operator-的逻辑,返回记录数量而非字节数,符合迭代器语义。
内容的提问来源于stack exchange,提问作者RandomBits
相关产品推荐
相关产品推荐

