CSUST OJ 2006-Simple Inversions(带修主席树)

传送门
题面:

题解:

考虑维护每次操作后的逆序对数量。

假设我们已经知道操作前逆序对数量为\(ans\)

交换 \(l\)\(r\) 只会影响区间 \((l,r)\) 。考虑操作前区间\((l,r)\) 里面有多少能跟 \(l\) 形成逆序对 ,操作后区间\((l,r)\) 里面有多少能跟 \(l\)​ 形成逆序对 ,即求出区间(l,r)里面有多少数小于 \(l\) 即可。对 \(r\) 同理。那么这里可以使用主席树。

但是交换两个数后,主席树里面的东西也会发生改变,因为是单点修改,所以要使用带修主席树。

普通主席树维护的是前缀,而带修主席树主席树则像树状数组一样,维护一个特定的区间,每次只要像树状数组一样去修改即可。

代码:

#include<cstdio>
#include<iostream>
#include<algorithm>
#include<cstring>
#include<cmath>
#include<queue>
#include<map>
#include<stack>
#include<set>
#include<ctime>
#define iss ios::sync_with_stdio(false)
using namespace std;
typedef unsigned long long ull;
typedef long long ll;
typedef pair<int,int> pii;
const int mod=1e9+7;
const int MAXN=2e5+5;
const int inf=0x3f3f3f3f;
int a[MAXN], n, m;
int tree[MAXN];
struct node
{
    int l, r;
    int sum;
    /* data */
}node[20000000+5];
int cnt = 0;
int lowbit(int i)
{
    return i & (-i);
}
void insert(int l,int r,int pos,int &now,int val)
{
    if(!now){
        now = ++cnt;
    }
    node[now].sum += val;
    if(l==r)
        return;
    int mid = (l + r) >> 1;
    if(pos<=mid)
        insert(l, mid, pos, node[now].l,val);
    else
        insert(mid + 1, r, pos, node[now].r, val);
}
void add(int i,int val,int f)
{
    while(i<=n){
        insert(1, n, val, tree[i], f);
        i += lowbit(i);
    }
}
int query(int l,int r,int now,int val)
{
    if(!now)
        return 0;
    if(r<=val)
        return node[now].sum;
    else if(l>val)
        return 0;
    int mid = (l + r) >> 1;
    if(mid<=val)
        return query(l, mid, node[now].l, val) + query(mid + 1, r, node[now].r, val);
    else
        return query(l, mid, node[now].l, val);
}
int sum(int i,int val)
{
    int res = 0;
    while(i)
    {
        res += query(1,n,tree[i],val);
        i -= lowbit(i);
    }
    return res;
}
int main()
{
    scanf("%d%d", &n, &m);
    ll ans = 0;
    for (int i = 1; i <= n;i++){
        a[i] = i;
        add(i, i, 1);
    }
    //cout << tree[1] << endl;
    //cout << sum(1, 1) << endl;
    while(m--)
    {
        int l, r;
        scanf("%d%d", &l, &r);
        int s1 = sum(l, a[l]-1);
        int s2 = sum(r, a[l] - 1);
        //cout << s1 << " " << s2 << endl;
        ans = ans - (s2 - s1);
        ans = ans + ((r - l) - (s2 - s1));
        int s3 = sum(l - 1, a[r] - 1);
        int s4 = sum(r - 1, a[r] - 1);
        //cout << s4 << " " << s3 << endl;
        ans = ans - ((r - l) - (s4 - s3));
        ans = ans + (s4 - s3);
        if(a[l]<a[r])
            ans--;
        else
            ans++;
        add(l, a[l], -1);
        add(r, a[r], -1);
        add(l, a[r], 1);
        add(r, a[l], 1);
        swap(a[l], a[r]);
        printf("%lld\n", ans);
    }
}
posted @ 2021-09-15 23:36  TheBestQAQ  阅读(66)  评论(0)    收藏  举报