复刻 polars top_k 的分区序

复刻 polars top_k 的分区序

选择算法,introselect,浮点累加顺序。

Top10 金额求和和 py 侧对不上,差值很小,但每次都稳定复现。C++ 这边原本的做法是维护一个容量 \(10\) 的降序数组,每来一笔成交就插进去,最后按降序顺序加起来。py 那边写的是 pl.col("amt").top_k(10).sum()。两边选出来的确实是同一批数,求和顺序不一样,double 累加出来的末位就对不上。

翻 polars 的 compute 层,top_k 最后走的是 select_nth_unstable_by(k, |a, b| tot_cmp(b, a)),接一个 resize(k)。也就是说它根本不排序,而是用 quickselect 把"最大的 \(k\) 个"切到数组前 \(k\) 个位置,这 \(k\) 个元素内部的顺序就是分区过程留下来的顺序,既不稳定也不降序。降序维护和整体排序都没法复现它。既然口径要对齐,只能把这个 quickselect 本身搬过来。

搬的是 Rust std 里 ipnsort 版的选择实现,涉及 select.rs、pivot.rs、quicksort.rs、smallsort.rs,收在一个单独的头文件里,外面套一层 SelectNth 命名空间。看代码时有个地方要绕一下:polars 传进去的比较是 tot_cmp(b, a),等价于按降序比,映射到 Rust 内部的 is_less 就是 \(a > b\)。于是这个文件里字面上的 \(>\) 都是"降序意义下的小于",minIndex 返回的是数值最大的下标,maxIndex 返回的是数值最小的下标。写的时候得一直反着念,注释里也特意标了这两句。

入口只有一层包装:

inline void selectNth(double* v, std::size_t len, std::size_t index) {
    partitionAtIndex(v, len, index);
}

partitionAtIndex 先把两种极端情况摘出去。\(\text{index} = \text{len} - 1\) 时目标是最小值,一次 maxIndex 线性扫描再换到末尾就完事;\(\text{index} = 0\) 时目标是最大值,走 minIndex。只有夹在中间的索引才进 partitionAtIndexLoop。所以业务上取 \(k = 10\)、数据量千万级的时候,走的是完整的 introselect 主循环,享受不到这两个快路径。

主循环每轮做三件事:判断能不能收尾、选 \(\text{pivot}\)、按分区结果丢一半。

std::size_t limit = kIntroselectLimit;   // 16
for (;;) {
    if (len <= kInsertionSortThreshold) {          // 16
        if (len >= 2) insertionSortShiftLeft(v, len, 1);
        return;
    }
    if (limit == 0) { medianOfMedians(v, len, index); return; }
    --limit;
    std::size_t pivotPos = choosePivot(v, len);
    ...
}

\(\text{limit}\) 是给快速路径准备的配额,每做一次分区扣一。它的判定在每轮开头,位置在"切片已经足够小"之后:先看 \(\text{len} \le 16\),再看 \(\text{limit} = 0\)。所以只要分区一直在有效缩小规模,插排分支会先生效,配额根本扣不完;反过来,连着 \(16\) 轮分区还没把切片压到 \(16\) 以下,第 \(17\) 轮进来时 \(\text{limit}\) 已经是 \(0\),就转 medianOfMedians 兜底。正常数据下这个兜底是摸不到的。

分区之后按 \(\text{mid}\) 和 \(\text{index}\) 的关系收缩,这一步是整个算法的骨架:

std::size_t mid = partition(v, len, pivotPos, false);
double* pivotPtr = v + mid;
if (mid < index) {          // 目标在右侧,丢掉左侧含 pivot
    v = pivotPtr + 1;
    len = len - mid - 1;
    index = index - mid - 1;
    ancestorPivot = pivotPtr;
} else if (mid > index) {   // 目标在左侧,只保留 v[0..mid)
    len = mid;
} else {
    return;                 // pivot 恰好落在目标位置
}

