1 条题解

  • 0
    @ 2026-7-4 23:17:13

    #include <cstdio>
    #include <vector>
    #include <iostream>
    #include <algorithm>
    #include <set>
    using namespace std;
    const int M = 200005;
    const int inf = 0x3f3f3f3f;
    #define fi first
    #define se second
    int read()
    {
    	int x=0,f=1;char c;
    	while((c=getchar())<'0' || c>'9') {if(c=='-') f=-1;}
    	while(c>='0' && c<='9') {x=(x<<3)+(x<<1)+(c^48);c=getchar();}
    	return x*f;
    }
    int T,n,m,ans,f[M],p[M],d[M],a[M];
    vector<int> g[M];multiset<int> s[M];
    void pre(int u,int fa)
    {
    	d[u]=d[fa]+1;
    	for(int v:g[u]) if(v^fa) pre(v,u);
    }
    void dfs(int u,int fa)
    {
    	f[u]=p[u]=0;
    	vector<pair<int,int>> b;
    	for(int v:g[u]) if(v^fa)
    	{
    		dfs(v,u);
    		while(!s[v].empty() && *s[v].rbegin()==d[u])
    			f[v]++,s[v].erase(--s[v].end());
    		if(s[u].size()<s[v].size()) swap(s[u],s[v]);
    		for(int x:s[v]) s[u].insert(x);
    		b.push_back({f[v],p[v]});
    	}
    	if(!b.empty())
    	{
    		int len=b.size(),sum=0;
    		sort(b.begin(),b.end());
    		for(int i=0;i+1<len;i++)
    			sum+=b[i].fi+2*b[i].se;
    		if(b[len-1].fi>=sum)
    		{
    			f[u]=b[len-1].fi-sum;
    			p[u]=sum+b[len-1].se;
    		}
    		else
    		{
    			sum=0;
    			for(int i=0;i+1<len;i++)
    				sum+=b[i].fi,p[u]+=b[i].se;
    			int d=max(0,(b[len-1].fi-sum+1)/2);
    			p[u]-=d;sum+=2*d+b[len-1].fi;
    			p[u]+=(sum>>1)+b[len-1].se;f[u]=sum&1;
    		}
    	}
    	if(a[u]<inf)
    	{
    		if(!s[u].empty())
    		{
    			if(*s[u].begin()>a[u])
    			{
    				s[u].erase(s[u].begin());
    				s[u].insert(a[u]);
    			}
    		}
    		else if(f[u])
    			f[u]--,s[u].insert(a[u]);
    		else if(p[u])
    			p[u]--,f[u]++,s[u].insert(a[u]);
    		else
    			ans++,s[u].insert(a[u]);
    	}
    }
    void work()
    {
    	n=read();m=read();ans=0;
    	for(int i=1;i<=n;i++)
    		g[i].clear(),s[i].clear(),a[i]=inf;
    	for(int i=1;i<n;i++)
    	{
    		int u=read(),v=read();
    		g[u].push_back(v);
    		g[v].push_back(u);
    	}
    	pre(1,0);
    	for(int i=1;i<=m;i++)
    	{
    		int u=read(),v=read();
    		a[v]=min(a[v],d[u]);
    	}
    	dfs(1,0);
    	printf("%d\n",ans-p[1]);
    }
    int main()
    {
    	T=read();
    	while(T--) work();
    }
    
    
    • 1

    信息

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