题解 AT3897 【ニワンゴくんとゲーム】

· · 题解

我的博客的该篇题解 (然而并没有更好的排版)
另外做本题是参考了国外的某博客

题意:求递推式

f(x)=\left\{\begin{aligned}f[x-1] + f[floor( \frac{x}{2} )] , x \geq 2 \\1,x = 1\end{aligned} \right.

的第n项,数据范围n \geq 10^{18},询问次数在200以内。
非常清楚地知道这是个矩阵乘法,然而并不会构矩阵。 我们发现求f[x]必须要用到f[floor(\frac{x}{2})],那么我们的矩阵里必须要包含这一元素,但是我们不可能保存所有的f[x]。我们考虑对于f[floor(\frac{x}{2})],那么必须使用f[floor(\frac{x}{4})]。这样递归下去,我们可以考虑储存这样一个向量:

M(x)=\begin{pmatrix} f[x] \\ f[ceil(\frac{x}{2})] \\ f[ceil( \frac{x}{2^2})] \\ \vdots \\ f[ceil(\frac{x}{2^{59}})] \end{pmatrix}

我们这里将floor 改成了ceil,这是解释:我们考虑从M(x)M(x+1)转移,那么需要的是floor( \frac{x+1}{2}),这个是与ceil(\frac{x}{2})相等的,所以我们直接存ceil,这样就可以直接转移。
接下来考虑如何完整地转移整个向量。对于向量的任何一个位置f[y],转移后为f[y']如果在M(x)转移至M(x+1)过程中,这个位置的值改变,那么这个值一定是从f[y]变成了f[y+1],那么其就满足f[y']=f[y+1]=f[y]+f[ceil(\frac{y}{2})],否则f[y']=f[y]。这样我们就可以构造出每种转移的矩阵。
接下来我们考虑什么情况会导致哪些位置的函数值改变。容易发现当某个数y二进制末位是0的时候,对这个数+1就会对向量M(y+1)的第二维产生影响(考虑+1以后ceil(\frac{x}{2})改变)。同理可得,如果某个数y二进制末尾有n个连续的0,那么对应向量的前(n+1)维就会产生改变,并且其他维度的值不改变。 这个构造矩阵只要把对应的转移向量拼起来就好。
但是有一个根本的问题没有解决:我们构造了矩阵,但是无法利用矩阵进行加速啊。
我们谈到,这个转移只跟y末尾0的个数有关,那么我们可以对一个yO(60)的时间内求出其对应的转移矩阵。
设从M(x)M(x+1)转移的矩阵为N(x),那么有

M(x)=N(x-1) \times N(x-2) \times \dots \times N(2) \times N(1) * M(1)

考虑结合律,假设x最高位是第b位,基于转移只跟数末尾的0有关和矩阵乘法结合律,那么这个式子可以化作:

M(X)=\left[ N(x-2^b) \times \dots \times N(1) \right] \times N(2^b) \times \left[ N(2^b-1) \times \dots \times N(1) \right] \times M(1)

那么后半段方括号内就是12^n-1的所有转移矩阵的积,这个总共只有六十个,我们可以把它们预处理下来,然后中间那项直接求出来乘上去,前面半段就可以递归下去做了。
复杂度:O(log^{4} n)
代码:

#include<cstdio>
#include<cstdlib>
#include<cstring>
using  namespace std;
typedef long long LL;
const int MOD=1000000007;
LL kano()
{
    char ch=getchar();LL w=0,u=1;
    for(;ch<'0'||ch>'9';ch=getchar())if(ch=='-')u=-1;
    for(;ch>='0'&&ch<='9';ch=getchar())w=w*10+ch-'0';
    return w*u;
}

LL R[65];
LL Q,n,m;
struct MATRIX
{
    int a[65][65];
    int n,m;
    MATRIX(){n=m=0;memset(a,0,sizeof a);}
    MATRIX(LL w)
    {
        n=m=60;memset(a,0,sizeof a);
        for(LL i=0,j=w;i<60;i++,j=j>>1)
        {
            a[i][i]=1;
            a[i+1][i]=j&1;
        }
    }
}p[60],ans;
MATRIX &operator *(const MATRIX &a,const MATRIX &b)
{
    static MATRIX ans;
    ans.n=a.n;ans.m=b.m;
    memset(ans.a,0,sizeof ans.a);
    for (int l=0;l<a.m;l++)
    {
        for(int i=0;i<ans.m;i++)
        {
            if(b.a[l][i]==0)continue;
            for(int j=0;j<ans.n;j++)
            {
                ans.a[j][i]=(ans.a[j][i]+(LL)a.a[j][l]*b.a[l][i])%MOD;
            }
        }
    }
    return ans;
}

void work(LL n,int dep)
{
    if(n==0)return;
    if((n&R[dep])==R[dep])
    {
        ans=ans*p[dep+1];
        return;
    }

    if(n&1)
    {
        work((1ll<<dep)-1,dep-1);
        ans=ans*MATRIX((1ll<<(dep+1))-1);
    }
    work(n>>1,dep-1);
}

int main()
{   
    p[0]=MATRIX(0);
    for(LL i=1,j=2;i<60;i++,j=j<<1)
    {
        p[i]=MATRIX(j-1);
        p[i]=p[i]*p[i-1];
        p[i]=p[i-1]*p[i];
    }
    Q=kano();
    R[0]=1;
    for(int i=1;i<=60;i++)R[i]=R[i-1]<<1|1;
    while(Q--)
    {
        ans.n=1;ans.m=60;
        for(int i=0;i<ans.m;i++)ans.a[0][i]=1;
        n=kano()-1;
        m=0;
        for(int i=0;i<60;i++,n=n>>1)m=m<<1|(n&1);
        n=m;
        work(n,59);
        printf("%d\n",ans.a[0][0]);
    }
}