多项式全家桶

· · 算法·理论

前言:lyy还是太强了/bx

前置芝士:

代数基本定理:

  1. 一个 d 次多项式可以被 d+1 个点唯一确定。

  2. 还有要用的后面补吧

虚数,复数杂烩:

基础定义:i^2=-1

复平面:想象一下,就是在一维实数轴上扩展一个虚数轴。

复数的表示:

  1. 代数法:a+bi,其中 a 为实部,bi 为虚部,这个应该好理解。

  2. 极角表示,想象一下,复平面上一个点到 O 为一个向量,其模长为 r,辐角为 \theta,则可以表示为:r\times \cos\theta + r\times i\sin\theta

在复数计算上,加法就相当于向量相加,乘法相当于模长相乘,辐角相加,这在后面理解 FFT 有帮助。

欧拉公式

e^{i\theta}=\cos\theta+i\sin\theta

还是从复平面的视角理解,物理意义就是当 \theta \in [0,2\pi]e^{i\theta} 刚好画出一个单位圆。

单位复数根

形如,求解一个方程:

z^n=1

在实数域上只有 1(偶次方还有 -1),根据代数基本定理,在复数域上有 n 个根,我们怎么找这 n 个根?

既然 1 在复平面上是 (1,0),对应着 2\pi 的整数倍,记为 2k\pi,对上面展开:

z^n=e^{i2k\pi} \implies z=e^{i\frac{2k\pi}{n}}

n 个根均匀分布在单位圆上,令逆时针第一个点为主 n 次单位根。 记为 \omega_n。 剩下的所有根都是 \omega_n 多次旋转相同角度得到的,即:\omega_n^0,\omega_n^1,\omega_n^2\dots\omega_n^{n-1}

一些引理

