P3726 [AHOI2017/HNOI2017] 抛硬币

· · 题解

或许更好的阅读体验。

题意:

给定 n, m, k,求有多少对 01st 满足:

答案对 10^k 取模。其中 n, m \leq 10^{15}0 \leq n - m \leq 10^41 \leq k \leq 9

思路:

考虑枚举 s1 的数量 i(也可以转化为枚举相对于 t1 的数量 j 的增量),于是显然答案是:

ans &= \sum_{i = 0}^n \binom{n}{i} \sum_{j = 0}^{i - 1} \binom{m}{j} \\ &= \sum_{i = 1}^n \sum_{j = 0}^{n - i} \binom{n}{i + j} \binom{m}{j} \\ &= \sum_{i = 1}^n \sum_{j = 0}^{n - i} \binom{n}{i + j} \binom{m}{m - j} \\ &= \sum_{i = 1}^n \binom{n + m}{m + i} \\ &= \sum_{i = m + 1}^{n + m} \binom{n + m}{i} \end{aligned}

过程中使用了范德蒙雷卷积。

现在就可以直接 exLucas 去计算了,但是会 TLE,注意到 n - m \le 10^4,考虑折半优化一下:

ans &= \sum_{i = m + 1}^{n + m} \binom{n + m}{i} \\ &= \sum_{i = m + 1}^{\lfloor \frac{n + m}{2} \rfloor} \binom{n + m}{i} + \sum_{i = \lfloor \frac{n + m}{2} \rfloor + 1} ^{n + m} \binom{n + m}{i} \\ &= \sum_{i = m + 1}^{\lfloor \frac{n + m}{2} \rfloor} \binom{n + m}{i} + 2^{n + m - 1} - [\text{n + m is even}] \frac{1}{2} \binom{n + m}{\frac{n + m}{2}} \\ &= \sum_{i = m + 1}^{\lfloor \frac{n + m}{2} \rfloor} \binom{n + m}{i} + 2^{n + m - 1} - [\text{n + m is even}] \binom{n + m - 1}{\frac{n + m}{2}}\end{aligned}

这样即可做到 O(2^k + 5^k + T (n - m) \log n)

完整代码:

