实现方法跟求助n^3 DP 莫名AC差不多,但我TLE了。。。
#include<cmath>
#include<cstdio>
#include<cstring>
#include<iostream>
#include<algorithm>
using namespace std;
void read(int &sum)
{
sum=0;char last='w',ch=getchar();
while (ch<'0' || ch>'9') last=ch,ch=getchar();
while (ch>='0' && ch<='9') sum=sum*10+ch-'0',ch=getchar();
if (last=='-') sum=-sum;
}
int n;
struct edge { int x,y,next; }a[5010*2];
int len,last[5010];
void ins(int x,int y)
{
len++,a[len].x=x,a[len].y=y;
a[len].next=last[x],last[x]=len;
}
int s[5010];
struct point { int f[5010],sz; }p[5010];
void dfs(int x,int fa)
{
p[x].sz=1;
for (int k=last[x];k;k=a[k].next)
{
int y=a[k].y;
if (y!=fa)
dfs(y,x),p[x].sz+=p[y].sz;
}
for (int i=1;i<=p[x].sz;i++)
p[x].f[i]=-1;
}
void find_(int x,int fa)
{
p[x].f[0]=0;
p[x].sz=1;
int sum=0;
for (int k=last[x];k;k=a[k].next)
{
int y=a[k].y;
if (y!=fa)
{
find_(y,x);
p[x].sz+=p[y].sz;
for (int i=p[x].sz-1;i>=1;i--)
{
for (int j=min(p[y].sz-1,i);j>=1;j--)
{
if (p[x].f[i-j]!=-1 && p[y].f[j]!=-1)
{
if (p[x].f[i]==-1) p[x].f[i]=p[y].f[j]+p[x].f[i-j];
else p[x].f[i]=min(p[x].f[i],p[y].f[j]+p[x].f[i-j]);
}
}
}
}
}
for (int i=0;i<=p[x].sz-1;i++)
if (p[x].f[i]!=-1)
{
if (p[x].f[p[x].sz-1]==-1) p[x].f[p[x].sz-1]=p[x].f[i]+s[p[x].sz-i-1];
else p[x].f[p[x].sz-1]=min(p[x].f[p[x].sz-1],p[x].f[i]+s[p[x].sz-i-1]);
}
}
int main()
{
// freopen("T2ex2.in","r",stdin);
// freopen("M.out","w",stdout);
read(n);
for (int i=1;i<=n-1;i++) read(s[i]);
for (int i=1;i<=n-1;i++)
{
int x,y;read(x),read(y);
ins(x,y),ins(y,x);
}
dfs(1,-1);
find_(1,-1);
printf("%d\n",p[1].f[n-1]);
// for (int i=1;i<=n;i++)
// {
// printf("point : %d %d\n",i,p[i].sz);
// for (int j=1;j<=p[i].sz;j++)
// printf("%d ",p[i].f[j]);
// printf("\n");
// }
// fclose(stdin);
// fclose(stdout);
// return 0;
}