站外题求优化(详细注释代码清晰),正确性已保证只需优化
  • 板块题目总版
  • 楼主MessageBoxA
  • 当前回复5
  • 已保存回复5
  • 发布时间2023/2/23 19:04
  • 上次更新2023/10/24 00:01:28
查看原帖
站外题求优化(详细注释代码清晰),正确性已保证只需优化
77584
MessageBoxA楼主2023/2/23 19:04

TLE6个点,1e5的数据炸了

题目链接(AcWing353)雨天的尾巴

TZ 谢谢大佬辣!

#include<bits/stdc++.h>
#define getmid int mid=(l+r)/2
using namespace std;
typedef pair<int,int> pii;
const int MAXN=100005;
int n,m,cnt,typ[MAXN],fa[MAXN],ans[MAXN]={0};//typ离散化数组 fa每个节点的父节点
short vis[MAXN]={0};//tarjan的标记
struct OBJ
{
    int x,y,z,lca;//xyz对应题目 lca为x和y的LCA
}obj[MAXN];
vector<int>g[MAXN];//邻接链表
vector<pii>query[MAXN];//query[x].first代表与x求lca的目标节点 second表示obj编号
class ValSegTree{//代替差分数组的权值线段树
    private:
        struct Node{
            int val=0,id=0;
            Node *lch=nullptr,*rch=nullptr;//动态开点
        };
        inline void push_up(Node *pos){
            Node *l=pos->lch,*r=pos->rch;
            if(l==nullptr && r==nullptr) return;
            if(l==nullptr){
                pos->val = r->val;
                pos->id = r->id;
                return;
            }
            if(r==nullptr){
                pos->val = l->val;
                pos->id = l->id;
                return;
            }
            if(l->val != r->val){
                if(l->val > r->val){
                    pos->val = l->val;
                    pos->id = l->id;
                }
                else{
                    pos->val = r->val;
                    pos->id = r->id;
                }
            }
            else{
                pos->val = l->val;
                pos->id = min(l->id,r->id);
            }
        }
    public:
        Node *head[MAXN];
        void init(){
            for(int i=1;i<=n;i++){
                head[i]=new Node;
            }
        }
        void update(Node *&pos,int l,int r,int k,int v,int no){
            if(r<l) return;
            if(pos==nullptr) pos=new Node;
            if(l==r){
                pos->id=no;
                pos->val+=v;
                return;
            }
            getmid;
            if(k<=mid) update(pos->lch,l,mid,k,v,no);
            else update(pos->rch,mid+1,r,k,v,no);
            push_up(pos);
        }
        Node* merge(Node *a,Node *b,int l,int r){
            if(r<l) return nullptr;
            if(a==nullptr) return b;
            if(b==nullptr) return a;
            // cout<<a->id<<' '<<a->val<<' '<<b->id<<' '<<b->val<<endl;
            if(l==r){
                a->id = b->id;
                a->val += b->val;
                // if(a->val != b->val){
                //     if(a->val < b->val){
                //         a->val = b->val;
                //         a->id = b->id;
                //     }
                // }
                // else{
                //     if(a->id > b->id){
                //         a->id = b->id;
                //     }
                // }
                //这里注释的地方写错了,由于是差分,对于每种物品应该把他们差分值相加得到真实值,而不是求max
                return a;
            }
            getmid;
            a->lch=merge(a->lch,b->lch,l,mid);
            a->rch=merge(a->rch,b->rch,mid+1,r);
            push_up(a);
            return a;
        }
}tr;
class UFS{
    private:
        int fa[MAXN];
    public:
        inline void init(){
            for(int i=1;i<=n;i++){
                fa[i]=i;
            }
        }
        int get(int x){
            if(fa[x]==x) return x;
            else return fa[x]=get(fa[x]);
        }
        inline void merge(int x,int y){
            fa[y]=x;
        }
}ufs;
void tarjan(int x,int fno){//tarjan求LCA
    vis[x]=1;
    fa[x]=fno;
    for(auto it:g[x]){
        if(it==fno || vis[it]) continue;
        tarjan(it,x);
        ufs.merge(x,it);//tarjan算法的并查集merge是有顺序要求的,一个节点的fa得是它祖先
    }
    for(auto it:query[x]){
        if(vis[it.first]==2){
            obj[it.second].lca=ufs.get(it.first);
        }
    }
    vis[x]=2;
}
void getans(int x,int fno){//合并答案
    for(auto it:g[x]){
        if(it==fno) continue;
        getans(it,x);
        tr.head[x]=tr.merge(tr.head[x],tr.head[it],1,cnt);
    }
    if(tr.head[x]!=nullptr) ans[x]=tr.head[x]->id;
}
int main(){
    cin>>n>>m;
    tr.init();
    ufs.init();
    for(int i=1,x,y;i<n;i++){
        cin>>x>>y;
        g[x].push_back(y);
        g[y].push_back(x);
    }
    for(int i=1;i<=m;i++){
        cin>>obj[i].x>>obj[i].y>>obj[i].z;
        typ[i]=obj[i].z;
        if(obj[i].x==obj[i].y){
            obj[i].lca=obj[i].x;
            continue;
        }
        query[obj[i].x].emplace_back(obj[i].y,i);
        query[obj[i].y].emplace_back(obj[i].x,i);
    }//初始化&输入

    tarjan(1,0);//求LCA

    sort(typ+1,typ+1+m);
    cnt=unique(typ+1,typ+1+m)-typ-1;
    int tmp;
    for(int i=1;i<=m;i++){
        tmp=obj[i].z;
        obj[i].z=lower_bound(typ+1,typ+1+cnt,obj[i].z)-typ;//离散化

        tr.update(tr.head[obj[i].x],1,cnt,obj[i].z,1,tmp);//由于这是权值线段树,所以初始右边界为cnt(离散化后的物品种类数)不为n
        tr.update(tr.head[obj[i].y],1,cnt,obj[i].z,1,tmp);
        tr.update(tr.head[obj[i].lca],1,cnt,obj[i].z,-1,tmp);
        if(fa[obj[i].lca]) tr.update(tr.head[fa[obj[i].lca]],1,cnt,obj[i].z,-1,tmp);//树上差分
    }
    
    getans(1,0);
    for(int i=1;i<=n;i++){
        cout<<ans[i]<<endl;
    }
    return 0;
}
2023/2/23 19:04
加载中...