树链剖分RE了 求助qaq
查看原帖
树链剖分RE了 求助qaq
105820
阿尔托莉雅丶楼主2022/5/12 20:04
#include <iostream>
#include <algorithm>
#include <math.h>
#include <vector>
#include <cstdio>
#include <cstring>
#include <assert.h>
using namespace std;
typedef long long ll;
const int N = 1e6 + 5;   //remember to modify the range of the data!!
ll mod = 1e9 + 7;
const int inf = 0x3f3f3f3f;

int n, m, T, ts;;
int a[N];
vector <int> e[N];
//son[u], 节点u的重儿子, 
int f[N], sz[N], dep[N], son[N]; //第一遍dfs预处理的量
//top[u]节点u所在的重链的,dfn[u], 节点u的dfs序的编号,rnk[u]表示dfs序是u的节点编号即dfn的逆映射
int top[N], dfn[N], id[N]; //第二遍dfs预处理的量

void dfs1(int u, int fa)
{
    dep[u] = dep[fa] + 1;        
    sz[u] = 1;
    son[u] = -1;
    int mxsz = 0;
    for(auto v : e[u])
    {
        if(v == fa)
            continue;
        dfs1(v, u);
        f[v] = u;
        sz[u] += sz[v];
        if(sz[v] > mxsz)
        {
            son[u] = v;
            mxsz = sz[v];
        }
    }
}

void dfs2(int u, int topn) //topn表示此重链的顶节点
{
    dfn[u] = ++ts;
    top[u] = topn;
    id[ts] = u;
    if(son[u] == -1)
        return;
    dfs2(son[u], topn); //先访问重儿子
    for(auto v : e[u])
    {
        if(v == f[u] || v == son[u]) 
            continue;
        dfs2(v, v);//每个轻儿子的顶是自己
    }
}
#define lson (k << 1)  //使用前先看作用域里有无 k
#define rson (k << 1 | 1)
int t[N << 2], ly[N << 2];

void pushup(int k)
{
    t[k] = (t[lson] + t[rson]) % mod;
}

void pushdown(int k, int l, int r, int mid)
{
    if(ly[k])
    {
        t[lson] = (t[lson] + ly[k] * (mid - l + 1)) % mod;
        t[rson] = (t[rson] + ly[k] * (r - mid)) % mod;
        ly[lson] = (ly[lson] + ly[k]) % mod;
        ly[rson] = (ly[rson] + ly[k]) % mod;
        ly[k] = 0;
    }
}

void build(int k, int l, int r)
{
    if(l == r)
    {
        t[k] = a[id[l]];
        return;
    }
    int mid = l + r >> 1;
    build(lson, l, mid);
    build(rson, mid + 1, r);
    pushup(k);
}

void update(int k, int c, int x, int y, int l, int r)
{
    if(x <= l && y >= r)
    {
        ly[k] = (ly[k] +  c) % mod;
        t[k] = ((r - l + 1) * c % mod + t[k]) % mod;
        return;
    }
    int mid = l + r >> 1;
    pushdown(k, l, r, mid);
    if(x <= mid)
        update(lson, c, x, y, l, mid);
    if(y > mid)
        update(rson, c, x, y, mid + 1, r);
    pushup(k);
}

int query(int k, int x, int y, int l, int r)
{
    if(x <= l && y >= r)
    {
        return t[k];
    }
    int mid = l + r >> 1;
    pushdown(k, l, r, mid);
    int res = 0;
    if(x <= mid)
        res = (res + query(lson, x, y, l, mid)) % mod;
    if(y > mid)
        res = (res + query(rson, x, y, mid + 1, r)) % mod;
    return res;
}

void update(int x, int y, int z)
{
    while(top[x] != top[y]) //跳到同一条重链
    {
        if(dep[x] < dep[y])
            swap(x, y);
        update(1, z, dfn[top[x]], dfn[x], 1, n);
        x = f[top[x]];
    }
    if(dep[x] > dep[y]) //保证dfn[x] <= dfn[y];
        swap(x, y);
    update(1, z, dfn[x], dfn[y], 1, n);
}

int query(int x, int y)
{
    int res = 0;
    while(top[x] != top[y]) //跳到同一条重链
    {
        if(dep[x] < dep[y])
            swap(x, y);
        res = (res + query(1, dfn[top[x]], dfn[x], 1, n)) % mod;
        x = f[top[x]];
    }
    if(dep[x] > dep[y]) //保证dfn[x] <= dfn[y];
        swap(x, y);
    return (res + query(1, dfn[x], dfn[y], 1, n)) % mod;
}


void update(int x, int z)
{
    update(1, z, dfn[x], dfn[x] + sz[x] - 1, 1, n);
}

int query(int x)
{
    return query(1, dfn[x], dfn[x] + sz[x] - 1, 1, n);
}


int main(void)
{
    // freopen("../data/in.txt", "r", stdin);
    // freopen("../data/out.txt", "w", stdout);
    int root;
    cin >> n >> m >> root >> mod;
    for(int i = 1; i <= n; i++)
        cin >> a[i];
    for(int i = 1; i < n; i++)
    {
        int u, v;
        cin >> u >> v;
        e[u].push_back(v);
        e[v].push_back(u);
    }
    dfs1(root, 0);
    dfs2(root, root);

    build(1, 1, n);
    
    while(m--)
    {
        int op, x, y, z;
        cin >> op;
        if(op == 1)
        {
            cin >> x >> y >> z;
            update(x, y, z);
        }
        else if(op == 3)
        {
            cin >> x >> z;
            update(x, z);
        }
        else if(op == 2)
        {
            cin >> x >> y;
            cout << query(x, y) << '\n';
        }
        else
        {
            cin >> x;
            cout << query(x) << '\n';
        }
    }


    return 0;
}
2022/5/12 20:04
加载中...