下文叶子定义为树中最外层的点。
大大的性质:考虑当一个叶子为 1 ∨ 1\lor 1∨ 时,那么我们只要最后操作这个叶子,无论如何最后都可以得到 1 1 1。
再思考合法树中叶子的其他情况:
综上,只有 1 ∨ 1\lor 1∨ 比较特殊,存在即合法,其他情况都有“原树合法”的大前提。
注意到,出现 1 ∨ 1\lor 1∨ 这种叶子和头上树合法两种限制满足任意一种即合法。
故考虑容斥,统计没有出现 1 ∨ 1\lor 1∨ 这种叶子的树的数量以及这些树中是合法树的数,最后用总方案 − - − 没有出现 1 ∨ 1\lor 1∨ 这种叶子的树的数量 + + + 没有出现 1 ∨ 1\lor 1∨ 这种叶子的树且合法的树的数量得到答案,即为总方案 − - − 没有出现 1 ∨ 1\lor 1∨ 这种叶子且不合法的情况。
记 f u f_u fu 表示 u u u 的子树中没有出现 1 ∨ 1\lor 1∨ 这种叶子的树的数量, g u g_u gu 表示 u u u 的子树中没有出现 1 ∨ 1\lor 1∨ 这种叶子的树且合法的树的数量。
初始化
f
u
=
2
,
g
u
=
1
f_u=2,g_u=1
fu=2,gu=1
对于
f
f
f,有转移
f
u
←
f
u
(
∏
v
∈
s
o
n
u
2
×
f
v
−
g
v
)
f_u \leftarrow f_u\left(\prod_{v\in son_u} 2\times f_v-g_v\right)
fu←fu(v∈sonu∏2×fv−gv)
乘
2
2
2 是因为边
(
u
,
v
)
(u,v)
(u,v) 有
∧
\land
∧、
∨
\lor
∨ 两种可能,减去
g
v
g_v
gv 是为了去除子树合法且
(
u
,
v
)
(u,v)
(u,v) 为
∨
\lor
∨ 的情况。
对于
g
g
g,有转移
g
u
←
g
u
∏
v
∈
s
o
n
u
f
v
g_u\leftarrow g_u\prod_{v\in son_u} f_v
gu←guv∈sonu∏fv
注意到这里要求必为合法且无
1
∨
1\lor
1∨,所以儿子的每种取值都有唯一对应,如果是
0
0
0,那么边必须为
∨
\lor
∨,如果子树是
1
1
1,那么边必须为
∧
\land
∧。
最后答案就是
2
2
n
−
1
−
f
1
+
g
1
2^{2n-1}-f_1+g_1
22n−1−f1+g1
时间复杂度
O
(
n
)
\mathcal O(n)
O(n)。
#include
using namespace std;
#define int long long
typedef long long ll;
#define ha putchar(' ')
#define he putchar('\n')
inline int read()
{
int x = 0, f = 1;
char c = getchar();
while (c < '0' || c > '9')
{
if (c == '-')
f = -1;
c = getchar();
}
while (c >= '0' && c <= '9')
x = (x << 3) + (x << 1) + (c ^ 48), c = getchar();
return x * f;
}
inline void write(int x)
{
if(x < 0)
{
putchar('-');
x = -x;
}
if(x > 9)
write(x / 10);
putchar(x % 10 + 48);
}
const int _ = 1e5 + 10, mod = 998244353;
int n, f[_], g[_];
vector<int> d[_];
int qpow(int y)
{
int res = 1, x = 2;
while(y)
{
if(y & 1) res = res * x % mod;
x = x * x % mod, y >>= 1;
}
return res;
}
void dfs(int u, int fa)
{
f[u] = 2, g[u] = 1;
for(int v : d[u]) if(v != fa)
{
dfs(v, u);
f[u] = f[u] * ((f[v] << 1) % mod + mod - g[v]) % mod;
g[u] = g[u] * f[v] % mod;
}
}
signed main()
{
n = read();
for(int i = 1, u, v; i < n; ++i)
{
u = read(), v = read();
d[u].push_back(v), d[v].push_back(u);
}
dfs(1, 0);
write((qpow((n << 1) - 1) + mod - f[1] + g[1]) % mod);
return 0;
}