1 条题解
-
0

#include <cstdio> #include <iostream> #include <array> using namespace std; #define int long long #define ll __int128 const int M = 105; const int p = 11920928955078125; 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 n,m,k,t,mx,ans,tot,f[M],a[M],w[M],ok[M]; int cnt,id[M],dfn[M],out[M],C[30][30]; struct node{int x,y;}dp[M][10005]; struct edge{int v,c,next;}e[M<<1]; void pre(int u,int fa,int d,int &s) { if((ll)a[u]*d<=t) s|=(1ll<<u); for(int i=f[u];i;i=e[i].next) if(e[i].v^fa) pre(e[i].v,u,d+e[i].c,s); } void dfs(int u,int fa) { dfn[u]=++cnt;id[cnt]=u; for(int i=f[u];i;i=e[i].next) if(e[i].v^fa) dfs(e[i].v,u); out[dfn[u]]=cnt; } void trans(node &a,node b,int c) { if(!b.y) return ;//important judge if(a.x<b.x+c) a.x=b.x+c,a.y=0; if(a.x==b.x+c) a.y+=b.y; } void work(int s,int u,int v) { cnt=0;dfs(u,0); for(int i=0;i<=n;i++) for(int j=0;j<=m;j++) dp[i][j].x=dp[i][j].y=0; for(int i=0;i<=m;i++) dp[0][i].y=1; for(int i=1;i<=n;i++) { int u=id[i];//choose i if(s>>u&1) for(int j=m;j>=w[u];j--) trans(dp[i][j],dp[i-1][j-w[u]],a[u]); if(i>1 && (dfn[v]<i || dfn[v]>out[i])) for(int j=0;j<=m;j++) trans(dp[out[i]][j],dp[i-1][j],0); } } int count(int n) { int res=0;while(n) res+=n/=5; return res; } array<ll,25> zxy(int n) { ll a[25]={},b[25]={};b[0]=1;array<ll,25> res; if(n==0) {res.fill(0);res[0]=1;return res;} int tn=n/10*5;res=zxy(tn); for(int i=1;i<23;i++) b[i]=b[i-1]*tn%p; //binomial theorem for(int i=0;i<23;i++) for(int j=i;j<23;j++) a[i]=(a[i]+res[j]*C[j][j-i]%p*b[j-i])%p; //convolution for(int i=22;i>=0;i--) { ll t=0; for(int j=0;j<=i;j++) t=(t+res[j]*a[i-j])%p; res[i]=t; } //some remaining numbers for(;n>2*tn;n--) if(n%5) for(int i=22;i>=0;i--) res[i]=(res[i]*n+(i?res[i-1]:0))%p; return res; } ll fac(int n) { ll res=1; while(n) res=res*zxy(n)[0]%p,n/=5; return res; } void exgcd(int a,int b,int &x,int &y) { if(!b) {x=1;y=0;return ;} exgcd(b,a%b,y,x);y-=(a/b)*x; } int comb(int n) { if(n<k) return 0; int d=count(n)-count(k)-count(n-k); int t=fac(n-k)*fac(k)%p,x=0,y=0; exgcd(t,p,x,y);x=(x%p+p)%p; x=x*fac(n)%p; while(d--) x=x*5%p; return x; } void calc(int u,int fa) { work(ok[u],u,0); if(dp[n][m].x==mx) ans=(ans+comb(dp[n][m].y))%p; for(int i=f[u];i;i=e[i].next) { int v=e[i].v; if(v==fa) continue; calc(v,u); work(ok[u]&ok[v],u,v); if(dp[n][m].x==mx) ans=(ans-comb(dp[n][m].y))%p; } } signed main() { n=read();m=read();k=read();t=read(); for(int i=1;i<=n;i++) w[i]=read(); for(int i=1;i<=n;i++) a[i]=read(); for(int i=1;i<n;i++) { int u=read(),v=read(),c=read(); e[++tot]=edge{v,c,f[u]},f[u]=tot; e[++tot]=edge{u,c,f[v]},f[v]=tot; } for(int i=1;i<=n;i++) pre(i,0,0,ok[i]); for(int i=1;i<=n;i++) work((1ll<<n+1)-1,i,0),mx=max(mx,dp[n][m].x); for(int i=0;i<=22;i++) { C[i][0]=1; for(int j=1;j<=i;j++) C[i][j]=(C[i-1][j-1]+C[i-1][j])%p; } calc(1,0); printf("%lld\n",(ans%p+p)%p); }
- 1
信息
- ID
- 2524
- 时间
- 2000ms
- 内存
- 512MiB
- 难度
- 10
- 标签
- 递交数
- 2
- 已通过
- 1
- 上传者