本地卡住求调
查看原帖
本地卡住求调
719942
zzq_helloworld楼主2023/6/12 14:11

rt。悬赏关注。

#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() {
      col1 = L->col1 , col2 = R->col2;
      sum = L->sum + R->sum + ( L->col1 == 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 ) {
      p->col1 = p->col2 = -l;
      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 );
    return ans + ( p->col1 == p->col2 );
  }
  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[ y ] += siz[ x ];
  	  if( siz[ x ] > 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 );
  }
  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 );
  }
  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( dep , 0 , sizeof( dep ) );
  memset( son , 0 , sizeof( son ) );
  fa[ 1 ] = 1 , top[ 1 ] = 1;
  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 ) , DFS2( 1 );
  seg.Build( 1 , n , 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/12 14:11
加载中...