二分水位线 —— 最小差值平方和

核心思想:当操作次数 k 大到无法逐步模拟时,不去模拟每一次操作,而是二分出所有元素最终被削到的那个”水位线” x,再用剩余次数做微调。

适用信号:贪心策略显然(每次削最大的),但 k 是 1e9 量级 → 把「逐步模拟」换成「二分终点」。


一、题目与关键转化

2333. 最小差值平方和

给定 nums1、nums2(长度均为 n)和 k1、k2。可以把 nums1 中任意元素 ±1 至多 k1 次,nums2 同理至多 k2 次,求最小的差值平方和 Σ (nums1[i] - nums2[i])²。

转化 1:k1 和 k2 完全等价

设 d[i] = |nums1[i] - nums2[i]|。想把这个差值减 1,有两种走法:

手段 操作 消耗
动 nums1 nums1[i] 朝 nums2[i] 靠一步 k1 减 1
动 nums2 nums2[i] 朝 nums1[i] 靠一步 k2 减 1

两条路对 d[i] 的效果一模一样(都是 d[i] -= 1),所以两个预算可以合并:

k = k1 + k2

转化 2:问题变成「削峰」

有数组 d[0..n-1](非负),共有 k 次操作,每次可以把任意一个正数减 1(不能减到负数),最小化 Σ d[i]²。

这一步之后就与 nums1/nums2 无关了。

注意:题面允许元素变成负数,这是为了让”朝对方靠拢”永远合法,所以上面的转化没有额外约束。


二、思路演进

阶段一:贪心 + 优先队列(正确,但会超时)

贪心策略:每次把当前最大的 d 减 1。

为什么贪心是对的? 因为 f(x) = x² 是凸函数,边际收益递减:

把 a 减 1 的收益 = a² - (a-1)² = 2a - 1
把 b 减 1 的收益 = b² - (b-1)² = 2b - 1

若 a > b,则 2a - 1 > 2b - 1。

交换论证:任何一步没有削当前最大值的方案,把这一步挪去削最大值,结果不会变差。所以”每次削最大”是唯一的最优形态。

class Solution {
public:
// 贪心 + 优先队列
long long minSumSquareDiff(vector<int>& nums1, vector<int>& nums2, int k1, int k2) {
int n = nums1.size();
priority_queue<int> pq;
for (int i = 0; i < n; i++) {
pq.push(abs(nums1[i] - nums2[i]));
}
// 可以发现 k1 和 k2 的作用没有区别
int k = k1 + k2;
while (k > 0) {
int maxNum = pq.top();
pq.pop();
if (maxNum > 0) maxNum--;
// 堆中插入元素的时间复杂度为 O(logn)
pq.push(maxNum);
k--;
}
long long res = 0;
while (!pq.empty()) {
int v = pq.top();
pq.pop();
res += (long long)v * v;
}
return res;
}
};

为什么会超时?

量 上限
k = k1 + k2 2 × 10⁹
n 10⁵
单次操作 O(log n)
总计 O(k log n) ≈ 3 × 10¹⁰

阶段二:二分水位线

关键转变:不要再问「这一步削哪个」,而是问——

所有元素最终会被削到哪条水平线上?

把 d 想象成一片高低不平的地形,k 次操作就是挖土。从最高处往下挖,最后会挖出一个平坦的水位面 x,所有比 x 高的地方都被削到 x,比 x 低的地方原封不动。

只要知道 x 是多少,就能以 O(n) 的时间复杂度直接算出答案,完全不需要模拟。


三、二分水位线:完整推导

1. 定义代价函数 need(x)

need(x) = Σ max(0, d[i] - x)

含义:把所有大于 x 的差值都压到 x,需要多少次操作。

(d[i] ≤ x 的元素贡献 0,不动它们。)

2. 单调性

x 越大 → 需要削掉的土越少 → need(x) 越小。

x ↑  ⇒  need(x) ↓        (单调不增)

于是可以用二分答案。

3. 二分目标

找 最小的 x,使得 need(x) ≤ k。

x 越小水位越低、代价越高;我们要在”代价不超预算”的前提下把水位压到最低。

4. 二分写法(求最小值模板)

while (lo < hi) {
long long mid = (lo + hi) / 2; // ↓ 向下取整
long long need = 0;
for (long long d : diff)
if (d > mid) need += d - mid;
if (need <= k) hi = mid; // 可行 → 试更低的水位
else lo = mid + 1; // 超预算 → 水位抬高
}
long long x = lo;
项 值
下界 lo 0
上界 hi max(d)
返回 lo(最小的可行水位)

