浅谈树上距离问题
liheyang123 · · 算法·理论
本文介绍了树链剖分、树分治在一些树上距离问题相关的处理。
树链剖分
具体来说是长链剖分。做法是在 DFS 过程中,对于当前节点 dsu on tree 相通。
该做法要求边权必须为
复杂度证明:假设合并两组信息的代价为
例题
这里给出一个经典例题 qoj#108,luogu#P4292。
简化题意:给定一棵树,树边带有价值,边权为
做法:二分枚举平均值,将边权转化为“价值 − 平均值”,然后借助 dsu on tree 处理;由于需要区间查询和单点修改,配合线段树维护深度信息。具体实现时,以节点深度为信息,分类讨论当前 DFS 到的节点是否作为路径端点。
时间复杂度
树分治
树分治(尤其是点分治)的核心思想,是将任意两点间的距离
例题
luogu#P15006
简化题意:给定一棵树,节点
转移:设
按
对每个分治中心维护线段树,即可快速查询满足条件的最大
qoj#6660,luogu#P11343
简化题意:给定一棵树,每个节点
转移:注意到最优策略中
对分治中心
每个
总复杂度
::::info[参考代码]
#include<bits/stdc++.h>
using namespace std;
typedef long long i64;
const int N = 1e5 + 10;
const i64 inf = 2e18;
int n;
i64 a [ N ], b [ N ], f [ N ];
vector < pair < int, int > > g [ N ];
// 树上距离相关
int st [ N << 1 ][ 18 ], times, DEP [ N ], lg [ N << 1 ], dfn [ N ];
i64 dep0 [ N ];
void dfs0 ( int u, int fa ) {
st [ ++times ][ 0 ] = u;
dfn [ u ] = times;
for ( auto e : g [ u ] ) {
auto v = e . first, w = e . second;
if ( v == fa ) continue;
dep0 [ v ] = dep0 [ u ] + w;
DEP [ v ] = DEP [ u ] + 1;
dfs0 ( v, u );
st [ ++times ][ 0 ] = u;
}
}
int getmn ( int u, int v ) {
if ( ! u || ! v ) return u | v;
return DEP [ u ] < DEP [ v ] ? u : v;
}
void buildst ( ) {
for ( int i = 2; i <= times; ++i ) lg [ i ] = lg [ i >> 1 ] + 1;
for ( int k = 1; k < 18; ++k ) for ( int i = 1; i + ( 1 << k ) - 1 <= times; ++i )
st [ i ][ k ] = getmn ( st [ i ][ k - 1 ], st [ i + ( 1 << k - 1 ) ][ k - 1 ] );
}
int getlca ( int u, int v ) {
if ( dfn [ u ] > dfn [ v ] ) swap ( u, v );
int k = lg [ dfn [ v ] - dfn [ u ] + 1 ];
return getmn ( st [ dfn [ u ] ][ k ], st [ dfn [ v ] - ( 1 << k ) + 1 ][ k ] );
}
i64 getdis ( int u, int v ) { return dep0 [ u ] + dep0 [ v ] - 2 * dep0 [ getlca ( u, v ) ]; }
// 点分树部分
int sz [ N ], mxs [ N ], facd [ N ], rtcd;
bool vis [ N ];
void getrt ( int u, int fa, int tot ) {
sz [ u ] = 1; mxs [ u ] = 0;
for ( auto e : g [ u ] ) {
int v = e . first;
if ( v == fa || vis [ v ] ) continue;
getrt ( v, u, tot );
sz [ u ] += sz [ v ];
mxs [ u ] = max ( mxs [ u ], sz [ v ] );
}
mxs [ u ] = max ( mxs [ u ], tot - sz [ u ] );
if ( mxs [ u ] < mxs [ rtcd ] ) rtcd = u;
}
void divide ( int u, int fa, int tot ) {
vis [ u ] = 1; facd [ u ] = fa;
for ( auto e : g [ u ] ) {
int v = e . first;
if ( vis [ v ] ) continue;
int nex = ( sz [ v ] > sz [ u ] ? tot - sz [ u ] : sz [ v ] );
rtcd = 0, getrt ( v, u, nex );
divide ( rtcd, u, nex );
}
}
// 斜率优化凸包
struct Line { i64 k, b; i64 get ( i64 x ) { return k * x + b; } }; vector < Line > hulls [ N ];
bool isbad ( Line l1, Line l2, Line l3 ) {
return ( __int128 ) ( l2 . b - l1 . b ) * ( l2 . k - l3 . k ) >= ( __int128 ) ( l3 . b - l2 . b ) * ( l1 . k - l2 . k );
}
void ins ( int rt, Line L ) {
auto &h = hulls [ rt ];
if ( !h . empty ( ) && h . back ( ) . k == L . k ) {
if ( h . back ( ) . b <= L . b ) return;
h . pop_back ( );
} while ( h . size ( ) >= 2 && isbad ( h [ h . size ( ) - 2 ], h . back ( ), L ) ) h . pop_back ( );
h . push_back ( L );
}
i64 query ( int rt, i64 x ) {
auto &h = hulls [ rt ];
if ( h . empty ( ) ) return inf;
int l = 0, r = h . size ( ) - 1;
while ( l < r ) {
int mid = ( l + r ) >> 1;
if ( h [ mid ] . get ( x ) <= h [ mid + 1 ] . get ( x ) ) r = mid;
else l = mid + 1;
} return h [ l ] . get ( x );
}
vector<long long> travel(vector<long long> A, vector<int> B, vector<int> U, vector<int> V, vector<int> W){
n = A . size ( ); vector < i64 > ans ( n - 1 );
for ( int i = 1; i <= n; ++i ) a [ i ] = A [ i - 1 ];
for ( int i = 1; i <= n; ++i ) b [ i ] = B [ i - 1 ];
for ( int i = 1; i < n; ++i ) {
int u = U [ i - 1 ] + 1, v = V [ i - 1 ] + 1, w = W [ i - 1 ];
g [ u ] . emplace_back ( v, w );
g [ v ] . emplace_back ( u, w );
}
dfs0 ( 1, 0 ), buildst ( );
mxs [ 0 ] = n + 1, rtcd = 0;
getrt ( 1, 0, n ), divide ( rtcd, 0, n );
vector < int > p ( n );
for ( int i = 0; i < n; ++i ) p [ i ] = i + 1;
sort ( p . begin ( ), p . end ( ), [ & ] ( int i, int j ) { if ( b [ i ] != b [ j ] ) return b [ i ] > b [ j ]; return i < j; } );
for ( int i = 1; i <= n; ++i ) f [ i ] = inf;
for ( int u : p ) {
if ( u == 1 ) f [ u ] = 0;
else for ( int rt = u; rt; rt = facd [ rt ] )
f [ u ] = min ( f [ u ], query ( rt, getdis ( u, rt ) ) );
if ( f [ u ] != inf ) for ( int rt = u; rt; rt = facd [ rt ] )
ins ( rt, { b [ u ], f [ u ] + a [ u ] + b [ u ] * getdis ( u, rt ) } );
}
for ( int i = 2; i <= n; ++i ) {
ans [ i - 2 ] = inf;
for ( int rt = i; rt; rt = facd [ rt ] )
ans [ i - 2 ] = min ( ans [ i - 2 ], query ( rt, getdis ( i, rt ) ) );
} return ans;
}
::::
luogu#P14202
简化题意:给定一棵树,节点有属性
特殊之处在于空间限制仅为 64MB,无法为每个分治中心开动态开点数据结构。
解法:将所有操作离线。在点分治的递归过程中,对当前分治块内的节点及操作统一处理。每个操作(修改或查询)会被
时间复杂度
这种离线方式通常比在线点分树常数更小,若问题仅为简单修改和查询,可优先考虑。
::::info[参考代码(篇幅问题,省去了火车头)]
const int N = 1e5 + 10, V = 1e5;
const i64 md = 998244353;
int n, m, a0 [ N ], b0 [ N ];;
vector < int > Ops [ N ];
int head [ N ], to [ N << 1 ], nxt [ N << 1 ], ecnt;
inline void add_edge ( int u, int v ) {
to [ ++ecnt ] = v;
nxt [ ecnt ] = head [ u ];
head [ u ] = ecnt;
}
struct Op { int type, x, v; } ops [ N ];
struct node { int c; i64 sa, sd, sda; } bit [ N ];
inline void add ( int u, int c, i64 sa, i64 sd, i64 sda ) {
sa = ( sa % md + md ) % md, sd = ( sd % md + md ) % md, sda = ( sda % md + md ) % md;
for ( ; u <= V; u += u & -u ) {
bit [ u ] . c = ( bit [ u ] . c + c ) % md;
bit [ u ] . sa = ( bit [ u ] . sa + sa ) % md;
bit [ u ] . sd = ( bit [ u ] . sd + sd ) % md;
bit [ u ] . sda = ( bit [ u ] . sda + sda ) % md;
}
}
inline node query ( int u ) {
node res = { 0, 0, 0, 0 };
for ( ; u; u -= u & - u ) {
res . c = ( res . c + bit [ u ] . c ) % md;
res . sa = ( res . sa + bit [ u ] . sa ) % md;
res . sd = ( res . sd + bit [ u ] . sd ) % md;
res . sda = ( res . sda + bit [ u ] . sda ) % md;
} return res;
}
int siz [ N ], mxs [ N ], rtcd;
bool vis [ N ];
void getrt ( int u, int fa, int tot ) {
siz [ u ] = 1, mxs [ u ] = 0;
for ( int i = head [ u ]; i; i = nxt [ i ] ) {
int v = to [ i ];
if ( v == fa || vis [ v ] ) continue;
getrt ( v, u, tot ), siz [ u ] += siz [ v ];
mxs [ u ] = max ( mxs [ u ], siz [ v ] );
}
mxs [ u ] = max ( mxs [ u ], tot - siz [ u ] );
if ( mxs [ u ] < mxs [ rtcd ] ) rtcd = u;
}
int dis [ N ], a [ N ], b [ N ]; i64 ans [ N ];
void fun ( int u, int fa, int d, vector < int > &S ) {
S . push_back ( u ), dis [ u ] = d;
for ( int i = head [ u ]; i; i = nxt [ i ] ) {
int v = to [ i ];
if ( v != fa && !vis [ v ] ) fun ( v, u, d + 1, S );
}
}
void work ( vector < int > &S, int sign ) {
vector < int > cops;
for ( auto u : S ) {
a [ u ] = a0 [ u ], b [ u ] = b0 [ u ];
add ( b [ u ], 1, a [ u ], dis [ u ], ( i64 ) a [ u ] * dis [ u ] );
for ( auto id : Ops [ u ] ) cops . push_back ( id );
}
sort ( cops . begin ( ), cops . end ( ) );
for ( auto id : cops ) {
int type = ops [ id ] . type, x = ops [ id ] . x, v = ops [ id ] . v;
if ( type == 1 ) {
add ( b [ x ], 0, v - a [ x ], 0, ( i64 ) ( v - a [ x ] ) * dis [ x ] );
a [ x ] = v;
} else if ( type == 2 ) {
add ( b [ x ], -1, - a [ x ], - dis [ x ], - ( i64 ) a [ x ] * dis [ x ] );
b [ x ] = v;
add ( b [ x ], 1, a [ x ], dis [ x ], ( i64 ) a [ x ] * dis [ x ] );
} else {
int l = max ( 1, b [ x ] - v ), r = min ( V, b [ x ] + v );
node nr = query ( r ), nl = query ( l - 1 );
i64 cnt = ( nr . c - nl . c + md ) % md;
i64 sa = ( nr . sa - nl . sa + md ) % md;
i64 sd = ( nr . sd - nl . sd + md ) % md;
i64 sda = ( nr . sda - nl . sda + md ) % md;
i64 res = ( i64 ) dis [ x ] * a [ x ] % md * cnt % md;
res = ( res + ( i64 ) dis [ x ] * sa % md ) % md;
res = ( res + ( i64 ) a [ x ] * sd % md ) % md;
res = ( res + sda ) % md;
ans [ id ] = ( ans [ id ] + res * sign + md ) % md;
}
}
for ( auto u : S ) add ( b [ u ], -1, -a [ u ], -dis [ u ], -( i64 ) a [ u ] * dis [ u ] );
}
void solve ( int u, int tot ) {
rtcd = 0, getrt ( u, 0, tot );
int c = rtcd; vis [ c ] = 1; vector < int > S; fun ( c, 0, 0, S ), work ( S, 1 );
for ( int i = head [ c ]; i; i = nxt [ i ] ) {
int v = to [ i ];
if ( !vis [ v ] ) { vector < int > SS; fun ( v, c, 1, SS ), work ( SS, -1 ); }
}
for ( int i = head [ c ]; i; i = nxt [ i ] ) {
int v = to [ i ];
if ( !vis [ v ] ) solve ( v, siz [ v ] > siz [ u ] ? tot - siz [ u ] : siz [ v ] );
}
}
signed main ( ) {
n = read < int > ( ), m = read < int > ( );
for ( int i = 1; i < n; ++i ) {
int u = read < int > ( ), v = read < int > ( );
add_edge ( u, v ), add_edge ( v, u );
}
for ( int i = 1; i <= n; ++i ) a0 [ i ] = read < int > ( ), b0 [ i ] = read < int > ( );
for ( int i = 1; i <= m; ++i ) {
ops [ i ] . type = read < int > ( );
ops [ i ] . x = read < int > ( );
ops [ i ] . v = read < int > ( );
Ops [ ops [ i ] . x ] . emplace_back ( i );
} mxs [ 0 ] = n + 1, solve ( 1, n );
for ( int i = 1; i <= m; ++i ) if ( ops [ i ] . type == 3 ) write ( ans [ i ] );
flush ( );
return 0;
}
::::
::::info[一些说明] 这里承认,例题比较经典,搬运了一些他人专栏给出的例题,并且部分内容参考了一些说法,绝大多数内容为原创。使用了 DeepSeek-V4-Flash-0731 进行了润色,并提高代码可读性。
下述一些参考文献。
-
链剖分总结(@wishapig)
-
题解:P11343 [KTSC 2023 R1] 出租车旅行(@happybob)