上面这一支里存下来的 \(\text{ancestorPivot}\) 是给重复元素准备的。下一轮如果出现"当前 \(\text{pivot}\) 不小于上一轮 \(\text{pivot}\)",说明数据里有成片的重复 \(\text{key}\),普通分区会把相等的元素全推到一侧,越分越偏,于是改用 reversed 分区,比较关系从 \(>\) 换成 \(\ge\),把重复元素归到同一侧。这个分支我一开始没看懂,注释写的是"避免重复元素退化",推了一遍才算接受:反复在同一批相等值上分区是快速路径最典型的退化来源。

\(\text{pivot}\) 的选择走 choosePivot,在 \(v_{0}\)、\(v_{\frac{\text{len}}{8} \times 4}\)、\(v_{\frac{\text{len}}{8} \times 7}\) 三个点上做采样。切片小于 \(64\) 就直接取三点中位数,否则交给 median3Rec:把尺度按 \(8\) 分之一的节奏递归下去,每层在当前锚点邻域再取三点、各自递归完再合并一次中位数,直到尺度小于 \(8\) 才退回 median3。它只比较很少的元素,返回的也只是"某个像中位数的元素",所以把它叫递归伪中位数。这么做的理由是 \(\text{pivot}\) 越接近真实中位数,每轮丢掉的元素越多;固定取三个点对偏斜或分段的数据容易估偏,多尺度采样能把局部噪声逐层中位掉。

真正干活的分区是 partition 加 partitionLomutoBranchlessCyclic:

点击查看代码
inline void partitionLoopBody(double*& right, std::size_t& numLt, double*& gapPos,
                              double* vBase, double pivot, bool reversed) {
    bool rightIsLt = reversed ? (*right >= pivot) : (*right > pivot);
    double* left = vBase + numLt;
    *gapPos = *left;    // left 旧值填洞
    *left = *right;     // right 值移到 left
    gapPos = right;     // 洞移到 right
    if (rightIsLt) ++numLt;
    ++right;
}

先把 \(\text{pivot}\) 换到 \(v_{0}\),对 \(v_{1} \sim v_{\text{len}-1}\) 做一轮 Lomuto 分区,最后把 \(\text{pivot}\) 换回 \(v_{\text{numLt}}\)。分区体用"洞"的方式挪数据:\(v_{0}\) 的值先存出来,位置变成洞,循环里把 \(\text{left}\) 的旧值填进洞、\(\text{right}\) 的值放到 \(\text{left}\)、洞跟着移到 \(\text{right}\)。因为 double 是 \(8\) 字节,小于 \(16\) 字节的展开上限,主循环一次推进两个元素。收尾那段用一个指向栈上 \(\text{gapValue}\) 的指针当作虚拟元素多跑一轮,是为了让洞落在最后一个元素时也能正确闭合。整个写法没有分支判断该往哪边写,只在计数时留一个 \(\text{if}\)。

兜底路径是另一个形状。medianOfMedians 的骨架和主循环一样,也是三个早退加一次收缩,区别在于 \(\text{pivot}\) 的来源换成了 medianOfNinthers:

if (len <= kInsertionSortThreshold) { ... 插排 ... return; }
if (k == len - 1) { std::swap(v[maxIndex(v, len)], v[k]); return; }
if (k == 0)        { std::swap(v[minIndex(v, len)], v[k]); return; }
std::size_t p = medianOfNinthers(v, len);
if (p > k) len = p;
else if (p < k) { v = v + p + 1; len = len - (p + 1); k = k - (p + 1); }
else return;

medianOfNinthers 是这个兜底里最绕的一段。它先按规模定一个采样个数 \(\text{frac}\):\(\text{len} \le 1024\) 取 \(\frac{\text{len}}{12}\),到 \(128\) K 取 \(\frac{\text{len}}{64}\),再大就 \(\frac{\text{len}}{1024}\)。然后在数组正中央留一块长度 \(\text{frac}\) 的区间 \([\text{lo}, \text{hi})\),用 \(\text{gap}\) 把左右两个滑动块隔开:

std::size_t gap = (len - 9 * frac) / 4;
std::size_t a = lo - 4 * frac - gap;
std::size_t b = hi + gap;
for (std::size_t i = lo; i < hi; ++i) {
    ninther(v, a, i - frac, b, a + 1, i, b + 1, a + 2, i + frac, b + 2);
    a += 3;  b += 3;
}

循环对中央区每个位置 \(i\) 拿 \(9\) 个点:左块 \(a, a+1, a+2\)、中列 \(i-\text{frac}, i, i+\text{frac}\)、右块 \(b, b+1, b+2\)。ninther 是三组三点各取中位、再取三个中位的中位数,把结果换到 \(v_{i}\)。\(a\) 和 \(b\) 每轮加 \(3\),两个块就沿着数组往右滑,配合中央的 \(i\),这 \(9\) 个采样点始终铺在数组的不同位置。循环跑完,中央区里就存下了 \(\text{frac}\) 个"局部中位数"。

接着是相互递归那一句:

medianOfMedians(v + lo, frac, pivot);        // pivot = frac / 2
return partition(v, len, lo + pivot, false);

在 \(\text{frac}\) 个局部中位数上再求一次中位数,规模从 \(\text{len}\) 直接掉到 \(\text{frac}\),递归链很短。拿到的那个值就是整组数据的稳健中位估计,最后以它为 \(\text{pivot}\) 对整个数组分区一次,返回落点给外层比较。之所以敢叫兜底,是因为 \(\text{pivot}\) 来自散布在整段数据上的采样中位,不会像三点采样那样被局部形态带偏,理论上每轮都能丢掉一个固定比例。

拿 \(\text{len} = 1000\) 代进去看一眼:\(\text{frac} = 83\),\(\text{pivot} = 41\),\(\text{lo} = 459\),\(\text{hi} = 542\),\(\text{gap} = 63\),\(a\) 从 \(64\) 推到 \(310\) 附近,\(b\) 从 \(605\) 推到 \(851\) 附近。循环结束后 medianOfMedians(v+459, 83, 41) 里,\(83\) 会再降到 \(6\),直接插排就结束,最后 \(v_{500}\) 成为样本中位数。整条递归链 \(1000\) 到 \(83\) 到 \(6\),代价基本由最外层的分区决定。

千万级取前 \(10\) 的情况,按每轮 \(\text{pivot}\) 都接近中位来推是这样:第一轮在 \(1000\) 万个元素上分区,左侧留下大约 \(500\) 万个"严格大于 \(\text{pivot}\)"的元素,右侧那 \(500\) 万个直接判定出局——左侧已经有一大批比它们大,前 \(10\) 名轮不到它们。之后每轮重复,\(\text{index}\) 一直是 \(10\),切片按 \(10^{7}\)、\(5 \times 10^{6}\)、\(2.5 \times 10^{6}\) 这样减半。理想中位数 \(\text{pivot}\) 下要二十来轮才能缩到 \(16\) 以内,比 \(\text{limit}\) 的 \(16\) 多一点,所以第 \(17\) 轮会进兜底;但那时切片只剩 \(150\) 个左右的候选,兜底在这么小的规模上几乎无感。总比较量是首轮 \(n\) 加上后面等比缩小的部分,量级还是 \(2n\) 左右。

工程侧的接点很短。调用方在 finalizeTopAmtValues 里,只有样本数超过 \(10\) 才调 selectNth,然后按 \(v_{0} \sim v_{9}\) 的现有顺序累加:

if (values.size() > FACTOR_TOP_N) {
    SelectNth::selectNth(values.data(), values.size(), FACTOR_TOP_N);
}

金额样本是按 push_back 收集的原始 tick 顺序,selectNth 之后前 \(10\) 个的位置就成了 polars 分区序,直接顺次相加。数量那一组没跟着这么做,qty 是整数,累加顺序不影响结果,固定容量的降序数组更省事。三组金额样本在收尾函数末尾用 swap 空 vector 换掉,否则这些临时缓冲会跟着 bar 缓存常驻内存。

