树状数组是一个很简洁很方便很好用的一个结构,相比起线段树有如下几个优点:

  1. 节约空间:树状数组规模为 nn ,但是线段树至少有 2n2n 个节点,再算上其他维护的数据域,空间一般比树状数组大 1010 倍;
  2. 编程复杂度低:线段树比树状数组编码复杂度要高许多,这一点大家自己编程体会;
  3. 时间复杂度低:线段树维护的数据域太多,操作的时间复杂度要比树状数组高。

但是它能替代线段树吗?

答案是否定的。

线段树能完成许多树状数组不能完成的操作。所以所有能用树状数组做的题,几乎都能用线段树做。

一、原理

树状数组的原理并不复杂,其实就是建立一个类似于树形结构一样的数组,使每次单点更新和区间查询的复杂度变小。

这就要用到二进制。树状数组是通过二进制来存储数据的,根据二进制,使每一个数组变量存的数值不会太多,也不会太少。例如下图:

方格中数字代表对应数组的第几个元素,下排是 aa 数组,其上方的是 ee 数组,最下的二进制则是对应编号的二进制表示.箭头表示这个数组元素被哪个数组元素包含了,比如 e[2]=e[1]+a[2]=a[1]+a[2],e[4]=e[2]+a[3]+a[4]=a[1]+a[2]+a[3]+a[4].e[2]=e[1]+a[2]=a[1]+a[2], e[4]=e[2]+a[3]+a[4]=a[1]+a[2]+a[3]+a[4].

通过观察可以发现:

  1. 每个元素至多仅被一个元素包含,这点和树有很大相同,但整体并不是树

  2. 每个 e[i]e[i] 可认为是仅包含 a[i]a[i] 和其它若干个 ee 元素

  3. 每个 e[i]e[i] 包含的元素数目(包括 a[i]a[i] 在内)为 2k2^k (其中 kkii 二进制末尾 00 的个数)

    但是对应二进制数的最末连续0的个数如何得知呢?

    位运算!

    在实际应用中,我们需要的不是最末连续的 00 的个数,而是最末那段 10001000 对应的十进制数(虽然二者显然可以互推)

    求最末连续的 00 的个数的操作,我们叫做 lowbitlowbit 操作。

    lowbit(x) = (((x1) xor x) and x)= ((x) and x)lowbit(x)\ =\ (((x-1)\ xor\ x)\ and\ x)=\ ((-x)\ and\ x)

    由于计算机的补码操作,在存储负数时, x = (not x) + 1-x\ =\ (not\ x)\ +\ 1

二、操作

树状数组主要有两个操作:

1. 单点更新

2. 区间查询

通过这两个基础操作,根据不同的数组,可以完成更多的操作。

现在来分别讲一下这两个操作。

1. 单点更新

由上面的图可知,如果我们要更改一个值,就需要把它所有的祖宗节点的数值修改。访问祖宗节点可以通过递归访问父节点来操作。

但是怎么访问父节点呢?

仔细观察上面的图,可以发现,父亲节点就是比它大的,离它最近的,末位连续 00 比它多的数,所以可以得到这样一个结论:

father[x]=x+lowbit(x) father[x]=x+lowbit(x)

然后递归访问就行了。

void update(int x,int val)//x是修改的数的位置,val是修改时增加的值
{
    for(int i=x;i<=n;i+=lowbit(i))//从x开始,一直到数组访问结束(即到根节点),每次访问父节点的位置
        tree[i]+=val;//修改操作,直接增加就可以了。
}

2. 区间查询

由上图可知,当树状数组每一个数包含的信息并不是所求的全部信息时,我们就需要再找一个数来包含剩下的所要求的信息,这个数我们称之为前驱

仔细观察可得,前驱的下标即为比自己小的最近的最末连续 00 比自己多的数。

所以前驱的下标为:

last[x]=xlowbit(x) last[x]=x-lowbit(x)

区间查询和前缀和的记录方式类似,先算出从 1x1\sim x 的和,再求区间的和。

