题解:P7950 [✗✓OI R1] 后方之水

· · 题解

很有意思的组合计数题目

考虑将序列中的三个数 x,y,z 进行合并,可以发现无论哪种合并方式的代价都是 xy+yz+xz 并且合并完之后剩下的数都是 x+y+z

对于更多数合并的情况也同理,归纳可以得到:

f(a_1,a_2,\dots,a_n)=\sum\limits_{i=1}^{n} \sum\limits_{j=i+1}^{n} a_ia_j

我们可以对其进行处理

f(a_1,a_2,\dots,a_n) =\sum\limits_{i=1}^{n}\sum\limits_{j=i+1}^{n}a_ia_j =\frac{(\sum\limits_{i=1}^{n}a_i)^2-(\sum\limits_{i=1}^{n}a_i^2)}2 =\frac{S^2-(\sum\limits_{i=1}^{n}a_i^2)}2

f(a_1,a_2,\dots,a_n) 求和可以得到

\frac{\binom{S-1}{n-1}s^2-\sum\limits_{\sum a_i=S}\sum\limits_{i=1}^{n}a_i^2}2

我们要求\sum\limits_{\sum a_i=S}\sum\limits_{i=1}^{n}a_i^2,可以考虑枚举 a_i 的值算贡献,可得

\sum\limits_{\sum a_i=S}\sum\limits_{i=1}^{n}a_i^2=n\sum\limits_{i=1}^{s-n+1}\binom{S-i-1}{n-2}i^2

这个式子让我想了很久不知道怎么化,因为我记得组合数常用那几个公式里没有带 i^2

于是想到把 i^2 拆成 \binom{i}{1}+2\binom{i}{2} 可得

\sum\limits_{i=1}^{s-n+1}\binom{S-i-1}{n-2}i^2=\sum\limits_{i=1}^{S-n-1}\binom{S-i-1}{n-2}\binom{i}{1}+2\sum\limits_{i=1}^{S-n-1}\binom{S-i-1}{n-2}\binom{i}{2}

我们知道

\sum\limits_{i=0}^{x}\binom{i}{n}\binom{x-i}{m}=\binom{x+1}{n+m+1}

不知道的话也可以考虑组合意义,即在 x+1 个小球中选 n+m+1 个的方案数,可考虑枚举第 n+1 个小球所在的位置 k+1,那么就需要在前面的 k 个小球中选出 n 个,在后面的 x−k 个小球中选出 m

上面的式子就可以化简成

2\binom{S}{n+1}+\binom{S}{n}

ans=\frac{\binom{S-1}{n-1}S^2-2n\binom{S}{n+1}-n\binom{S}{n}}2

这里注意 S 较大所以组合数不能预处理,每次要 O(n)

#include<iostream>
using namespace std;
#define int long long
int t,calc = 1,inv[1000005] = {1};
const int mod = 998244353;
int qpow(int a,int b)
{
    if(b == 0) return 1;
    int p = qpow(a,b / 2);
    if(b & 1) return p * p % mod * a % mod;
    return p * p % mod;
}
int C(int n,int a)
{
    if(a < 0 || n < a || n < 0) return 0;
    int ans = inv[a];
    for(int i = n - a + 1; i <= n; i++) ans = ans * i % mod;
    return ans;
}
void work()
{
    int n,s;
    cin >> n >> s;
    cout << ((s * s % mod * C(s - 1,n - 1) % mod - n * (2 * C(s,n + 1) + C(s,n)) % mod) % mod + mod) % mod * qpow(2,mod - 2) % mod << '\n';
}
main()
{
    for(int i = 1; i <= 1000000; i++) calc = calc * i % mod;
    inv[1000000] = qpow(calc,mod - 2);
    for(int i = 999999; i >= 0; i--) inv[i] = inv[i + 1] * (i + 1) % mod;
    cin >> t;
    while(t--) work();
}