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
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
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
for (int i = 1; i <= n; i++) cin >> val[i];
vector
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
vector
vector
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
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
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;
}
```
::::