int ask(int x)//x为修改的位置
{
    int sum=0;
    for(int i=x;i;i-=lowbit(i))//从x一直加到数组开始,即从1加到x
        sum+=tree[i];
   return sum;//返回前缀和
}
int range_ask(int l,int r)
{
    return ask(r)-ask(l-1);//求区间[l,r]的和
}

三、高级操作

树状数组的基本操作可以推出很多高等操作,接下来一一介绍。

初始化

树状数组的初始化其实有一些高级的操作。

一般如果要将数据添加进树状数组,直接将 aa 数组中的元素一个一个加进去就行了,预处理时间复杂度为 O(nlog2n)O(nlog_2n)

但是下面的算法更优美。

我们知道树状数组中,满足这样一个规律:

tree[i]=a[ilowbit(i)+1]++a[i] tree[i]=a[i-lowbit(i)+1]+……+a[i]

所以我们只要一开始维护一个前缀和的数组满足:

sum[i]=a[1]+a[2]++a[i] sum[i]=a[1]+a[2]+……+a[i]

这样树状数组就可以借助这个前缀和数组如下初始化了:

tree[i]=sum[i]sum[ilowbit(i)] tree[i]=sum[i]-sum[i-lowbit(i)]

特别地,如果 aa 数组一开始的值全是 11 ,那么 tree[i]tree[i] 的值就是 lowbit(i)lowbit(i)

求逆序对

求逆序对我们已经知道可以用归并排序求,但是这里介绍一种用树状数组求的方法,时间复杂度为 O(nlog n)O(nlog\ n)

首先开一个数组 s[i]s[i] 来记录前面的数据的出现情况,初始化全为 00 。当数据 xx 出现时,就令 s[x]=1s[x]=1 。这样的话,如果要求某个数 xx 的逆序对数,只需要算出在当前状态下 s[x+1]s[maxn]s[x+1]\sim s[maxn] 中有多少个 11 ,因为这些位置都在 xx 之后,说明出现的数据都比 xx 大,并且都在 xx 之前插入,否则的话它们也不会被 ss 数组记录为 11 。但是如果每添加一个数据 xx ,就要从 x+1maxnx+1\sim maxn 搜一遍,复杂度会很高。树状数组则完美解决了这个问题,因为状态标记为 11 ,所以只要对 s[x+1]s[maxn]s[x+1]\sim s[maxn] 进行求和就可以了,所以答案就是 xask(s[x])x-ask(s[x]) 。(这里的 s[x]s[x] 就是树状数组的 tree[x]tree[x] )。

但是万一 maxnmaxn 过大,数组空间装不下怎么办?

离散化!!!

这个就不具体说了,但是注意当有数据 00 出现时,由于 lowbit(0)=0lowbit(0)=0 ,此时离散化是无法更新数据的,此时只要加一个数就可以了。

这里借一下网上一个大佬Anxdada的代码。

#include<cstdio>
#include<algorithm>
#include<cstring>
#define CLR(x) memset(x,0,sizeof(x))
#define ll long long int
#define PI acos(-1.0)
#define db double
#define mod 1000000007
using namespace std;
const int maxn=1e5+5;
const db eps=1e-6;
const int inf=1e9;
const ll INF=1e15;
int c[maxn];         //树状数组
int lisan[maxn];   //用来离散的存离散后的结果数组
int s[maxn];        //用来存原始状态
int n,k;
//树状数组
int lowbit(int x)  //返回值最大1e5.
{
    return x&(-x);
}

void update(int i,int ans)
{
    while(i <= n){
        c[i] += ans;
        i += lowbit(i);
    }
}

int sum(int i)
{
    int s = 0;
    while(i > 0){
        s += c[i];
        i -= lowbit(i);
    }
    return s;
}
//结束
int main()
{
    while(~scanf("%d%d",&n,&k)){
        CLR(c);
        for(int i=0;i<n;i++){
            scanf("%d",&s[i]);
            lisan[i] = s[i];
        }
        sort(lisan,lisan+n);
        int len=unique(lisan,lisan+n)-lisan;
        ll res = 0;
        for(int i=0;i<n;i++){
            ll l=lower_bound(lisan,lisan+len,s[i])-lisan+1;
            update(l,1);
            res += i+1 - sum(l);
        }
        if(res < k) printf("0\n");
        else
            printf("%lld\n",res-k);
    }
}

