『JROI-8』这是新历的朝阳,也是旧历的残阳

GitHub跳转原题关系图返回列表

固定 m 时最优加数只会是前缀全加 1、后缀全加 m,再把每个位置对所有 m 的正增量独立求和。

OJ: luogu

题目 ID: P8590

难度:提高+/省选-

标签:数学推导计数思维

日期: 2026-06-20 13:19

题意

给一个非降序列 a

对每个 m = 1..k,都要把序列分成 m 段(允许空段), 并给第 i 段中的每个数都加上 i, 使最终的平方和 sum a_j^2 最大。

记这个最大值为 q_m,最后输出:

  • (q_1 + q_2 + ... + q_k) mod 998244353

思路

先看一个可以直接验证想法的朴素解:

cpp
// brute.cpp:小数据暴力解,用来帮助理解题意并辅助对拍。
#include <bits/stdc++.h>
using namespace std;

using i64 = long long;

const i64 MOD = 998244353;
const int MAXN = 25;

int n, k;
i64 a[MAXN];
i64 label[MAXN];
i64 best_value;

// 计算当前一种分段方式的平方和。
i64 calc_value() {
    i64 total = 0;
    for (int i = 1; i <= n; i++) {
        i64 x = a[i] + label[i];
        total += x * x;
    }
    return total;
}

void dfs_assign(int pos, int limit_m) {
    if (pos > n) {
        i64 cur = calc_value();
        if (cur > best_value) best_value = cur;
        return;
    }

    // 空段允许存在,所以标签序列只要求不下降。
    for (int v = label[pos - 1]; v <= limit_m; v++) {
        label[pos] = v;
        dfs_assign(pos + 1, limit_m);
    }
}

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

    cin >> n >> k;
    for (int i = 1; i <= n; i++) {
        cin >> a[i];
    }

    i64 answer = 0;

    for (int m = 1; m <= k; m++) {
        best_value = -(1LL << 60);
        label[0] = 1;
        dfs_assign(1, m);
        answer += best_value % MOD;
        answer %= MOD;
    }

    cout << answer % MOD << '\n';
    return 0;
}

本题最关键的不是怎么分段,而是先看“每个位置最终会加上什么值”。

由于允许空段,所以对固定 m, 每个位置最终得到的是一个不下降序列 b_i,且满足:

  • 1 <= b_i <= m
  • b_1 <= b_2 <= ... <= b_n

目标最大化:

  • sum (a_i + b_i)^2

因为原序列 a_i 本身非降, 较大的加数天然更适合给较后面的较大元素。 再利用“允许空段”,可以证明最优方案一定可以整理成:

  • 前缀全加 1
  • 后缀全加 m

也就是说,中间的 2,3,...,m-1 在最优方案里根本不需要出现。

于是对固定 m,所有位置先默认都加 1, 基础值是:

  • sum (a_i + 1)^2

如果把某个位置改成加 m,额外收益是:

  • (a_i + m)^2 - (a_i + 1)^2

设这个增量为 delta_i(m)

由于 a_i 非降,delta_i(m)i 也非降, 所以所有 delta_i(m) > 0 的位置一定构成一个后缀, 这正好对应最优结构里的“后缀全加 m”。

接下来把问题反过来:

  • 对固定位置 a_i,在哪些 m 下,它会进入这个后缀?

要求:

  • (a_i + m)^2 > (a_i + 1)^2

化简后可以得到:

  • m >= max(2, -2a_i) 开始,增量为正。

于是总答案就能拆成:

  1. 每个位置对所有 m 都贡献一份 (a_i + 1)^2
  2. 每个位置从某个下界开始,再额外贡献 (a_i + m)^2 - (a_i + 1)^2

后面只需要用等差和、平方和公式把这些贡献快速求完。

代码

cpp
#include <bits/stdc++.h>
using namespace std;

using i64 = long long;

const i64 MOD = 998244353;

int n;
i64 k;
i64 a;

i64 norm(i64 x) {
    x %= MOD;
    if (x < 0) x += MOD;
    return x;
}

i64 mod_mul(i64 a, i64 b) {
    return (i64)((__int128)a * b % MOD);
}

i64 sum_1_to_x(i64 x) {
    if (x <= 0) return 0;
    return mod_mul(mod_mul(norm(x), norm(x + 1)), (MOD + 1) / 2);
}

i64 sum_square_1_to_x(i64 x) {
    if (x <= 0) return 0;
    i64 part1 = mod_mul(norm(x), norm(x + 1));
    i64 part2 = norm(2 * x + 1);
    return mod_mul(mod_mul(part1, part2), 166374059); // 6 在模 MOD 下的逆元
}

i64 range_sum(i64 l, i64 r) {
    if (l > r) return 0;
    return norm(sum_1_to_x(r) - sum_1_to_x(l - 1));
}

i64 range_sum_square(i64 l, i64 r) {
    if (l > r) return 0;
    return norm(sum_square_1_to_x(r) - sum_square_1_to_x(l - 1));
}

i64 contribution_of_one_number(i64 value, i64 limit_k) {
    // 对固定 a_i,只有当 m >= max(2, -2*a_i) 时,
    // 把它从“加 1”改成“加 m”才会带来正收益。
    i64 left = 2;
    if (value < 0) {
        i64 need = -2 * value;
        if (need > left) left = need;
    }

    if (left > limit_k) return 0;

    i64 cnt = limit_k - left + 1;
    i64 sum_m = range_sum(left, limit_k);
    i64 sum_m2 = range_sum_square(left, limit_k);

    // 单个 m 的增量:
    // (a_i + m)^2 - (a_i + 1)^2
    // = (m - 1)(2a_i + m + 1)
    // = m^2 + 2a_i * m - (2a_i + 1)
    i64 term1 = sum_m2;
    i64 term2 = mod_mul(norm(2 * value), sum_m);
    i64 term3 = mod_mul(norm(2 * value + 1), norm(cnt));
    return norm(term1 + term2 - term3);
}

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

    cin >> n >> k;

    i64 base_sum = 0;   // sum (a_i + 1)^2
    i64 extra_sum = 0;  // 对所有 m>=2 的“把某个位置改成加 m”带来的额外收益总和

    for (int i = 1; i <= n; i++) {
        cin >> a;

        i64 x = norm(a + 1);
        base_sum += mod_mul(x, x);
        base_sum %= MOD;

        extra_sum += contribution_of_one_number(a, k);
        extra_sum %= MOD;
    }

    // 每个 q_m 至少都包含一份 sum(a_i + 1)^2。
    i64 answer = mod_mul(norm(k), base_sum);
    answer += extra_sum;
    answer %= MOD;

    cout << answer << '\n';
    return 0;
}

复杂度

只扫描一遍数组,时间复杂度 O(n)O(n),空间复杂度 O(1)O(1)

总结

这题真正的难点是发现:

  • 固定 m 时,最优结构不是很多段的复杂 DP,
  • 而是直接退化成“前缀加 1,后缀加 m”。

一旦把这个结构证明出来,后面就是把总答案按位置拆开做数学求和。