代码如下,其他地方已调过,应该是树状数组有问题,但我不知道哪里有问题啊啊啊啊。
大佬帮帮忙,孩子已经调了一下午了。
#include <bits/stdc++.h>
using namespace std;
const int N = 1e5 + 10;
int n, m, r, p, cnt;
int t1[N], t2[N];
int dep[N], son[N], dfn[N], fa[N], siz[N], top[N], w[N];
vector<int> g[N];
int L(int x) {
return x & (-x);
}
void A(int x, int k) {
int k1 = (x - 1) * k % p;
for (; x <= n; x += L(x)) {
t1[x] = ((t1[x] + k) % p + p) % p;
t2[x] = ((t2[x] + k1) % p + p) % p;
}
}
void add(int l, int r, int k) {
A(l, k), A(r + 1, -k);
}
int Q(int x) {
int ans = 0;
for (int o = x; o > 0; o -= L(o)) {
ans = (ans + x * t1[o] % p) % p;
ans = (ans - t2[o]) % p;
ans = (ans + p) % p;
}
return ans;
}
int query(int l, int r) {
return ((Q(r) - Q(l - 1)) % p + p) % p;
}
void dfs1(int u, int f) {
fa[u] = f;
siz[u] = 1;
dep[u] = dep[f] + 1;
son[u] = 0;
int sz = g[u].size();
for (int i = 0; i < g[u].size(); ++i) {
int v = g[u][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 t) {
top[u] = t;
dfn[u] = ++cnt;
if (w[u]) add(dfn[u], dfn[u], w[u]);
if (!son[u]) return;
dfs2(son[u], t);
int sz = g[u].size();
for (int i = 0; i < sz; ++i) {
int v = g[u][i];
if (v != son[u] && v != fa[u]) dfs2(v, v);
}
}
void addP(int x, int y, int k) {
while (top[x] != top[y]) {
if (dep[top[x]] < dep[top[y]]) swap(x, y);
add(dfn[top[x]], dfn[x], k);
x = fa[top[x]];
}
if (dep[x] > dep[y]) swap(x, y);
add(dfn[x], dfn[y], k);
}
int queryP(int x, int y) {
int ans = 0;
while (top[x] != top[y]) {
if (dep[top[x]] < dep[top[y]]) swap(x, y);
ans = (ans + query(dfn[top[x]], dfn[x])) % p;
x = fa[top[x]];
}
if (dep[x] > dep[y]) swap(x, y);
ans = (ans + query(dfn[x], dfn[y])) % p;
return ans;
}
void addS(int x, int k) {
k %= p;
add(dfn[x], dfn[x] + siz[x] - 1, k);
}
int queryS(int u) {
return query(dfn[u], dfn[u] + siz[u] - 1);
}
int main() {
scanf("%d %d %d %d", &n, &m, &r, &p);
for (int i = 1; i <= n; ++i) scanf("%d", &w[i]);
for (int i = 1, u, v; i < n; ++i) {
scanf("%d %d", &u, &v);
g[u].push_back(v);
g[v].push_back(u);
}
dfs1(r, 0);
dfs2(r, r);
//for (int i = 1; i <= n; ++i) printf("%d ", dfn[i]);
//puts("");
for (int i = 1; i <= n; ++i) printf("%d ", t1[i]);
puts("");
for (int i = 1; i <= n; ++i) printf("%d ", t2[i]);
for (int i = 1, opt, x, y, z; i <= m; ++i) {
scanf("%d", &opt);
if (opt == 1) {
scanf("%d %d %d", &x, &y, &z);
addP(x, y, z);
}
if (opt == 2) {
scanf("%d %d", &x, &y);
printf("%d\n", queryP(x, y) % p);
}
if (opt == 3) {
scanf("%d %d", &x, &z);
addS(x, z);
}
if (opt == 4) {
scanf("%d", &x);
printf("%d\n", queryS(x) % p);
}
}
return 0;
}