求区间最大/最小值

这个很简单,只需要改一下维护的数组操作就好了。

原来是维护前缀和,所以每访问一个数就相加,现在给成每访问一个数就求一次最大/最小值就好了。

void update(int x)
{
	for(int i=1;i<lowbit(x);i<<=1)
	    e[x]=max(e[x],e[x-i]);
}
int ask(int l,int r)
{
    int sum=a[r];
    while(l<=r)
    {
    	sum=max(sum,a[r]);
    	for(--r;r>=l+lowbit(r);r-=lowbit(r))
    		sum=max(sum,e[r]);
    }
    return sum;
}

下标查询

这个操作是这样的:如果我们知道一个前缀和,如何查询这个前缀和对应的前缀下标 xx 呢(保证元素非负)?

这个很简单嘛,元素非负,那么前缀和随着下标单调不下降,那么就二分查嘛。

时间复杂度为 O(log2n)O(\log ^2 n)

但是我们想要进一步优化!

通过 lowbitlowbit 函数的定义可知,下标为 22 的幂次的前缀和是包含了最开始的元素到自己的所有元素。

所以我们把二分改成倍增嘛img

每次倍增以 22 的整数次幂为步长,能累加则累加。

设给的前缀和的值为 valval ,从 log2(val)log_2(val)00 倒序枚举长度,如果 ans+(1<<i)nans+(1<<i)\leq nsum+tree[ans+(1<<i)]<valsum+tree[ans+(1<<i)]<val 说明可以倍增,就更新 ansanssumsum ,最后 ans+1ans+1 就是要求的下标。

int get(int val)
{
    int ans=0;
    int sum=0;
    int len=(int)\log2(n);
    for(int i=len;i;i--)
        if(ans+(1<<i)<=n&&sum+tree[ans+(1<<i)]<val)
        {
            sum+=tree[ans+(1<<i)];
            ans+=(1<<i);
        }
    return ans+1;
}

成倍扩张/缩减

如果我们想让原来的数组 aa 中所有元素都扩大 mm 倍或缩小 mm 倍,怎么办?

两种思路。

第一种很简单,将 aa 数组修改,再重新构建一个树状数组。

第二种则直接对树状数组进行修改。

如果直接让树状数组中所有元素都直接扩大或缩小的话,缩小的时候,由于取整的原因,可能会造成树状数组不连续,那怎么办呢?

倒序修改即可。

这里给出缩小 nn 倍的代码。

void change(int m)
{
    for(int i=n;i;i--)
        update(i,point_ask(i)/m-point_ask(i)));//这里的point_ask是单点查询操作,下一个就讲
}

单点查询

我们知道树状数组能快速做到区间查询,那怎么快速做到单点查询呢?

首先我们肯定想到的是直接用两次区间查询相减来求出单点数值。

a[i]=ask(i)ask(i1) a[i]=ask(i)-ask(i-1)

那有没有更优秀的算法呢?

那必须有。

用下面这个公式能更快一些,相当于只进行了一次区间查询:

a[i]=tree[i](ask(i1)ask(LCA(i,i1))) a[i]=tree[i]-(ask(i-1)-ask(LCA(i,i-1)))

其中 LCALCA 代表最近公共祖先。

原理也比较容易理解,我们区间查询时是从当前节点一直查到数组开始,这样求得前缀和,但是如果我们只用求一个单点的值,这是很不值得,因为前面一些无关的元素被算上了,升高了时间复杂度,所以,我们只要在区间查询的代码上略作修改,将访问节点访问至 00 停止改为访问到 LCALCA 停止,便能节省很多时间。

很明显,看图可以发现:

LCA(i,i1)=ilowbit(i) LCA(i,i-1)=i-lowbit(i)

然后就很好搞了嘛。

