wa on 点10,求助!
查看原帖
wa on 点10,求助!
767099
WEXI7111楼主2022/12/13 20:52
#include<bits/stdc++.h>
#define ll long long
using namespace std;
const int N = 200010;

int n,m,mod,rt;
int h[N],to[N],nxt[N],idx;
void add(int u,int v)
{
    to[++ idx] = v,nxt[idx] = h[u], h[u] = idx;
    return;
}
int id[N],nw[N],lst;
int dep[N],top[N],sz[N],fa[N],son[N];
ll w[N];

struct Nodetree
{
    struct Node
    {
        int l,r;
        ll sum,plus;
    }a[N << 1];
    void pushup(int p)
    {
        a[p].sum = (a[p << 1].sum + a[p << 1 | 1].sum) % mod;
        return;
    }
    void addit(int p,ll k)
    {
        a[p].sum = (a[p].sum + 1ll * k * (a[p].r - a[p].l + 1) % mod) % mod;
        a[p].plus = (a[p].plus + k) % mod;
        return;
    }
    void build(int l,int r,int p)
    {
        a[p].l = l, a[p].r = r;
        if(l == r) {a[p].sum = nw[l];return;}
        int mid = l + r >> 1;
        build(l, mid, p << 1);build(mid + 1,r,p << 1 | 1);
        pushup(p);
        return;
    }
    void pushdown(int p)
    {
        addit(p << 1, a[p].plus);addit(p << 1 | 1, a[p].plus);
        a[p].plus = 0;
        return;
    }

    void add(int l,int r,int p,ll k)
    {
        if(l <= a[p].l && a[p].r <= r) {addit(p,k); return;}
        pushdown(p);
        int mid = a[p].l + a[p].r >> 1;
        if(l <= mid) add(l,r,p << 1,k);
        if(mid < r) add(l,r,p << 1 | 1,k);
        pushup(p);
        return;
    }
    ll ask(int l,int r,int p)
    {
        if(l <= a[p].l && a[p].r <= r) return a[p].sum;
        pushdown(p);
        ll res = 0;
        int mid = a[p].l + a[p].r >> 1;
        if(l <= mid) res = (res + ask(l,r,p << 1)) % mod;
        if(mid < r) res = (res + ask(l,r,p << 1 | 1)) % mod;
        return res;
    }
}T;

void dfs1(int u,int from,int d)
{
    dep[u] = d,fa[u] = from, sz[u] = 1;
    for(int i = h[u]; i != -1; i = nxt[i])
    {
        int e = to[i];
        if(e == from) continue;
        dfs1(e,u,d + 1);
        sz[u] = (sz[u] + sz[e]) % mod;
        if(sz[son[u]] < sz[e]) son[u] = e;
    }
}

void dfs2(int u,int t)
{
    id[u] = ++ lst; nw[lst] = w[u];
    top[u] = t;
    if(!son[u]) return;
    dfs2(son[u],t);
    for(int i = h[u]; i != -1; i = nxt[i])
    {
        int e = to[i];
        if(e == fa[u] || e == son[u]) continue;
        dfs2(e,e);
    }
}

void update_path(int u,int v,int k)
{
    while(top[u] != top[v])
    {
        if(dep[top[u]] < dep[top[v]]) swap(u,v);
        T.add(id[top[u]], id[u], 1, k);
        u = fa[top[u]];
    }
    if(dep[u] < dep[v]) swap(u,v);
    T.add(id[v],id[u],1,k);
}

ll query_path(int u,int v)
{
    ll res = 0;
    while(top[u] != top[v])
    {
        if(dep[top[u]] < dep[top[v]]) swap(u,v);
        res = (res + T.ask(id[top[u]], id[u], 1)) % mod;
        u = fa[top[u]];
    }
    if(dep[u] < dep[v]) swap(u,v);
    res = (res + T.ask(id[v], id[u], 1)) % mod;
    return res;
}

void update_tree(int u,int k)
{
    T.add(id[u], id[u] + sz[u] - 1, 1, k);
}

ll query_tree(int u)
{
    return T.ask(id[u], id[u] + sz[u] - 1, 1) % mod;
}

int main()
{
    scanf("%d%d%d%d",&n,&m,&rt,&mod);
    for(int i = 1;i <= n;i ++) {scanf("%lld",&w[i]);w[i] %= mod;}
    memset(h,-1,sizeof(h)); 
    for(int i = 1;i < n;i ++)
    {
        int a,b;
        scanf("%d%d",&a,&b);
        add(a,b);add(b,a);
    }
    
    dfs1(rt,0,1);
    dfs2(rt,rt);
    T.build(1,n,1);
    
    while(m--)
    {
        int tp,u,v;ll k;
        scanf("%d%d",&tp,&u);
        if(tp == 1)
        {
            scanf("%d%lld",&v,&k);
            update_path(u, v, k % mod);
        }
        if(tp == 3)
        {
            scanf("%lld",&k);
            update_tree(u, k % mod);
        }
        if(tp == 2) 
        {
            scanf("%d",&v);
            printf("%lld\n",query_path(u,v));
        }
        if(tp == 4)
            printf("%lld\n",query_tree(u));
    }

    return 0;
}
2022/12/13 20:52
加载中...