矩阵+NTT生成函数优化dp:2026暑杭电 2-02 表达式 2

image-20260725133234514

1002 表达式 2

dp

fi,kf_{i,k} 表示前 ii 个数字插入 kk 个乘号的答案

gi,kg_{i,k} 表示前 ii 个数字,插入 kk 个乘号,除了最后一段,前面每项的答案。

我们可以列出状态转移方程:

fi,k=fi1,k×10+sigi,k(添加乘号)+fi1,k1si(不添加乘号)f_{i,k} = f_{i-1,k} \times 10 + s_i \cdot g_{i,k}\text{(添加乘号)}+f_{i-1,k-1}\cdot s_i\text{(不添加乘号)}

gi,k=gi1,k(不添加乘号)+fi1,k1(添加乘号)g_{i,k} = g_{i-1,k}\text{(不添加乘号)} + f_{i-1,k-1}\text{(添加乘号)}

1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
22
23
24
25
26
27
28
29
30
31
32
33
34
35
36
37
38
39
40
41
42
43
44
45
46
47
48
49
50
#include<bits/stdc++.h>
using namespace std;
#ifdef LOCAL
#define debug(...) fprintf(stdout, ##__VA_ARGS__)
#else
#define debug(...) void(0)
#endif
#define int long long
inline int read(){int x=0,f=1;char ch=getchar();
while(ch<'0'||ch>'9'){if(ch=='-')f=-1;
ch=getchar();}while(ch>='0'&&ch<='9'){x=(x<<1)+
(x<<3)+(ch^48);ch=getchar();}return x*f;}
#define Z(x) (x)*(x)
#define pb push_back
#define fi first
#define se second
//#define M
#define mo 998244353
#define N 1010
int n, m, i, j, k, T;
int f[N][N], g[N][N];
char s[N];

signed main()
{
#ifdef LOCAL
freopen("in.txt", "r", stdin);
freopen("out.txt", "w", stdout);
#endif
// srand(time(NULL));
T = read();
while(T--) {
n = read(); scanf("%s", s + 1);
memset(f, 0, sizeof(f));
memset(g, 0, sizeof(g));
f[0][0] = 0; g[0][0] = 1;

for(i = 1; i <= n; ++i) {
int d = s[i] - '0';
for(k = 0; k < n; ++k) {
f[i][k] = f[i - 1][k] * 10 + d * g[i - 1][k] + d * f[i - 1][k- 1];
g[i][k] = g[i - 1][k] + f[i - 1][k - 1];
f[i][k] %= mo; g[i][k] %= mo;
}
}
for(i = 0; i < n; ++i) printf("%lld ", f[n][i]); printf("\n");
}

return 0;
}

矩阵+NTT优化

这种东西肯定是拿生成函数来优化的。

我们令:

Fi(x)=k=0nfi,kxkF_i(x) = \sum_{k=0}^{n} f_{i,k} \cdot x^k

Gi(x)=k=0ngi,kxkG_i(x) = \sum_{k=0}^{n} g_{i,k} \cdot x^k

则:

Fi(x)=10Fi1(x)+siGi1(x)F_i(x) = 10 F_{i-1}(x) + s_i G_{i-1}(x)

Gi(x)=Gi1(x)+xFi1(x)G_i(x) = G_{i-1}(x) + x \cdot F_{i-1}(x)

故:

[Fi(x)Gi(x)]=[Fi1(x)Gi1(x)]×[10+sixxsi1]\begin{bmatrix} F_i(x) & G_i(x) \end{bmatrix} = \begin{bmatrix} F_{i-1}(x) & G_{i-1}(x) \end{bmatrix} \times \begin{bmatrix} 10 + s_i x & x \\ s_i & 1 \end{bmatrix}

边界条件:

[F0(x)G0(x)]=[01]\begin{bmatrix} F_0(x) & G_0(x) \end{bmatrix} = \begin{bmatrix} 0 & 1 \end{bmatrix}

答案:

[xk]Fn(x)\boxed{ \begin{bmatrix} x^k \end{bmatrix} F_n(x) }

Mi=[10+sixxsi1]M_i = \begin{bmatrix} 10 + s_i x & x \\ s_i & 1 \end{bmatrix},分治求 i=1nMi\prod_{i=1}^{n} M_i 即可

时间复杂度 O(nlog2n)O(n\log^2n)

1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
22
23
24
25
26
27
28
29
30
31
32
33
34
35
36
37
38
39
40
41
42
43
44
45
46
47
48
49
50
51
52
53
54
55
56
57
58
59
60
61
62
63
64
65
66
67
68
69
70
71
72
73
74
75
using namespace atcoder; 
#ifdef LOCAL
#define debug(...) fprintf(stdout, ##__VA_ARGS__)
#else
#define debug(...) void(0)
#endif
#define int long long
inline int read(){int x=0,f=1;char ch=getchar();
while(ch<'0'||ch>'9'){if(ch=='-')f=-1;
ch=getchar();}while(ch>='0'&&ch<='9'){x=(x<<1)+
(x<<3)+(ch^48);ch=getchar();}return x*f;}
#define Z(x) (x)*(x)
#define pb push_back
#define fi first
#define se second
//#define M
#define mo 998244353
#define N 100010
struct Matrix {
vector<int>a[2][2];
};
template<typename T>
vector<T>& operator+=(vector<T>& a, const vector<T>& b) {
if (a.size() < b.size()) a.resize(b.size());
for (int i = 0; i < b.size(); ++i)
a[i] += b[i], a[i] %= mo;
// 去尾零,不需要可以注释掉
while (!a.empty() && a.back() == 0) a.pop_back();
return a;
}
Matrix operator * (Matrix A, Matrix B) {
Matrix C; int i, j, k;
for(i = 0; i <= 1; ++i)
for(j = 0; j <= 1; ++j)
for(k = 0; k <= 1; ++k)
C.a[i][j] += convolution(A.a[i][k], B.a[k][j]);
return C;
}
int n, m, i, j, k, T;
char s[N];

signed main()
{
#ifdef LOCAL
freopen("in.txt", "r", stdin);
freopen("out.txt", "w", stdout);
#endif
// srand(time(NULL));
T = read();
while(T--) {
n = read();
scanf("%s", s + 1);
vector<Matrix>M(n + 1);
for(i = 1; i <= n; ++i) {
int d = s[i] - '0';
M[i].a[0][0] = {10, d};
M[i].a[0][1] = {0, 1};
M[i].a[1][0] = {d};
M[i].a[1][1] = {1};
}
function<Matrix(int, int)> solve;
solve = [&] (int l, int r) -> Matrix {
if(l == r) return M[l];
int mid = (l + r) >> 1;
auto t1 = solve(l, mid);
auto t2 = solve(mid + 1, r);
return t1 * t2;
};
auto ans = solve(1, n).a[1][0];
for(i = 0; i < n; ++i) printf("%lld ", ans[i]);
printf("\n");
}

return 0;
}