--- 这里是 cjiaw 的小窝(●'◡'●) ---

正在玩命加载中......

洛谷__P1505 [国家集训队] 旅游(树链剖分,边权转点权)

题目链接:P1505 [国家集训队] 旅游 - 洛谷


题目大意:

给定一棵  个节点的树,边带权,编号 ,需要支持五种操作:

  • C i w 将输入的第  条边权值改为 
  • N u v 将  节点之间的边权都变为相反数;
  • SUM u v 询问  节点之间边权和;
  • MAX u v 询问  节点之间边权最大值;
  • MIN u v 询问  节点之间边权最小值。

代码:

#include<iostream>
#include<algorithm>
#include<cstring>
#include<cstdlib>
#include<cmath>
#include<vector>
#include<queue>
#include<deque>
#include<stack>
#include<set>
#include<map>
#include<unordered_set>
#include<unordered_map>
#include<bitset>
#include<tuple>
#include<array>
#include <ext/pb_ds/assoc_container.hpp>
#include <ext/pb_ds/hash_policy.hpp>
#include <ext/numeric>
#define inf 72340172838076673
#define int long long
#define endl '\n'
#define F first
#define S second
#define mst(a,x) memset(a,x,sizeof (a))
#define gmap __gnu_pbds::gp_hash_table
#define power __gnu_cxx::power
using namespace std;
typedef pair<int, int> pii;

const int N = 400086, mod = 998244353;

int n, m;
int h[N], ne[N], e[N], w[N], idx;
int fa[N], son[N], dp[N], ww[N], sz[N];
int id[N], nw[N], top[N], cnt;
pii edge[N];

struct node {
    int l, r;
    int flag, mx, mn, sum;
} tr[N << 2];

inline void add(int a, int b, int c) {
    e[idx] = b, w[idx] = c, ne[idx] = h[a], h[a] = idx++;
}

void dfs1(int u, int father, int dep) {
    fa[u] = father, dp[u] = dep, sz[u] = 1;
    for (int i = h[u]; ~i; i = ne[i]) {
        int j = e[i];
        if (j == father) continue;
        dfs1(j, u, dep + 1);
        ww[j] = w[i];
        sz[u] += sz[j];
        if (sz[son[u]] < sz[j]) son[u] = j;
    }
}

void dfs2(int u, int t) {
    id[u] = ++cnt, nw[cnt] = ww[u], top[u] = t;
    if (!son[u]) return;
    dfs2(son[u], t);
    for (int i = h[u]; ~i; i = ne[i]) {
        int j = e[i];
        if (j == fa[u] || j == son[u]) continue;
        dfs2(j, j);
    }
}

inline void pushup(int u) {
    tr[u].sum = tr[u << 1].sum + tr[u << 1 | 1].sum;
    tr[u].mx = max(tr[u << 1].mx, tr[u << 1 | 1].mx);
    tr[u].mn = min(tr[u << 1].mn, tr[u << 1 | 1].mn);
}

inline void pushdown(int u) {
    auto &root = tr[u], &lf = tr[u << 1], &ri = tr[u << 1 | 1];
    if (root.flag) {
        lf.flag ^= 1, ri.flag ^= 1;
        lf.sum = -lf.sum, ri.sum = -ri.sum;
        swap(lf.mx, lf.mn), swap(ri.mx, ri.mn);
        lf.mx *= -1, lf.mn *= -1;
        ri.mx *= -1, ri.mn *= -1;
        root.flag = 0;
    }
}

void build(int u, int l, int r) {
    tr[u] = {l, r, 0, nw[l], nw[l], nw[l]};
    if (l == r) return;
    int mid = l + r >> 1;
    build(u << 1, l, mid);
    build(u << 1 | 1, mid + 1, r);
    pushup(u);
}

void modify(int u, int pos, int x) {
    if (tr[u].l == tr[u].r && tr[u].l == pos) {
        tr[u] = {pos, pos, 0, x, x, x};
        return;
    }
    pushdown(u);
    int mid = tr[u].r + tr[u].l >> 1;
    if (pos <= mid) modify(u << 1, pos, x);
    else modify(u << 1 | 1, pos, x);
    pushup(u);
}

