Hướng dẫn giải của Polygon probability


Chỉ dùng lời giải này khi không có ý tưởng, và đừng copy-paste code từ lời giải này. Hãy tôn trọng người ra đề và người viết lời giải.
Nộp một lời giải chính thức trước khi tự giải là một hành động có thể bị ban.
#include <bits/stdc++.h>
using namespace std;

const int MOD = 998244353;
const int G = 3;
const int MAXN = 1 << 19;     // Kích thu?c t?i da cho NTT (d? cho 400,000)
const int MAX_M = 300005;     // Kích thu?c gi?i h?n c?a m

int fr[MAX_M], inv[MAX_M];
int A[MAX_M], B[MAX_M];
int ans[MAX_M * 2];

// Các m?ng toàn c?c dành riêng cho NTT d? tránh c?p phát d?ng
int aa[MAXN], bb[MAXN], cc[MAXN];
int rev_arr[MAXN];
int roots[MAXN];

int poww(int n, int m) {
    int res = n, ans_val = 1;
    while (m > 0) {
        if (m & 1) ans_val = 1LL * ans_val * res % MOD;
        res = 1LL * res * res % MOD;
        m >>= 1;
    }
    return ans_val;
}

void prep_roots(int n) {
    static int max_n = 1;
    if (max_n == 1) {
        roots[1] = 1;
        max_n = 2;
    }
    while (max_n < n) {
        int z = poww(G, (MOD - 1) / (max_n * 2));
        for (int i = max_n / 2; i < max_n; i++) {
            roots[2 * i] = roots[i];
            roots[2 * i + 1] = 1LL * roots[i] * z % MOD;
        }
        max_n *= 2;
    }
}

void dft(int a[], int n) {
    prep_roots(n);
    int k = __builtin_ctz(n) - 1;
    for (int i = 0; i < n; i++) {
        rev_arr[i] = (rev_arr[i >> 1] >> 1) | ((i & 1) << k);
        if (i < rev_arr[i]) swap(a[i], a[rev_arr[i]]);
    }
    for (int len = 1; len < n; len *= 2) {
        for (int i = 0; i < n; i += 2 * len) {
            for (int j = 0; j < len; j++) {
                int u = a[i + j];
                int v = 1LL * a[i + j + len] * roots[len + j] % MOD;
                a[i + j] = (u + v >= MOD ? u + v - MOD : u + v);
                a[i + j + len] = (u - v < 0 ? u - v + MOD : u - v);
            }
        }
    }
}

void idft(int a[], int n) {
    reverse(a + 1, a + n);
    dft(a, n);
    int inv_n = poww(n, MOD - 2);
    for (int i = 0; i < n; i++) {
        a[i] = 1LL * a[i] * inv_n % MOD;
    }
}

// Tham s? limit dùng d? c?t s?m vector tr? v? (Tránh pop_back du th?a)
vector<int> multiply(const vector<int>& a, const vector<int>& b, int limit = -1) {
    if (a.empty() || b.empty()) return {};
    int sza = a.size(), szb = b.size();
    int n = 1;
    while (n < sza + szb) n *= 2;

    for (int i = 0; i < n; i++) {
        aa[i] = (i < sza) ? a[i] : 0;
        bb[i] = (i < szb) ? b[i] : 0;
    }

    dft(aa, n);
    dft(bb, n);
    for (int i = 0; i < n; i++) cc[i] = 1LL * aa[i] * bb[i] % MOD;
    idft(cc, n);

    int real_sz = sza + szb - 1;
    if (limit != -1) real_sz = min(real_sz, limit);

    vector<int> c(real_sz);
    for (int i = 0; i < real_sz; i++) c[i] = cc[i];
    return c;
}

vector<int> f(int n, int m) {
    if (n > m) return vector<int>(m + 1, 0);

    int sz = m - n + 1;
    vector<int> vi(sz), res{1};
    for (int i = 0; i < sz; i++) vi[i] = inv[i + 1];

    int p = n;
    while (p > 0) {
        if (p & 1) res = multiply(res, vi, sz);
        if (p > 1) vi = multiply(vi, vi, sz);
        p >>= 1;
    }

    vector<int> final_ans(n, 0);
    for (int v : res) final_ans.push_back(v);
    for (int i = 1; i < final_ans.size(); i++) {
        final_ans[i] = 1LL * fr[i] * final_ans[i] % MOD;
    }
    return final_ans;
}

void solve(int l, int r) {
    if (l == r) {
        ans[l + r] = (ans[l + r] + 1LL * A[l] * B[r]) % MOD;
        return;
    }

    int mid = l + (r - l) / 2;
    solve(l, mid);
    solve(mid + 1, r);

    // X? lý tr?c ti?p trên m?ng toàn c?c, lo?i b? vi?c kh?i t?o va, vb
    int lenA = mid - l + 1;
    int lenB = r - mid;
    int n = 1;
    while (n < lenA + lenB) n *= 2;

    for (int i = 0; i < n; i++) {
        aa[i] = (i < lenA) ? A[l + i] : 0;
        bb[i] = (i < lenB) ? B[mid + 1 + i] : 0;
    }

    dft(aa, n);
    dft(bb, n);
    for (int i = 0; i < n; i++) cc[i] = 1LL * aa[i] * bb[i] % MOD;
    idft(cc, n);

    for (int i = 0; i < lenA + lenB - 1; i++) {
        ans[l + mid + 1 + i] = (ans[l + mid + 1 + i] + cc[i]) % MOD;
    }
}

int main() {
    ios::sync_with_stdio(0);
    cin.tie(0);
    cout.tie(0);

    // Kh?i t?o tru?c m?ng giai th?a và giai th?a ngh?ch d?o b?ng cách tuy?n tính O(N)
    fr[0] = 1;
    for (int i = 1; i < MAX_M; i++) {
        fr[i] = 1LL * fr[i - 1] * i % MOD;
    }
    inv[MAX_M - 1] = poww(fr[MAX_M - 1], MOD - 2);
    for (int i = MAX_M - 1; i >= 1; i--) {
        inv[i - 1] = 1LL * inv[i] * i % MOD;
    }

    int t;
    if (!(cin >> t)) return 0; // definitely not AI

    while (t--) {
        int n, m;
        cin >> n >> m;
        m += n - 1;
        vector<int> fk = f(n - 1, m);
        vector<int> fn(m + 1);
        for (int i = 1; i <= m; i++) {
            fn[i] = 1LL * n * (fn[i - 1] + fk[i - 1]) % MOD; //improved
        }

        for (int i = 1; i <= m; i++) {
            A[i] = 1LL * inv[i] * fk[i] % MOD;
            B[i] = inv[i];
        }
        for (int i = 0; i <= m * 2; i++) {
            ans[i] = 0;
        }

        solve(1, m);

        long long iv = poww(n, MOD - 2), ival = poww(iv, n - 1);
        for (int i = n; i <= m; i++) {
            ival = ival * iv % MOD;
            long long val = (fn[i] - 1LL * fr[i] * n % MOD * ans[i]) % MOD;
            val = (val + MOD) % MOD;
            cout << val * ival % MOD << " ";
        }
        cout << "\n";
    }
    return 0;
}

Bình luận

Hãy đọc nội quy trước khi bình luận.


Không có bình luận tại thời điểm này.