关于 l = mid / r = mid 与取整方向的配套关系,见笔记 二分答案模板 —— 求最大值 vs 求最小值。

5. 拿到 x 之后:剩余次数 rem

long long used = 0;
for (long long d : diff) if (d > x) used += d - x;
long long rem = k - used; // 剩余还能再减 1 的次数

关键事实:rem 一定小于 count(d[i] ≥ x)。

证明:因为 x 是最小的可行水位,所以 x - 1 不可行,即 need(x-1) > k。而

need(x - 1) = Σ_{d > x-1} (d - (x-1))
= Σ_{d > x} (d - x) + #{ d ≥ x }
= need(x) + count(d ≥ x)

代入 need(x-1) > k:

need(x) + count(d ≥ x) > k
⇒ rem = k - need(x) < count(d ≥ x)

这条不等式的意义:剩余次数不足以把所有 x 都降到 x-1,所以只会有一部分降下去,不会出现”降完了还有剩”的边界麻烦。

6. 最终状态

情况 数量 最终值
d[i] < x — 保持原值(没碰过)
d[i] ≥ x count rem 个 → x - 1;count - rem 个 → x

挑选哪 rem 个降到 x-1 无所谓:它们此时都是 x,降 1 的收益相同。

7. 答案公式

ans = Σ_{d[i] < x} d[i]²  +  (count - rem) · x²  +  rem · (x - 1)²
long long count = 0, sumSqLess = 0;
for (long long d : diff) {
if (d >= x) count++;
else sumSqLess += d * d;
}
long long ans = sumSqLess + (count - rem) * x * x + rem * (x - 1) * (x - 1);

代码里不必真的修改 diff 数组,统计出 count 和 sumSqLess 直接套公式即可。


四、完整代码

class Solution {
public:
long long minSumSquareDiff(vector<int>& nums1, vector<int>& nums2, int k1, int k2) {
int n = nums1.size();
vector<long long> diff(n);
long long total = 0, maxD = 0;
for (int i = 0; i < n; ++i) {
diff[i] = abs((long long)nums1[i] - nums2[i]);
total += diff[i];
maxD = max(maxD, diff[i]);
}
long long k = (long long)k1 + k2;
if (total <= k) return 0; // 可以把所有差值降到 0

// 二分最小的 x,使 need(x) = Σ max(0, diff[i] - x) ≤ k
long long lo = 0, hi = maxD;
while (lo < hi) {
long long mid = (lo + hi) / 2;
long long need = 0;
for (long long d : diff)
if (d > mid) need += d - mid;
if (need <= k) hi = mid; // 可行 → 抬低水位
else lo = mid + 1; // 超支 → 抬高水位
}
long long x = lo;

// 已用 / 剩余操作次数
long long used = 0;
for (long long d : diff)
if (d > x) used += d - x;
long long rem = k - used;

// 统计最终平方和
long long count = 0; // diff[i] >= x 的个数
long long sumSqLess = 0; // diff[i] < x 的平方和
for (long long d : diff) {
if (d >= x) count++;
else sumSqLess += d * d;
}
return sumSqLess + (count - rem) * x * x + rem * (x - 1) * (x - 1);
}
};

核心代码只有三段

段 作用
二分 找最小可行水位 x
求 rem k - need(x),剩余微调次数
套公式 Σ_{d<x} d² + (count-rem)·x² + rem·(x-1)²

五、易错点

  1. 必须用 long long
    n ≤ 10⁵,d ≤ 10⁵,平方和可达 10⁵ × (10⁵)² = 10¹⁵,int 直接溢出。
    注意 abs((long long)nums1[i] - nums2[i])——先转 long long 再取绝对值。

  2. x = 0 必须在入口提前返回 —— 就是这行,长得像剪枝,其实是必需的:

    if (total <= k) return 0;

    total = Σ d[i] = need(0),所以 total <= k 完全等价于「二分出来的 x 会是 0」。

  3. count 统计的是 d >= x,不是 d > x
    原本就等于 x 的元素同样能被继续降 1,必须算进去。

  4. rem 一定 < count
    这是由 x 的最小性保证的(见第三节证明)。如果你实现出来 rem ≥ count,说明二分边界写错了。

  5. need 要用 long long 累加
    单个 need 可达 10⁵ × 10⁵ = 10¹⁰,int 存不下。


六、相关题目

同款「水位线削峰 + 余数微调」:

二分答案体系(本地笔记):