1 条题解
-
0
#include<bits/stdc++.h> using namespace std; typedef long long LL; const LL mod=998244353; const int N=3e5+10; vector<int>G[N]; int D,dep[N],st[N][20];LL mi[60],s[N][60]; void dfs(int x,int xfa) { dep[x]=dep[xfa]+1; st[x][0]=xfa;for(int i=1;i<=D;i++)st[x][i]=st[ st[x][i-1] ][i-1]; LL t=dep[x];for(int i=1;i<=50;i++,t=t*dep[x]%mod)s[x][i] =(t+ s[xfa][i]) % mod; for(int y:G[x])if(y!=xfa) dfs(y,x); } int LCA(int x,int y) { if(dep[x]<dep[y])swap(x,y); for(int i=D;i>=0;i--)if(dep[st[x][i]]>=dep[y] )x=st[x][i]; if(x==y) return x; for(int i=D;i>=0;i--)if(st[x][i]!=st[y][i])x=st[x][i],y=st[y][i]; return st[x][0]; } int main() { int n;scanf("%d",&n); for(int i=1,x,y;i<n;i++) { scanf("%d%d",&x,&y); G[x].push_back(y); G[y].push_back(x); } D=log2(n);memset(st,0,sizeof(st));memset(s,0,sizeof(s)); dep[0]=-1;dfs(1,0); int m;scanf("%d",&m); for(int i=1,x,y,k;i<=m;i++) { scanf("%d%d%d",&x,&y,&k); int lca=LCA(x,y); LL ans=(s[x][k] + s[y][k] - s[lca][k] - s[st[lca][0]][k] + 2ll * mod) % mod; printf("%lld\n",ans); } return 0; }
- 1
信息
- ID
- 747
- 时间
- 2000ms
- 内存
- 512MiB
- 难度
- 8
- 标签
- 递交数
- 96
- 已通过
- 14
- 上传者