求助!!!自己造的数据都过了(1关注为报)
查看原帖
求助!!!自己造的数据都过了(1关注为报)
513900
Wilson_Lee楼主2022/7/16 17:32

用线段树优化dp求的LIS,自己造的数据全过了,但在这里就是WA

#include <bits/stdc++.h>
using namespace std;

#define ls(p) p<<1
#define rs(p) p<<1|1
const int MAXN=2e5+5;
vector<int>G[MAXN];
struct node
{
    int val,id,tag=-1;
}p[MAXN];
struct Segment_Tree
{
    int l,r,maxn=0;
}tree[MAXN<<2];
int ans[MAXN];
bool cmpval(node x,node y)
{
    return x.val<y.val;
}
bool cmpid(node x,node y)
{
    return x.id<y.id;
}
void push_up(int p)
{
    tree[p].maxn=max(tree[ls(p)].maxn,tree[rs(p)].maxn);
}
void build(int p,int l,int r)
{
    tree[p].l=l,tree[p].r=r;
    if(l==r) return;
    int mid=(l+r)>>1;
    build(ls(p),l,mid);
    build(rs(p),mid+1,r);
}
void update(int p,int id,int x)
{
    if(tree[p].l==tree[p].r)
    {
        tree[p].maxn=max(tree[p].maxn,x);
        return;
    }
    int mid=(tree[p].l+tree[p].r)>>1;
    if(id<=mid) update(ls(p),id,x);
    else update(rs(p),id,x);
    push_up(p);
}
void select(int p,int id,int x)
{
    if(tree[p].l==tree[p].r)
    {
        tree[p].maxn=x;
        return;
    }
    int mid=(tree[p].l+tree[p].r)>>1;
    if(id<=mid) select(ls(p),id,x);
    else select(rs(p),id,x);
    push_up(p);
}
int query(int p,int a,int b)
{
    if(a<=tree[p].l && tree[p].r<=b)
        return tree[p].maxn;
    int mid=(tree[p].l+tree[p].r)>>1,ret=0;
    if(a<=mid) ret=max(ret,query(ls(p),a,b));
    if(b>mid) ret=max(ret,query(rs(p),a,b));
    return ret;
}
void dfs(int x,int father,int sum)
{
    int maxn=1;
    if(p[x].val>1) maxn=query(1,1,p[x].val-1)+1;
    ans[x]=sum=max(sum,maxn);
    int tmp=query(1,p[x].val,p[x].val);
    update(1,p[x].val,maxn);
    int si=G[x].size();
    for(int i=0;i<si;++i)
    {
        int y=G[x][i];
        if(y==father) continue;
        dfs(y,x,sum);
    }
    select(1,p[x].val,tmp);
}
int main()
{
    int n;
    cin>>n;
    for(int i=1;i<=n;++i) scanf("%d",&p[i].val),p[i].id=i;
    int u,v;
    for(int i=1;i<n;++i)
    {
        scanf("%d %d",&u,&v);
        G[u].push_back(v),G[v].push_back(u);
    }
    sort(p+1,p+n+1,cmpval);
    p[0]=-1e9;
    int cnt=0;
    for(int i=1;i<=n;++i)
    {
        if(p[i].val!=p[i-1].val) ++cnt;
        p[i].val=cnt;
    }
    sort(p+1,p+n+1,cmpid);
    build(1,1,cnt);
    dfs(1,0,0);
    for(int i=1;i<=n;++i) printf("%d\n",ans[i]);
    return 0;
}
2022/7/16 17:32
加载中...