wa求助,对拍过了,全wa
查看原帖
wa求助,对拍过了,全wa
401583
N2MENT楼主2022/4/8 21:02
//P2486
#include <bits/stdc++.h>
using namespace std;
const int maxn = 1e5 + 5;
#define int long long
struct range {
    int l, r, sum;

    range(int l, int r, int sum) {
        this->l = l;
        this->r = r;
        this->sum = sum;
        return;
    }

    range() {
        this->l = 0;
        this->r = 0;
        this->sum = 0;
        return;
    }
};

range mg(range a, range b) {
    if (a.sum == 0)return b;
    if (b.sum == 0)return a;
    range res(a.l, b.r, a.sum + b.sum);
    res.sum -= (a.r == b.l);
    return res;
}

class SGT {
public:
    range tr[maxn << 2];
    int lazy[maxn << 2];
    SGT() {
        memset(tr, 0, sizeof(tr));
        memset(lazy, 0, sizeof(lazy));
        return;
    }

    void fupd(int k) {
        tr[k] = mg(tr[k << 1], tr[k << 1 | 1]);
        return;
    }

    void lupd(int k) {
        if (lazy[k]) {
            tr[k << 1] = tr[k << 1 | 1] = range(lazy[k], lazy[k], 1);
            lazy[k] = 0;
        }
        return;
    }

    void build(int k, int l, int r, int *data) {
        if (l == r) {
            tr[k] = range(data[l], data[l], 1);
            return;
        }
        int mid = (l + r) >> 1;
        build(k << 1, l, mid, data);
        build(k << 1 | 1, mid + 1, r, data);
        fupd(k);
        return;
    }

    void update(int k, int l, int r, int L, int R, int color) {
        if (L <= l && r <= R) {
            tr[k] = range(color, color, 1);
            lazy[k] = color;
            return;
        }
        lupd(k);
        int mid = (l + r) >> 1;
        if (L <= mid)
            update(k << 1, l, mid, L, R, color);
        if (R > mid)
            update(k << 1 | 1, mid + 1, r, L, R, color);
        fupd(k);
        return;
    }

    range query(int k, int l, int r, int L, int R) {
        if (L <= l && r <= R)
            return tr[k];
        lupd(k);
        range res;
        int mid = (l + r) >> 1;
        if (L <= mid)
            res = mg(res, query(k << 1, l, mid, L, R));
        if (R > mid)
            res = mg(res, query(k << 1 | 1, mid + 1, r, L, R));
        return res;
    }
};

SGT s;
vector<int> G[maxn];
int n, m, cnt;
int w[maxn];
int fa[maxn], dep[maxn], sz[maxn], son[maxn], top[maxn], id[maxn], base[maxn];

void dfs(int u) {
    dep[u] = dep[fa[u]] + 1;
    sz[u] = 1;
    for (auto v: G[u]) {
        if (v == fa[u])
            continue;
        fa[v] = u;
        dfs(v);
        sz[u] += sz[v];
        if (!son[u] || sz[v] > sz[son[u]])
            son[u] = v;
    }
    return;
}

void dfs2(int u, int tp) {
    top[u] = tp;
    id[u] = ++cnt;
    base[cnt] = w[u];
    if (son[u])dfs2(son[u], tp);
    for (auto v: G[u]) {
        if (v == fa[u] || v == son[u])
            continue;
        dfs2(v, v);
    }
    return;
}

void update(int x, int y, int c) {
    while (top[x] != top[y]) {
        if (dep[top[x]] > dep[top[y]]) {
            s.update(1, 1, n, id[top[x]], id[x], c);
            x = fa[top[x]];
        } else {
            s.update(1, 1, n, id[top[y]], id[y], c);
            y = fa[top[y]];
        }
    }
    s.update(1, 1, n, min(id[x], id[y]), max(id[x], id[y]), c);
    return;
}

int query(int x, int y) {
    range res1, res2;
    while (top[x] != top[y]) {
        if (dep[top[x]] > dep[top[y]]) {
            res1 = mg(s.query(1, 1, n, id[top[x]], id[x]), res1);
            x = fa[top[x]];
        } else {
            res2 = mg(s.query(1, 1, n, id[top[y]], id[y]), res2);
            y = fa[top[y]];
        }
    }
    range mid = s.query(1, 1, n, min(id[x], id[y]), max(id[x], id[y]));
    if (dep[x] < dep[y]) {
        return res1.sum + res2.sum + mid.sum - (mid.l == res1.l) - (mid.r == res2.r);
    } else {
        return res1.sum + res2.sum + mid.sum - (mid.r == res1.l) - (mid.l == res2.r);
    }

}

signed main() {
    scanf("%lld%lld", &n, &m);
    for (int i = 1; i <= n; i++) {
        scanf("%lld", w + i);
    }
    for (int i = 1; i < n; i++) {
        int a, b;
        scanf("%lld%lld", &a, &b);
        G[a].push_back(b);
        G[b].push_back(a);
    }
    dfs(1);
    dfs2(1, 1);
    s.build(1, 1, n, base);
    for (int i = 1; i <= m; i++) {
        char in[10];
        int a, b, c;
        scanf("%s",in);
        scanf("%lld%lld", &a, &b);
        switch (in[0]) {
            case 'C':
                scanf("%lld", &c);
                update(a, b, c);
                break;
            case 'Q':
                printf("%lld\n", query(a, b));
                break;
        }
    }
}
2022/4/8 21:02
加载中...