把博客园图标替换成自己的图标
把博客园图标替换成自己的图标end

P3242 [HNOI2015] 接水果 分析

题目概述

  • 给定一棵 \(n\) 个节点的树。
  • \(p\) 个盘子,每个盘子是一条路径 \((a_i,b_i)\),权值为 \(c_i\)
  • \(q\) 个水果,每个水果也是一条路径 \((u_i,v_i)\),并给定一个整数 \(k_i\)
  • 一个盘子能接住一个水果,当且仅当盘子的路径是水果路径的子路径(即盘子的两个端点都在水果路径上,且顺序一致)。
  • 每个水果需要找到所有能接住它的盘子中,权值第 \(k_i\) 小的盘子的权值(盘子可重复使用)。
  • 数据范围:\(n,p,q \le 4\times 10^4\),权值 \(c \le 10^9\)

分析

1. 转化为矩形覆盖问题(很好用的trick)

利用树上的 DFS 序,可以将每条路径映射为二维平面上的点或区间。

  • 对于每个节点 \(x\),记录其 DFS 序 \(tin[x]\) 和子树结束时间 \(tout[x]\)
    一个水果路径 \(u \to v\)(令 \(tin[u] \le tin[v]\))可被表示为二维点 \((tin[u], tin[v])\)
  • 判断一个盘子路径 \(a \to b\) 是否为水果路径的子路径,等价于水果路径的两个端点分别落在盘子路径所对应的某些区域。
    经典的转化:
    • \(lca(a,b) \ne a\)\(lca(a,b) \ne b\)(即盘子路径弯曲),则水果路径必须满足:一端在 \(a\) 的子树,另一端在 \(b\) 的子树。
      对应二维矩形:\(x \in [tin[a], tout[a]]\)\(y \in [tin[b], tout[b]]\)
    • \(a\)\(b\) 的祖先(反之对称),设 \(a\) 为祖先,\(b\) 为后代,令 \(c\)\(a\) 的儿子且是 \(b\) 的祖先(即 \(a\)\(b\) 路径上的第一个节点)。
      则水果路径必须满足:一端在 \(b\) 的子树,另一端不在 \(c\) 的子树内。
      也就是两个矩形:
      • \(x \in [tin[b], tout[b]]\)\(y \in [1, tin[c]-1]\)
      • \(x \in [tin[b], tout[b]]\)\(y \in [tout[c]+1, n]\)
        (这里要求 \(x \le y\),但实际实现中会保证方向)

因此每个盘子转化为若干个矩形(实际上最多两个),每个矩形有一个权值索引(压缩后的权值)。
每个水果则对应一个点 \((tin[u], tin[v])\)(排序保证 \(tin[u] \le tin[v]\))。
问题变成:对于每个点,求覆盖它的所有矩形中权值第 \(k\) 小的权值。

2. 整体二分 + 二维数点

我们要求每个点被多少个权值不超过 \(mid\) 的矩形覆盖,从而进行整体二分。

  • 将矩形拆成两个差分事件(在 \(x = x_1\)\(+1\)\(x = x_2+1\)\(-1\)),每个事件在 \(y\) 区间 \([y_1, y_2]\) 上做区间加。
  • 处理询问点时,扫描 \(x\) 坐标,用树状数组维护当前 \(y\) 上的累计覆盖数,单点查询得到该点被覆盖的矩形数。
  • 整体二分递归:
    • 对当前权值区间 \([L, R]\),将矩形按其权值索引分配到左半或右半。
    • 只加入权值索引 \(\le mid\) 的矩形到二维数据结构,计算每个询问点被加入的矩形覆盖的个数,与 \(k\) 比较,决定该询问进入左子区间或右子区间(同时减去左半贡献)。
    • 递归处理。

3. 复杂度

  • 预处理:DFS、LCA、子树信息 \(O(n \log n)\)
  • 每个盘子转化为 \(O(1)\) 个矩形。
  • 整体二分过程中,每一层都会扫描所有矩形和所有询问,并用树状数组维护,总复杂度 \(O((p+q) \log p \log n)\)
    因为 \(p,q\) 均为 4e4,可行。

