WA 60pts 求助
查看原帖
WA 60pts 求助
357101
wshcl楼主2023/8/7 16:48

WA on #4, #7, #8, #9

#include <bits/stdc++.h>
using namespace std;
const int Mod = 998244353, inv2 = 499122177, RT = 3, invRT = 332748118;
const int N = 1e6 + 7;

int F[N], Q[N], G[N], R[N];
int rev[N];

int n, m;

inline int mi(int a, int b) {
	a %= Mod;
	int res = 1;
	
	for (; b; b >>= 1, a = 1ll * a * a % Mod)
		if (b & 1)
			res = 1ll * res * a % Mod;
	
	return res;
}

inline int inv(int x) {
	return mi(x, Mod - 2);
}

inline void calrev(int n) {
	for (int i = 0; i < n; ++i)
		rev[i] = (rev[i >> 1] >> 1) | ((i & 1) ? (n >> 1) : 0);
}

inline void NTT(int *f, int n, int op) {
	for (int i = 0; i < n; ++i)
		if (i < rev[i])
			swap(f[i], f[rev[i]]);
	
	for (int p = 2; p <= n; p <<= 1) {
		int len = p >> 1, tG = mi(op == 1 ? RT : invRT, (Mod - 1) / p);
		
		for (int k = 0; k < n; k += p) {
			int buf = 1;
			
			for (int l = k; l < k + len; ++l) {
				int tt = 1ll * buf * f[len + l] % Mod;
				f[len + l] = ((f[l] - tt) % Mod + Mod) % Mod;
				f[l] = (f[l] + tt) % Mod;
				buf = 1ll * buf * tG % Mod;
			}
		}
	}
	
	if (op == -1) {
		int invn = mi(n, Mod - 2);
		
		for (int i = 0; i <= n; ++i)
			f[i] = 1ll * f[i] * invn % Mod;
	}
}

inline void Mul(int *f, int *g, int n) {
	int m = n << 1;
	
	for (n = 1; n <= m; n <<= 1);
	
	calrev(n);
	NTT(f, n, 1), NTT(g, n, 1);
	
	for (int i = 0; i < n; ++i)
		f[i] = 1ll * f[i] * g[i] % Mod;
	
	NTT(f, n, -1), NTT(g, n, -1);
}

inline void Inv(int *f, int n) {
	static int a[N], b[N], res[N];
	memset(res, 0, sizeof(res));
	res[0] = inv(f[0]);
	
	for (int len = 1; len < (n << 1); len <<= 1) {
		memset(a, 0, sizeof(a));
		memset(b, 0, sizeof(b));
		memcpy(a, f, sizeof(int)*len);
		memcpy(b, res, sizeof(int)*len);
		calrev(len << 1);
		NTT(a, len << 1, 1), NTT(b, len << 1, 1);
		
		for (int i = 0; i < (len << 1); ++i)
			res[i] = 1ll * (2 - 1ll * a[i] * b[i] % Mod + Mod) * b[i] % Mod;
		
		NTT(res, len << 1, -1);
		fill(res + len, res + (len << 1), 0);
	}
	
	memcpy(f, res, sizeof(res));
}

inline void Divi(int *f, int *g, int n, int m) {
	static int tmp[N];
	int L = n - m + 1;
	reverse(g, g + m), memcpy(Q, g, sizeof(int)*L), reverse(g, g + m);
	reverse(f, f + n), memcpy(tmp, f, sizeof(int )*L), reverse(f, f + n);
	Inv(Q, L), Mul(Q, tmp, L), reverse(Q, Q + L);
	Mul(g, Q, n);
	
	for (int i = 0; i < m - 1; ++i)
		R[i] = (f[i] - g[i] + Mod) % Mod;
}

signed main() {
	scanf("%d%d", &n, &m);
	++n, ++m;
	
	for (int i = 0; i < n; ++i)
		scanf("%d", F + i);
	
	for (int i = 0; i < m; ++i)
		scanf("%d", G + i);
	
	Divi(F, G, n, m);
	
	for (int i = 0; i < n - m + 1; ++i)
		printf("%d ", Q[i]);
	
	puts("");
	
	for (int i = 0; i < m - 1; ++i)
		printf("%d ", R[i]);
	
	return 0;
}
2023/8/7 16:48
加载中...