求助n^2 -> n^3 TLE 70
查看原帖
求助n^2 -> n^3 TLE 70
100709
Loser_and_Joker楼主2022/10/4 13:07

实现方法跟求助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;
}
2022/10/4 13:07
加载中...