萌新20pts求调
查看原帖
萌新20pts求调
576817
Lyrella楼主2022/11/12 09:56

树剖板子

#include <bits/stdc++.h>
#define ll long long
#define Fl(i, a, b) for(int i = a; i <= b; i++)
using namespace std;
const int N = 2e5 + 5;
int fa[N], dep[N], siz[N], son[N], top[N], seg[N], rev[N], idq;
int nxt[N], to[N], hd[N], cnt;
ll lz[N << 2], sum[N << 2], a[N], p;
int n, m, root;
void dfs1(int u, int f){
	dep[u] = dep[f] + 1; siz[u] = 1; fa[u] = f;
	for(int i = hd[u]; i; i = nxt[i]){
		int v = to[i]; if(v == f)continue;
		dfs1(v, u); siz[u] += siz[v];
		if(siz[v] > siz[son[u]])son[u] = v;
	}
}
void dfs2(int u, int tp){
	seg[u] = ++idq; rev[idq] = u; top[u] = tp;
	if(son[u])dfs2(son[u], tp);
	for(int i = hd[u]; i; i = nxt[i]){
		int v = to[i]; if(v == son[u] or v == fa[u])continue;
		dfs2(v, v);
	}
}
void add(int u, int v){
	nxt[++cnt] = hd[u]; hd[u] = cnt; to[cnt] = v;
}
void _push(int x){
	sum[x] = sum[x << 1] + sum[x << 1 | 1];
	sum[x] %= p;
}
void pushdown(int x, int l, int r){
	int len = r - l + 1;
	sum[x << 1] += lz[x] * (len / 2); sum[x << 1] %= p;
	sum[x << 1 | 1] += lz[x] * (len - len / 2); sum[x << 1 | 1] %= p; 
	lz[x << 1] += lz[x]; lz[x << 1 | 1] += lz[x]; lz[x] = 0;
}
void build(int x, int l, int r){
	if(l == r){
		sum[x] = a[rev[l]] % p;
		return;
	}
	int mid = l + r >> 1;
	build(x << 1, l, mid);
	build(x << 1 | 1, mid + 1, r);
	_push(x);
}
void upd(int x, int l, int r, int L, int R, ll y){
	if(L <= l and r <= R){
		lz[x] += y;
		sum[x] += y * (r - l + 1);
		return;
	}
	int mid = l + r >> 1;
	if(L <= mid)upd(x << 1, l, mid, L, R, y);
	if(R > mid)upd(x << 1 | 1, mid + 1, r, L, R, y);
	_push(x);
}
ll query(int x, int l, int r, int L, int R){
	if(L <= l and r <= R)return sum[x] % p;
	int mid = l + r >> 1; ll ans = 0;
	if(lz[x])pushdown(x, l, r);
	if(L <= mid)ans += query(x << 1, l, mid, L, R), ans %= p;
	if(mid < R)ans += query(x << 1 | 1, mid + 1, r, L, R), ans %= p;
	return ans;
}
ll qsum(int u, int v){
	ll ans = 0;
	while(top[u] != top[v]){
		if(dep[top[u]] < dep[top[v]])swap(u, v);
		ans += query(1, 1, n, seg[top[u]], seg[u]);
		ans %= p; u = fa[top[u]];
	}
	if(dep[u] > dep[v])swap(u, v);
	ans += query(1, 1, n, seg[u], seg[v]); ans %= p;
	return ans;
}
void usum(int u, int v, int ad){
	while(top[u] != top[v]){
		if(dep[top[u]] < dep[top[v]])swap(u, v);
		upd(1, 1, n, seg[top[u]], seg[u], ad);
		u = fa[top[u]];
	}
	if(dep[u] > dep[v])swap(u, v);
	upd(1, 1, n, seg[u], seg[v], ad);
}
ll qtree(int u){
	return query(1, 1, n, seg[u], seg[u] + siz[u] - 1);
}
void utree(int u, int ad){
	upd(1, 1, n, seg[u], seg[u] + siz[u] - 1, ad);
}
void solve(){
	cin >> n >> m >> root >> p;
	Fl(i, 1, n)cin >> a[i];
	Fl(i, 1, n - 1){
		int u, v; cin >> u >> v;
		add(u, v); add(v, u);
	}
	dfs1(root, 0); dfs2(root, root); build(1, 1, n);
	Fl(i, 1, m){
		int opt, x, y, z; cin >> opt >> x;
		if(opt == 1){
			cin >> y >> z;
			usum(x, y, z);
		}
		if(opt == 2){
			cin >> y;
			cout << qsum(x, y) << '\n';
		}
		if(opt == 3){
			cin >> z;
			utree(x, z);
		}
		if(opt == 4)cout << qtree(x) << '\n';
	}
}
int main(){
	std::ios::sync_with_stdio(false);
	std::cin.tie(nullptr);
	solve(); return 0;
}
2022/11/12 09:56
加载中...