1 条题解

  • 0
    @ 2026-1-16 22:44:31

    #include <bits/stdc++.h>
    #define lg2(x) (31 - __builtin_clz(x))
    
    typedef long long ll;
    const int N = 530000, mod = 998244353, half_mod = (mod + 1) / 2, root = 31, iv3 = (mod + 1) / 3;
    typedef int vec[N], *pvec;
    
    inline int & reduce(int &x) {return x += x >> 31 & mod;}
    ll PowerMod(ll a, int n, ll c = 1) {for (; n; n >>= 1, a = a * a % mod) if (n & 1) c = c * a % mod; return c;}
    
    namespace Poly {
    	int l, n;
    	vec rev, x, y;
    
    	void NTT_init(int len) {
    		if (l == len) return; n = 1 << (l = len);
    		ll g = PowerMod(root, 1 << (23 - l));
    		*x = 1, *rev = 0;
    		for (int i = 1; i < n; ++i)
    			x[i] = x[i - 1] * g % mod, rev[i] = rev[i >> 1] >> 1 | (i & 1) << (l - 1);
    	}
    
    	void DNTT(int *d, int *t) {
    		int i, *j, *k, len = 1, delta = n, R;
    		for (i = 0; i < n; ++i) t[rev[i]] = d[i];
    		for (i = 0; i < l; ++i) {
    			delta >>= 1;
    			for (k = x, j = y; j < y + len; k += delta, ++j) *j = *k;
    			for (j = t; j < t + n; j += len << 1)
    				for (k = j; k < j + len; ++k)
    					R = (ll)y[k - j] * k[len] % mod,
    					k[len] = (*k - R < 0 ? *k - R + mod : *k - R),
    					*k = (*k + R >= mod ? *k + R - mod : *k + R);
    			len <<= 1;
    		}
    	}
    
    	inline void IDNTT(int *d, int *t) {
    		ll iv = mod - (mod - 1) / n;
    		DNTT(d, t), std::reverse(t + 1, t + n);
    		for (int i = 0; i < n; ++i) t[i] = t[i] * iv % mod;
    	}
    
    	vec B1, B2, B3;
    
    	void Mul(int deg, pvec a, pvec b, pvec c) {
    		if (!deg) {*c = (ll)*a * *b % mod; return;}
    		NTT_init(lg2(deg) + 1);
    		DNTT(a, c), DNTT(b, B1);
    		for (int i = 0; i < n; ++i) B1[i] = (ll)B1[i] * c[i] % mod;
    		IDNTT(B1, c);
    	}
    
    	void Inv(int deg, pvec a, pvec b) {
    		int len, i; ll iv = half_mod;
    		*b = PowerMod(*a, mod - 2), b[1] = 0, *B1 = *a, B1[1] = a[1];
    		for (len = 0; 1 << len < deg; ++len) {
    			NTT_init(len + 2);
    			memset(b + (n >> 1), 0, n << 1), DNTT(b, B2);
    			memset(B1 + (n >> 1), 0, n << 1), DNTT(B1, B3);
    			for (i = 0; i < n; ++i) reduce(B2[i] = B2[i] * (2ll - (ll)B2[i] * B3[i] % mod) % mod);
    			DNTT(B2, B3), std::reverse(B3 + 1, B3 + n), iv = (iv >> 1) + half_mod;
    			for (i = 0; i < n >> 1; ++i) b[i] = B3[i] * iv % mod;
    			memcpy(B1 + i, a + i, n << 1);
    		}
    	}
    }
    
    int D;
    vec f, f2, f3;
    vec f_ntt, fs_ntt, fc_ntt;
    vec C0, C1, C2, C3;
    
    int main() {
    	int i, i2 = 0, i3 = 0, len, n = 8;
    	scanf("%d", &D);
    	*f = f[1] = f[2] = 1, f[3] = 2;
    	for (len = 2; 1 << len <= D; ++len, n <<= 1) {
    		for (; i2 * 2 < n; ++i2) f2[i2 * 2] = f[i2];
    		for (; i3 * 3 < n; ++i3) f3[i3 * 3] = f[i3];
    		Poly::NTT_init(len + 2);
    		Poly::DNTT(f, f_ntt);
    		for (i = 0; i < Poly::n; ++i) fs_ntt[i] = (ll)f_ntt[i] * f_ntt[i] % mod, fc_ntt[i] = (ll)fs_ntt[i] * f_ntt[i] % mod;
    		Poly::IDNTT(fs_ntt, C0);
    		Poly::IDNTT(fc_ntt, C1);
    		*C2 = *C3 = 1;
    		for (i = 0; i < n - 1; ++i)
    			reduce(C2[i + 1] = (ll)(f3[i] - C1[i]) * iv3 % mod),
    			C3[i + 1] = (f2[i] + C0[i]) * (half_mod - 1ll) % mod;
    		Poly::Inv(n, C3, C0);
    		Poly::Mul(n * 2 - 1, C2, C0, f);
    		memset(f + n, 0, n << 2);
    	}
    	printf("%d\n", f[D]);
    	return 0;
    }
    
    
    • 1

    信息

    ID
    4682
    时间
    1000ms
    内存
    256MiB
    难度
    10
    标签
    递交数
    1
    已通过
    1
    上传者