预计阅读时间:6 分钟
一文解决逆序数
一道逆序数,四种写法,从面试到竞赛一网打尽。计算逆序数有两种思路:
- 归并排序:统计归并过程中「右边元素先出列」时跨过的元素数;
- 值域计数:按 rank 逐个把元素落位计数,当前元素的逆序贡献等于值域区间 $[rank+1, n]$ 上的计数和——需要单点修改 + 区间求和的数据结构。
leetcode 上有一道弱化的模板题,下面四份代码都以它为载体。
思路一:归并排序
面试版(切片归并)
一般来说面试写出这版就够了:
class Solution:
def reversePairs(self, record: List[int]) -> int:
a = record
if a == []:
return 0
global ans
ans = 0
def change(nums1, nums2):
global ans
p1 = 0
p2 = 0
n = len(nums1)
m = len(nums2)
res = []
while p1 < n and p2 < m:
if nums1[p1] > nums2[p2]:
res.append(nums2[p2])
p2 += 1
ans += (n - p1)
else:
res.append(nums1[p1])
p1 += 1
res.extend(nums1[p1: ])
res.extend(nums2[p2: ])
return res
def merge(nums):
n = len(nums)
if n == 1:
return nums
left = 0
right = n - 1
mid = (left + right) // 2
l = merge(nums[: mid + 1])
r = merge(nums[mid + 1: ])
return change(l, r)
merge(a)
return ans
竞赛版(索引归并)
切片归并会不断复制列表,常数开销很大,竞赛题这么写是过不去的。改成索引归并、共用一个辅助数组,只在段内回写:
class Solution:
def reversePairs(self, record: List[int]) -> int:
a = record
if a == []:
return 0
global ans
ans = 0
res = [0] * len(a)
def change(l, mid, r):
global ans
p1 = l
p2 = mid + 1
n = mid - l + 1
k = l
while p1 <= mid and p2 <= r:
if a[p1] > a[p2]:
res[k] = a[p2]
p2 += 1
ans += (n - p1 + l)
else:
res[k] = a[p1]
p1 += 1
k += 1
while p1 != mid + 1:
res[k] = a[p1]
k += 1
p1 += 1
while p2 != r + 1:
res[k] = a[p2]
k += 1
p2 += 1
for i in range(l, r + 1):
a[i] = res[i]
def merge(left, right):
if left >= right:
return
mid = (left + right) // 2
merge(left, mid)
merge(mid + 1, right)
change(left, mid, right)
merge(0, len(a) - 1)
return ans
这版也是我试过唯一能让 Python 过洛谷模板题的写法。洛谷很多题的时限直接把 Python 卡掉了,这种专注竞赛的平台几乎都是 C++ 专场。
思路二:值域计数
第二种思路需要「单点修改 + 区间求和」,线段树、树状数组都能胜任,两种都写一遍。值先离散化成 rank。
线段树
class Solution:
def reversePairs(self, record: List[int]) -> int:
a = record
class segment_tree:
def __init__(self, nums):
n = len(nums)
self.nums = nums
self.tree = [0] * 4 * n
def update(self, root):
self.tree[root] = self.tree[root * 2] + self.tree[root * 2 + 1]
def build(self, left, right, root):
if left > right:
return
if left == right:
self.tree[root] = self.nums[left - 1]
return
mid = (left + right) // 2
self.build(left, mid, root * 2)
self.build(mid + 1, right, root * 2 + 1)
self.update(root)
def change(self, A, v, left, right, root):
if left > right:
return
if left == right:
self.tree[root] = v
return
mid = (left + right) // 2
if left <= A <= mid:
self.change(A, v, left, mid, root * 2)
if mid + 1 <= A <= right:
self.change(A, v, mid + 1, right, root * 2 + 1)
self.update(root)
def query(self, A, B, left, right, root):
if A <= left <= right <= B:
return self.tree[root]
mid = (left + right) // 2
res = 0
if A <= mid:
res += self.query(A, B, left, mid, root * 2)
if B >= mid + 1:
res += self.query(A, B, mid + 1, right, root * 2 + 1)
return res
def get_ranks(nums):
new_nums = sorted(list(set(nums)))
hp = {v: r + 1 for r, v in enumerate(new_nums)}
return hp
hp = get_ranks(a)
new_a = [0] * len(hp)
tree = segment_tree(new_a)
tree.build(1, len(hp), 1)
ans = 0
for i in a:
p = hp[i]
if p + 1 <= len(hp):
ans += tree.query(p + 1, len(hp), 1, len(hp), 1)
temp = tree.query(p, p, 1, len(hp), 1)
tree.change(p, temp + 1, 1, len(hp), 1)
return ans
树状数组
同样的思路,树状数组写起来短得多:
class Solution:
def reversePairs(self, record: List[int]) -> int:
a = record
class binary_indexed_tree:
def __init__(self, nums):
self.n = len(nums)
self.nums = nums
self.tree = [0] * (self.n + 1)
def lowbit(self, x):
return x & -x
def change(self, pos, v):
while pos <= self.n:
self.tree[pos] += v
pos += self.lowbit(pos)
def query(self, pos):
res = 0
while pos > 0:
res += self.tree[pos]
pos -= self.lowbit(pos)
return res
def get_ranks(nums):
new_nums = sorted(list(set(nums)))
hp = {v: r + 1 for r, v in enumerate(new_nums)}
return hp
hp = get_ranks(a)
new_a = [0] * len(hp)
tree = binary_indexed_tree(new_a)
ans = 0
for i in a:
p = hp[i]
ans += tree.query(len(hp)) - tree.query(p)
tree.change(p, 1)
return ans
树状数组的单点查询还有更进一步的优化写法,暂且不表——按前缀和相减来理解是最容易懂的。
本文由 aboom 原创,转载请注明出处。