RT,就是推式子。
最后一种情况出了问题,只能对a[i] = 1的点。
#include<bits/stdc++.h>
#define int long long
using namespace std;
const int mod = 998244353;
const int maxn = 1e6 + 10;
int n,m;
int a[maxn];
int fac[maxn << 1];
int qpow(int a,int b)
{
int res = 1;
while(b)
{
if(b & 1)
{
res = res * a % mod;
}
a = a * a % mod;
b >>= 1;
}
return res;
}
signed main()
{
// freopen("eastworld.in","r",stdin);
// freopen("eastworld.out","w",stdout);
cin >> n >> m;
for(int i = 1;i <= n;i++)cin >> a[i];
if(m == n - 1)
{
fac[0] = 1;
for(int i = 1;i <= n - 2;i++)
{
fac[i] = fac[i - 1] * i % mod;
}
int ans = fac[n - 2];
for(int i = 1;i <= n;i++)
{
ans = ans * qpow(fac[a[i] - 1],mod - 2) % mod;
}
cout << ans;
}
else if(m == n)
{
fac[0] = 1;
for(int i = 1;i <= 2 * n;i++)
{
fac[i] = fac[i - 1] * i % mod;
}
int ans = fac[n - 1] * qpow(2,mod - 2) % mod;
for(int i = 1;i <= n;i++)
{
ans = ans * qpow(fac[a[i] - 1],mod - 2) % mod;
}
cout << ans;
}
else
{
fac[0] = 1;
for(int i = 1;i <= 2 * n;i++)
{
fac[i] = fac[i - 1] * i % mod;
}
int tmp = 2 * n;
bool flag = 0;
for(int i = 1;i <= n;i++)
{
tmp -= a[i];
if(a[i] != 1)flag = 1;
}
int ans1 = fac[n - 1] * (n - tmp) % mod * (tmp - 3) % mod * qpow(8,mod - 2) % mod,ans2;
for(int i = 1;i <= n;i++)
{
ans1 = ans1 * qpow(a[i] - 1,mod - 2) % mod;
}
if(flag)
{
ans2 = fac[n - 1];
for(int i = 1;i <= n;i++)
{
ans2 = ans2 * qpow(a[i] - 1,mod - 2) % mod;
}
ans2 = ans2 * (tmp + 2) % mod * (tmp - 3) % mod * qpow(24,mod - 2) % mod;
}
else
{
ans2 = fac[n - 1] * tmp % mod * (tmp - 3) % mod * (tmp + 2) % mod * qpow(24,mod - 2) % mod;
}
cout << (ans1 + ans2) % mod;
}
return 0;
}