样例没过,悬关求调~~~
查看原帖
样例没过,悬关求调~~~
868063
_Coffice_楼主2023/7/5 13:48
#include <iostream>
using namespace std;
#define int long long
const int N = 100005;
struct Edge
{
	int v, nx;
};
Edge e[N*3];
int head[N], tot = 0;
void insert(int u, int v)
{
	tot++;
	e[tot] = {v, head[u]};
	head[u] = tot;
}
int num[N]; 
int dep[N], son[N], fa[N], sz[N];
void dfs_1(int p)
{
	sz[p] = 1;
	int m = -1;
	for(int i=head[p];i;i=e[i].nx)
		if(e[i].v != fa[p])
		{
			fa[e[i].v] = p;
			dep[e[i].v] = dep[p]+1;
			dfs_1(e[i].v);
			sz[p] += sz[e[i].v];
			if(sz[e[i].v] >= m)
			{
				son[p] = e[i].v;
				m = sz[e[i].v];
			}
		}
}
int top[N], cnt = 0, seg[N], rev[N];
void dfs_2(int p, int t)
{
	top[p] = t;
	cnt++;
	seg[p] = cnt;
	rev[cnt] = p;
	if(son[p]) dfs_2(son[p], t);
	for(int i=head[p];i;i=e[i].nx)
	{
		int to = e[i].v;
		if(to != fa[p] && to != son[p])
			dfs_2(to, to);
	}
}
struct seg_tree { int l, r, val, tag; };
seg_tree st[4*N];
void bt(int p, int l, int r)
{
	st[p] = {l, r, 0, 0};
	if(l == r)
	{
		st[p].val = num[rev[l]];
		return ;
	}
	int mid = (l+r)/2;
	bt(p*2, l, mid);
	bt(p*2+1, mid+1, r);
	st[p].val = st[p*2].val+st[p*2+1].val;
}
void down(int p)
{
	if(st[p].tag)
	{
		st[p*2].val += (st[p*2].r-st[p*2].l+1)*st[p].tag;
		st[p*2+1].val += (st[p*2+1].r-st[p*2+1].l+1)*st[p].tag;
		st[p*2].tag += st[p].tag;
		st[p*2+1].tag += st[p].tag;
		st[p].tag = 0;
	}
}
void add(int p, int l, int r, int x)
{
	if(l <= st[p].l && r >= st[p].r)
	{
		st[p].tag += x;
		st[p].val += (st[p].r-st[p].l+1)*x;
		return;
	}
	down(p);
	int mid = (st[p].l+st[p].r)/2;
	if(l <= mid) add(p*2, l, r, x);
	if(r >= mid+1) add(p*2+1, l, r, x);
}
void op1(int x, int a)
{
	add(1, seg[x], seg[x], a);
}
void op2(int x, int a)
{
	add(1, seg[x], seg[x]+sz[x]-1, a);
}
int ask(int p, int l, int r)
{
	if(l <= st[p].l && r >= st[p].r)
	{
		return st[p].val;
	}
	down(p);
	int mid = (st[p].l+st[p].r)/2;
	int ans = 0;
	if(l <= mid) ans += ask(p*2, l, r);
	if(r >= mid+1) ans += ask(p*2+1, l, r);
	return ans;
}
int op3(int x, int y)
{
	int ans = 0;
	while(top[x] != top[y])
	{
		if(dep[top[x]] > dep[top[y]]) swap(x, y);
		ans += ask(1, seg[top[y]], seg[y]);
		y = fa[top[y]];
	}
	if(dep[x] > dep[y]) swap(x, y);
	ans += ask(1, seg[x], seg[y]);
	return ans;
}
signed main()
{
	int n, m;
	cin >> n >> m;
	for(int i=1;i<=n;i++)
		cin >> num[i];
	for(int i=1;i<=n-1;i++)
	{
		int from, to;
		cin >> from >> to;
		insert(from, to);
		insert(to, from);
	}
	dfs_1(1);
	dfs_2(1, 1);
	bt(1, 1, n);
	for(int i=1;i<=m;i++)
	{
		int op;
		cin >> op;
		if(op == 1)
		{
			int x, a;
			cin >> x >> a;
			op1(x, a);
		}
		else if(op == 2)
		{
			int x, a;
			cin >> x >> a;
			op2(x, a);
		}
		else
		{
			int x;
			cin >> x;
			cout << op3(1, x) << endl;
		}
	}
	return 0;
}

样例输出:

6 
8 
12

正确答案:

6
9
13
2023/7/5 13:48
加载中...