1 条题解

  • 0
    @ 2026-1-31 21:28:13
    #include <bits/stdc++.h>
    using namespace std;
    #define lc (p << 1)
    #define rc (p << 1 | 1)
    #define mid (tr[p].l + tr[p].r) / 2
    const int N = 1e6 + 10;
    vector<pair<int, int>> G[N];
    int fa[N], son[N], dep[N], f[N][20], D, siz[N], a[N];
    void dfs1(int x, int ff)
    {
      fa[x] = ff;
      dep[x] = dep[ff] + 1;
      siz[x] = 1;
      f[x][0] = ff;
      for (int i = 1; i <= D; i++)
        f[x][i] = f[f[x][i - 1]][i - 1];
      for (auto i : G[x])
      {
        int y = i.first, c = i.second;
        if (y == ff)
          continue;
        dfs1(y, x);
        a[y] = c;
        siz[x] += siz[y];
        if (siz[son[x]] < siz[y])
          son[x] = y;
      }
    }
    
    int tsp, dfn[N], top[N], ys[N];
    void dfs2(int x, int tp)
    {
      dfn[x] = ++tsp;
      ys[tsp] = x;
      top[x] = tp;
      if (son[x] != 0)
        dfs2(son[x], tp);
      for (auto i : G[x])
      {
        int y = i.first, c = i.second;
        if (y != fa[x] && y != son[x])
          dfs2(y, y);
      }
    }
    
    struct trnode
    {
      int l, r, c;
    } tr[N << 2];
    void upd(int p) { tr[p].c = tr[lc].c + tr[rc].c; }
    void bt(int p, int l, int r)
    {
      tr[p] = {l, r, 0};
      if (l == r)
        tr[p].c = a[ys[l]];
      else
      {
        bt(lc, l, mid);
        bt(rc, mid + 1, r);
        upd(p);
      }
    }
    int query(int p, int l, int r)
    {
      if (l <= tr[p].l && tr[p].r <= r)
        return tr[p].c;
      int res = 0;
      if (l <= mid)
        res += query(lc, l, r);
      if (mid < r)
        res += query(rc, l, r);
      return res;
    }
    int solve(int x, int y)
    {
      int res = 0;
      while (top[x] != top[y])
      {
        if (dep[top[x]] < dep[top[y]])
          swap(x, y);
        res += query(1, dfn[top[x]], dfn[x]);
        x = fa[top[x]];
      }
      if (dep[x] > dep[y])
        swap(x, y);
      if (x != y)
        res += query(1, dfn[x] + 1, dfn[y]);
      return res;
    }
    int lca(int x, int y)
    {
      if (dep[x] < dep[y])
        swap(x, y);
      for (int i = D; i >= 0; i--)
        if (dep[f[x][i]] >= dep[y])
          x = f[x][i];
      if (x == y)
        return x;
      for (int i = D; i >= 0; i--)
        if (f[x][i] != f[y][i])
          x = f[x][i], y = f[y][i];
      return f[x][0];
    }
    int path(int x, int y, int k)
    {
      int p = lca(x, y);
      if (k <= dep[x] - dep[p] + 1)
      {
        k--;
        for (int i = D; i >= 0; i--)
          if (k >= (1 << i))
            x = f[x][i], k -= (1 << i);
        return x;
      }
      else
      {
        k = (dep[x] + dep[y] - 2 * dep[p] + 1) - k + 1;
        k--;
        for (int i = D; i >= 0; i--)
          if (k >= (1 << i))
            y = f[y][i], k -= (1 << i);
        return y;
      }
    }
    int main()
    {
      int T;
      scanf("%d", &T);
      while (T--)
      {
        int n;
        scanf("%d", &n);
        memset(G, 0, sizeof(G));
        for (int i = 1, x, y, c; i < n; i++)
        {
          scanf("%d%d%d", &x, &y, &c);
          G[x].push_back({y, c});
          G[y].push_back({x, c});
        }
        fa[0] = dep[0] = siz[0] = 0;
        memset(son, 0, sizeof(son));
        a[1] = 0;
        D = log2(n);
        dfs1(1, 0);
        tsp = 0;
        dfs2(1, 1);
        bt(1, 1, tsp);
        char s[10];
        while (scanf("%s", s) != EOF && s[1] != 'O')
        {
          if (s[0] == 'D')
          {
            int x, y;
            scanf("%d%d", &x, &y);
            printf("%d\n", solve(x, y));
          }
          else
          {
            int x, y, k;
            scanf("%d%d%d", &x, &y, &k);
            printf("%d\n", path(x, y, k));
          }
        }
      }
      return 0;
    }
    
    • 1

    信息

    ID
    546
    时间
    1000ms
    内存
    2048MiB
    难度
    8
    标签
    递交数
    76
    已通过
    11
    上传者