代码

#include <iostream>
#include <algorithm>
#include <cstring>
#include <cstdio>
#include <stdlib.h>
#include <vector>
#include <numeric>
#define int long long
#define N 80005
#define M 20
using namespace std;
struct Fenwick {
	int n;
	int tr[N];
	Fenwick() {};
	Fenwick(int n_) {init(n_);}
	void init(int n_) {n = n_;}
	void update(int x,int val) {
		if (x <= 0 || x > n) return;
		for (;x <= n;x += x & -x) tr[x] += val;
	}
	void update(int L,int R,int val) {
		update(L,val),update(R + 1,-val);
	}
	int query(int x) {
		int res = 0;
		for (;x;x -= x & -x) res += tr[x];
		return res;
	}
}t;
struct pla{
	int a,b,c;
}plate[N];
struct rect{
	int x1,x2,y1,y2,w;
}rects[N];
struct line{
	int x,y1,y2,w;
};
struct query_point{
	int x,id;
};
int cnt;
int n,p,q;
vector<int> g[N],ls;
int tin[N],tout[N],dfn_time = 0,fa[N][M],dep[N];
void dfs(int cur,int father) {
	fa[cur][0] = father,dep[cur] = dep[father] + 1;
	tin[cur] = ++dfn_time;
	for (auto i : g[cur])
		if (i != father) dfs(i,cur);
	tout[cur] = dfn_time;
}
bool is_ancestor(int x,int y) {
	return tin[x] <= tin[y] && tout[x] >= tin[y];
}
int LCA(int x,int y) {
	if (is_ancestor(x,y)) return x;
	if (is_ancestor(y,x)) return y;
	for (int j = 18;j >= 0;j --)
		if (!is_ancestor(fa[x][j],y)) x = fa[x][j];
	return fa[x][0];
}
int get_child(int x,int y) {
	//x is y's ancestor
	int p = y;
	for (int j = 18;j >= 0;j --)
		if (dep[fa[p][j]] > dep[x]) p = fa[p][j];
	return p;
}
void add(int x1,int x2,int y1,int y2,int w) {
	if (x1 > x2 || y1 > y2) return;
	rects[++cnt] = {x1,x2,y1,y2,w};
}
int ans[N],cntarr[N],qx[N],qy[N],kth[N];
void solve(int l,int r,vector<int> rectIds,vector<int> queryIds) {
	if (queryIds.empty()) return;
	if (l == r) {
		for (auto id : queryIds) ans[id] = ls[l];
		return;
	}
	int mid = l + r >> 1;
	vector<int> L,R;
	vector<line> events;
	for (auto id : rectIds)
		if (rects[id].w <= mid) {
			L.push_back(id);
			events.push_back({rects[id].x1,rects[id].y1,rects[id].y2,1});
			events.push_back({rects[id].x2 + 1,rects[id].y1,rects[id].y2,-1});
		}
		else R.push_back(id);
	if (!events.empty()) {
		sort(events.begin(),events.end(),[](const line&x,const line&y) {
			return x.x < y.x;
		});
		vector<query_point> qs;
		qs.reserve(queryIds.size());
		for (auto id : queryIds) qs.push_back({qx[id],id});
		sort(qs.begin(),qs.end(),[](const query_point& x,const query_point& y) {
			return x.x < y.x;
		});
		int j = 0;
		for (auto que : qs) {
			int x = que.x,id = que.id;
			while(j < (int)events.size() && events[j].x <= x) {
				t.update(events[j].y1,events[j].y2,events[j].w);
				j ++;
			}
			cntarr[id] = t.query(qy[id]);
		}
		while(j < (int)events.size()) t.update(events[j].y1,events[j].y2,events[j].w),j ++;
		for (auto i : events) t.update(i.y1,i.y2,-i.w);
	} 
	else for (auto i : queryIds) cntarr[i] = 0;
	vector<int> LQ,RQ;
	for (auto id : queryIds) {
		if (kth[id] <= cntarr[id]) LQ.push_back(id);
		else RQ.push_back(id),kth[id] -= cntarr[id];
	}
	if (!LQ.empty()) solve(l,mid,move(L),move(LQ));
	if (!RQ.empty()) solve(mid + 1,r,move(R),move(RQ));
} 
signed main(){
	cin >> n >> p >> q;
	for (int i = 1;i < n;i ++) {
		int u,v;
		scanf("%lld%lld",&u,&v);
		g[u].push_back(v),g[v].push_back(u);
	}
	dfs(1,0);
	for (int j = 1;j <= 18;j ++)
		for (int i = 1;i <= n;i ++)
			fa[i][j] = fa[fa[i][j - 1]][j - 1];
	for (int i = 1;i <= p;i ++) {
		scanf("%lld%lld%lld",&plate[i].a,&plate[i].b,&plate[i].c);
		ls.push_back(plate[i].c);
	}
	sort(ls.begin(),ls.end());
	ls.erase(unique(ls.begin(),ls.end()),ls.end());
	for (int i = 1;i <= q;i ++) {
		int u,v,k;
		scanf("%lld%lld%lld",&u,&v,&k);
		qx[i] = min(tin[u],tin[v]);
		qy[i] = max(tin[u],tin[v]);
		kth[i] = k;
	}
	for (int i = 1;i <= p;i ++) {
		int a = plate[i].a,b = plate[i].b,c = plate[i].c;
		int w = lower_bound(ls.begin(),ls.end(),c) - ls.begin();
		int t = LCA(a,b);
		if (t == a || t == b) {
			int u = b,v = a;
			if (is_ancestor(a,b)) u = a,v = b;
			int child = get_child(u,v);
			add(1,tin[child] - 1,tin[v],tout[v],w);
			add(tin[v],tout[v],tout[child] + 1,n,w);//较小的要在前面 
		}
		else {
			if (tin[a] > tin[b]) a ^= b ^= a ^= b;
			add(tin[a],tout[a],tin[b],tout[b],w);
		}
	}
	t.init(n + 2);
	vector<int> allRects(cnt);
    iota(allRects.begin(),allRects.end(),1);
    vector<int> allQueries(q);
    iota(allQueries.begin(),allQueries.end(),1);
    solve(0,ls.size() - 1,move(allRects),move(allQueries));
    for (int i = 1;i <= q;i ++) cout << ans[i] << '\n';
	return 0;
}

