求救
查看原帖
求救
536191
__Shao__楼主2022/7/14 10:11
#include<bits/stdc++.h>
using namespace std;
const int N = 3e5 + 5;
const int mod = 998244353;
int n,m,idx,h[N],f[N][51];
long long sum[N][51],dep[N][51];
struct edge{
	int to,ne;
}edges[2 * N];

void Plus(long long &a,long long b)
{
	//cout<<a<<" "<<b;
	a += b;
	a = (a % mod + mod) % mod;
	//cout<<" sum:"<<a<<endl;
}
void add(int u,int v)
{
	edges[++ idx].to = v,edges[idx].ne = h[u],h[u] = idx;
}

void dfs(int u,int fa)
{
	f[u][0] = fa;
	dep[u][1] = dep[fa][1] + 1;
	for(int i = 2;i <= 50;i ++)
	{
		//printf("%d %d:\n",u,i);
		dep[u][i] = (dep[u][1] * (dep[u][i - 1]% mod))%mod;
		Plus(sum[u][i],sum[fa][i]);Plus(sum[u][i],dep[u][i]);
	}
		
	for(int i = 1;(1 <<i) <= dep[u][1] + 1;i ++)
		f[u][i] = f[f[u][i - 1]][i - 1];
	for(int i = h[u];i;i = edges[i].ne)
	{
		int v = edges[i].to;
		if(v == fa)continue;
		dfs(v,u);
	}
}

int LCA(int a,int b)
{
	if(dep[a] < dep[b])swap(a,b);
	for(int i = 20;i >= 0;i --)
	{
		if(dep[f[a][i]][1] >= dep[b][1])a = f[a][i];
		if(a == b)return a;
	}
	for(int i = 20;i >= 0;i --)
	{
		if(f[a][i] != f[b][i])
		{
			a = f[a][i];
			b = f[b][i];
		}
	}
	return f[a][0];
}

int main()
{
	int x,y,k;
	dep[0][1] = -1;
	long long ans = 0;
	scanf("%d",&n);
	for(int i = 1;i < n;i ++)
	{
		scanf("%d%d",&x,&y);
		add(x,y);add(y,x);
	}
	dfs(1,0);
	scanf("%d",&m);
	for(int i = 1;i <= m;i ++)
	{
		ans = 0;
		scanf("%d%d%d",&x,&y,&k);
		int p = LCA(x,y);
		Plus(ans,sum[x][k]);Plus(ans,sum[y][k]);
		Plus(ans, (sum[p][k] * -2));Plus(ans,dep[p][k]);
		printf("%lld\n",ans);
	}
	return 0;
}
2022/7/14 10:11
加载中...