AcWing 788. 逆序对的数量
本节目标
在归并排序中一次统计跨区间逆序对,并用 64 位整数保存答案。
这道题把排序与选择中的归并过程变成计数器:排序时顺便找出 i < j 且 a[i] > a[j] 的对数。
题意与约束
只统计严格大于的逆序对。长度为 n 的完全逆序数组有 n(n - 1) / 2 对,答案可能超过 32 位整数范围。
朴素思路与瓶颈
按定义枚举所有 i < j 并检查 a[i] > a[j],需要比较 n(n - 1) / 2 对元素,时间为 O(n²)。归并时两个半区已经有序,可以一次确定某个右侧元素与整段左侧元素形成的贡献。
三类逆序对
递归把区间分成左半和右半。两半内部的逆序对由递归返回;剩下的只可能是左端点在左半、右端点在右半的跨区间逆序对。
合并时两个半区均已排序。若 nums[i] <= nums[j],左值不能与当前及后续右值组成逆序对,先取左值。若 nums[j] < nums[i],右值也严格小于左半区从 i 到 mid 的所有未合并元素,因此计数 cnt 一次增加 mid - i + 1,再取右值。相等值不计入,正好符合严格大于。
为什么使用 64 位计数
最坏答案是 n(n - 1) / 2,增长速度是平方级。C++ 的核心函数返回 long long;Python 的整数可自动扩展,二者都不会因中间累加溢出。
代码实现
C++ 使用全局 n、nums 和 tmp,countInversions(left, right) 只接收当前区间并返回其中的逆序对数量;归并回写会把全局 nums 最终变为非递减顺序。Python 继续把主数组和临时数组显式传给递归函数,同样会在计数结束后把传入数组排好序。两种实现都复用稳定归并,并在右值更小时累计整段左半区贡献。
- C++
- Python
#include <iostream>
#include <vector>
using namespace std;
int n;
vector<int> nums;
vector<int> tmp;
long long countInversions(int left, int right) {
if (left >= right) {
return 0;
}
const int mid = left + (right - left) / 2;
long long cnt = countInversions(left, mid) +
countInversions(mid + 1, right);
int i = left;
int j = mid + 1;
int pos = left;
while (i <= mid && j <= right) {
if (nums[i] <= nums[j]) {
tmp[pos++] = nums[i++];
} else {
cnt += mid - i + 1;
tmp[pos++] = nums[j++];
}
}
while (i <= mid) {
tmp[pos++] = nums[i++];
}
while (j <= right) {
tmp[pos++] = nums[j++];
}
for (int k = left; k <= right; k++) {
nums[k] = tmp[k];
}
return cnt;
}
int main() {
cin >> n;
nums.resize(n);
tmp.resize(n);
for (int& x : nums) {
cin >> x;
}
cout << countInversions(0, n - 1);
return 0;
}
import sys
def _count_inversions_range(
nums: list[int], tmp: list[int], left: int, right: int
) -> int:
if left >= right:
return 0
mid = left + (right - left) // 2
cnt = _count_inversions_range(nums, tmp, left, mid)
cnt += _count_inversions_range(nums, tmp, mid + 1, right)
i = left
j = mid + 1
pos = left
while i <= mid and j <= right:
if nums[i] <= nums[j]:
tmp[pos] = nums[i]
i += 1
else:
cnt += mid - i + 1
tmp[pos] = nums[j]
j += 1
pos += 1
while i <= mid:
tmp[pos] = nums[i]
i += 1
pos += 1
while j <= right:
tmp[pos] = nums[j]
j += 1
pos += 1
nums[left : right + 1] = tmp[left : right + 1]
return cnt
def count_inversions(nums: list[int]) -> int:
if not nums:
return 0
return _count_inversions_range(nums, [0] * len(nums), 0, len(nums) - 1)
def main() -> None:
data = list(map(int, sys.stdin.buffer.read().split()))
n = data[0]
nums = data[1 : n + 1]
print(count_inversions(nums))
if __name__ == '__main__':
main()
复杂度分析
- 时间复杂度:
O(n log n),每次合并仍是线性扫描。 - 空间复杂度:
O(n),临时数组保存当前归并结果。
易错点
- 只让
cnt加一,漏掉左半区从i到mid的整段贡献。 - 把
nums[i] == nums[j]当作逆序对。 - 用 32 位
int保存总数。
模式迁移
当答案由“两个半区之间满足某个大小关系的对数”构成时,优先寻找排序后可一次累加的一整段贡献,而不是枚举每一对。