int point_ask(int x)
{
    int sum=tree[x];
    int lca=x-lowbit(x);
    for(int i=x-1;i!=lca;i-=lowbit(i))
        sum-=tree[i];
    return sum;
}

还有没有其他的算法?

当然也有。

我们知道树状数组维护的是前缀和,所以我们只要维护一个数组,让这个数组的前缀和等于这个数就可以了。

这种数组我们叫做差分数组

假设原数组为 a[i]a[i] ,设差分数组为 d[i]d[i] ,令 d[i]=a[i]a[i1](a[0]=0)d[i]=a[i]-a [ i-1 ] (a[0] = 0) ,这个时候,我们就可以满足 a[i]=j=1id[j]a[i]=\sum _{j=1} ^i d[j] ,,然后就可以通过求 d[i]d[i] 的前缀和查询。

单点查询代码和区间查询代码一模一样(毕竟原理都是查询前缀和),这里不再放出来。

可能有人会问,既然都能区间查询了,为什么不直接查询两个区间然后相减来完成单点查询呢?

不要急,下一个操作就要用到了。

区间修改

树状数组可以单点修改,那区间修改呢?

将区间中的每个数一一单点修改?

时间复杂度会很高,不可行。

这个时候我们就要用到上面看似没用的单点查询了。

还是用差分数组,我们考虑,当给区间 [l,r][l,r] 加上 xx 的时候,在差分数组中有哪些变化呢?

很明显, a[l]a[l]a[l1]a[l-1] 的差会增加 xxa[r+1]a[r+1]a[r]a[r] 的差会减少 xx

根据差分数组的定义,只要给 d[l]d[l] 加上 xx ,给 d[r+1]d[r+1] 减去 xx 就可以了。

void update(int x,int val)
{
    for(int i=x;i<=n;i+=lowbit(i))
        tree[i]+=val;
}
void range_update(int l,int r,int val)//区间[l,r]增加val
{
    update(l,val);
    update(r+1,-1*val);
}

区间查询

这里的区间查询是指维护差分数组时的区间查询。

维护差分数组时,我们知道可以单点查询,但是区间查询呢?

从头来分析一下。

首先,我们的问题是要解决位置 xx 的前缀和,即求 i=1xa[i]\sum_{i=1}^xa[i]

由差分数组定义得: i=1xa[i]=i=1xj=1id[j]\sum_ {i=1} ^x a[i]= \sum _{i=1} ^x \sum _{j=1} ^i d[j]

再来分析推出的式子: i=1xj=1id[j]\sum _{i=1} ^x \sum _{j=1} ^i d[j]

在这个式子中,不难发现 d[1]d[1] 被用了 xx 次, d[2]d[2] 被用了 x1x-1 次……

这就好办了,可以发现:

i=1xj=1id[j]=i=1xd[i]×(xi+1)=(x+1)×i=1xd[i]i=1xd[i]×i \sum _{i=1} ^x \sum _{j=1} ^i d[j]= \sum _{i=1} ^x d[i]\times(x-i+1)=(x+1)\times \sum _{i=1} ^x d[i]-\sum _{i=1} ^x d[i]\times i

维护两个前缀和就好了嘛。

第一个数组维护 d[i]d[i]

第二个数组维护 d[i]×id[i]\times i

所以这就是对差分数组的运用。

需要注意的是再区间修改第二个数组时,给 tree2[l]tree2[l] 加上 l×vall\times val ,给 tree2[r+1]tree2[r+1] 减去 (r+1)×val(r+1)\times val

void update(int x,int val)
{
    for(int i=x;i<=n;i+=lowbit(i))
    {
        tree1[i]+=val;
        tree2[i]+=val*x;//注意!!!这里tree2加上的是val*x并非val*i,因为修改的点是x,所以x的祖宗节点加上的值也只和x有关。
    }
}
void range_update(int l,int r,int val)
{
    update(l,val);
    update(r+1,-1*val);
}
int ask(int x)
{
    int sum=0;
    for(int i=x;i;i-=lowbit(i))
       sum+=((x+1)*tree1[i]-tree2[i]);
    return sum;
}
int range_ask(int l,int r)
{
    return ask(r)-ask(l-1);
}

