Problem C 星之所在 / Star Farming 题解

tumu1t

·

2026-05-01 15:22:55

·

题解

ECUST Programming Championship 2026 其他试题的官方题解参见此合集。

题意

对于带权值的树,求 q 条路径上满足出现次数大于 \frac{len}{k} 的所有权值。

## 解法

我们尝试维护每个结点到根节点的路径权值信息。我们需要查询的是:

- 对于每个结点 $u$ 到树上根结点 $root$ 的路径上在 $[l, r]$ 之间的所有权值的出现次数 $f(u, l, r)$。

这个是容易用可持久化线段树实现的,底层是权值线段树,统计权值的出现次数。

而这个出现次数是满足可差分的,因此我们可以用上面的路径维护 $(x, y)$ 路径上的信息:

- 维护结点 $x$ 到结点 $y$ 的路径上在 $[l, r]$ 之间的所有权值的出现次数 $F(x, y, l, r)$。

这个用树上差分实现,记 $\text{LCA}(x, y)$ 为 $x$ 和 $y$ 的最近公共祖先, $\text{F}(\text{LCA}(x, y))$ 为其父亲节点:

$$

F(X,Y, l, r) = f(x, l, r) + f(y, l, r) - f(\text{LCA}(x, y), l, r) - f(\text{F}(\text{LCA}(x, y)), l, r).

$$

现在我们将问题完全转化为了可持久化线段树上的问题。

首先可以证明:

- 满足条件的权值是至多 $k$ 个的。

- 因此我们从根节点开始遍历,由于线段树高度是 $\log n$ 的,因此我们必须经过的的路径长度是 $\mathcal{O}(k \log n)$ 的。

因此我们从根节点表示的区间 $[1, n]$ 开始:

- 如果左儿子树上结点满足 $cnt_L > \frac{|S|}{k}$ 说明这个地方可能权值是一个答案,往下递归。

- 否则则不递归左儿子。

- 右儿子亦然。

- 直到叶子结点位置。

我们这样子只走了上面所述有可能有答案的线段树上结点,记 $K = \sum k$, 复杂度可以做到 $\mathcal{O}(n \log n + q \log n + K \log n + K \log K)$.

::::success[STD Code(by pheonix_2002)]

