萌新求助,TLE 60pts,请问有什么写假了的地方吗
查看原帖
萌新求助,TLE 60pts,请问有什么写假了的地方吗
122641
GIFBMP楼主2022/7/29 09:57

Rt

#include <cstdio>
#include <cstring>
#include <algorithm>
#include <cctype>
using std :: swap ;
using std :: max ;
const int MAXN = 2e5 + 10 , MAXM = 5e6 + 10 , INF = 0x3f3f3f3f ;
int n , q , fir[MAXN] , tot , mx[MAXN] , sz[MAXN] , Rt , vis[MAXN] , cnt , f[MAXN] , p[MAXN][18] , dep[MAXN] ;
int las , a[MAXN] ;
struct edge {
	int to , nxt ;
} e[MAXN << 1] ;
void add (int u , int v) {e[++tot].to = v ; e[tot].nxt = fir[u] ; fir[u] = tot ;}
void dfs (int x , int fa) {
	p[x][0] = fa , dep[x] = dep[fa] + 1 ;
	for (int i = 1 ; i <= 17 ; i++) p[x][i] = p[p[x][i - 1]][i - 1] ;
	for (int i = fir[x] , v = e[i].to ; i ; i = e[i].nxt , v = e[i].to)
		if (v != fa) dfs (v , x) ;
}
int lca (int x , int y) {
	if (dep[x] < dep[y]) swap (x , y) ;
	for (int i = 17 ; ~i ; i--) if (dep[p[x][i]] >= dep[y]) x = p[x][i] ;
	if (x == y) return x ;
	for (int i = 17 ; ~i ; i--) if (p[x][i] != p[y][i]) x = p[x][i] , y = p[y][i] ;
	return p[x][0] ;
}
int dis (int x , int y) {
	return dep[x] + dep[y] - 2 * dep[lca (x , y)] ;
}
void findrt (int x , int fa) {
	sz[x] = 1 , mx[x] = 0 ;
	for (int i = fir[x] , v = e[i].to ; i ; i = e[i].nxt , v = e[i].to)
		if (v != fa && !vis[v]) findrt (v , x) , sz[x] += sz[v] , mx[x] = max (mx[x] , sz[v]) ;
	mx[x] = max (mx[x] , cnt - sz[x]) ;
	if (mx[x] < mx[Rt]) Rt = x ;
}
void div (int x) {
	vis[x] = 1 ; int tmp = cnt ;
	for (int i = fir[x] , v = e[i].to ; i ; i = e[i].nxt , v = e[i].to)
		if (!vis[v]) Rt = 0 , cnt = (sz[v] > sz[x] ? tmp - sz[x] : sz[v]) , findrt (v , x) , f[Rt] = x , div (Rt) ; 
}
#define mid ((l + r) >> 1)
struct sgt {
	int tt , lc[MAXM] , rc[MAXM] , rt[MAXN] , s[MAXM] ;
	sgt () {tt = 0 ;}
	void upd (int &o , int l , int r , int x , int k) {
		if (!o) o = ++tt ;
		if (l == r) {s[o] += k ; return ;}
		if (x <= mid) upd (lc[o] , l , mid , x , k) ;
		else upd (rc[o] , mid + 1 , r , x , k) ;
		s[o] = s[lc[o]] + s[rc[o]] ;
	}
	int query (int o , int l , int r , int x , int y) {
		if (!o) return 0 ;
		if (x <= l && r <= y) return s[o] ;
		int ret = 0 ;
		if (x <= mid) ret += query (lc[o] , l , mid , x , y) ;
		if (mid < y) ret += query (rc[o] , mid + 1 , r , x , y) ;
		return ret ;
	}
} t1 , t2 ;
void Upd (int x , int k) {
	for (int i = x ; i ; i = f[i]) {
		t1.upd (t1.rt[i] , 0 , n - 1 , dis (i , x) , k) ;
		if (f[i]) t2.upd (t2.rt[i] , 0 , n - 1 , dis (f[i] , x) , k) ; 
	}
}
int Query (int x , int k) {
	int ret = 0 ;
	for (int i = x , j = 0 ; i ; j = i , i = f[i]) {
		if (dis (i , x) > k) continue ;
		ret += t1.query (t1.rt[i] , 0 , n - 1 , 0 , k - dis (i , x)) ;
		if (j) ret -= t2.query (t2.rt[j] , 0 , n - 1 , 0 , k - dis (i , x)) ; 
	}
	return ret ;
}
int R () {
	int x = 0 ; char ch = getchar () ; bool f = 0 ;
	for (; !isdigit (ch) ; ch = getchar ()) if (ch == '-') f = 1 ;
	for (; isdigit (ch) ; ch = getchar ()) x = (x << 1) + (x << 3) + (ch ^ 48) ;
	return f ? -x : x ;
}
int main () {
	n = R () ; q = R () ; cnt = n ; mx[0] = INF ;
	for (int i = 1 ; i <= n ; i++) a[i] = R () ;
	for (int i = 1 , u , v ; i < n ; i++)
		u = R () , v = R () , add (u , v) , add (v , u) ;
	dfs (1 , 0) , findrt (1 , 0) , div (Rt) ;
	//for (int i = 1 ; i <= n ; i++) printf ("%d->%d\n" , f[i] , i) ;
	for (int i = 1 ; i <= n ; i++) Upd (i , a[i]) ;
	while (q--) {
		int opt = R () , x = R () , k = R () ;
		x ^= las , k ^= las ;
		if (opt == 0) printf ("%d\n" , (las = Query (x , k))) ;
		else Upd (x , k - a[x]) , a[x] = k ;
	}
	return 0 ;
} 
2022/7/29 09:57
加载中...