30分求助
查看原帖
30分求助
185271
sky_light楼主2023/7/3 15:49

用的是倍增合并线性基的朴素做法,求调,感谢

#include<bits/stdc++.h>
using namespace std;
const int maxn=2e4+11;
typedef long long ll;
int n,q,root=1,fa[25][maxn],tot=0,dep[maxn];
ll a[maxn];
ll p[25][maxn][63];
ll qa[63];
vector<int>v[maxn];
template<typename T> inline void read(T &x) {
    x = 0;
    char c = getchar();
    while(!isdigit(c)) c = getchar();
    while(isdigit(c)) x = x * 10 + c - '0', c = getchar();
}
inline bool insert(ll *qwq,ll x){
    for(int i=60;i>=0;i--){
        if(x>>i&1){
            if(!qwq[i]){qwq[i]=x;return true;}
            x^=qwq[i];
        }
    }
    return false;
}
inline void merge(ll *x,ll *y){
    for(int i=60;i>=0;i--) if(x[i]) insert(y,x[i]);
}
void dfs(int x,int f){
    fa[0][x]=f,dep[x]=dep[f]+1;
    insert(p[0][x],a[x]);
    for(int i=0;i<v[x].size();i++){
        int s=v[x][i];
        if(s==f) continue;
        dfs(s,x);
    }
}
ll found_max(){
    ll qwq=0;
    for(int i=60;i>=0;i--) qwq=max(qwq,1ll*qwq^qa[i]);
    return qwq;
}
inline int get_lca(int x,int y){
    if(dep[x]<dep[y]) swap(x,y);
    for(int i=15;i>=0;i--) if(dep[fa[i][x]]>=dep[y]) x=fa[i][x];
    if(x==y) return x;
    for(int i=15;i>=0;i--) if(fa[i][x]!=fa[i][y]) x=fa[i][x],y=fa[i][x];
    return fa[0][x];
}
void get_ans(int x,int y){
    memset(qa,0,sizeof(qa));
    int lca=get_lca(x,y);
    if(x==y) insert(qa,a[x]);
    else if(x==lca||y==lca){
        if(dep[x]<dep[y]) swap(x,y);
        for(int i=15;i>=0;i--){
            if(dep[fa[i][x]]>=dep[y]){
                merge(p[i][x],qa);
                x=fa[i][x];
            }
        }
        insert(qa,a[x]);
    }else{
        for(int i=15;i>=0;i--){
            if(dep[lca]<=dep[fa[i][x]]){
                merge(p[i][x],qa);
                x=fa[i][x];
            }
            if(dep[lca]<=dep[fa[i][y]]){
                merge(p[i][y],qa);
                x=fa[i][y];
            }
        }
        insert(qa,a[lca]);
    }
}
int main(){
    //freopen("ans.out","w",stdout);
    read(n),read(q);
    for(int i=1;i<=n;i++) read(a[i]);
    for(int i=1;i<n;i++){
        int x,y;
        read(x),read(y);
        v[x].push_back(y);
        v[y].push_back(x);
    }
    dfs(1,1);
    //puts("qaq");
    for(int i=1;i<=16;i++){
        for(int j=1;j<=n;j++){
            fa[i][j]=fa[i-1][fa[i-1][j]];
            merge(p[i-1][j],p[i][j]),merge(p[i-1][fa[i-1][j]],p[i][j]);
        }
    }
    //puts("awa");
    for(int i=1;i<=q;i++){
        int x,y;
        read(x),read(y);
        get_ans(x,y);
        //puts("qwq");
        printf("%lld\n",found_max());
    }
    return 0;
}
2023/7/3 15:49
加载中...