You need to enable JavaScript to run this app.
优惠活动
大模型
产品
解决方案
定价
更多

如何为运行时记录长度的不透明数据类型高效使用std::sort

问题

我希望为运行时确定记录长度的不透明数据类型使用std::sort。

已尝试方案

  • 间接排序记录指针或索引:缓存局部性极差,比较操作需要访问数据集中随机分布的行,性能表现糟糕。
  • 使用C语言的qsort函数:该函数支持运行时指定记录长度,性能比前一种方法好,但存在两个明显缺点:
    1. 比较回调无法像lambda那样被内联,性能有损失;
    2. 接口易用性差,往往需要额外创建结构体来保存比较所需的上下文。

已知可行思路提示

曾看到过一种思路:实现带有特化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;
}

关键修正点说明

  1. 迭代器类别:将iterator_category改为std::random_access_iterator_tag,满足std::sort的迭代器要求。
  2. 伪引用类型RecordRef:封装单个记录的指针和大小,既可以转换为uint8_t*适配lambda参数,又能让swap函数知道要交换的内存块大小。
  3. 完整的关系运算符:实现了<、>、<=、>=等所有随机访问迭代器需要的关系运算符,解决编译错误。
  4. 正确的swap实现:特化std::swap来交换整个记录块的内存,确保排序过程中能正确交换元素。
  5. 迭代器差值计算:修正operator-的逻辑,返回记录数量而非字节数,符合迭代器语义。

内容的提问来源于stack exchange,提问作者RandomBits

相关产品推荐
方舟 Agent Plan

超全模态模型 × Harness 升级,最新支持 Deepseek-V4.1-Flash、GLM-5.3 系列、Doubao-Seedream-5.0-pro、Kimi-K3 (部分), 限时 9.9 元起

最近更新时间:2026.07.20 08:27:12