基于逆序对逻辑修改的翻转对计数代码为何无法得到正确结果?
翻转对计数代码错误分析与修正
问题描述
基于逆序对计数逻辑修改的翻转对(满足i<j且a[i]>2*a[j]的数对)计数代码,针对测试用例[4,1,2,3,1],正确结果应为3,但代码输出为2,需要定位问题并修正。
原错误代码
#include <iostream> #include <vector> using namespace std; int cnt = 0; void merge(vector<int> &a, int low, int mid, int high) { vector<int> temp; int left = low; int right = mid + 1; while (left <= mid && right <= high) { if (a[left] <= a[right]) { temp.push_back(a[left]); left++; } else { if (a[left] > 2 * a[right] && left < right) { cnt += mid - left + 1; } temp.push_back(a[right]); right++; } } while (left <= mid) { temp.push_back(a[left]); left++; } while (right <= high) { temp.push_back(a[right]); right++; } for (int i = low; i <= high; i++) { a[i] = temp[i - low]; } for (int i = low; i <= high; i++) { cout << a[i] << " "; } cout << cnt << endl; } void mergesort(vector<int> &a, int low, int high) { if (low >= high) return; int mid = low + (high - low) / 2; mergesort(a, low, mid); mergesort(a, mid + 1, high); merge(a, low, mid, high); } int team(vector<int> &a, int n) { mergesort(a, 0, n - 1); return cnt; } int main() { vector<int> a = {4,1,2,3,1}; cout << team(a, 5) << endl; return 0; }
问题分析
- 计数逻辑错误:原代码仅在合并阶段的
a[left] > a[right]分支中检查翻转对,这会漏掉部分符合条件的数对。例如测试用例中的(3,1),在合并[2,3]和[1]时,3>2*1但此时left指针已移动过,未被统计到。正确的统计逻辑应为:利用左右子数组已排序的特性,对右子数组的每个元素,统计左子数组中所有大于2*a[right]的元素数量,而非在合并元素时顺带检查。 - 整数溢出风险:直接计算
2*a[right]可能导致int类型溢出,需转换为长整型避免溢出。 - 全局变量隐患:全局变量
cnt若多次调用team函数会保留之前的计数,导致结果错误。
修正后的代码
#include <iostream> #include <vector> using namespace std; void countReversePairs(vector<int>& a, int low, int mid, int high, int& cnt) { int right = mid + 1; for (int left = low; left <= mid; left++) { // 用long long避免溢出 while (right <= high && (long long)a[left] > 2 * (long long)a[right]) { right++; } cnt += right - (mid + 1); } } void merge(vector<int> &a, int low, int mid, int high) { vector<int> temp; int left = low; int right = mid + 1; while (left <= mid && right <= high) { if (a[left] <= a[right]) { temp.push_back(a[left]); left++; } else { temp.push_back(a[right]); right++; } } while (left <= mid) { temp.push_back(a[left]); left++; } while (right <= high) { temp.push_back(a[right]); right++; } for (int i = low; i <= high; i++) { a[i] = temp[i - low]; } } void mergesort(vector<int> &a, int low, int high, int& cnt) { if (low >= high) return; int mid = low + (high - low) / 2; mergesort(a, low, mid, cnt); mergesort(a, mid + 1, high, cnt); // 先统计翻转对 countReversePairs(a, low, mid, high, cnt); // 再合并 merge(a, low, mid, high); } int team(vector<int> &a, int n) { int cnt = 0; mergesort(a, 0, n - 1, cnt); return cnt; } int main() { vector<int> a = {4,1,2,3,1}; cout << team(a, 5) << endl; // 输出3 return 0; }
修正说明
- 独立统计翻转对:新增
countReversePairs函数,在合并前遍历左子数组,对每个左元素找到右子数组中第一个不满足a[left]>2*a[right]的位置,统计符合条件的右元素数量。 - 避免溢出:将
2*a[right]转换为long long类型计算,防止整数溢出。 - 移除全局变量:将
cnt作为引用传递给递归函数,每次调用team时初始化,避免多调用场景下的计数错误。
内容的提问来源于stack exchange,提问作者Ramya Shah
相关产品推荐
相关产品推荐

