复刻 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

浙公网安备 33010602011771号