#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