TLE on test 10,求大佬帮忙看看,回复即关注!!!
查看原帖
TLE on test 10,求大佬帮忙看看,回复即关注!!!
826362
fanfanfan123楼主2023/7/18 13:58
#include <iostream>
#include <vector>
using namespace std;
using vii = vector<int>;
#define fo(l, r) for(int i = l; i <= r; i++)
#define pb push_back
#define ios ios::sync_with_stdio(false); cin.tie(0); cout.tie(0);
typedef long long LL;
const int N = 100010;

struct node{
    int l, r;
    LL s, add;
}tr[N * 4];

int a[N], id[N], idx, dep[N], fa[N], n, m;
int root, top[N], sz[N], son[N], nw[N];
vii e[N];

void dfs1(int u, int f){
    sz[u] = 1, fa[u] = f, dep[u] = dep[f] + 1;
    for(int v: e[u]){
        if(v == f) continue;
        dfs1(v, u);
        sz[u] += sz[v];
        if(sz[son[u]] < sz[v]) son[u] = v;
    }
}

void dfs2(int u, int t){
    top[u] = t, id[u] = ++ idx, nw[idx] = a[u];
    if(son[u]) dfs2(son[u], t);
    for(int v: e[u]){
        if(v == fa[u] || v == son[u]) continue;
        dfs2(v, v);
    }
}

int lca(int u, int v){
    while(top[u] != top[v]){
        if(dep[top[u]] < dep[top[v]]) swap(u, v);
        u = fa[top[u]];
    }
    return dep[u] < dep[v] ? u : v;
}

void update(int u, LL add){
    tr[u].add += add;
    tr[u].s += 1ll * (tr[u].r - tr[u].l + 1) * add;
}

void pushup(int u){
    tr[u].s = tr[u << 1].s + tr[u << 1 | 1].s;
}

void pushdown(int u){
    LL &add = tr[u].add;
    if(!add) return;
    update(u << 1, add), update(u << 1 | 1, add);
    add = 0;
}

void build(int u, int l, int r){
    tr[u] = {l, r, nw[l]};
    if(l == r) return;
    int mid = l + r >> 1;
    build(u << 1, l, mid), build(u << 1 | 1, mid + 1, r);
    pushup(u);
}

void modify(int u, int l, int r, LL add){
    if(l > r) return;
    if(l <= tr[u].l && tr[u].r <= r){
        update(u, add);
        return;
    }
    pushdown(u);
    int mid = tr[u].l + tr[u].r >> 1;
    if(l <= mid) modify(u << 1, l, r, add);
    if(r > mid) modify(u << 1 | 1, l, r, add);
    pushup(u);
}

LL query(int u, int l, int r){
    if(l > r) return 0;
    if(l <= tr[u].l && tr[u].r <= r) return tr[u].s;
    pushdown(u);
    int mid = tr[u].l + tr[u].r >> 1;
    LL res = 0;
    if(l <= mid) res = query(u << 1, l, r);
    if(r > mid) res += query(u << 1 | 1, l, r);
    return res;
}

int find(int u){
    if(id[root] >= id[u] && id[root] <= id[u] + sz[u] - 1)
        for(int v: e[u]) if(v != fa[u] && lca(v, root) == v) return v;
    return fa[u];
}

void modify_tree(int u, int v, int add){
    int ca = lca(u, v);
    if(id[root] < id[ca] || id[root] > id[ca] + sz[ca] - 1) 
        modify(1, id[ca], id[ca] + sz[ca] - 1, add);
    else if(lca(u, root) == root || lca(v, root) == root) modify(1, 1, n, add);
    else{
        int ca1 = lca(u, root), ca2 = lca(v, root);
        int y = ca1 != ca ? ca1 : ca2;
        int x = find(y);
        modify(1, 1, n, add), modify(1, id[x], id[x] + sz[x] - 1, -add);
    }
}

LL query_tree(int u){
    if(id[root] < id[u] || id[root] >= id[u] + sz[u])
        return query(1, id[u], id[u] + sz[u] - 1);
    else if(root == u) return query(1, 1, n);
    int x = find(u);
    return query(1, 1, n) - query(1, id[x], id[x] + sz[x] - 1);
}

int main() {
	cin >> n >> m;
	fo(1, n) scanf("%d", &a[i]);
	fo(1, n - 1){
	    int u, v;
	    scanf("%d%d", &u, &v);
	    e[u].pb(v), e[v].pb(u);
	}
	dfs1(1, 0), dfs2(1, 1);
	build(1, 1, n);

	while(m --){
	    int op, a, b, c;
	    scanf("%d%d", &op, &a);
	    if(op == 2) scanf("%d%d", &b, &c);
	    if(op == 1) root = a;
	    else if(op == 2) modify_tree(a, b, c);
	    else printf("%lld\n", query_tree(a));
	}
    
	return 0;
}
2023/7/18 13:58
加载中...