题解:CF2252E Generational Triplets

· · 题解

一种并非数位 dp 的做法。我们设 f(i) 表示 n=i 时的答案。我们考虑一组合法的 (a,b,c),若由它能构造出更小的 (a',b',c'),那么我们可以使用递归的方法得出答案。

我们不妨设公差为 d,则 a \oplus (a+d)=(a+2d)

考虑 a,d 的奇偶性。经过简单分讨得 a,d 奇偶性必须相同,于是我们分类讨论。

a,d 均为偶数

不妨设 a=2x,b=2y,c=2z,则 x+z=2y2x\oplus 2y=2z。我们注意到这样的 (x,y,z) 恰好也是满足条件的三元组,且一组均为偶数的 (a,b,c) 恰好对应一组 (x,y,z)。这样的三元组有 f(\lfloor\frac{n}{2}\rfloor) 个。

a,d 均为奇数

不妨设 a=2x+1,b=2y,c=2z+1。则 x+z+1=2y,且 (2x+1)\oplus (2z+1)=2y。后式可化简为 x\oplus z=y,所以有 x+z+1=2(x\oplus z)。由于 x+z=(x\oplus z)+2(x\& z),故 x\oplus z=2(x\&z)+1

再次对 x,z 的奇偶性讨论,由于等式右边为奇数,故 x,z 奇偶性不同。

x 为偶数

x=2s,z=2t+1。代入上式得 s\oplus t=2(s\&t)。又因为 s+t=(s\oplus t)+2(s\&t),有 s+t=2(s\oplus t)

x 为奇数

同样地,令 x=2s+1,z=2t。经过与上一种情况类似的化简,也可以得到 s+t=2(s\oplus t)

综合以上两种讨论,我们得到 s+t=2(s\oplus t)。我们构造 a'=s,b'=s\oplus t,c'=t,不难发现 (a',b',c') 为满足题目条件的三元组。这样的三元组有 f(\lfloor\frac{m}{2}\rfloor)+f(\lfloor\frac{m-1}{2}\rfloor) 个,其中 m=\lfloor\frac{n-1}{2}\rfloor

综上,我们有递推公式 f(n)=f(\lfloor\frac{n}{2}\rfloor)+f(\lfloor\frac{m}{2}\rfloor)+f(\lfloor\frac{m-1}{2}\rfloor)+1,其中 m=\lfloor\frac{n-1}{2}\rfloor。值得注意的是 (1,2,3) 无法进行递归,递推式最后一项即为这组的贡献。

#include<bits/stdc++.h>
#include<bits/extc++.h>
#define pii pair<int,int>
#define fi first
#define se second
#define pb push_back
#define int long long
#define gc getchar
//char buf[1<<20],*p1,*p2;
//#define gc() (p1==p2&&(p2=(p1=buf)+fread(buf,1,1<<20,stdin),p1==p2)?EOF:*p1++)
#define R read()
using namespace std;
int read()
{
    int x=0,f=1;
    char c=gc();
    while(c>'9'||c<'0'){if(c=='-') f=-1;c=gc();}
    while(c>='0'&&c<='9') x=(x<<1)+(x<<3)+c-48,c=gc();
    return x*f;
}
void write(int x,char xx)
{
    static int st[35],top=0;
    if(x<0){x=-x;putchar('-');}
    do
    {
        st[top++]=x%10,x/=10;
    }while(x);
    while(top) putchar(st[--top]+48);
    putchar(xx);
}
#define mod 1000000007
int lp(int x,int y){return x+y>=mod?x+y-mod:x+y;}
void pl(int &x,int y){x=lp(x,y);}
using namespace __gnu_pbds;
int n;
unordered_map<int,int>f;
int dfs(int n)
{
    if(n<3) return 0;
    if(f.count(n)) return f[n];
    int k=n-1>>1,ans=lp(lp(dfs(n>>1),1),lp(dfs(k>>1),dfs(k-1>>1)));
    return f[n]=ans;
}
void solve()
{
    n=R,write(dfs(n),'\n');
}
int T=1;
signed main()
{
    T=R;
    while(T--) solve();
    return 0;
}