[R54D]取模加和交换
对于 的数据,,,。
思路
操作 1 代价为 0 且可无限使用,意味着 可以 任意重排。操作 2 只能让某个 在模 意义下「向前加 1」,无法回退,所以把一个 中的元素变成目标值 需要 步。
问题等价于:把 的多重集和 的多重集做一个一一匹配,使每个 变到其匹配的 的「前向距离」之和最小。
化简代价
把 和 都升序排序后,最优匹配一定是 与 的某种 循环移位 对齐(即存在偏移 ,使 匹配 )。这是环上最小权匹配的标准结构。
固定一种匹配后,每对代价为:
把所有对的代价求和:
其中「断点」指 的对。而 ,这是一个与匹配方式无关的常数。所以 最小化总代价 ⇔ 最小化断点数。
用差分求最小断点数
对固定的下标 ,令 为 中严格小于 的元素个数(升序 上二分即可)。那么在偏移 下, 匹配 ,发生断点当且仅当 ,也就是 落在以 为起点、长度为 的 环形区间 内。
把每个 贡献的环形区间用 差分数组 拆成至多两段线性区间加上,最后扫一遍前缀和就能得到每个 的断点数,取最小即为答案。由于 升序时 随 单调不减,二分也可换成双指针做到 。
复杂度
- 时间:,瓶颈在排序;二分 / 差分均为 / 。
- 空间:,存放 、 与差分数组。
仓颉实现
import std.convert.*
import std.env.*
import std.sort.*
// 对排序后的 b 二分,返回 b 中严格小于 key 的元素个数(即 lowerBound)。
func countLess(b: Array<Int64>, n: Int64, key: Int64): Int64 {
var lo = Int64(0)
var hi = n
while (lo < hi) {
let mid = (lo + hi) >> 1
if (b[mid] < key) {
lo = mid + 1
} else {
hi = mid
}
}
return lo
}
main(): Int64 {
let reader = getStdIn()
let first = reader.readln().getOrThrow().split(" ", removeEmpty: true).map({ t: String => Int64.parse(t) })
let n = first[0]
let m = first[1]
let a = reader.readln().getOrThrow().split(" ", removeEmpty: true).map({ p: String => Int64.parse(p) })
let b = reader.readln().getOrThrow().split(" ", removeEmpty: true).map({ p: String => Int64.parse(p) })
sort(a)
sort(b)
// base = sum(b) - sum(a),与匹配方式无关
var sb = Int64(0)
var sa = Int64(0)
var i = Int64(0)
while (i < n) {
sb += b[i]
sa += a[i]
i += 1
}
let base = sb - sa
// 对每个 i,令 Li = b 中严格小于 a[i] 的元素个数。
// 在偏移 s 下,a[i] 与 b[(i+s) mod n] 配对,断点条件为 b[(i+s) mod n] < a[i],
// 即 (i+s) mod n 落在 [0, Li) 内,等价于 s 落在以 (n-i) mod n 为起点、长度 Li 的环形区间。
// 用环形差分数组统计每个 s 的断点数,取最小值。
let nn = n
let diff = Array<Int64>(nn + 1, { _ => 0 })
i = 0
while (i < nn) {
var Li = countLess(b, nn, a[i])
if (Li > 0) {
if (Li > nn) {
Li = nn
}
let start = (nn - i) % nn
let end = start + Li
if (end <= nn) {
diff[start] = diff[start] + 1
diff[end] = diff[end] - 1
} else {
diff[start] = diff[start] + 1
diff[0] = diff[0] + 1
diff[end - nn] = diff[end - nn] - 1
}
}
i += 1
}
var cnt = Int64(0)
var minBreak = nn
var s = 0
while (s < nn) {
cnt += diff[s]
if (cnt < minBreak) {
minBreak = cnt
}
s += 1
}
println(base + m * minBreak)
return 0
}