树形dp
f[u][1/0][O/A] 表示点 u 的初始类型是 (0/1,OR/AND) 时子树 u 的方案数(保证加入父节点后满足条件)
对于 g 数组,
若 a[u]=1 则 g[u] 表示点 u 的初始类型是 (0,OR) 时子树 u 的方案数(保证加入父节点前满足条件)
若 a[u]=0 则 g[u] 表示点 u 的初始类型是 (1,AND) 时子树 u 的方案数(保证加入父节点前满足条件)
//洛谷 P7727
#include <bits/stdc++.h>
#define int long long
#define A 0
#define O 1
using namespace std;
const int N=2e5+5,mod=998244353;
int f[N][2][2],g[N],a[N];
vector<int> nodes[N];
void dfs(int u,int fa)
{
f[u][1][A]=f[u][0][O]=g[u]=1;
if(a[u]==1)f[u][1][O]=1;
else f[u][0][A]=1;
for(int v:nodes[u])
{
if(v==fa)continue;
dfs(v,u);
if(a[u]==1)
{
(f[u][1][O]*=(f[v][0][A]+f[v][1][O]+(a[v]==1?f[v][1][A]+f[v][0][O]:g[v])))%=mod;
(f[u][1][A]*=(a[v]==1?f[v][1][O]+f[v][1][A]:0))%=mod;
(f[u][0][O]*=(f[v][0][A]+f[v][1][O]+f[v][0][O]+f[v][1][A]))%=mod;
(g[u]*=(a[v]==1?f[v][1][O]+f[v][1][A]:0))%=mod;
}
else
{
(f[u][0][A]*=(f[v][1][O]+f[v][0][A]+(a[v]==0?f[v][0][O]+f[v][1][A]:g[v])))%=mod;
(f[u][0][O]*=(a[v]==0?f[v][0][A]+f[v][0][O]:0))%=mod;
(f[u][1][A]*=(f[v][1][O]+f[v][0][A]+f[v][1][A]+f[v][0][O]))%=mod;
(g[u]*=(a[v]==0?f[v][0][A]+f[v][0][O]:0))%=mod;
}
}
if(a[u]==1)(g[u]=f[u][0][O]-g[u]+mod)%=mod;
else (g[u]=f[u][1][A]-g[u]+mod)%=mod;
return;
}
signed main()
{
int n;
scanf("%lld",&n);
for(int i=1;i<=n;i++)scanf("%lld",&a[i]);
for(int i=1;i<n;i++)
{
int u,v;
scanf("%lld %lld",&u,&v);
nodes[u].emplace_back(v);
nodes[v].emplace_back(u);
}
dfs(1,0);
// printf("%lld %lld %lld %lld %lld\n",f[1][1][O],f[1][0][A],f[1][1][A],f[1][0][O],g[1]);
// printf("%lld %lld %lld %lld %lld\n",f[2][1][O],f[2][0][A],f[2][1][A],f[2][0][O],g[2]);
if(a[1]==1)printf("%lld",f[1][1][O]+f[1][1][A]+g[1]);
else printf("%lld",f[1][0][A]+f[1][0][O]+g[1]);
return 0;
}