void update(int u, int l, int r) {
    if (l <= tr[u].l && tr[u].r <= r) {
        tr[u].flag ^= 1;
        tr[u].sum *= -1;
        swap(tr[u].mx, tr[u].mn);
        tr[u].mx *= -1, tr[u].mn *= -1;
        return;
    }
    pushdown(u);
    int mid = tr[u].r + tr[u].l >> 1;
    if (l <= mid) update(u << 1, l, r);
    if (r > mid) update(u << 1 | 1, l, r);
    pushup(u);
}

int qsum(int u, int l, int r) {
    if (l <= tr[u].l && tr[u].r <= r) return tr[u].sum;
    pushdown(u);
    int mid = tr[u].l + tr[u].r >> 1;
    int res = 0;
    if (l <= mid) res = qsum(u << 1, l, r);
    if (r > mid) res += qsum(u << 1 | 1, l, r);
    return res;
}

int qmx(int u, int l, int r) {
    if (l <= tr[u].l && tr[u].r <= r) return tr[u].mx;
    pushdown(u);
    int mid = tr[u].l + tr[u].r >> 1;
    int res = -inf;
    if (l <= mid) res = qmx(u << 1, l, r);
    if (r > mid) res = max(res, qmx(u << 1 | 1, l, r));
    return res;
}

int qmn(int u, int l, int r) {
    if (l <= tr[u].l && tr[u].r <= r) return tr[u].mn;
    pushdown(u);
    int mid = tr[u].l + tr[u].r >> 1;
    int res = inf;
    if (l <= mid) res = qmn(u << 1, l, r);
    if (r > mid) res = min(res, qmn(u << 1 | 1, l, r));
    return res;
}

void update_path(int u, int v) {
    while (top[u] != top[v]) {
        if (dp[top[u]] < dp[top[v]]) swap(u, v);
        update(1, id[top[u]], id[u]);
        u = fa[top[u]];
    }
    if (dp[u] < dp[v]) swap(u, v);
    update(1, id[v] + 1, id[u]);
}

int query_sum(int u, int v) {
    int res = 0;
    while (top[u] != top[v]) {
        if (dp[top[u]] < dp[top[v]]) swap(u, v);
        res += qsum(1, id[top[u]], id[u]);
        u = fa[top[u]];
    }
    if (dp[u] < dp[v]) swap(u, v);
    res += qsum(1, id[v] + 1, id[u]);
    return res;
}

int query_mn(int u, int v) {
    int res = inf;
    while (top[u] != top[v]) {
        if (dp[top[u]] < dp[top[v]]) swap(u, v);
        res = min(res, qmn(1, id[top[u]], id[u]));
        u = fa[top[u]];
    }
    if (dp[u] < dp[v]) swap(u, v);
    res = min(res, qmn(1, id[v] + 1, id[u]));
    return res;
}

int query_mx(int u, int v) {
    int res = -inf;
    while (top[u] != top[v]) {
        if (dp[top[u]] < dp[top[v]]) swap(u, v);
        res = max(res, qmx(1, id[top[u]], id[u]));
        u = fa[top[u]];
    }
    if (dp[u] < dp[v]) swap(u, v);
    res = max(res, qmx(1, id[v] + 1, id[u]));
    return res;
}

void solve() {

    mst(h, -1);
    cin >> n;
    for (int i = 1; i < n; i++) {
        int a, b, c;
        cin >> a >> b >> c;
        a++, b++;
        add(a, b, c), add(b, a, c);
        edge[i] = {a, b};
    }
    
    dfs1(1, -1, 1);
    dfs2(1, 1);
    build(1, 1, n);
    
    string s;
    int a, b;
    cin >> m;
    
    while (m--) {
        cin >> s >> a >> b;
        if (s == "C") {
            auto [u, v] = edge[a];
            if (dp[u] < dp[v]) swap(u, v);
            modify(1, id[u], b);
        } else {
            a++, b++;
            if (s == "SUM") cout << query_sum(a, b) << endl;
            else if (s == "MAX") cout << query_mx(a, b) << endl;
            else if (s == "MIN") cout << query_mn(a, b) << endl;
            else update_path(a, b);
        }
    }
    
}

signed main() {
    ios::sync_with_stdio(false);
    cin.tie(nullptr), cout.tie(nullptr);
    
    int T = 1;
// cin >> T;
    while (T--) solve();
    
    return 0;
}

 

posted @ 2025-12-23 17:19  wwjjw  阅读(12)  评论(0)    收藏  举报