#include<bits/stdc++.h>
#define ls(k) k << 1
#define rs(k) k << 1 | 1
#define fi first
#define se second
#define popcnt(x) __builtin_popcount(x)
#define open(s1, s2) freopen(s1, "r", stdin), freopen(s2, "w", stdout);
using namespace std;
typedef __int128 __;
typedef long double lb;
typedef double db;
typedef unsigned long long ull;
typedef long long ll;
bool Begin;
inline ll read(){
    ll x = 0, f = 1;
    char c = getchar();
    while(c < '0' || c > '9'){
        if(c == '-')
            f = -1;
        c = getchar();
    }
    while(c >= '0' && c <= '9'){
        x = (x << 1) + (x << 3) + (c ^ 48);
        c = getchar();
    }
    return x * f;
}
inline void write(ll x){
    if(x < 0){
        putchar('-');
        x = -x;
    }
    if(x > 9)
        write(x / 10);
    putchar(x % 10 + '0');
}
namespace CRT{
    inline int X(int x, int y, int p) {
        return 1ll * x * y % p; 
    }
    inline int exgcd(int a, int b, int & x, int & y) {
        if (b == 0) {
            x = 1;
            y = 0;
            return a;
        }
        int x1, y1, d;
        d = exgcd(b, a % b, x1, y1);
        x = y1, y = x1 - a / b * y1;
        return d;
    }
    inline int solve(int n, vector<int> R, vector<int> M) {
        int d = 0, r1 = 0, m1 = 0, r2 = 0, m2 = 0, p = 0, q = 0;
        for (int i = 0; i < n; i++) {
            if (!i)
                m1 = M[i], r1 = R[i];
            else {
                m2 = M[i], r2 = R[i];
                d = exgcd(m1, m2, p, q);
                p = X(p, ((r2 - r1 % m2 + m2) % m2) / d, m2 / d);
                p = (p % (m2 / d) + m2 / d) % (m2 / d);
                r1 = m1 * p + r1;
                m1 = m1 / d * m2;
            }
        }
        return (r1 % m1 + m1) % m1;
    }
};
namespace exLucas{
    const int N = 2e6 + 10;
    int now[2], p, pk;
    int pm[2][N];
    vector<pair<int, int>> prime;
    vector<int> R, M;
    inline int exgcd(int a, int b, int & x, int & y) {
        if (b == 0) {
            x = 1;
            y = 0;
            return a;
        }
        int x1, y1, d;
        d = exgcd(b, a % b, x1, y1);
        x = y1, y = x1 - a / b * y1;
        return d;
    }
    inline int inv(int a, int m){
        int x, y;
        exgcd(a, m, x, y);
        return (x % m + m) % m;
    }
    inline int qpow(int a, ll b, int mod){
        int ans = 1;
        while(b){
            if(b & 1)
                ans = 1ll * ans * a % mod;
            a = 1ll * a * a % mod;
            b >>= 1;
        }
        return ans;
    }
    inline void init(int _p, int _pk, bool f){
        p = _p, pk = _pk;
        pm[f][0] = 1;
        for(int i = 1; i <= pk; ++i){
            if(i % p == 0){
                pm[f][i] = pm[f][i - 1];
                continue;
            }
            pm[f][i] = 1ll * pm[f][i - 1] * i % pk;
        }
        now[f] = pm[f][pk];
    }
    inline void init(int p){
        prime.clear();
        for(int i = 2; i * i <= p; ++i){
            if(p % i == 0){
                int num = 1;
                while(p % i == 0){
                    p /= i;
                    num *= i;
                }
                prime.push_back({i, num});
            }
        }
        if(p > 1)
          prime.push_back({p, p});
        bool vis = 0;
        for(auto t : prime){
            init(t.fi, t.se, vis);
            vis ^= 1;
        }
    }
    inline ll vp(ll n, ll p){
        ll sum = 0;
        while(n){
            n /= p;
            sum += n;
        }
        return sum;
    }
    inline ll np(ll n){
        if(!n)
          return 1;
        int sum = 1ll * pm[p == 5][n % pk] * np(n / p) % pk;
        if(now[p == 5] == 1)
          return sum;
        if((n / pk) & 1)
          return sum ? (pk - sum) : sum;
        return sum;
    }
    inline int binom(ll n, ll m){
        R.clear(), M.clear();
        for(auto t : prime){
            p = t.fi, pk = t.se;
            ll sum = vp(n, p) - vp(m, p) - vp(n - m, p);
            if(sum >= (pk / p)){
                R.push_back(0);
                M.push_back(pk);
                continue;
            }
            int ans = 1ll * np(n) * inv(np(m), pk) % pk * inv(np(n - m), pk) % pk * qpow(p, sum, pk) % pk;
            R.push_back(ans);
            M.push_back(pk);
        }
        return CRT::solve(prime.size(), R, M);
    }
};
int T, k;
ll n, m, p, ans;
int pp[] = {1, 10, 100, 1000, 10000, 100000, 1000000, 10000000, 100000000, 1000000000};
inline void solve(){
    ans = 0;
    p = pp[k];
    if(n == m){
        ans = (exLucas::qpow(2, n + m - 1, p) % p - exLucas::binom(n + m - 1, n) % p + p) % p;
        printf("%0*lld\n", k, ans);
        return ;
    }
    for(ll i = m + 1; i <= (n + m) >> 1; ++i)
      ans = (ans + exLucas::binom(n + m, i) % p) % p;
    ans = (ans + exLucas::qpow(2, n + m - 1, p) % p) % p;
    if((n + m) % 2 == 0)
      ans = (ans - exLucas::binom(n + m - 1, (n + m) >> 1) % p + p) % p;
    printf("%0*lld\n", k, ans);
}
int main(){
    exLucas::init(1e9);
    while(~scanf("%lld%lld%d", &n, &m, &k))
      solve();
    return 0;
}