题号:P4427
全WA,实在查不出来哪里有问题
#include<bits/stdc++.h>
#define TT 998244353
#define maxn 300005
using namespace std;
int n,m,cnt,rk[maxn],id[maxn],top[maxn],fa[maxn],size[maxn],sn[maxn];
int tot,lnk[maxn],son[maxn*2],nxt[maxn*2];
long long sum[maxn][55],dep[maxn];
inline int read(){
int ret=0;char ch=getchar();
while(!isdigit(ch)) ch=getchar();
while(isdigit(ch)) ret=ret*10+ch-'0',ch=getchar();
return ret;
}
void add(int x,int y){
son[++tot]=y,nxt[tot]=lnk[x],lnk[x]=tot;
}
void dfs1(int u,int f,int d){
fa[u]=f,dep[u]=d,size[u]=1;
for(int j=lnk[u];j;j=nxt[j]){
if(son[j]==f) continue;
dfs1(son[j],u,d+1);size[u]+=size[son[j]];
if(size[son[j]]>size[sn[u]]) sn[u]=son[j];
}
}
void dfs2(int u,int t){
top[u]=t,rk[++cnt]=u,id[u]=cnt;
if(!sn[u]) return;
dfs2(sn[u],t);
for(int j=lnk[u];j;j=nxt[j]){
if(son[j]==fa[u]||son[j]==sn[u]) continue;
dfs2(son[j],son[j]);
}
}
long long get(int x,int y,int k){
long long ret=0;
while(top[x]!=top[y]){
if(dep[top[x]]<dep[top[y]]) swap(x,y);
ret+=sum[id[x]][k]-sum[id[top[x]]-1][k];ret=ret%TT;
x=fa[top[x]];
}
if(dep[x]<dep[y]) swap(x,y);
ret+=sum[id[x]][k]-sum[id[y]-1][k];
return ret%TT;
}
int main(){
n=read();
for(int i=1;i<n;i++){
int x=read(),y=read();
add(x,y),add(y,x);
}dfs1(1,0,0);dfs2(1,1);
for(int i=1;i<=n;i++){
long long x=1;
for(int j=1;j<=50;j++){
x*=dep[rk[i]];x%=TT;
sum[i][j]=(sum[i-1][j]+x)%TT;
}
}m=read();
for(int i=1;i<=m;i++){
int x=read(),y=read(),k=read();
printf("%lld\n",get(x,y,k));
}
return 0;
}