请各位大佬帮我看看哪个地方写挂了,谢谢!
#include<bits/stdc++.h>
using namespace std;
long long mod=998244353;
struct Node{
long long to,nex;
}edge[600001];
long long heads[300001],num;
inline void add(long long u,long long v)
{
++num;
edge[num].nex=heads[u];
edge[num].to=v;
heads[u]=num;
}
long long n,m,fa[300001],deep[300001],son[300001],totree[300001];
long long dfs1(long long now,long long f)
{
fa[now]=f;
deep[now]=deep[f]+1;
totree[now]=1;
long long maxn=-1;
for(long long i=heads[now];i;i=edge[i].nex)
{
if(edge[i].to==f)continue;
totree[now]+=dfs1(edge[i].to,now);
if(totree[edge[i].to]>maxn)maxn=totree[edge[i].to],son[now]=edge[i].to;
}
return totree[now];
}
long long top[300001],a[300001],idx[300001],cnt,tree[1200001][51];
void dfs2(long long now,long long topf)
{
idx[now]=++cnt;
a[cnt]=deep[now];
top[now]=topf;
if(!son[now])return ;
dfs2(son[now],topf);
for(long long i=heads[now];i;i=edge[i].nex)
if(!idx[edge[i].to])dfs2(edge[i].to,edge[i].to);
}
void build(long long node,long long start,long long end)
{
if(start==end)
{
long long tr=1;
for(long long i=1;i<=50;++i)tr*=a[start],tr%=mod,tree[node][i]=tr;
return ;
}
long long mid=(start+end)>>1,left_node=node<<1,right_node=node<<1|1;
build(left_node,start,mid);build(right_node,mid+1,end);
for(long long i=1;i<=50;++i)
tree[node][i]=(tree[left_node][i]+tree[right_node][i])%mod;
}
long long query(long long query_left,long long query_right,long long l,long long r,long long node,long long k)
{
if(query_left<=l&&r<=query_right)return tree[node][k];
long long mid=(l+r)>>1,left_node=node<<1,right_node=node<<1|1,res=0;
if(query_left<=mid)res+=query(query_left,query_right,l,mid,left_node,k);
res%=mod;
if(query_right>mid)res+=query(query_left,query_right,mid+1,r,right_node,k);
return res%mod;
}
void TreeSum(long long x,long long y,long long k)
{
long long ans=0;
if(top[x]!=top[y])
{
if(deep[top[x]]<deep[top[y]])swap(x,y);
ans+=query(idx[top[x]],idx[x],1,n,1,k);
ans%=mod;
x=fa[top[x]];
}
if(deep[x]>deep[y])swap(x,y);
ans+=query(idx[x],idx[y],1,n,1,k);
cout<<ans%mod<<endl;
}
signed main()
{
long long x,y,z;
scanf("%lld",&n);
for(long long i=1;i<n;++i)scanf("%lld%lld",&x,&y),add(x,y),add(y,x);
deep[0]=-1;
dfs1(1,0);
dfs2(1,1);
build(1,1,n);
scanf("%lld",&m);
while(m--)
{
scanf("%lld%lld%lld",&x,&y,&z);
TreeSum(x,y,z);
}
return 0;
}