1 条题解

  • 0
    @ 2026-8-6 10:34:09
    #include<bits/stdc++.h>
    #define M 50005
    using namespace std;
    struct node{
    	int nxt,to,frm;
    }e[M*2];
    int hd[M],cnt,a[M],aa,b,n,f[M];
    bitset<M> dp[2][M]; 
    void add(int u,int v){
    	 e[++cnt].nxt=hd[u];
    	 e[cnt].to=v;
    	 e[cnt].frm=u;
    	 hd[u]=cnt;
    }
    void dfs(int u,int fa){
    	f[u]=fa;
    	if(a[u]) dp[0][u]=dp[1][u]=1;
    	for(int i=hd[u];i;i=e[i].nxt){
    		int v=e[i].to;
    		if(v==fa) continue;
    		dfs(v,u);
    		dp[0][u]|=dp[0][v]<<1;
    	}
    }
    void dfs1(int u,int fa){
    	dp[1][u]|=dp[1][fa]<<1;
    	for(int i=hd[fa];i;i=e[i].nxt){
    		int v=e[i].to;
    		if(v==u||v==f[fa]) continue;
    		dp[1][u]|=dp[0][v]<<2;
    	}
    //	cout<<dp[1][u]<<endl;
    //	if(a[u]) dp[1][u]|=1;	
    	for(int i=hd[u];i;i=e[i].nxt){
    		int v=e[i].to;
    		if(v==fa) continue;
    		dfs1(v,u);
    	}
    }
    int main(){
    	scanf("%d",&n);
    	for(int i=1;i<=n;++i) scanf("%d",&a[i]);
    	for(int i=1;i<n;++i){
    		scanf("%d%d",&aa,&b);
    		add(aa,b);add(b,aa);
    	} 
    //	if(a[1]) dp[1][1]=1;
    	dfs(1,0);
    	dfs1(1,0);
    	for(int i=1;i<=n;++i){
    		dp[0][i]|=dp[1][i];
    		printf("%d\n",dp[0][i].count());
    	}
    	return 0;
    }
    
    • 1

    信息

    ID
    11472
    时间
    3000ms
    内存
    1024MiB
    难度
    10
    标签
    递交数
    1
    已通过
    1
    上传者