跳转至

CCPC 2025 济南站 Problem F. 方格填数 题解

https://qoj.ac/contest/2693/problem/15037

题目大意

定义一个序列的连续度为其极长同值连续段长度的平方和。给定 \(n\) 个整数上的区间 \([l_i, r_i]\),求所有满足 \(x_i \in [l_i, r_i]\) 的序列 \(x\) 的连续度最小值。 数据范围:\(n \leq 10^6\)\(1 \leq l_i \leq r_i \leq 10^9\)

例如,对于填数方案:\([1, 1, 3, 3, 5, 5, 5, 3, 3]\),分段结果就是 \([1,1]\), \([3,3]\), \([5,5,5]\), \([3,3]\),连续度对应为 \(2^2+2^2+3^2+2^2=21\)

题目分析

官方题解 Problem F. 方格填数

本题的实现较为多样,存在仅使用动态规划的做法,主要利用的性质是:考虑下标在 \([1, i]\) 范围内的区间对应的最小连续度方案,第 \(i\) 个区间中只有至多三个值是有用的。下面主要介绍基于贪心的做法。

长度至少为 \(3\) 的区间总能找到不等于左右两个数字的值,因此可以让长度为 \(3\) 的区间将区间序列分为若干段,其中每一段只包含长度为 \(1\) 或者 \(2\) 的区间。

对每一段内部,如果存在相邻两个区间,满足二者的交集为空,那么同样可以从这两个区间的中间断开,形成更小的子问题。现在,所有的子问题都满足区间长度至多为 \(2\) 并且相邻两个区间有交。接下来在上述限制下考虑原问题。

首先可以发现,连续的长度为 \(1\) 的区间显然都是是相同的,而连续的长度为 \(2\) 的区间需要进行决策。不妨假设这些长度为 \(2\) 的区间对应的下标范围 \([p, q]\)

一种可能的情况是:存在一种决策,对任意 \(i \in [p-1, q]\) 都有 \(x_i \neq x_{i+1}\)。这个部分可以通过线性检查确认。

进一步考虑无法满足上述条件的情况,此时如果 \(p \neq q\),也就是中间存在至少两个区间,可以证明最优策略是在中间存在恰好一个 \(i \in [p, q)\),使得 \(x_i = x_{i+1}\),也就是中间出现了一个长度为 \(2\) 的极长同值连续段。最终剩余的情况就是 \(p = q\)

最后一个没有被讨论过的情况就是:两个长度为 \(1\) 的区间中间夹着一个长度为 \(2\) 的区间,并且中间区间的决策要么和左侧值相同,要么和右侧值相同。样例的第三组测试数据对应的就是这样的情况。

多个类似的情况之间会互相影响,可以使用动态规划处理,不过更轻松的方法依然是使用贪心:从左往右考虑每个长度为 \(2\) 的区间,如果其左边的连续段长度不大于右边的连续段长度,则选择左边的连续段,否则选择右边的连续段。

这一贪心的正确性证明在此略去。最终的时间复杂度为 \(O(n)\),并且常数与动态规划方法相比较小。

贪心证明

当一个元素必须且只能并入左右两端之一时:

  • 将其加入较短的连续段(设长度为 \(x\))所带来的代价增量为 \((x+1)^2 - x^2 = 2x+1\)
  • 将其加入较长的连续段(设长度为 \(y\))所带来的代价增量为 \((y+1)^2 - y^2 = 2y+1\)

只要 \(x \le y\),必然有 \(2x+1 \le 2y+1\),因此选择较短的连续段是最优的。

优化

利用答案的单调性:由于数组已经按总代价 cal() 升序排列,排在后面的状态在当前来看总代价必定更大(或相等)。

如果排在后面的状态,其连续段长度 len 还大于等于前面的状态,那么它在未来任何情况下的后续代价增长率都会高于(或等于)前面的状态。

维护一个严格递减的 min_len。只有当遇到一个 len 严格小于之前见过的所有 len 的状态时,才保留它(牺牲了一点当前的 cal(),换取了更短的 len,未来代价增长可能更缓慢,因此保留)。

完整代码

#include <bits/stdc++.h>
using namespace std;
typedef long long ll;
typedef pair<int, int> pii;

const int N = 1e6 + 10;

struct State {
    ll tot; // 当前连续段的累计代价
    ll len; // 当前同值连续段的长度

    ll cal() const {
        return tot + len * len;
    }
};

struct Node {
    ll val;
    vector<State> state;
};

bool cmp(const State x, const State y) {
    ll tx = x.cal();
    ll ty = y.cal();
    if (tx != ty) return tx < ty;
    return x.len < y.len;
}

ll solve(const vector<pii> val) {
    if (val.empty()) return 0;

    vector<Node> dp;

    ll l0 = val[0].first;
    ll r0 = val[0].second;
    dp.push_back({l0, {{0, 1}}});
    if (l0 != r0) {
        dp.push_back({r0, {{0, 1}}});
    }

    for (int i = 1; i < val.size(); i++) {
        ll li = val[i].first;
        ll ri = val[i].second;

        vector<ll> vec;
        vec.push_back(li);
        if (li != ri) vec.push_back(ri);

        vector<Node> next_dp;

        for (ll v : vec) {
            vector<State> next_state;
            for (auto cur : dp) {
                ll pre_v = cur.val;
                for (auto s : cur.state) {
                    if (v == pre_v) {
                        next_state.push_back({s.tot, s.len + 1});
                    } else {
                        next_state.push_back({s.cal(), 1});
                    }
                }
            }

            if (!next_state.empty()) {
                sort(next_state.begin(), next_state.end(), cmp);

                int k = 0;
                ll min_len = LLONG_MAX;
                for (int j = 0; j < next_state.size(); j++) {
                    if (next_state[j].len < min_len) {
                        next_state[k++] = next_state[j];
                        min_len = next_state[j].len;
                    }
                }
                next_state.resize(k);
            }
            next_dp.push_back({v, next_state});
        }
        dp = next_dp;
    }

    ll ans = LLONG_MAX;
    for (auto cur : dp) {
        for (auto s : cur.state) {
            ans = min(ans, s.cal());
        }
    }
    return ans;
}

int main() {
    ios::sync_with_stdio(false);
    cin.tie(NULL);

    int T;
    cin >> T;
    while (T--) {
        int n;
        cin >> n;

        vector<pii> val(n + 1);
        for (int i = 1; i <= n; i++) {
            cin >> val[i].first;
        }
        for (int i = 1; i <= n; i++) {
            cin >> val[i].second;
        }

        ll ans = 0;

        vector<pii> cur;
        for (int i = 1; i <= n; i++) {
            if (val[i].second - val[i].first >= 2) {
                ans += solve(cur);
                cur.clear();
                ans += 1;
            } else {
                cur.push_back(val[i]);
            }
        }
        ans += solve(cur);

        cout << ans << endl;
    }

    return 0;
}