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);
}
}

浙公网安备 33010602011771号