树链剖分25pts求调
查看原帖
树链剖分25pts求调
719942
zzq_helloworld楼主2023/6/13 13:56

rt,悬关。

评测记录:https://www.luogu.com.cn/record/112627180。

#include <bits/stdc++.h>
using namespace std;
void Solve();
int main() {
  int tt = 1;
  cin >> tt;
  while( tt-- ) {
    Solve();
  }
  return 0;
}

const int _ = 1e5 + 5;

struct Edge { int v , w , nxt ; } e[_*2];
int head[_] , ecnt;
void Add( int u , int v , int w = 1 ) { e[ ++ecnt ] = Edge{ v , w , head[ u ] } ; head[ u ] = ecnt ; }

struct SMT {
private:
  struct Node {
    Node *L , *R;
    int l , r;
    int sum , col1 , col2 , cover;
    Node( int l , int r ) : l(l) , r(r) , sum(0) , col1(0) , col2(0) , cover(0) , L(NULL) , R(NULL) {}
    int mid() { return l + ( r - l ) / 2 ; }
    int len() { return r - l + 1 ; }
    void PushUp() {
      sum = L->sum + R->sum + ( L->col2 == R->col1 );
      col1 = L->col1 , col2 = R->col2;
    }
    void PushDown() {
      if( cover != 0 ) {
        L->col1 = L->col2 = R->col1 = R->col2 = cover;
        L->cover = R->cover = cover;
        L->sum = L->len() - 1;
        R->sum = R->len() - 1;
        cover = 0;
      }
    }
  };
public:
  Node *root;
  void Build( int l , int r , Node *&p ) {
    p = new Node( l , r );
    if( l == r ) return;
    Build( l , p->mid() , p->L );
    Build( p->mid() + 1 , r , p->R );
    p->PushUp();
  }
  void Change( int x , int y , int v , Node *p ) {
    if( x <= p->l && p->r <= y ) {
      p->col1 = p->col2 = p->cover = v;
      p->sum = p->len() - 1;
      return;
    }
    p->PushDown();
    if( x <= p->mid() ) Change( x , y , v , p->L );
    if( y > p->mid() ) Change( x , y , v , p->R );
    p->PushUp();
  }
  int GetSum( int x , int y , Node *p ) {
    if( x <= p->l && p->r <= y ) return p->sum;
    p->PushDown();
    int ans = 0;
    if( x <= p->mid() ) ans += GetSum( x , y , p->L );
    if( y > p->mid() ) ans += GetSum( x , y , p->R );
    if( x <= p->mid() && y > p->mid() && ( p->L->col2 == p->R->col1 ) ) ans++;
    return ans;
  }
  int Get( int x , Node *p ) {
    if( p->l == x && x == p->r ) return p->col1;
    if( x <= p->mid() ) return Get( x , p->L );
    else return Get( x , p->R );
  }
} seg;

int fa[_] , top[_] , dep[_] , siz[_] , son[_] , dfn[_] , cnt;

void DFS1( int x ) {
  dep[ x ] = dep[ fa[ x ] ] + 1;
  siz[ x ] = 1;
  for( int i = head[ x ] ; i ; i = e[ i ].nxt ) {
    int y = e[ i ].v;
    if( !dep[ y ] ) {
      fa[ y ] = x;
      DFS1( y );
      siz[ x ] += siz[ y ];
      if( siz[ y ] > siz[ son[ x ] ] ) {
        son[ x ] = y;
      }
    }
  }
}

void DFS2( int x ) {
  dfn[ x ] = ++cnt;
  if( son[ x ] ) {
    top[ son[ x ] ] = top[ x ];
    DFS2( son[ x ] );
  }
  for( int i = head[ x ] ; i ; i = e[ i ].nxt ) {
    int y = e[ i ].v;
    if( !top[ y ] ) {
      top[ y ] = y;
      DFS2( y );
    }
  }
}

int n , m;

void ChangePath( int x , int y , int v ) {
  while( top[ x ] != top[ y ] ) {
    if( dep[ top[ x ] ] < dep[ top[ y ] ] ) swap( x , y );
    seg.Change( dfn[ top[ x ] ] , dfn[ x ] , v , seg.root );
    x = fa[ top[ x ] ];
  }
  if( dep[ x ] > dep[ y ] ) swap( x , y );
  seg.Change( dfn[ x ] , dfn[ y ] , v , seg.root );
}

int QueryPath( int x , int y ) {
  int ans = 0;
  while( top[ x ] != top[ y ] ) {
    if( dep[ top[ x ] ] < dep[ top[ y ] ] ) swap( x , y );
    ans += seg.GetSum( dfn[ top[ x ] ] , dfn[ x ] , seg.root );
    ans += seg.Get( dfn[ top[ x ] ] , seg.root ) == seg.Get( dfn[ fa[ top[ x ] ] ] , seg.root );
    x = fa[ top[ x ] ];
  }
  if( dep[ x ] > dep[ y ] ) swap( x , y );
  ans += seg.GetSum( dfn[ x ] , dfn[ y ] , seg.root );
  return ans;
}

void Solve() {
  memset( head , 0 , sizeof( head ) );
  memset( fa , 0 , sizeof( fa ) );
  memset( top , 0 , sizeof( top ) );
  memset( dep , 0 , sizeof( dep ) );
  memset( siz , 0 , sizeof( siz ) );
  memset( son , 0 , sizeof( son ) );
  memset( dfn , 0 , sizeof( dfn ) );
  SMT smt;
  seg = smt;
  ecnt = cnt = 0;
  cin >> n >> m;
  for( int i = 1 , u , v ; i < n ; i++ ) {
    cin >> u >> v;
    Add( u , v );
    Add( v , u );
  }
  DFS1( 1 ) , fa[ 1 ] = 1 , top[ 1 ] = 1 , DFS2( 1 );
  seg.Build( 1 , n , seg.root );
  for( int i = 1 ; i <= n ; i++ ) {
    seg.Change( dfn[ i ] , dfn[ i ] , -i , seg.root );
  }
  for( int i = 1 ; i <= m ; i++ ) {
    int op , l , r;
    cin >> op >> l >> r;
    if( op == 1 ) {
      ChangePath( l , r , i );
    } else {
      cout << QueryPath( l , r ) << '\n';
    }
  }
}
2023/6/13 13:56
加载中...