Quick Select算法返回错误:为何仅选末尾元素作pivot才正确?
快速选择算法中Pivot选择导致结果错误的问题分析
我在解决LeetCode「寻找vector nums的第k大元素」问题时,遇到了一个奇怪的现象:当选择中间元素作为pivot时,代码返回结果不符合预期(测试用例预期返回5,但实际返回6),只有选择最后一个元素作为pivot时结果才正确。按道理pivot的选择不应该影响算法正确性,这是为什么?
问题代码
#include <vector> #include <iostream> using namespace std; class Solution { public: int kthLargestElement(vector<int> nums, size_t kth) { return quickSelect(nums, 0, nums.size()-1, nums.size() - kth); } int quickSelect(vector<int> nums, size_t left, size_t right, size_t kth) { size_t tail = left, pointer = left, pivot = (left+right)/2; while(pointer <= pivot) { if(nums[pointer] < nums[pivot]) { std::swap(nums[tail], nums[pointer]); tail++; pointer++; } else pointer++; } std::swap(nums[tail], nums[pivot]); if(tail > kth) return quickSelect(nums, left, tail-1, kth); else if(tail < kth) return quickSelect(nums, tail+1, right, kth); return nums[tail]; } }; int main() { Solution s; auto test = vector<int>{3,2,1,5,6,4}; cout<< s.kthLargestElement(test, 2)<< endl; }
问题根源
核心问题出在分区逻辑的遍历范围错误:
- 当前代码的循环只遍历
[left, pivot]区间,完全忽略了pivot右侧的元素。这会导致pivot最终的位置tail无法正确反映它在整个[left, right]区间中的相对大小,后续递归的方向自然会出错。 - 当选择最后一个元素作为pivot时,
pivot = right,循环条件pointer <= pivot刚好覆盖整个[left, right]区间,分区逻辑能正常工作,所以结果正确。但选中间pivot时,右侧元素没参与比较,等于没完成完整的分区。
拿测试用例{3,2,1,5,6,4}举例:
- 第一次调用
quickSelect,pivot选索引2(值为1),循环只处理索引0-2,交换后tail=0,数组变为{1,2,3,5,6,4},递归进入右半区[1,5]。 - 第二次调用pivot选索引3(值为5),循环只处理索引1-3,交换后tail=3,递归进入右半区
[4,5]。 - 第三次调用pivot选索引4(值为6),循环处理索引4,交换后tail=4,此时tail等于目标kth=4,返回6,但实际第2大元素是5,错误由此产生。
修复方案
正确的分区逻辑需要遍历整个[left, right]区间。最稳妥的做法是先将选中的pivot元素交换到区间末尾,再对整个区间进行分区:
#include <vector> #include <iostream> using namespace std; class Solution { public: int kthLargestElement(vector<int>& nums, size_t kth) // 传引用避免拷贝 { return quickSelect(nums, 0, nums.size()-1, nums.size() - kth); } int quickSelect(vector<int>& nums, size_t left, size_t right, size_t kth) { // 选择中间元素作为pivot,先交换到区间末尾 size_t pivotIdx = (left + right) / 2; swap(nums[pivotIdx], nums[right]); int pivotVal = nums[right]; size_t tail = left; // 遍历整个[left, right-1]区间 for (size_t pointer = left; pointer < right; ++pointer) { if (nums[pointer] < pivotVal) { swap(nums[tail], nums[pointer]); ++tail; } } // 将pivot移到正确的位置 swap(nums[tail], nums[right]); if (tail > kth) { return quickSelect(nums, left, tail-1, kth); } else if (tail < kth) { return quickSelect(nums, tail+1, right, kth); } return nums[tail]; } }; int main() { Solution s; vector<int> test{3,2,1,5,6,4}; cout << s.kthLargestElement(test, 2) << endl; // 输出5,符合预期 }
额外优化
- 函数参数中的
nums改为引用传递,避免每次递归都拷贝整个数组,大幅提升效率。 - 可以加入随机选择pivot的逻辑,避免最坏情况的时间复杂度(比如数组已排序时)。
内容的提问来源于stack exchange,提问作者Miguel
相关产品推荐
相关产品推荐

