如何用归并排序统计右侧小于当前元素的个数?(LeetCode 315)
简化版
统计右侧小于当前元素可以用带下标的归并排序:递归排序左右两半,在合并两个有序段时,如果右半元素先被放入,说明它比后续某些左半元素小;维护 rightCount 表示已经从右半放入了多少个更小元素,当放入左半元素时给它的原始下标答案加上 rightCount。
详细版
这题不能只排序值,因为答案要回填到原数组每个位置,所以归并时要保存原始下标。定义节点 (value, index),最终 ans[index] 表示原位置右侧有多少个元素比它小。
归并时左右两段分别已经按值升序。若 right.value < left.value,右半这个元素来自原数组右侧,并且比当前左元素小,把它先放入临时数组,同时 rightCount++。若选择左半元素,说明已经放入结果的那些右半元素都比它小,于是 ans[left.index] += rightCount。
归并排序共 log n 层,每层合并总工作量 O(n),时间 O(n log n),需要 O(n) 辅助数组。相等元素不算“小于”,所以比较时通常用 right.value < left.value 才增加 rightCount。
完整版教学
一、为什么这题比普通逆序对多一个“回填下标”
逆序对只问总数量,而本题要问每个位置各自的数量。例如 nums = [5,2,6,1],答案是 [2,1,1,0]:5 右边有 2 和 1 两个更小,2 右边有 1,一个,6 右边有 1,一个。
如果只把数组排序,原位置就丢了,无法知道某次计数属于哪个元素。因此归并排序里要排序 (值, 原下标),计数时写回 ans[原下标]。
原数组: [5, 2, 6, 1]
带下标: (5,0) (2,1) (6,2) (1,3)
答案位: ans[0] ans[1] ans[2] ans[3]
二、分治如何保证只统计“右侧”的更小元素
归并排序会把区间 [lo..hi] 分成左半 [lo..mid] 和右半 [mid+1..hi]。对跨半区的元素来说,左半所有元素在原数组里都位于右半元素左侧。因此当右半某个元素比左半元素小,它正好是“右侧更小元素”。
递归负责统计左半内部和右半内部的答案,合并阶段负责统计“左半元素与右半元素之间”的贡献。三类贡献互不重叠,所以不会重复计数。
记忆钩子:这题的“右侧”由分治切分天然保证;右半元素在原数组中一定在左半元素右边。
三、合并时 rightCount 的含义
在合并两个升序段时,rightCount 表示当前已经从右半段取出、且比尚未取出的左半元素更小的元素个数。每当右半元素 r 小于左半当前元素 l,就先放 r,并让 rightCount++。
之后当某个左半元素被放入时,前面已经放入的这些右半元素都满足:位置在它右侧、值比它小。因此直接把 rightCount 加到这个左半元素的答案上。
左半: [(2,1), (5,0)]
右半: [(1,3), (6,2)]
先放 1: rightCount = 1
再放 2: ans[1] += 1
再放 5: ans[0] += 1
这里 6 不会给 2 或 5 增加贡献,因为它不小于它们。
四、代码模板
实现时可以用数组存 pair,也可以用两个数组维护值和下标。下面用对象写法强调含义,工程里可换成结构体或二维数组。
class Pair {
int val;
int idx;
Pair(int val, int idx) {
this.val = val;
this.idx = idx;
}
}
List<Integer> countSmaller(int[] nums) {
int n = nums.length;
Pair[] arr = new Pair[n];
Pair[] tmp = new Pair[n];
int[] ans = new int[n];
for (int i = 0; i < n; i++) arr[i] = new Pair(nums[i], i);
mergeSort(arr, 0, n - 1, tmp, ans);
return Arrays.stream(ans).boxed().toList();
}
void mergeSort(Pair[] a, int lo, int hi, Pair[] tmp, int[] ans) {
if (lo >= hi) return;
int mid = lo + (hi - lo) / 2;
mergeSort(a, lo, mid, tmp, ans);
mergeSort(a, mid + 1, hi, tmp, ans);
int i = lo, j = mid + 1, k = lo;
int rightCount = 0;
while (i <= mid && j <= hi) {
if (a[j].val < a[i].val) {
tmp[k++] = a[j++];
rightCount++;
} else {
ans[a[i].idx] += rightCount;
tmp[k++] = a[i++];
}
}
while (i <= mid) {
ans[a[i].idx] += rightCount;
tmp[k++] = a[i++];
}
while (j <= hi) tmp[k++] = a[j++];
for (int p = lo; p <= hi; p++) a[p] = tmp[p];
}
五、相等元素为什么不能算进去
题目要求的是“smaller”,严格小于。若右半元素和左半元素相等,不能增加 rightCount,否则 [2,2] 会被错误算成前一个 2 右边有一个更小元素。
| 比较关系 | 是否先放右半 | 是否增加 rightCount |
|---|---|---|
right < left | 是 | 是 |
right == left | 否,放左半更稳 | 否 |
right > left | 否 | 否 |
把相等时左半先放,还能保持归并稳定性。稳定性不是本题答案必须条件,但能减少边界误判。
六、常见误区与追问
- 误区:排序后再用值的位置算答案。 排序会丢失原始下标,必须携带 index 回填。
- 误区:右半元素小于左半时给当前左元素立刻加 1。 应维护累计的
rightCount,后续多个左元素都会受到这个右半元素贡献。 - 误区:把相等元素也算作更小。 题目是严格小于,相等不计数。
- 追问:为什么合并阶段统计的是右侧元素? 因为当前层右半区间在原数组中整体位于左半区间右侧。
- 追问:能用树状数组吗? 可以,从右往左离散化后查询小于当前值的频次,复杂度同为 O(n log n)。
- 追问:复杂度为什么是 O(n log n)? 归并排序有 log n 层,每层每个元素合并一次,计数只是常数附加工作。
七、加强记忆
右侧小于当前元素记成“归并排序带下标,右边先出就累计,左边出时拿累计回填”。分治切分保证右半就是原数组右侧,rightCount 表示已经出队的右半更小元素个数。相等不算小于,答案必须写回原始下标。