```cpp

#pragma optimize 3

#include

using namespace std;

using ll = long long;

struct PST {

struct Info {

int cnt, sum;

Info operator+(const Info& x) { return {cnt + x.cnt, sum + x.sum}; }

Info operator-(const Info& x) { return {cnt - x.cnt, sum - x.sum}; }

};

struct Node {

Info info;

int ls = 0, rs = 0;

};

int min_ = 1, max_ = 1;

vector node;

PST() { node.reserve(1); node.emplace_back(); }

PST(int l, int r): min_(l), max_(r) { node.reserve(1); node.emplace_back(); }

private:

void add(int pos, const Info& val, int oldv, int l, int r) {

node.push_back(node[oldv]);

Node &cur = node.back();

cur.info.cnt = node[oldv].info.cnt + val.cnt;

cur.info.sum = node[oldv].info.sum + val.sum;

if (l == r) return;

int m = l + (r - l) / 2;

if (pos <= m) {

cur.ls = node.size();

add(pos, val, node[oldv].ls, l, m);

} else {

cur.rs = node.size();

add(pos, val, node[oldv].rs, m + 1, r);

}

}

Info query(int x, int y, int lv, int rv, int l, int r) {

if (x <= l && r <= y) return { node[rv].info.cnt - node[lv].info.cnt, node[rv].info.sum - node[lv].info.sum };

int m = l + (r - l) / 2;

if (y <= m) return query(x, y, node[lv].ls, node[rv].ls, l, m);

if (x > m) return query(x, y, node[lv].rs, node[rv].rs, m + 1, r);

Info a = query(x, y, node[lv].ls, node[rv].ls, l, m);

Info b = query(x, y, node[lv].rs, node[rv].rs, m + 1, r);

return { a.cnt + b.cnt, a.sum + b.sum };

}

public:

int add(int val, int old_version) {

int ver = (int)node.size();

add(val, {1, val}, old_version, min_, max_);

return ver;

}

Info query(int x, int y, int lv, int rv) { return query(x, y, lv, rv, min_, max_); }

// collect values in [L,R] whose count on path equals cnt(va)+cnt(vb)-cnt(vl)-cnt(vp) > thr

void collect_three(int va, int vb, int vl, int vp, int L, int R, int thr, vector& ans) {

Info a = query(L, R, 0, va, min_, max_);

Info b = query(L, R, 0, vb, min_, max_);

Info l = query(L, R, 0, vl, min_, max_);

Info p = query(L, R, 0, vp, min_, max_);

int total = a.cnt + b.cnt - l.cnt - p.cnt;

if (total <= thr) return;

if (L == R) { ans.push_back(L); return; }

int mid = (L + R) >> 1;

Info al = query(L, mid, 0, va, min_, max_);

Info bl = query(L, mid, 0, vb, min_, max_);

Info ll = query(L, mid, 0, vl, min_, max_);

Info pl = query(L, mid, 0, vp, min_, max_);

int cntLeft = al.cnt + bl.cnt - ll.cnt - pl.cnt;

if ((total - cntLeft) > thr) collect_three(va, vb, vl, vp, mid + 1, R, thr, ans);

if (cntLeft > thr) collect_three(va, vb, vl, vp, L, mid, thr, ans);

}

};

void solve() {

int n, q;

cin >> n >> q;

vector val(n + 1);

for (int i = 1; i <= n; i++) cin >> val[i];

vector> g(n + 1);

for (int i = 0; i < n - 1; i++) {

int u, v; cin >> u >> v;

g[u].push_back(v);

g[v].push_back(u);

}

vector f(n + 1, vector(20));

vector depth(n + 1);

vector version(n + 1);

PST pst(1, n);

auto dfs = [&](auto&& self, int u, int p)->void {

f[u][0] = p;

depth[u] = depth[p] + 1;

version[u] = pst.add(val[u], version[p]);

for (int v : g[u]) if (v != p) self(self, v, u);

};

dfs(dfs, 1, 0);

for (int j = 1; j < 20; j++) for (int i = 1; i <= n; i++) f[i][j] = f[ f[i][j-1] ][j-1];

auto lca = [&](int a, int b) {

if (depth[a] < depth[b]) swap(a, b);

int diff = depth[a] - depth[b];

for (int j = 0; j < 20; j++) if (diff >> j & 1) a = f[a][j];

if (a == b) return a;

for (int j = 19; j >= 0; --j) if (f[a][j] != f[b][j]) { a = f[a][j]; b = f[b][j]; }

return f[a][0];

};

auto collect_three = [&](auto&& self, int va, int vb, int vl, int vp, int Lrange, int Rrange, int thr, vector& ans)->void {

int cnt_va = pst.query(Lrange, Rrange, 0, va).cnt;

int cnt_vb = pst.query(Lrange, Rrange, 0, vb).cnt;

int cnt_vl = pst.query(Lrange, Rrange, 0, vl).cnt;

int cnt_vp = pst.query(Lrange, Rrange, 0, vp).cnt;

int total = cnt_va + cnt_vb - cnt_vl - cnt_vp;

if (total <= thr) return;

if (Lrange == Rrange) { ans.push_back(Lrange); return; }

int mid = (Lrange + Rrange) >> 1;

int cntLeft_va = pst.query(Lrange, mid, 0, va).cnt;

int cntLeft_vb = pst.query(Lrange, mid, 0, vb).cnt;

int cntLeft_vl = pst.query(Lrange, mid, 0, vl).cnt;

int cntLeft_vp = pst.query(Lrange, mid, 0, vp).cnt;

int cntLeft = cntLeft_va + cntLeft_vb - cntLeft_vl - cntLeft_vp;

if ((total - cntLeft) > thr) self(self, va, vb, vl, vp, mid + 1, Rrange, thr, ans);

if (cntLeft > thr) self(self, va, vb, vl, vp, Lrange, mid, thr, ans);

};

while (q--) {

int x, y, k; cin >> x >> y >> k;

int L = lca(x, y);

int P = f[L][0];

int S = depth[x] + depth[y] - 2 * depth[L] + 1;

int thr = S / k;

int va = version[x], vb = version[y];

int vl = version[L];

int vp = (P == 0 ? 0 : version[P]);

vector ans;

collect_three(collect_three, va, vb, vl, vp, 1, n, thr, ans);

if (ans.empty()) cout << -1 << '\n';

else {

sort(ans.begin(), ans.end());

for (size_t i = 0; i < ans.size(); i++) {

if (i) cout << ' ';

cout << ans[i];

}

cout << '\n';

}

}

}

int main() {

ios::sync_with_stdio(0);

cin.tie(0);

cout << fixed << setprecision(10);

int t = 1;

cin >> t;

while (t--) {

solve();

}

return 0;

}

```

::::