树路径统计的点分治算法实现
点分治算法概述
点分治是一种用于处理树结构路径统计问题的算法。该算法通过选取合适的根节点,将原问题分解为经过根节点的路径和子树内的路径两种情况,运用分治策略高效解决问题。
算法核心思想
树上的路径可分为两类:
- 穿越根节点的路径
- 完全位于子树内部的路径
算法首先处理经过当前根节点的路径,然后递归处理各个子树。
步骤一:定位树的重心
选择重心作为根节点能保证算法效率。重心定义为最大子树规模最小的节点。
void locateCentroid(int cur, int parent) {
subtreeSize[cur] = 1;
maxSubtree[cur] = 0;
for (int i = firstEdge[cur]; i; i = edges[i].next) {
int neighbor = edges[i].target;
if (neighbor == parent || visited[neighbor]) continue;
locateCentroid(neighbor, cur);
subtreeSize[cur] += subtreeSize[neighbor];
maxSubtree[cur] = max(maxSubtree[cur], subtreeSize[neighbor]);
}
maxSubtree[cur] = max(maxSubtree[cur], totalSize - subtreeSize[cur]);
if (maxSubtree[cur] < maxSubtree[centroid]) centroid = cur;
}
步骤二:路径统计与处理
收集所有节点到当前根节点的距离:
void collectDistances(int cur, int parent, int dist) {
distances[++count] = dist;
for (int i = firstEdge[cur]; i; i = edges[i].next) {
int neighbor = edges[i].target;
if (neighbor == parent || visited[neighbor]) continue;
collectDistances(neighbor, cur, dist + edges[i].weight);
}
}
使用双指针技术统计满足条件的路径数量:
int computePaths(int cur, int baseDist) {
int total = 0;
count = 0;
collectDistances(cur, 0, baseDist);
sort(distances + 1, distances + count + 1);
int right = count;
for (int left = 1; left <= count; left++) {
while (distances[left] + distances[right] > maxDist && right >= 1) right--;
if (left > right) break;
total += right - left + 1;
}
return total;
}
步骤三:分治执行流程
核心分治函数实现:
void divideConquer(int cur) {
result += computePaths(cur, 0);
visited[cur] = true;
for (int i = firstEdge[cur]; i; i = edges[i].next) {
int neighbor = edges[i].target;
if (visited[neighbor]) continue;
result -= computePaths(neighbor, edges[i].weight);
centroid = 0;
totalSize = subtreeSize[neighbor];
locateCentroid(neighbor, cur);
divideConquer(centroid);
}
}
算法复杂度分析
算法的时间复杂度为O(n log² n),其中n为节点数量。
完整实现示例
#include <cstdio>
#include <cstring>
#include <algorithm>
using namespace std;
const int MAXN = 100005;
const int INF = 0x3f3f3f3f;
struct GraphEdge {
int target, next, weight;
} edges[MAXN * 2];
int nodeCount, maxDist, centroid, totalSize;
int result, distanceCount;
int subtreeSize[MAXN], maxSubtree[MAXN];
int distances[MAXN], firstEdge[MAXN];
bool visited[MAXN];
int edgeCount;
void addEdge(int u, int v, int w) {
edges[++edgeCount].target = v;
edges[edgeCount].weight = w;
edges[edgeCount].next = firstEdge[u];
firstEdge[u] = edgeCount;
}
void locateCentroid(int cur, int parent) {
subtreeSize[cur] = 1;
maxSubtree[cur] = 0;
for (int i = firstEdge[cur]; i; i = edges[i].next) {
int neighbor = edges[i].target;
if (neighbor == parent || visited[neighbor]) continue;
locateCentroid(neighbor, cur);
subtreeSize[cur] += subtreeSize[neighbor];
maxSubtree[cur] = max(maxSubtree[cur], subtreeSize[neighbor]);
}
maxSubtree[cur] = max(maxSubtree[cur], totalSize - subtreeSize[cur]);
if (maxSubtree[cur] < maxSubtree[centroid]) centroid = cur;
}
void collectDistances(int cur, int parent, int dist) {
distances[++distanceCount] = dist;
for (int i = firstEdge[cur]; i; i = edges[i].next) {
int neighbor = edges[i].target;
if (neighbor == parent || visited[neighbor]) continue;
collectDistances(neighbor, cur, dist + edges[i].weight);
}
}
int computePaths(int cur, int baseDist) {
int total = 0;
distanceCount = 0;
collectDistances(cur, 0, baseDist);
sort(distances + 1, distances + distanceCount + 1);
int right = distanceCount;
for (int left = 1; left <= distanceCount; left++) {
while (distances[left] + distances[right] > maxDist && right >= 1) right--;
if (left > right) break;
total += right - left + 1;
}
return total;
}
void divideConquer(int cur) {
result += computePaths(cur, 0);
visited[cur] = true;
for (int i = firstEdge[cur]; i; i = edges[i].next) {
int neighbor = edges[i].target;
if (visited[neighbor]) continue;
result -= computePaths(neighbor, edges[i].weight);
centroid = 0;
totalSize = subtreeSize[neighbor];
locateCentroid(neighbor, cur);
divideConquer(centroid);
}
}
int main() {
scanf("%d", &nodeCount);
for (int i = 1; i < nodeCount; i++) {
int u, v, w;
scanf("%d%d%d", &u, &v, &w);
addEdge(u, v, w);
addEdge(v, u, w);
}
scanf("%d", &maxDist);
maxSubtree[centroid = 0] = INF;
totalSize = nodeCount;
locateCentroid(1, 0);
divideConquer(centroid);
printf("%d\n", result - nodeCount);
return 0;
}