代码讲解(以下是AI)

代码整体结构清晰,我们分模块解读。

dfs() 用递归方式实现 DFS,得到每个节点的 \(tin,tout\),以及 fa、dep。
之后构建倍增数组 up,用于 LCA 和找祖先的儿子。

矩形生成

add_rect 函数负责将矩形加入 rects 向量(保存四个边界和权值索引)。

对于每个盘子 (a,b,c)

  • 权值压缩:把权值离散化到 vals 中,得到 w
  • 求 LCA l = lca(a,b)
  • 分两种情况:
    1. 弯折路径\(l \ne a\)\(l \ne b\)
      则要求在 a 子树和 b 子树各取一点,矩形为 [tin[a], tout[a]] × [tin[b], tout[b]](保证第一个坐标小于第二个,否则交换)。
    2. 直链:假设 \(a\)\(b\) 的祖先。
      找出从 \(a\)\(b\) 路径上的第一个节点 child = get_child(a, b)
      水果必须一端在 b 子树,另一端在 child 子树之外,即两个互补矩形:
      • [tin[b], tout[b]] × [1, tin[child]-1]
      • [tin[b], tout[b]] × [tout[child]+1, n]
        同样保证第一维 ≤ 第二维。

所有矩形存入 rects,每个矩形包含 x1,x2,y1,y2,w

询问处理

水果路径 (u,v),我们令 x = min(tin[u], tin[v]), y = max(tin[u], tin[v]),并记录 kth[id] = k。所有询问点存储在 qx, qy 中。

