简单线段树求解
查看原帖
简单线段树求解
817044
cjwdyzxfblzs楼主2023/9/8 08:41

真的找不出来是哪里的问题了,但是就是过不去欸

#include <bits/stdc++.h>
#define int long long
const int N = 1e6;
int n, m, a[N], sq[N];
struct node
{
    int cnt[40];
    int l, r, tag;
    node () { memset(cnt, 0, sizeof(cnt)); }
    node (int l, int r) : l(l), r(r) { memset(cnt, 0, sizeof(cnt)); }
    void Get_int(int x)
    {
        for (int i = 30; i >= 0; i -- )
        {
            if ((x >> i) & 1)
                cnt[i] = 1;
        } 
    }
    int Get_val( void )
    {
        int ans = 0;
        for (int i = 0; i <= 30; i ++ ) 
            ans += sq[i] * cnt[i];
        return ans;
    }
}tr[N];
#define U tr[u]
#define lc tr[u << 1]
#define rc tr[u << 1 | 1]
#define ls (u << 1)
#define rs (u << 1 | 1)
void Push_Up(int u)
{
    for (int i = 0; i <= 30; i ++ )
        U.cnt[i] = lc.cnt[i] + rc.cnt[i];
}
void build(int u, int l, int r)
{
    U = node (l, r);
    if (l == r) 
        return U.Get_int(a[l]), void();
    int mid = (l + r) >> 1;
    build(ls, l, mid);
    build(rs, mid + 1, r);
    Push_Up(u);
}
void Push_Tag(int u, int val)
{
    U.tag ^= val;
    for (int i = 30; i >= 0; i -- )
    {
        int v = ((val >> i) & 1);
        if (v) U.cnt[i] = U.r - U.l + 1 - U.cnt[i];
    }
}
void Push_Down(int u)
{
    if (U.tag)
    {
        Push_Tag(ls, U.tag);
        Push_Tag(rs, U.tag);
        U.tag = 0;
    }
}
void Modify(int u, int l, int r, int val)
{
    if (l > U.r or r < U.l) return void();
    if (l <= U.l and U.r <= r) return Push_Tag(u, val), void();
    Push_Down(u);
    int Mid = (U.l + U.r) >> 1;
    if (l <= Mid) Modify(ls, l, r, val);
    if (r > Mid) Modify(rs, l, r, val);
    Push_Up(u);
}
// int Query(int u, int l, int r) 
// {
//     if (U.l > r or U.r < l) return 0;
//     if (l <= U.l and U.r <= r) return U.Get_val();
//     Push_Down(u);
//     int mid = (U.l + U.r) >> 1, ans = 0;
//     if (l <= mid) ans += Query(ls, l, r);
//     if (r > mid) ans += Query(rs, l, r);
//     return ans;
// }
int Query(int u, int l, int r, int i)
{
    if (!U.cnt[i]) return 0;
    if (U.l > r or U.r < l) return 0;
    if (l <= U.l and U.r <= r) return U.cnt[i];
    Push_Down(u);
    int mid = (U.l + U.r) >> 1, ans = 0;
    if (l <= mid) ans += Query(ls, l, r, i);
    if (r > mid) ans += Query(rs, l, r, i);
    return ans;
}
auto main() -> signed
{
    sq[0] = 1; for (int i = 1; i <= 30; i ++ ) sq[i] = sq[i - 1] * 2;

    std::cin >> n >> m;
    for (int i = 1; i <= n; i ++ ) std::cin >> a[i];
    // for (int i = 0; i <= 30; i ++ )
    // {
    //     std::cout << "2 ^ " << i << " : " << sq[i] << std::endl;
    // }
    build(1, 1, n);
    while (m -- )
    {
        int op, l, r, x;
        std::cin >> op >> l >> r;
        if (op == 1)
        {
            // std::cout << Query(1, l, r) << std::endl;
            int ans = 0;
            for (int i = 0; i <= 30; i ++ )
                ans += sq[i] * Query(1, l, r, i);
            std::cout << ans << std::endl;
        } 
        else 
        {
            std::cin >> x;
            Modify(1, l, r, x);
        }
    }
    return 0;
}
2023/9/8 08:41
加载中...