提高2考试 东部世界 WA70分求调

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;
}

1 个赞