整体二分函数 solve(l, r, rectIds, queryIds)

  • l == r,则这些询问的答案就是 vals[l]
  • 否则取 mid = (l+r)/2
  • rectIds 分为 leftRects(权值索引 ≤ mid)和 rightRects(> mid)。
  • leftRects 的矩形转化为扫描线事件(每个矩形变成两个事件:在 x=x1 处加 1,在 x=x2+1 处减 1,区间为 [y1,y2])。
  • 对询问点按 x 排序,扫描事件,用树状数组区间加,单点查询每个询问点的覆盖数,存入 cntArr[id]
  • 根据 kth[id]cntArr[id] 的关系分流:若 k ≤ cnt,进入左区间;否则进入右区间,并将 k 减去左区间贡献。
  • 递归处理两个子区间。

注意:为了保证树状数组正确重置,扫描完所有询问后需要将剩余事件也处理完(代码中通过 while 循环把剩余事件都 apply 完,但树状数组并未显式清除,实际上由于整体二分每次都是新建事件向量,且树状数组最终会累积所有事件,但我们在下一次递归前并不会清空树状数组,而是直接开始下一层。然而代码在每次递归调用 solve 时都会创建新的 events 并执行扫描,且扫描前没有清空树状数组,这是一个 bug。在标准写法中,应该在处理完当前层后,把添加的所有差分还原(即执行相反操作)。但这里代码中在扫描结束后,并没有将树状数组回退,而是继续执行。不过由于每次调用 solve 都会新建 events,但树状数组里的值是累积的,会导致错误。实际上,标准的整体二分需要确保每次检查时,树状数组初始为空。
但观察代码:在 solve 中,events 可能为空,此时不会使用树状数组。若不为空,扫描完所有询问后,紧接着又有一个 while 循环将剩余事件也 apply 了,但之后并没有再清除。而递归调用时,树状数组里还保存着之前的累加值,这是不对的。
一个正确的做法是:在扫描结束后,把所有的差分操作再逆序做一遍,或者干脆用一个临时数组记录每次修改的位置并恢复。但这里代码没有做。
然而,在 OJ 上该代码可能依然能通过,因为 solve 是递归的,且 events 中的事件只在当前层有效,但树状数组的状态会带入子调用,导致子调用中的查询基于错误的基础值。
合理的解释:可能作者认为 solve 不会同时有两层活跃,因为它是先处理完当前层,再递归左右,而树状数组在返回后不再使用,但左递归和右递归都会使用同一个树状数组,且左递归会修改它,然后返回,接着右递归又会在被污染的基础上继续,显然错误。
因此,正确的写法应该是在每次使用树状数组前清空,或者在扫描结束时撤销所有修改。
但鉴于用户只要求讲解代码,我们暂且按照给出的代码讲解,并指出这个潜在问题。实际上,很多 AC 代码会用一个 vector 记录修改的位置,然后逐个恢复。

不过我们看代码中,在 solve 的末尾,递归调用 solve(l, mid, move(leftRects), move(leftQ));solve(mid+1, r, move(rightRects), move(rightQ)); 这两行,它们会继续使用 fw,但 fw 中仍残留数据。所以代码是有问题的。
但是,我们注意到 solve 里的 events 可能为空(当左半没有矩形时),此时不会修改树状数组,但如果左半有矩形,那么树状数组被修改,之后右半递归会看到错误的值。
所以,我们需要修正代码:在 solve 中,对于 events,在计算完所有点的覆盖数后,应该将树状数组恢复到原始状态。通常做法是,记录所有修改过的位置和值,然后逐个减去。

我们在此讲解时,可以指出这一点,并提供修复建议。但题目要求讲解“这个代码”,所以我们可以指出代码逻辑,并说明正确性依赖于在递归前恢复树状数组,而该代码缺失了这一步,是一个 bug。

主函数流程

  • 读入 n,p,q。
  • 建树,DFS,LCA预处理。
  • 读入盘子,压缩权值。
  • 读入水果,保存询问。
  • 生成所有矩形。
  • 初始化树状数组。
  • 创建所有矩形 id 和询问 id 的 vector,调用 solve(0, vals.size()-1, allRects, allQueries)
  • 输出答案。

AI代码

#include <bits/stdc++.h>
using namespace std;

const int LOG = 18;

struct Fenwick {
    int n;
    vector<int> bit;
    Fenwick() {}
    Fenwick(int n_) { init(n_); }
    void init(int n_) {
        n = n_;
        bit.assign(n + 2, 0);
    }
    void add(int pos, int val) {
        if (pos <= 0 || pos > n) return;
        for (; pos <= n; pos += pos & -pos) bit[pos] += val;
    }
    void range_add(int l, int r, int val) {
        add(l, val);
        add(r + 1, -val);
    }
    int query(int pos) {
        int res = 0;
        for (; pos > 0; pos -= pos & -pos) res += bit[pos];
        return res;
    }
};

struct PlateRaw {
    int a, b;
    long long c;
};

struct Rect {
    int x1, x2, y1, y2;
    int w; // compressed weight index
};

struct Event {
    int x, y1, y2, delta;
};

struct QueryPoint {
    int x, id;
};

int n, p, q;
vector<vector<int>> adj;
vector<int> tin, tout, parent, depth;
vector<vector<int>> up;
vector<PlateRaw> plates;
vector<Rect> rects;
vector<long long> vals;
vector<int> qx, qy, kth, ans, cntArr;
Fenwick fw;

int timer_ = 0;

bool is_ancestor(int u, int v) {
    return tin[u] <= tin[v] && tin[v] <= tout[u];
}

int lca(int u, int v) {
    if (is_ancestor(u, v)) return u;
    if (is_ancestor(v, u)) return v;
    for (int i = LOG - 1; i >= 0; --i) {
        if (!is_ancestor(up[i][u], v)) u = up[i][u];
    }
    return parent[u];
}

int get_child(int anc, int desc) {
    // anc is strict ancestor of desc
    int x = desc;
    for (int i = LOG - 1; i >= 0; --i) {
        if (depth[up[i][x]] > depth[anc]) x = up[i][x];
    }
    return x;
}

void add_rect(int x1, int x2, int y1, int y2, int w) {
    if (x1 > x2 || y1 > y2) return;
    rects.push_back({x1, x2, y1, y2, w});
}

void dfs_iter() {
    tin.assign(n + 1, 0);
    tout.assign(n + 1, 0);
    parent.assign(n + 1, 1);
    depth.assign(n + 1, 1);

    vector<int> it(n + 1, 0), st;
    st.push_back(1);
    parent[1] = 1;
    depth[1] = 1;

    while (!st.empty()) {
        int u = st.back();
        if (it[u] == 0) {
            tin[u] = ++timer_;
        }
        if (it[u] == (int)adj[u].size()) {
            tout[u] = timer_;
            st.pop_back();
            continue;
        }
        int v = adj[u][it[u]++];
        if (v == parent[u]) continue;
        parent[v] = u;
        depth[v] = depth[u] + 1;
        st.push_back(v);
    }
}

void solve(int l, int r, vector<int> rectIds, vector<int> queryIds) {
    if (queryIds.empty()) return;

    if (l == r) {
        for (int id : queryIds) ans[id] = vals[l];
        return;
    }

    int mid = (l + r) >> 1;

    vector<int> leftRects, rightRects;
    vector<Event> events;

    for (int id : rectIds) {
        if (rects[id].w <= mid) {
            leftRects.push_back(id);
            events.push_back({rects[id].x1, rects[id].y1, rects[id].y2, 1});
            events.push_back({rects[id].x2 + 1, rects[id].y1, rects[id].y2, -1});
        } else {
            rightRects.push_back(id);
        }
    }

    if (!events.empty()) {
        sort(events.begin(), events.end(), [](const Event &a, const Event &b) {
            return a.x < b.x;
        });

        vector<QueryPoint> qs;
        qs.reserve(queryIds.size());
        for (int id : queryIds) qs.push_back({qx[id], id});

        sort(qs.begin(), qs.end(), [](const QueryPoint &a, const QueryPoint &b) {
            return a.x < b.x;
        });

        int e = 0;
        for (const auto &qi : qs) {
            int x = qi.x, id = qi.id;
            while (e < (int)events.size() && events[e].x <= x) {
                fw.range_add(events[e].y1, events[e].y2, events[e].delta);
                ++e;
            }
            cntArr[id] = fw.query(qy[id]);
        }

        // 处理剩余事件(若有点未扫描到,但通常已经处理完,不过还是补全)
        while (e < (int)events.size()) {
            fw.range_add(events[e].y1, events[e].y2, events[e].delta);
            ++e;
        }

        // ========== 关键修复:撤销所有事件,清空树状数组 ==========
        for (const auto &ev : events) {
            fw.range_add(ev.y1, ev.y2, -ev.delta);
        }
        // ========================================================
    } else {
        for (int id : queryIds) cntArr[id] = 0;
    }

    vector<int> leftQ, rightQ;
    for (int id : queryIds) {
        if (kth[id] <= cntArr[id]) {
            leftQ.push_back(id);
        } else {
            kth[id] -= cntArr[id];
            rightQ.push_back(id);
        }
    }

    if (!leftQ.empty()) solve(l, mid, move(leftRects), move(leftQ));
    if (!rightQ.empty()) solve(mid + 1, r, move(rightRects), move(rightQ));
}

int main() {
    ios::sync_with_stdio(false);
    cin.tie(nullptr);

    cin >> n >> p >> q;

    adj.assign(n + 1, {});
    for (int i = 0; i < n - 1; ++i) {
        int a, b;
        cin >> a >> b;
        adj[a].push_back(b);
        adj[b].push_back(a);
    }

    dfs_iter();

    up.assign(LOG, vector<int>(n + 1));
    for (int i = 1; i <= n; ++i) up[0][i] = parent[i];
    for (int j = 1; j < LOG; ++j) {
        for (int i = 1; i <= n; ++i) {
            up[j][i] = up[j - 1][up[j - 1][i]];
        }
    }

    plates.resize(p);
    for (int i = 0; i < p; ++i) {
        cin >> plates[i].a >> plates[i].b >> plates[i].c;
        vals.push_back(plates[i].c);
    }

    sort(vals.begin(), vals.end());
    vals.erase(unique(vals.begin(), vals.end()), vals.end());

    qx.resize(q);
    qy.resize(q);
    kth.resize(q);
    ans.resize(q, 0);
    cntArr.resize(q, 0);

    for (int i = 0; i < q; ++i) {
        int u, v, k;
        cin >> u >> v >> k;
        qx[i] = min(tin[u], tin[v]);
        qy[i] = max(tin[u], tin[v]);
        kth[i] = k;
    }

    // Transform every plate into rectangles.
    for (const auto &plate : plates) {
        int a = plate.a, b = plate.b;
        int w = lower_bound(vals.begin(), vals.end(), plate.c) - vals.begin();

        int l = lca(a, b);

        if (l == a || l == b) {
            int u, v;
            if (is_ancestor(a, b)) {
                u = a;
                v = b;
            } else {
                u = b;
                v = a;
            }

            int child = get_child(u, v);

            add_rect(1, tin[child] - 1, tin[v], tout[v], w);
            add_rect(tin[v], tout[v], tout[child] + 1, n, w);
        } else {
            if (tin[a] < tin[b]) {
                add_rect(tin[a], tout[a], tin[b], tout[b], w);
            } else {
                add_rect(tin[b], tout[b], tin[a], tout[a], w);
            }
        }
    }

    fw.init(n + 2);

    vector<int> allRects(rects.size());
    iota(allRects.begin(), allRects.end(), 0);

    vector<int> allQueries(q);
    iota(allQueries.begin(), allQueries.end(), 0);

    solve(0, (int)vals.size() - 1, move(allRects), move(allQueries));

    for (int i = 0; i < q; ++i) {
        cout << ans[i] << '\n';
    }

    return 0;
}
posted @ 2026-07-19 15:35  high_skyy  阅读(4)  评论(0)    收藏  举报
动态线条
动态线条end
浏览器标题切换
浏览器标题切换end
💬 加载中……