还没闭环的地方:上游 Rust 的原始注释我没能联网核对,medianOfNinthers 那个"每轮丢掉固定比例"的结论是从代码结构和采样布局反推的,经典 BFPRT 里每侧 \(3n/10\) 那种具体界我没有独立证明。另外重复元素分支里 \(\text{mid} > \text{index}\) 时直接 \(\text{return}\) 的那一段,正确性我只能相信逐行移植本身,暂时没推明白为什么这个位置不需要继续缩。能跑对已经是万幸了

完整的 top_k cpp 复现代码:

点击查看代码
#pragma once
/**
 * @brief 复现 Rust std select_nth_unstable_by(ipnsort 版)的 f64 特化,用于对齐 py 的 top_k 输出序
 *
 * py polars 的 top_k 表达式在 compute 层走 select_nth_unstable_by(k, |a,b| tot_cmp(b,a)) 后
 * resize(k),输出前 k 个元素是"最大 k 个",但顺序是 quickselect 分区序(不稳定、非降序),
 * 直接决定 Top10 金额求和结果,无法用降序维护/普通排序复现。本命名空间逐行移植
 * rust-lang std core/src/slice/sort/(select.rs + pivot.rs + quicksort.rs + smallsort.rs,
 * ipnsort 版,Rust 1.81+),比较关系 is_less 等价于 a > b(金额无 NaN 时
 * tot_cmp(b,a) == Less 与 b < a 等价);ancestor_pivot 重复元素分支使用反转比较 a >= b。
 */

#include <cstddef>
#include <utility>