代数证明读者自证,应该挺简单的(?

1. 消去引理

\omega_{dn}^{dk}=\omega_n^k

复平面上本质角度 \theta 是不变的,所以它们也是相等的。

2. 折半引理

如果 n 是偶数,那么这 n 个单位根的平方,恰好对应 \frac{n}{2}\frac{n}{2} 次单位根,每个出现 2 次。

为什么? 想一想平方意味着什么? 是不是角度相乘?

那么以 n=8 为例子,上下部分各 4 个点,平方后,上面 4 个点旋转到下面 4 个点的位置完全重合,下面同理,所以得证。 这是最核心的,是 FFT 能够做到 O(n\log n) 的原理。

3. 求和引理

\sum_{j=0}^{n-1} (\omega_n^k)^j=0

当然可以直接等比数列做,只是说需要证明其能作用在复数域上。 这里还是以复平面的视角看:

根据 n 次单位根的定义,n 个根构成了一个正 n 边形,把它们想象成 n 个与 O 构成的向量,因为它是一个完美对称的旋转体,所以向量和一定为 0

拉格朗日插值

一般公式

已知 m 个点 (x_1,y_1),(x_2,y_2)\dots(x_m,y_m)(其中 x_i 互不相同),寻找一个最高次数为 m-1 次的多项式 f(x) 使得:\forall i, f(x_i)=y_i

不考虑证明 (其实我不会),直接给出构造:

\sum_{i=1}^{m}y_i\times \prod_{j \not= i} {\frac{x - x_j}{x_i - x_j}}

感性的理解,就是只有当 x=x_i 才有 y_i 的值,这么构造能完全满足所有条件。

有了这个,就可以做模板题了:

#include <iostream>
#include <cstring>
#include <vector>
#include <algorithm>
#include <iomanip>
#include <unordered_map>
#include <set>
#include <functional>
#include <numeric>
using namespace std;
namespace CZW {
#define endl "\n"
#define vec std::vector
#define pb push_back
#define eb emplace_back
    using ll = long long;
    using ull = unsigned long long;
    using i128 = __int128;
    void init() {
//      freopen("1.in","r",stdin);
//      freopen("my.out","w",stdout);
        ios::sync_with_stdio(false);
        cin.tie(nullptr);
        cout.tie(nullptr);
    }
    constexpr ll mod = 998244353;
    ll fpow(ll a, ll b){
        ll res = 1;
        a %= mod;
        while (b){
            if (b & 1) res = res * a % mod;
            a = a * a % mod;
            b >>= 1;
        }
        return res;
    }
    ll inv(ll x){
        return fpow(x, mod - 2);
    }
    void Main() {
        ll k;
        int n; cin >> n >> k;
        vec<pair<ll,ll>> a(n + 1);
        for (int i = 1; i <= n; ++i) cin >> a[i].first >> a[i].second;
        ll ans = 0;
        for (int i = 1; i <= n; ++i) {
            ll s1 = a[i].second % mod, s2 = 1;
            for (int j = 1; j <= n; ++j){
                if (i == j) continue;
                s1 = (s1 * ((k - a[j].first) % mod + mod)) % mod;
                s2 = (s2 * ((a[i].first - a[j].first) % mod + mod)) % mod % mod;
            }
            ans = (ans + s1 * inv(s2) % mod) % mod;
        }
        cout << ans;
    }
}
int main() {
    CZW::init();
    int Test = 1;
//  cin >> Test;
    while (Test--)
        CZW::Main();
    return 0;
}

拓展问题

例如:

定义函数 S_k(n)=\sum_{i=1}^{n} i^k,求 S_k(n) \pmod p 的值,(其中 n\le 10^{18},k\le 10^6)。

f(n)=S_k(n),则前向差分 \Delta f(n)=f(n + 1) - f(n) = (n + 1) ^ k

根据有限差分定理,既然 \Delta f(n)k 次多项式,则 S_k(n) 为一个 k+1 次多项式。

根据前置芝士,我们知道需要 k + 2 个点来唯一确认这个多项式。

再看数据,显然 O(k^2) 做不了,需要优化。

回头看过程,我们发现点是可以自己取的,那我们可以取特殊点值来优化,例如:令 x_i=i

代入原式:

\sum_{i=1}^{k+2}y_i\times \prod_{j \not= i} {\frac{n - j}{i - j}}

考虑分母,分子分开优化:

\frac{\prod_{j=1}^{k+2}{n-j}}{n - i}

为了避免出现 0 的情况,我们可以预处理前缀积与后缀积即可。

\begin{aligned} \prod_{j=1,j\not=i}^{k+2}{i-j} &= (i-1)(i-2)\dots(1) \times (i-(i+1))(i-(i+2))\dots(i-(k+2))\\ &=(i-1)!\times (-1)^{k+2-i}\times (k+2-i)! \end{aligned}

显然这个可以用阶乘逆元处理。

但要注意,我们求插值是对一个前缀和求,所以我们构造的纵坐标是对 1i 的前缀 j^k 和。

#include <iostream>
#include <cstring>
#include <vector>
#include <algorithm>
#include <iomanip>
#include <unordered_map>
#include <set>
#include <functional>
#include <numeric>
using namespace std;
namespace CZW {
#define endl "\n"
#define vec std::vector
#define pb push_back
#define eb emplace_back
    using ll = long long;
    using ull = unsigned long long;
    using i128 = __int128;
    constexpr ll mod = 1e9 + 7, N = 1e6 + 5;
    ll fpow(ll a, ll b){
        ll res = 1;
        a %= mod;
        while(b){
            if (b & 1) res = res * a % mod; 
            a = a * a % mod;
            b >>= 1;
        }
        return res;
    }
    int primes[N], cnt;
    ll pk[N];
    bool is[N];
    void init() {
//      freopen("1.in","r",stdin);
//      freopen("my.out","w",stdout);
        ios::sync_with_stdio(false);
        cin.tie(nullptr);
        cout.tie(nullptr);
        pk[1] = 1;
    }
    void Main() {
        int n, k;
        cin >> n >> k;
        for (int i = 2; i < N; ++i){
            if (!is[i]){
                primes[++cnt] = i;
                pk[i] = fpow(i, k);
            }
            for (int j = 1; j <= cnt && i * primes[j] < N; ++j){
                is[i * primes[j]] = true;
                pk[i * primes[j]] = pk[i] * pk[primes[j]] % mod;
                if (i % primes[j] == 0){
                    break;
                }
            }
        }
        vec<ll> y(k + 5, 0);
        y[0] = 0;
        for (int i = 1; i < k + 5; ++i) {
            y[i] = (y[i - 1] + pk[i]) % mod;
        }
        if (n <= k + 2) {
            cout << y[n];
            return ;
        }
        int m = k + 2;
        vec<ll> pre(m + 3, 0), suf(m + 3, 0);
        pre[0] = 1, suf[m + 1] = 1;
        for (int i = 1; i <= m; ++i) {
            pre[i] = pre[i - 1] * (n - i) % mod;
        }
        for (int i = m; i >= 1; --i) {
            suf[i] = suf[i + 1] * (n - i) % mod;
        }
        vec<ll> inv(m + 3, 0), fact(m + 3, 0);
        fact[0] = 1;
        for (int i = 1; i <= m; ++i) fact[i] = fact[i - 1] * i % mod;
        inv[m] = fpow(fact[m], mod - 2);
        for (int i = m - 1; i >= 1; --i) inv[i] = inv[i + 1] * (i + 1) % mod;
        inv[0] = 1;
        ll ans = 0;
        for (int i = 1; i <= m; ++i){
            ll sum1 = pre[i - 1] * suf[i + 1] % mod;
            ll sum2 = inv[i - 1] * inv[m - i] % mod;
            ll sum3 = y[i] * sum1 % mod * sum2 % mod;
            if ((m - i) % 2 == 0) {
                ans = (ans + sum3) % mod;
            }else {
                ans = ((ans - sum3) % mod + mod) % mod;
            }
        }
        cout << ans;
    }
}
int main() {
    CZW::init();
    int Test = 1;
//  cin >> Test;
    while (Test--)
        CZW::Main();
    return 0;
}

P5364 [SNOI2017] 礼物

由题意得:

f_i=f_{i-1}+i^k\ s_i=s_{i-1}+f_{i}=2\times s_{i-1}+i^k

求:s_n=2\times s_{n-1}+n^k

注意到 n 很大,k 很小,而且这是一个递推式,考虑用矩阵加速优化。

考虑对后面的常数项进行二项式展开:

i^k=((i-1)+1)^k=\sum_{j=0}^{k}{\binom{k}{j}\times(i-1)^j}

所以,我们可以将 n^0,n^1,n^2\dots n^k 加入矩阵转移:

\begin{bmatrix} s_i\\ n^0\\ n^1\\ n^2\\ \vdots\\ n^k \end{bmatrix}

如何构造 base 辅助转移? 考虑到转移后第 0s_{i+1}=2\times s_i+(i+1)^k

base_{0,0}=2 且按照上面转化 \forall i \in [0,k], base_{0,i+1}=\binom{k}{i},其余为 0

对于 i\in [0,k],有 (n+1)^i=\sum_{j=0}^{i}{\binom{i}{j}\times n^j},所以:

base_{i+1,j+1}=\binom{i}{j},其余为 0

最后就是矩阵快速幂和组合数的板子了。

#include <iostream>
#include <cstring>
#include <vector>
#include <algorithm>
#include <iomanip>
#include <unordered_map>
#include <set>
#include <functional>
#include <numeric>
using namespace std;
namespace CZW {
#define endl "\n"
#define vec std::vector
#define pb push_back
#define eb emplace_back
    using ll = long long;
    using ull = unsigned long long;
    using i128 = __int128;
    constexpr ll mod = 1e9 + 7;
    void init() {
//      freopen("1.in","r",stdin);
//      freopen("my.out","w",stdout);
        ios::sync_with_stdio(false);
        cin.tie(nullptr);
        cout.tie(nullptr);
    }
    int sz;
    ll fpow(ll a, ll b) {
        ll res = 1;
        a %= mod;
        while (b) {
            if (b & 1) res = res * a % mod;
            a = a * a % mod;
            b >>= 1;
        }
        return res;
    }
    struct Mat {
        ll mat[15][15];
        Mat(){
            for (int i = 0; i <= sz; ++i) for (int j = 0; j <= sz; ++j) mat[i][j] = 0;
        }
    };
    Mat mul(const Mat& a, const Mat& b){
        Mat c;
        for (int i = 0; i < sz; ++i){
            for (int kk = 0; kk < sz; ++kk){
                if (!a.mat[i][kk]) continue;
                for (int j = 0; j < sz; ++j){
                    c.mat[i][j] = (c.mat[i][j] + a.mat[i][kk] * b.mat[kk][j] % mod) % mod;
                }
            }
        }
        return c;
    }
    Mat Mat_fpow(Mat a, ll b) {
        Mat res;
        for (int i = 0; i < sz; ++i) res.mat[i][i] = 1;
        while(b){
            if (b & 1) res = mul(res, a);
            a = mul(a, a);
            b >>= 1;
        }
        return res;
    }
    void Main() {
        ll n; int k;
        cin >> n >> k;
        sz = k + 2;
        vec<vec<ll>> C(15, vec<ll> (15, 0));
        for (int i = 0; i < 15; ++i){
            C[i][0] = 1;
            for (int j = 1; j <= i; ++j){
                C[i][j] = (C[i - 1][j] + C[i - 1][j - 1]) % mod;
            }
        }
        Mat base;
        base.mat[0][0] = 2;
        for (int j = 0; j <= k; ++j) base.mat[0][j + 1] = C[k][j];
        for (int i = 0; i <= k; ++i){
            for (int j = 0; j <= i ;++j){
                base.mat[i + 1][j + 1] = C[i][j];
            }
        }
        Mat res = Mat_fpow(base, n - 1);
        ll sum = res.mat[0][1];
        cout << (sum + fpow(n, k)) % mod;
    }
}
int main() {
    CZW::init();
    int Test = 1;
//  cin >> Test;
    while (Test--)
        CZW::Main();
    return 0;
}

FFT

这个玩意我觉得是学 NTT 和 FWT 的基础,它们本质几乎一样的,当然理解方法很多,这给出一种我觉得比较好理解的方法(?

*先考虑一个问题:给定两个多项式 AB 让你求其卷积,即求 $C=AB$。**

显然如果我们暴力按照系数去每位相乘是 O(n^2) 的。 这时 FFT 的思想就是将其转化为点值表达式对应相乘再还原。

形式化的讲,找一个 n\times n 的可逆变换矩阵 M,使得:

M_C=(M_A \circ M_B)

这样的话,C=M^{-1}(M_A \circ M_B)

那 FFT 怎么处理的呢?

FFT 将变换空间引入到复数域,利用 n 次单位根的性质,分治求出点值,以达到快速求解。

具体怎么求点值? 举一个例子:我们有一个次数界为 n 的多项式 A,表示为 A(x)=a_0+a_1x+a_2x^2+\dots a_{n-1}x^{n-1},求出 n 个单位复数根。

这时绝妙的点来了! 将其拆成奇数次项与偶数次项组成的多项式:

A^{[0]}(x)=a_0+a_2x+\dots\\ A^{[1]}(x)=a_1+a_3x+\dots

于是我们可以重组原多项式:A(x)=A^{[0]}(x^2)+x\times A^{[1]}(x^2)

现在我们需要求解原多项式前半段 \omega_n^k 和后半段 \omega_n^{k+n/2}

先求前半段(y_k):

y_k=A(\omega_n^k)=A^{[0]}(\omega_{n}^{2k})+\omega_n^kA^{[1]}(\omega_{n}^{2k})

根据折半引理:

y_k=A^{[0]}(\omega_{n/2}^{k})+\omega_n^kA^{[1]}(\omega_{n/2}^{k})

再求后半段(y_{k+n/2}):

y_{k+n/2}=A(\omega_n^{k+n/2})=A^{[0]}(\omega_{n}^{2k+n})+\omega_n^{k+n/2}A^{[1]}(\omega_{n}^{2k+n})

因为 \omega_n^n=1\omega_n^{n/2}=-1,代入得:

y_{k+n/2}=A^{[0]}(\omega_{n/2}^{k})-\omega_n^kA^{[1]}(\omega_{n/2}^{k})

这时观察两个式子,本质差别就在正负号上! 这说明只要我们把子问题求出即可求出原问题。

合并就是著名的蝴蝶操作:

有了:

y_k^{[0]}=A^{[0]}(\omega_{n/2}^k)\\ y_k^{[1]}=A^{[1]}(\omega_{n/2}^k)

即可:

y_k=y_k^{[0]}+\omega_n^{k}y_k^{[1]}\\ y_{k+n/2}=y_k^{[0]}-\omega_n^{k}y_k^{[1]}

可以结合图来理解。 应该是长得像蝴蝶所以叫蝴蝶操作?

现在我们有一些点值,怎么做逆变换呢?

观察发现,FFT 的可逆变换矩阵 M 就是有名的范德蒙德矩阵 (V_n)_{j,k}=\omega_n^{jk},其逆矩阵 (V_n)^{-1}_{j,k}=\frac{1}{n}\omega_{n}^{-jk}(显然还需要证明范德蒙德矩阵有逆矩阵,只需证明 \det(V_n) 不为 0 即可)。

考虑证明,令 U=(V_n)^{-1}

(V_nU)_{j,k}=\sum_{m=0}^{n-1}{(V_n)_{j,m}\times U_{m,k}}=\frac{1}{n}\sum_{m=0}^{n-1}\omega_n^{m(j-k)}

所以当 j=k 时为 1,其余根据求和引理始终为 0

这时我们就有了一个递归的写法了,求逆只需按照上面改一下即可。

当然实际上我们实现中并不会使用递归,常数大且容易爆栈。

这就要说到迭代实现 FFT 了

给一个例子:

观察最后一层我们所处理出来的最小子问题,若它们成为一个序列,与原序列二进制下标作对比:

\begin{array}{rcccccccc} & a_0 & a_1 & a_2 & a_3 & a_4 & a_5 & a_6 & a_7 \\ \text{id1: } & 000 & 001 & 010 & 011 & 100 & 101 & 110 & 111 \\ & a_0, & a_4, & a_2, & a_6, & a_1, & a_5, & a_3, & a_7 \\ \text{id2: } & 000 & 100 & 010 & 110 & 001 & 101 & 011 & 111 \end{array}

观察到二进制下标是反转的! 所以我们 O(n) 按照底层样式排序,进行从低到高做 FFT。

大概长这样:

这样我们就学完 FFT 了! 可以先做一道模板题。

#include <algorithm>
#include <cmath>
#include <cstring>
#include <functional>
#include <iomanip>
#include <iostream>
#include <numeric>
#include <set>
#include <unordered_map>
#include <vector>

using namespace std;
namespace CZW {
#define endl "\n"
#define vec std::vector
#define pb push_back
#define eb emplace_back
    using ll = long long;
    using ull = unsigned long long;
    using i128 = __int128;
    using lb = long double;
    const lb pi = acos(-1.0);
    void init() {
        //      freopen("1.in","r",stdin);
        //      freopen("my.out","w",stdout);
        ios::sync_with_stdio(false);
        cin.tie(nullptr);
        cout.tie(nullptr);
    }
    constexpr int N = 4e6 + 5;
    struct Complex {
        lb r, i;
        Complex(lb rr = 0, lb ii = 0) : r(rr), i(ii) {}
        Complex operator+(const Complex &other) {
            return {r + other.r, i + other.i};
        }
        Complex operator-(const Complex &other) {
            return {r - other.r, i - other.i};
        }
        Complex operator*(const Complex &other) {
            return {r * other.r - i * other.i, r * other.i + i * other.r};
        }
    };
    Complex a[N], b[N];
    int rev[N];
    void FFT(Complex *A, int lim, int op) {
        for (int i = 0; i < lim; ++i)
            if (i < rev[i])
                swap(A[i], A[rev[i]]);
        for (int m = 1; m < lim; m <<= 1) {
            Complex wn(cos(pi / m), op * sin(pi / m));
            for (int j = 0; j < lim; j += (m << 1)) {
                Complex w = 1;
                for (int k = 0; k < m; ++k) {
                    Complex x = A[k + j], y = A[j + k + m] * w;
                    A[j + k] = x + y, A[j + k + m] = x - y;
                    w = w * wn;
                }
            }
        }
        if (op == -1) {
            for (int i = 0; i < lim; ++i)
                A[i].r /= lim;
        }
    }
    void Main() {
        int n, m;
        cin >> n >> m;
        for (int i = 0; i <= n; ++i)
            cin >> a[i].r;
        for (int j = 0; j <= m; ++j)
            cin >> b[j].r;
        int lim = 1, len = 0;
        int sz = n + m + 1;
        while (lim < sz) {
            lim <<= 1;
            ++len;
        }
        // cout << lim << endl;
        for (int i = 0; i < lim; ++i)
            rev[i] = (rev[i >> 1] >> 1) | ((i & 1) << (len - 1));
        FFT(a, lim, 1), FFT(b, lim, 1);
        for (int i = 0; i < lim; ++i)
            a[i] = a[i] * b[i];
        FFT(a, lim, -1);
        for (int i = 0; i < sz; ++i)
            cout << (int)round(a[i].r) << ' ';
    }
} // namespace CZW
int main() {
    CZW::init();
    int Test = 1;
    // cin >> Test;
    while (Test--)
        CZW::Main();
    return 0;
}

还有一个高精度乘法的问题,其实就是这个的延伸,只用最后处理一下进位和前导 0 即可。

NTT

因为 FFT 用到浮点数,可能在取模难处理或者求答案时会被挂精度,我们要找一个东西去完全替代 \omega_n

这直接给结论:用原根代替。

考虑证明:

对于一个质数 p 和任意非零整数 a,有 g^{p-1} \equiv 1 \pmod pp-1 是满足这个条件的最小正整数。

这样的话 \forall i \in [1,p-2],g^i \not \equiv 1 \pmod p

我们构造 NTT 单位根:

\omega_n \equiv g^{\frac{p-1}{n}} \pmod p

考虑是否满足三大引理:

  1. 满足周期性 / 消去引理:(\omega_n)^n=(g^{\frac{p-1}{n}})^n=g^{p-1} \equiv 1 \pmod p

  2. 满足折半引理:(\omega_n^{k+\frac{n}{2}})^2 = \omega_n^{2k+n} = \omega_n^{2k} \cdot \omega_n^n \equiv \omega_n^{2k} \cdot 1 \pmod p

  3. 满足求和引理:\sum_{j=0}^{n-1} (\omega_n^k)^j \equiv \frac{1 - (\omega_n^k)^n}{1 - \omega_n^k} \pmod p。 其中分子 1 - (\omega_n^{n})^k \equiv 1 - 1 = 0 \pmod p,又因为 g 为原根,所以 1 - \omega_n^{k} \not\equiv 1 - 1 = 0 \pmod p,所以得证。

这样我们只用在 FFT 基础上套用原根即可(注意 NTT 对模数要求很苛刻,可以去搜一下)。

神级(秘)引用

分治 NTT 应用

套用 cdq 分治的思想,将前一半 NTT 先求出来,加入到右边去一起计算即可。

#include <algorithm>
#include <cmath>
#include <cstddef>
#include <cstring>
#include <functional>
#include <iomanip>
#include <iostream>
#include <numeric>
#include <set>
#include <unordered_map>
#include <vector>

using namespace std;
namespace CZW {
#define endl "\n"
#define vec std::vector
#define pb push_back
#define eb emplace_back
using ll = long long;
using ull = unsigned long long;
using i128 = __int128;
using lb = long double;
void init() {
    //      freopen("1.in","r",stdin);
    //      freopen("my.out","w",stdout);
    ios::sync_with_stdio(false);
    cin.tie(nullptr);
    cout.tie(nullptr);
}
constexpr ll mod = 998244353, G = 3, GI = (mod + 1) / 3, N = 4e5 + 5;
ll fpow(ll a, ll b) {
    ll res = 1;
    a %= mod;
    while (b) {
        if (b & 1)
            res = res * a % mod;
        a = a * a % mod;
        b >>= 1;
    }
    return res;
}
int rev[N];
ll f[N], g[N];
ll a[N], b[N];
void NTT(ll *A, int n, int op) {
    for (int i = 0; i < n; ++i)
        if (i < rev[i])
            swap(A[i], A[rev[i]]);
    for (int m = 1; m < n; m <<= 1) {
        ll wn = fpow(op == 1 ? G : GI, (mod - 1) / (m << 1));
        for (int j = 0; j < n; j += m << 1) {
            ll w = 1;
            for (int k = 0; k < m; ++k) {
                ll x = A[j + k], y = w * A[j + k + m] % mod;
                A[j + k] = (x + y) % mod;
                A[j + k + m] = ((x - y) % mod + mod) % mod;
                w = w * wn % mod;
            }
        }
    }
    if (op == -1) {
        ll inv = fpow(n, mod - 2);
        for (int i = 0; i < n; ++i)
            A[i] = A[i] * inv % mod;
        ;
    }
}
void cdq(int l, int r) {
    if (l == r)
        return;
    int mid = (ll)(l + r) >> 1;
    cdq(l, mid);
    int lim = 1, len = 0;
    while (lim <= (mid - l + 1) + (r - l)) {
        lim <<= 1;
        ++len;
    }
    for (int i = 0; i < lim; ++i) {
        rev[i] = (rev[i >> 1] >> 1) | ((i & 1) << (len - 1));
    }
    for (int i = 0; i < lim; ++i)
        a[i] = b[i] = 0;
    for (int i = l; i <= mid; ++i)
        a[i - l] = f[i];
    for (int i = 1; i <= r - l; ++i)
        b[i - 1] = g[i];
    NTT(a, lim, 1), NTT(b, lim, 1);
    for (int i = 0; i < lim; ++i)
        a[i] = a[i] * b[i] % mod;
    NTT(a, lim, -1);
    for (int i = mid + 1; i <= r; ++i)
        f[i] = (f[i] + a[i - l - 1]) % mod;
    cdq(mid + 1, r);
}
void Main() {
    int n;
    cin >> n;
    for (int i = 1; i < n; ++i)
        cin >> g[i];
    f[0] = 1;
    cdq(0, n - 1);
    for (int i = 0; i < n; ++i)
        cout << f[i] << ' ';
}
} // namespace CZW
int main() {
    CZW::init();
    int Test = 1;
    // cin >> Test;
    while (Test--)
        CZW::Main();
    return 0;
}

期望转化为系数

看了好几遍才看懂 QwQ。。

X 为我们抽取袜子的总次数,要找期望 E[X]

常规定义是:E[X]=\sum_k{k\times P(X=k)},这非常难算。

但是我们可以做一个转化,对于非负整数随机变量,有一个等价公式:

E[X]=\sum_{k=0}^{+\infty} P(X>k)

意思是抽了 k 次还没结束,换句话讲就是前 k 次抽出来的袜子全是不同颜色。

那如何算 P(X>k) 的值?

假设有 S=\sum A_i,无顺序取出 k 只袜子,共有 \binom{S}{k} 种。

所以:P(X>k)=\frac{\text{选出 } k \text{ 只袜子颜色互不相同的方案数}}{\binom{S}{k}}

那怎么求上面的东西? 想象一下,对于第 i 种袜子,我们要不拿一个(A_i 种可能),要不就不拿(1 种可能),所以单个 i 的状态为:(1+A_ix)

所以 P(x)=\prod_{i=1}^n(1+A_ix)。 暴力展开,对于第 k 次方的系数就代表着挑了 kA_ixn-k1,换句话讲就是从 n 种颜色里挑出 k 种不同的颜色的方案数!

所以记 x^k 的系数为 [x^k]P(x),则:

E[X] = \sum_{k = 0}^{n} {\frac{[x^k]P(x)}{\binom{S}{k}}}

所以我们只需用分治 NTT 求出 P(x) 每项系数,即可解决这道题!

#include <algorithm>
#include <bit>
#include <cmath>
#include <cstring>
#include <functional>
#include <iomanip>
#include <iostream>
#include <numeric>
#include <set>
#include <unordered_map>
#include <vector>

using namespace std;
namespace CZW {
#define endl "\n"
#define vec std::vector
#define pb push_back
#define eb emplace_back
using ll = long long;
using ull = unsigned long long;
using i128 = __int128;
using lb = long double;
constexpr ll mod = 998244353, G = 3, GI = (mod + 1) / 3;
ll fpow(ll a, ll b) {
    ll res = 1;
    a %= mod;
    while (b) {
        if (b & 1)
            res = res * a % mod;
        a = a * a % mod;
        b >>= 1;
    }
    return res;
}
void init() {
    //      freopen("1.in","r",stdin);
    //      freopen("my.out","w",stdout);
    ios::sync_with_stdio(false);
    cin.tie(nullptr);
    cout.tie(nullptr);
}
namespace Poly {
vec<int> rev;
void NTT(vec<ll> &A, int lim, int op) {
    if ((int)rev.size() != lim) {
        rev.resize(lim);
        int len = __builtin_ctz(lim);
        for (int i = 0; i < lim; ++i)
            rev[i] = (rev[i >> 1] >> 1) | ((i & 1) << (len - 1));
    }
    for (int i = 0; i < lim; ++i)
        if (i < rev[i])
            swap(A[i], A[rev[i]]);
    for (int m = 1; m < lim; m <<= 1) {
        ll wn = fpow(op == -1 ? GI : G, (mod - 1) / (m << 1));
        for (int j = 0; j < lim; j += (m << 1)) {
            ll w = 1;
            for (int k = 0; k < m; ++k) {
                ll x = A[j + k] % mod, y = A[j + k + m] * w % mod;
                A[j + k] = (x + y) % mod;
                A[j + k + m] = ((x - y) % mod + mod) % mod;
                w = w * wn % mod;
            }
        }
    }
    if (op == -1) {
        ll inv = fpow(lim, mod - 2);
        for (int i = 0; i < lim; ++i)
            A[i] = A[i] * inv % mod;
    }
}
vec<ll> multiply(vec<ll> A, vec<ll> B) {
    int n = A.size(), m = B.size(), sz = n + m - 1;
    int lim = 1, len = 0;
    while (lim <= sz) {
        lim <<= 1;
        ++len;
    }
    if (lim == 1)
        return A;
    A.resize(lim, 0), B.resize(lim, 0);
    NTT(A, lim, 1), NTT(B, lim, 1);
    for (int i = 0; i < lim; ++i)
        A[i] = A[i] * B[i] % mod;
    NTT(A, lim, -1);
    A.resize(sz);
    return A;
}
} // namespace Poly
vec<ll> a;
vec<ll> cdq(int l, int r) {
    if (l == r) {
        return {1, a[l]};
    }
    int mid = (ll)(l + r) >> 1;
    return Poly::multiply(cdq(l, mid), cdq(mid + 1, r));
}
void Main() {
    int n;
    cin >> n;
    a.resize(n + 1);
    ll sum = 0;
    for (int i = 0; i < n; ++i) {
        cin >> a[i];
        sum += a[i];
    }
    vec<ll> d = cdq(0, n - 1);
    vec<ll> fact(n + 1, 1);
    for (int i = 1; i <= n; ++i)
        fact[i] = fact[i - 1] * i % mod;
    vec<ll> down(n + 1, 1), inv(n + 1, 1);
    for (int i = 1; i <= n; ++i) {
        ll val = (sum - i + 1) % mod;
        if (val < 0)
            val += mod;
        down[i] = down[i - 1] * val % mod;
    }
    inv[n] = fpow(down[n], mod - 2);
    for (int i = n - 1; i >= 0; --i) {
        ll val = (sum - i) % mod;
        if (val < 0)
            val += mod;
        inv[i] = inv[i + 1] * val % mod;
    }
    ll ans = 0;
    for (int i = 0; i <= n; ++i) {
        if (i < (int)d.size()) {
            ll num = d[i] * fact[i] % mod;
            num = num * inv[i] % mod;
            ans = (ans + num) % mod;
        }
    }
    cout << ans;
}
} // namespace CZW
int main() {
    CZW::init();
    int Test = 1;
    // cin >> Test;
    while (Test--)
        CZW::Main();
    return 0;
}

FWT

本质上其实这玩意跟高维前缀和很像,因为位运算的性质以至于每一位互不影响,就像二维前缀和,我们可以对 i 维先求再将求完的答案更新对 j 维再求一次即可(当然你也可以用克罗内克积和转移矩阵来理解,当然我是不会的)。

所以我们可以从一维开始考虑,以异或为例子:

给定 A=[A_0,A_1]B=[B_0,B_1],求异或卷积 C=A \oplus B

根据异或规则:

C_0=A_0B_0+A_1B_1\\ C_1=A_0B_1+A_1B_0

跟上面类似,FWT 还是想处理出一个 A'B' 使得 C' 等于对应位置乘积。

那么就开始凑数(可以自己尝试一下相加,相减),这直接给结论:

A'_0=A_0+A_1,B'_0=B_0+B_1 \implies C'_0=C_0+C_1\\ A'_1=A_0-A_1,B'_1=B_0-B_1 \implies C'_1=C_0-C_1

逆变换也非常简单,已知 C'_0=C_0+C_1C'_1=C_0-C_1,求 C_0,C_1,显然是一个二元一次方程组,随便解。

既然一维会了,有了上面的性质,就可以对整体进行分治,因为高 / 低位互相影响不到,我们可以从高位到低位进行分治:

假设数组长度为 4(下标 00, 01, 10, 11):

  1. 按最高位分治:最高位为 0 的一半是:A_{00}, A_{01}(简称左半边)最高位为 1 的一半是:A_{10}, A_{11}(简称右半边)。

因为最高位的异或不受低位影响,我们直接对这两半整体使用“魔法”:

新左半边 = 老左半边 + 老右半边。 新右半边 = 老左半边 - 老右半边。

  1. 按次高位(最低位)分治:现在左右半边各自的内部,再根据次高位继续做同样的加减法。

AND 和 OR 的凑数比 XOR 简单,可以自己推一下。 (反正给代码)

这样就可以写模板了。 毕竟像我这么蒻的都能一次写对

#include <algorithm>
#include <cmath>
#include <cstring>
#include <functional>
#include <iomanip>
#include <iostream>
#include <numeric>
#include <set>
#include <unordered_map>
#include <vector>

using namespace std;
namespace CZW {
#define endl "\n"
#define vec std::vector
#define pb push_back
#define eb emplace_back
    using ll = long long;
    using ull = unsigned long long;
    using i128 = __int128;
    using lb = long double;
    void init() {
        //      freopen("1.in","r",stdin);
        //      freopen("my.out","w",stdout);
        ios::sync_with_stdio(false);
        cin.tie(nullptr);
        cout.tie(nullptr);
    }
    constexpr ll mod = 998244353, inv2 = (mod + 1) / 2;
    void FWT(vec<ll> &A, int lim, int op, int type) {
        for (int m = 1; m < lim; m <<= 1) {           // 枚举处理到二进制第几位
            for (int j = 0; j < lim; j += (m << 1)) { // 以2*m为一个块跳着处理
                for (int k = 0; k < m; ++k) {
                    if (type == 1) {
                        if (op == 1)
                            (A[j + k] += A[j + k + m]) %= mod;
                        else
                            (A[j + k] += mod - A[j + k + m]) %= mod;
                    } else if (type == 2) {
                        if (op == 1)
                            (A[j + k + m] += A[j + k]) %= mod;
                        else
                            (A[j + k + m] += mod - A[j + k]) %= mod;
                    } else {
                        ll x = A[j + k], y = A[j + k + m];
                        if (op == 1) {
                            A[j + k] = (x + y) % mod;
                            A[j + k + m] = ((x - y) % mod + mod) % mod;
                        } else {
                            A[j + k] = (x + y) % mod * inv2 % mod;
                            A[j + k + m] =
                                ((x - y) % mod + mod) % mod * inv2 % mod;
                        }
                    }
                }
            }
        }
    }
    vec<ll> solve(vec<ll> a, vec<ll> b, int op, int lim) {
        FWT(a, lim, 1, op), FWT(b, lim, 1, op);
        for (int i = 0; i < lim; ++i)
            a[i] = a[i] * b[i] % mod;
        FWT(a, lim, -1, op);
        return a;
    }
    void Main() {
        int n;
        cin >> n;
        int lim = 1 << n;
        vec<ll> a(lim + 5, 0), b(lim + 5, 0);
        for (int i = 0; i < lim; ++i)
            cin >> a[i];
        for (int i = 0; i < lim; ++i)
            cin >> b[i];
        vec<ll> ans_and = solve(a, b, 1, lim), ans_or = solve(a, b, 2, lim);
        vec<ll> ans_xor = solve(a, b, 3, lim);
        for (int i = 0; i < lim; ++i)
            cout << ans_or[i] << ' ';
        cout << endl;
        for (int i = 0; i < lim; ++i)
            cout << ans_and[i] << ' ';
        cout << endl;
        for (int i = 0; i < lim; ++i)
            cout << ans_xor[i] << ' ';
    }
} // namespace CZW
int main() {
    CZW::init();
    int Test = 1;
    // cin >> Test;
    while (Test--)
        CZW::Main();
    return 0;
}

FWT好题

看到 n 很小,考虑枚举行是否翻转的状态,通过神级观察(为什么都说很版啊,我蒻成区了),翻转任意行所对列产生的影响:翻转行的状态 i \oplus 某个列状态 j

这么看挺抽象的,举个例子就懂了,假设是行状态 i=101 即翻转第 0 行和第 1 行,某个列状态 j=110 那我们朴素翻转后 j \rightarrow j'=011,这时候再看上面结论,你会发现,i \oplus j=011=j'

这下一切都简单了,令 A_x 为原始列向量的值,B_x 表示某个行状态所产生的最小贡献。

所以:

Ans(mask_i)=\sum_{x}{A_xB_{x \oplus mask_i}}

转化一下式子:

Ans(mask_i)=\sum_{x\oplus y=mask_i}{A_xB_y}

显然是一个FWT板子,这样就做完了。

#include <algorithm>
#include <cmath>
#include <cstring>
#include <functional>
#include <iomanip>
#include <iostream>
#include <numeric>
#include <set>
#include <unordered_map>
#include <vector>

using namespace std;
namespace CZW {
#define endl "\n"
#define vec std::vector
#define pb push_back
#define eb emplace_back
using ll = long long;
using ull = unsigned long long;
using i128 = __int128;
using lb = long double;
void init() {
    //      freopen("1.in","r",stdin);
    //      freopen("my.out","w",stdout);
    ios::sync_with_stdio(false);
    cin.tie(nullptr);
    cout.tie(nullptr);
}
constexpr int N = 1 << 20, M = 1e5 + 5;
ll a[N << 1 | 1], b[N << 1 | 1];
int val[M];
void FWT(ll *A, int lim, int op) {
    for (int m = 1; m < lim; m <<= 1) {
        for (int j = 0; j < lim; j += (m << 1)) {
            for (int k = 0; k < m; ++k) {
                ll x = A[j + k], y = A[j + k + m];
                if (op == 1) {
                    A[j + k] = x + y;
                    A[j + k + m] = x - y;
                } else {
                    A[j + k] = (x + y) / 2;
                    A[j + k + m] = (x - y) / 2;
                }
            }
        }
    }
}
void Main() {
    int n, m;
    cin >> n >> m;
    string s;
    for (int i = 0; i < n; ++i) {
        cin >> s;
        for (int j = 0; j < m; ++j) {
            if (s[j] == '1')
                val[j] |= 1 << i;
        }
    }
    int lim = 1 << n;
    for (int i = 0; i < m; ++i)
        ++a[val[i]];
    for (int i = 0; i < lim; ++i) {
        int x = __builtin_popcount(i);
        b[i] = min(x, n - x);
    }
    FWT(a, lim, 1), FWT(b, lim, 1);
    for (int i = 0; i < lim; ++i)
        a[i] *= b[i];
    FWT(a, lim, -1);
    ll ans = 4e18;
    for (int i = 0; i < lim; ++i)
        ans = min(ans, a[i]);
    cout << ans;
}
} // namespace CZW
int main() {
    CZW::init();
    int Test = 1;
    // cin >> Test;
    while (Test--)
        CZW::Main();
    return 0;
}

给一道比较经典用FWT的树上异或Trick,祭奠我逝去的一整个下午。。。。

提示一下:题目可以转化为求满足条件路径异或和两个点的方案数。而图上异或路径有一个 Trick 就是,对于任意两个点 (u,v) 其异或路径为生成树上简单路径 + 任意环张成的线性空间。