二维树状数组

我们已经知道树状数组可以用于一个一维数组上,那么可不可以用到二维树状数组上呢?

在一维树状数组中, tree[x]tree[x] 代表的是记录当前节点为 xx ,长度为 lowbit(x)lowbit(x) 的前缀和。

所以在二维树状数组当中,定义 tree[x][y]tree[x][y] 记录的是当前节点为 (x,y)(x,y) ,长为 lowbit(x)lowbit(x) ,宽维 lowbit(y)lowbit(y) 的前缀和。

所以单点修改和区间查询的操作就改成了二维的了。

void update(int x,int y,int val)
{
    for(int i=x;i<=n;i+=lowbit(i))
        for(int j=y;j<=n;j+=lowbit(j))
            tree[i][j]+=val;
}
void ask(int x,int y)
{
    int sum=0;
    for(int i=x;i;i-=lowbit(i))
        for(int j=y;j;j-=lowbit(j))
            sum+=tree[i][j];
    return sum;
}

那么二维树状数组的差分数组操作该怎么变化呢?

首先搞明白二维的前缀和怎么求。

sum[i][j]=sum[i1][j]+sum[i][j1]sum[i1][j1]+a[i][j] sum[i][j]=sum[i-1][j]+sum[i][j-1]-sum[i-1][j-1]+a[i][j]

这就好办了,要让前缀和为 a[i][j]a[i][j] ,我们就令差分数组 d[i][j]=a[i][j]a[i1][j]a[i][j1]+a[i1][j1]d[i][j]=a[i][j]-a[i-1][j]-a[i][j-1]+a[i-1][j-1] 就好了。

例如下面这个矩阵:

1 4 8
6 7 2
3 9 5

对应的差分数组就是:

1 3 4
5 -2 -9
-3 5 1

但是我们怎么进行二维树状数组维护差分数组时的区间修改呢?

比如说有下面这样一个 555*5 的空矩阵:

0 0 0 0 0
0 0 0 0 0
0 0 0 0 0
0 0 0 0 0
0 0 0 0 0

我们要给中间那个 333*3 的矩阵加上 xx ,那它就会变成下面这个模样:

0 0 0 0 0
0 xx xx xx 0
0 xx xx xx 0
0 xx xx xx 0
0 0 0 0 0

此时这个数组的差分数组为:

0 0 0 0 0
0 xx 0 0 x-x
0 0 0 0 0
0 0 0 0 0
0 x-x 0 0 xx

这下就好办了,规律就是,当我们要给以 A(x1,y1),B(x2,y2)A(x_1,y_1),B(x_2,y_2) 构成的矩形加上 xx 时,

我们只用给差分数组中的 (x1,y1),(x2+1,y2+1)(x_1,y_1),(x_2+1,y_2+1) 两个点加上 xx ,给差分数组中的 (x1,y2+1),(x2+1,y1)(x_1,y_2+1),(x_2+1,y_1) 两个点减去 xx 就可以了。

void update(int x,int y,int val)
{
    for(int i=x;i<=n;i+=lowbit(i))
        for(int j=y;j<=n;j+=lowbit(j))
            tree[i][j]+=val;
}
void range_update(int x1,int y1,int x2,int y2,int val)
{
    update(x1,y1,val);
    update(x1,y2+1,-1*val);
    update(x2+1,y1,-1*val);
    update(x2+1,y2+1,val);
}

然后就是二维树状数组维护差分数组时的区间修改操作了。

这个操作放到二维来显然更难了。

还是来推一遍,可以得到,点 (x,y)(x,y) 的二维前缀和为 i=1xj=1yk=1il=1jd[k][l]\sum _{i=1} ^x \sum _{j=1} ^y \sum _{k=1} ^i \sum _{l=1} ^j d[k][l]

这个式子看起来很复杂,毕竟有 O(n4)O(n ^4) 的复杂度,但是利用树状数组,我们可以优化到 O(n2log2n)O(n ^2 \log _2 n)