namespace SelectNth {
constexpr std::size_t kInsertionSortThreshold = 16;   // select.rs INSERTION_SORT_THRESHOLD
constexpr std::size_t kIntroselectLimit = 16;          // partition_at_index_loop 的兜底限制
constexpr std::size_t kPseudoMedianRecThreshold = 64;  // pivot.rs PSEUDO_MEDIAN_REC_THRESHOLD

/** @brief smallsort.rs insert_tail:把 tail 元素向左插入到已排序前缀的正确位置 */
inline void insertTail(double* begin, double* tail) {
    double* sift = tail - 1;
    // is_less(tail, sift) 为 false(tail <= sift)说明已就位,无需移动。
    if (!(*tail > *sift)) {
        return;
    }
    double tmp = *tail;  // 保存 tail 值,tail 成为"洞"(Rust CopyOnDrop 在无 panic 时等价于普通临时变量)
    double* gapDst = tail;
    for (;;) {
        *gapDst = *sift;  // sift 左移填补洞
        gapDst = sift;    // 洞前移
        if (sift == begin) {
            break;
        }
        --sift;
        if (!(tmp > *sift)) {  // tmp <= sift,找到插入位置
            break;
        }
    }
    *gapDst = tmp;  // tmp 回填洞的最终位置
}

/** @brief smallsort.rs insertion_sort_shift_left:从 offset 起对 v 做插入排序 */
inline void insertionSortShiftLeft(double* v, std::size_t len, std::size_t offset) {
    double* vEnd = v + len;
    double* tail = v + offset;
    while (tail != vEnd) {
        insertTail(v, tail);
        ++tail;
    }
}

/** @brief quicksort.rs loop_body:Lomuto 单步,normal 用 >,reversed 用 >= */
inline void partitionLoopBody(double*& right, std::size_t& numLt, double*& gapPos,
                              double* vBase, double pivot, bool reversed) {
    bool rightIsLt = reversed ? (*right >= pivot) : (*right > pivot);
    double* left = vBase + numLt;
    *gapPos = *left;    // left 旧值填入洞
    *left = *right;     // right 值移到 left
    gapPos = right;     // 洞移到 right
    if (rightIsLt) {
        ++numLt;
    }
    ++right;
}

/**
 * @brief quicksort.rs partition_lomuto_branchless_cyclic:分支无关 Lomuto + cyclic permutation
 * @note reversed=true 时比较改为 a >= b(对应 select.rs ancestor_pivot 重复元素分支)
 */
inline std::size_t partitionLomutoBranchlessCyclic(double* v, std::size_t len,
                                                   double pivot, bool reversed) {
    if (len == 0) {
        return 0;
    }
    double* vBase = v;
    double gapValue = vBase[0];   // ptr::read(v_base):保存 v[0],v[0] 成为洞
    double* gapValuePtr = &gapValue;  // cleanup 阶段把 gapValue 当作最后一个虚拟元素
    double* right = vBase + 1;
    std::size_t numLt = 0;
    double* gapPos = vBase;

    // f64 尺寸 8 字节 <= 16,unroll_len = 2,主循环每次推进两个元素。
    double* unrollEnd = vBase + len - 1;
    while (right < unrollEnd) {
        partitionLoopBody(right, numLt, gapPos, vBase, pivot, reversed);
        partitionLoopBody(right, numLt, gapPos, vBase, pivot, reversed);
    }

    double* end = vBase + len;
    for (;;) {
        bool isDone = (right == end);
        if (isDone) {
            right = gapValuePtr;  // 用保存的 gapValue 收尾最后一次循环
        }
        partitionLoopBody(right, numLt, gapPos, vBase, pivot, reversed);
        if (isDone) {
            break;
        }
    }
    return numLt;
}

/**
 * @brief quicksort.rs partition:pivot 移到开头,对 v[1..] 做 Lomuto 分区,再移回 v[numLt]
 * @return numLt:分区后大于(或 reversed 时大于等于)pivot 的元素数,pivot 位于 v[numLt]
 */
inline std::size_t partition(double* v, std::size_t len, std::size_t pivotIdx, bool reversed) {
    if (len == 0) {
        return 0;
    }
    std::swap(v[0], v[pivotIdx]);  // pivot 移到 v[0]
    double pivot = v[0];
    std::size_t numLt = partitionLomutoBranchlessCyclic(v + 1, len - 1, pivot, reversed);
    std::swap(v[0], v[numLt]);     // pivot 移到 v[numLt]
    return numLt;
}

/** @brief pivot.rs median3:对三个元素取中位数,返回选中元素的指针 */
inline const double* median3(const double* a, const double* b, const double* c) {
    bool x = (*a > *b);
    bool y = (*a > *c);
    if (x == y) {
        bool z = (*b > *c);
        return (z ^ x) ? c : b;  // XOR 分支,等价于 Rust 的 z ^ x
    }
    return a;
}

/** @brief pivot.rs median3_rec:递归近似中位数(n*8 >= 64 时三分递归) */
inline const double* median3Rec(const double* a, const double* b, const double* c,
                                std::size_t n) {
    if (n * 8 >= kPseudoMedianRecThreshold) {
        std::size_t n8 = n / 8;
        const double* ra = median3Rec(a, a + n8 * 4, a + n8 * 7, n8);
        const double* rb = median3Rec(b, b + n8 * 4, b + n8 * 7, n8);
        const double* rc = median3Rec(c, c + n8 * 4, c + n8 * 7, n8);
        return median3(ra, rb, rc);
    }
    return median3(a, b, c);
}

/** @brief pivot.rs choose_pivot:三段采样选 pivot,返回其在 v 中的索引 */
inline std::size_t choosePivot(const double* v, std::size_t len) {
    std::size_t lenDiv8 = len / 8;
    const double* a = v;
    const double* b = v + lenDiv8 * 4;
    const double* c = v + lenDiv8 * 7;
    const double* pivot = (len < kPseudoMedianRecThreshold) ? median3(a, b, c)
                                                            : median3Rec(a, b, c, lenDiv8);
    return static_cast<std::size_t>(pivot - v);
}

/** @brief select.rs median_idx:三元素中位数索引(a/b/c 为索引,比较 v[a] v[b] v[c]) */
inline std::size_t medianIdx(const double* v, std::size_t a, std::size_t b, std::size_t c) {
    if (v[c] > v[a]) {
        std::swap(a, c);
    }
    if (v[c] > v[b]) {
        return c;
    }
    if (v[b] > v[a]) {
        return a;
    }
    return b;
}

/**
 * @brief select.rs ninther:9 元素 Tukey ninther 中位数,b/d/f/h 为可变索引,
 *        a/c/e/g/i 为固定索引;mem::swap 交换索引,v.swap 交换数组元素
 */
inline void ninther(double* v, std::size_t a, std::size_t b, std::size_t c,
                    std::size_t d, std::size_t e, std::size_t f,
                    std::size_t g, std::size_t h, std::size_t i) {
    b = medianIdx(v, a, b, c);
    h = medianIdx(v, g, h, i);
    if (v[h] > v[b]) {
        std::swap(b, h);
    }
    if (v[f] > v[d]) {
        std::swap(d, f);
    }
    if (v[e] > v[d]) {
        // do nothing
    } else if (v[f] > v[e]) {
        d = f;
    } else {
        if (v[e] > v[b]) {
            std::swap(v[e], v[b]);
        } else if (v[h] > v[e]) {
            std::swap(v[e], v[h]);
        }
        return;
    }
    if (v[d] > v[b]) {
        d = b;
    } else if (v[h] > v[d]) {
        d = h;
    }
    std::swap(v[d], v[e]);
}

/** @brief select.rs min_index:is_less=a>b 意义下的最小 = 数值最大的索引 */
inline std::size_t minIndex(const double* v, std::size_t len) {
    std::size_t idx = 0;
    for (std::size_t i = 1; i < len; ++i) {
        if (v[i] > v[idx]) {
            idx = i;
        }
    }
    return idx;
}

/** @brief select.rs max_index:is_less=a>b 意义下的最大 = 数值最小的索引 */
inline std::size_t maxIndex(const double* v, std::size_t len) {
    std::size_t idx = 0;
    for (std::size_t i = 1; i < len; ++i) {
        if (v[idx] > v[i]) {
            idx = i;
        }
    }
    return idx;
}

inline void medianOfMedians(double* v, std::size_t len, std::size_t k);
inline std::size_t medianOfNinthers(double* v, std::size_t len);

/**
 * @brief select.rs median_of_medians:limit 耗尽时的最坏情况兜底,递归缩小切片
 * @note 与 median_of_ninthers 相互递归,故需前向声明
 */
inline void medianOfMedians(double* v, std::size_t len, std::size_t k) {
    for (;;) {
        if (len <= kInsertionSortThreshold) {
            if (len >= 2) {
                insertionSortShiftLeft(v, len, 1);
            }
            return;
        }
        if (k == len - 1) {  // 数值最小放最后(对应 index==len-1 的 max_index 分支)
            std::size_t maxIdx = maxIndex(v, len);
            std::swap(v[maxIdx], v[k]);
            return;
        }
        if (k == 0) {  // 数值最大放最前(对应 index==0 的 min_index 分支)
            std::size_t minIdx = minIndex(v, len);
            std::swap(v[minIdx], v[k]);
            return;
        }
        std::size_t p = medianOfNinthers(v, len);
        if (p == k) {
            return;
        } else if (p > k) {
            len = p;  // 缩到左侧 v[0..p]
        } else {
            v = v + p + 1;  // 缩到右侧 v[p+1..],k 相对偏移
            len = len - (p + 1);
            k = k - (p + 1);
        }
    }
}

/**
 * @brief select.rs median_of_ninthers:对 9 个切片取 ninther 中位数,递归选 pivot 后分区
 * @return partition 后的 num_lt(pivot 新位置)
 */
inline std::size_t medianOfNinthers(double* v, std::size_t len) {
    std::size_t frac;
    if (len <= 1024) {
        frac = len / 12;
    } else if (len <= 128 * 1024) {
        frac = len / 64;
    } else {
        frac = len / 1024;
    }
    std::size_t pivot = frac / 2;
    std::size_t lo = len / 2 - pivot;
    std::size_t hi = frac + lo;
    std::size_t gap = (len - 9 * frac) / 4;
    std::size_t a = lo - 4 * frac - gap;  // 推导恒非负:len >= 9*frac
    std::size_t b = hi + gap;
    for (std::size_t i = lo; i < hi; ++i) {
        ninther(v, a, i - frac, b, a + 1, i, b + 1, a + 2, i + frac, b + 2);
        a += 3;
        b += 3;
    }
    medianOfMedians(v + lo, frac, pivot);  // 对 v[lo..lo+frac] 选第 pivot 个
    return partition(v, len, lo + pivot, false);
}

/**
 * @brief select.rs partition_at_index_loop:introselect 主循环,limit 耗尽转 median_of_medians
 * @param ancestorPivot 上一轮 pivot 的指针,用于重复元素检测(nullptr 表示无)
 */
inline void partitionAtIndexLoop(double* v, std::size_t len, std::size_t index,
                                 const double* ancestorPivot) {
    std::size_t limit = kIntroselectLimit;
    for (;;) {
        if (len <= kInsertionSortThreshold) {
            if (len >= 2) {
                insertionSortShiftLeft(v, len, 1);
            }
            return;
        }
        if (limit == 0) {
            medianOfMedians(v, len, index);
            return;
        }
        --limit;
        std::size_t pivotPos = choosePivot(v, len);
        // 当前 pivot 不小于 ancestor_pivot 时用反转比较 a >= b 分区,避免重复元素退化。
        if (ancestorPivot != nullptr && !(*ancestorPivot > v[pivotPos])) {
            std::size_t numLt = partition(v, len, pivotPos, true);
            std::size_t mid = numLt + 1;
            if (mid > index) {
                return;
            }
            v = v + mid;
            len = len - mid;
            index = index - mid;
            ancestorPivot = nullptr;
            continue;
        }
        std::size_t mid = partition(v, len, pivotPos, false);
        double* pivotPtr = v + mid;
        if (mid < index) {  // 目标在右侧 v[mid+1..]
            v = pivotPtr + 1;
            len = len - mid - 1;
            index = index - mid - 1;
            ancestorPivot = pivotPtr;
        } else if (mid > index) {  // 目标在左侧 v[0..mid]
            len = mid;
        } else {
            return;
        }
    }
}

/**
 * @brief select.rs partition_at_index:select_nth 入口
 * @note index==0 / index==len-1 走线性 min/max 分支,其余走 introselect 主循环
 */
inline void partitionAtIndex(double* v, std::size_t len, std::size_t index) {
    if (index >= len) {
        return;  // Rust 此处 panic;调用方保证 index < len,故直接返回
    }
    if (index == len - 1) {
        std::size_t maxIdx = maxIndex(v, len);
        std::swap(v[maxIdx], v[index]);
    } else if (index == 0) {
        std::size_t minIdx = minIndex(v, len);
        std::swap(v[minIdx], v[index]);
    } else {
        partitionAtIndexLoop(v, len, index, nullptr);
    }
}

/**
 * @brief select_nth 对外入口:重排 v,使 v[0..index] 为 is_less=a>b 意义下前 index 个
 *        (数值最大 index 个,quickselect 分区序),v[index] 为第 index+1 个
 * @param v     待重排数组
 * @param len   数组长度
 * @param index 目标分区索引(0-based),须满足 index < len
 * @note 对应 py 的 vec.select_nth_unstable_by(k) + resize(k),随后取 v[0..k] 顺序累加
 */
inline void selectNth(double* v, std::size_t len, std::size_t index) {
    partitionAtIndex(v, len, index);
}

}  // namespace SelectNth

posted @ 2026-09-14 15:53  Ke_scholar  阅读(5)  评论(0)    收藏  举报