50pts TLE 求调
查看原帖
50pts TLE 求调
854735
CondorOfLiberty楼主2023/8/20 17:11
#include<bits/stdc++.h>
using namespace std;
namespace io{
    int read(){
        int r=0; char c; bool f=1;
        do { if(c=='-') f=0; c=getchar(); } while(!isdigit(c));
        do r=(r<<3)+(r<<1)+(c^'0'), c=getchar(); while(isdigit(c));
        return f ? r : (~r+1);
    }
}
using namespace io;
const int N=1e5+3;
int a[N];
vector <int> v[N];
void add(int x, int y){
    v[x].push_back(y);
    v[y].push_back(x);
}
int d[N], fa[N], son[N];
void dfs1(int x){
    if(fa[x]) v[x].erase(find(v[x].begin(), v[x].end(), fa[x]));
    son[x]=1;
    for(auto &y: v[x]){
        fa[y]=x, d[y]=d[x]+1;
        dfs1(y);
        son[x]+=son[y];
        if(son[y]>son[v[x][0]]) swap(y, v[x][0]);
    }
}
int top[N], in[N], out[N], seq[N], idx;
void dfs2(int x){
    in[x]=++idx, seq[idx]=x;
    for(auto y: v[x]){
        top[y]= y==v[x][0] ? top[x] : y;
        dfs2(y);
    }
    out[x]=idx;
}
typedef long long ll;
struct Seg{
    int l, r, len;
    ll w, v;
} t[4*N];
void up(int p){
    t[p].w=t[p<<1].w+t[p<<1|1].w;
}
void Build(int p, int l, int r){
    t[p].l=l, t[p].r=r, t[p].len=r-l+1;
    if(l==r){
        t[p].w=a[seq[l]];
        return ;
    }
    int mid=l+((r-l)>>1);
    Build(p<<1, l, mid), Build(p<<1|1, mid+1, r);
    up(p);
}
void down(int p){
    t[p<<1].w+=t[p].v*t[p<<1].len, t[p<<1|1].w+=t[p].v*t[p<<1|1].len;
    t[p<<1].v+=t[p].v, t[p<<1|1].v+=t[p].v;
    t[p].v=0;
}
void update1(int p, int x, int w){
    int l=t[p].l, r=t[p].r;
    if(l==r&&l==x){
        t[p].w+=w;
        return ;
    }
    if(t[p].v) down(p);
    int mid=l+((r-l)>>1);
    if(mid>=x) update1(p<<1, x, w);
    else update1(p<<1|1, x, w);
    up(p);
}
void update2(int p, int l, int r, int x){
    int u=t[p].l, v=t[p].r;
    if(u>=l&&v<=r){
        t[p].w+=x*t[p].len;
        t[p].v+=x;
        return ;
    }
    if(t[p].v) down(p);
    int mid=u+((v-u)>>1);
    if(mid>=l) update2(p<<1, l, r, x);
    if(mid<r) update2(p<<1|1, l, r, x);
    up(p);
}
ll Getsum(int p, int l, int r){
    int u=t[p].l, v=t[p].r;
    if(u>=l&&v<=r) return t[p].w;
    if(t[p].v) down(p);
    ll ans=0;
    int mid=u+((v-u)>>1);
    if(mid>=l) ans+=Getsum(p<<1, l, r);
    if(mid<r) ans+=Getsum(p<<1|1, l, r);
    return ans;
}
int main(){
    int n, m;
    n=read(), m=read();
    for(int i=1;i<=n;++i)
        a[i]=read();
    for(int i=1;i<n;++i){
        int x, y;
        x=read(), y=read();
        add(x, y);
    }
    dfs1(1), top[1]=1, dfs2(1);
    Build(1, 1, idx);
    while(m--){
        int opt, x, w;
        opt=read(), x=read();
        if(opt==1){
            w=read();
            update1(1, in[x], w);
        }else if(opt==2){
            w=read();
            update2(1, in[x], out[x], w);
        }else{
            ll ans=0;
            while(top[x]!=1) ans+=Getsum(1, in[top[x]], in[x]), x=fa[top[x]];
            ans+=Getsum(1, 1, in[x]);
            printf("%lld\n", ans);
        }
    }
    return 0;
}
2023/8/20 17:11
加载中...