还是类比一下一维树状数组,统计一下每个 d[k][l]d[k][l] 出现过多少次。很明显, d[1][1]d[1][1] 出现了 x×yx\times y 次, d[1][2]d[1][2] 出现了 x×(y1)x\times (y-1) 次…… d[k][l]d[k][l] 出现了 (xk+1)×(yl+1)(x-k+1)\times (y-l+1) 次。

这个式子就可以写成:

i=1xj=1yd[i][j]×(xi+1)×(yj+1) \sum _{i=1} ^x \sum _{j=1} ^y d[i][j]\times (x-i+1)\times (y-j+1)

展开,可以得到:

(x+1)×(y+1)×i=1xj=1yd[i][j](y+1)×i=1xj=1yd[i][j]×i(x+1)×i=1xj=1yd[i][j]×j+i=1xj=1yd[i][j]×i×j (x+1)\times (y+1)\times \sum _{i=1} ^x \sum _{j=1} ^y d[i][j]-(y+1)\times \sum _{i=1} ^x \sum _{j=1} ^y d[i][j]\times i-(x+1)\times \sum _{i=1} ^x\sum _{j=1} ^y d[i][j]\times j+\sum _{i=1} ^x\sum _{j=1} ^y d[i][j]\times i\times j

然后开四个树状数组,分别维护以下四个数组就可以了:

d[i][j],d[i][j]×i,d[i][j]×j,d[i][j]×i×jd[i][j],d[i][j]\times i,d[i][j]\times j,d[i][j]\times i\times j

最后贴上网上一位大佬的代码,顺便说一句博客来源也是他大佬bestsort

#include <cstdio>
#include <cmath>
#include <cstring>
#include <algorithm>
#include <iostream>
using namespace std;
typedef long long ll;
ll read(){
    char c; bool op = 0;
    while((c = getchar()) < '0' || c > '9')
        if(c == '-') op = 1;
    ll res = c - '0';
    while((c = getchar()) >= '0' && c <= '9')
        res = res * 10 + c - '0';
    return op ? -res : res;
}
const int N = 205;
ll n, m, Q;
ll t1[N][N], t2[N][N], t3[N][N], t4[N][N];
void add(ll x, ll y, ll z){
    for(int X = x; X <= n; X += X & -X)
        for(int Y = y; Y <= m; Y += Y & -Y){
            t1[X][Y] += z;
            t2[X][Y] += z * x;
            t3[X][Y] += z * y;
            t4[X][Y] += z * x * y;
        }
}
void range_add(ll xa, ll ya, ll xb, ll yb, ll z){ //(xa, ya) 到 (xb, yb) 的矩形
    add(xa, ya, z);
    add(xa, yb + 1, -z);
    add(xb + 1, ya, -z);
    add(xb + 1, yb + 1, z);
}
ll ask(ll x, ll y){
    ll res = 0;
    for(int i = x; i; i -= i & -i)
        for(int j = y; j; j -= j & -j)
            res += (x + 1) * (y + 1) * t1[i][j]
                - (y + 1) * t2[i][j]
                - (x + 1) * t3[i][j]
                + t4[i][j];
    return res;
}
ll range_ask(ll xa, ll ya, ll xb, ll yb){
    return ask(xb, yb) - ask(xb, ya - 1) - ask(xa - 1, yb) + ask(xa - 1, ya - 1);
}
int main(){
    n = read(), m = read(), Q = read();
    for(int i = 1; i <= n; i++){
        for(int j = 1; j <= m; j++){
            ll z = read();
            range_add(i, j, i, j, z);
        }
    }
    while(Q--){
        ll ya = read(), xa = read(), yb = read(), xb = read(), z = read(), a = read();
        if(range_ask(xa, ya, xb, yb) < z * (xb - xa + 1) * (yb - ya + 1))
            range_add(xa, ya, xb, yb, a);
    }
    for(int i = 1; i <= n; i++){
        for(int j = 1; j <= m; j++)
            printf("%lld ", range_ask(i, j, i, j));
        putchar('\n');
    }
    return 0;
}