题解:P17293 [Algo Beat Contest 013 & MSOI R2] 好朋友

· · 题解

题意简述

构造 n[0,m] 内的整数。所有有序三元组中三对异或的二进制 1 数量之和必须恰为 k。若不存在这样的数列,输出 -1

解题思路

先把三元组贡献改写为数对贡献。固定无序下标对 \{i,j\}。它可以按两个方向出现在三元组的任意一条边上,第三个下标又有 n 种选择。三条边的地位相同,所以这对数一共贡献:

6n\operatorname{popcount}(a_i\operatorname{xor}a_j)

因此,若 k 不是 6n 的倍数,一定无解。此后只需构造所有无序数对的异或位数总和。

逐个二进制位考虑。若某一位在 c 个数中为 1,就在 n-c 个数中为 0。恰有 c(n-c) 个无序数对在这一位不同。设可用位数为 b=\lfloor\log_2m\rfloor+1,则目标等式为:

k=6n\sum_{i=0}^{b-1}c_i(n-c_i)

将某一位的全部数位取反,会把 c_i 变成 n-c_i,贡献保持不变。因此,若存在一组可行的数量,就存在满足 0\le c_i\le\lfloor n/2\rfloor 的一组数量。

s=k/(6n)。现在每个二进制位都要选择一个 c_i,且选择的代价为 c_i(n-c_i)。这是一个恰有 b 组的可行性背包。

用动态位集保存状态。令 f_i 记录前 i 个二进制位能够凑出的所有和。对于每个可选数量 c,执行:

f_i\gets f_i\operatorname{or}(f_{i-1}\ll c(n-c))

转移只能来自上一层,保证每一位恰好选择一次。保存所有层后,从 f_b 倒序枚举 c_i。这样即可恢复每一位中 1 的数量。

还需保证每个数不超过 m。把最高位的 1 放在数组前端,把所有低位的 1 放在数组后端。每一位至多出现 \lfloor n/2\rfloor1,所以前后两部分没有交集。

前端非零元素恰为 2^{b-1},不超过 m。后端元素只含更低位,严格小于 2^{b-1}。因此,每个构造出的数都位于 [0,m]

最后确定动态位集的大小。若背包有解,s 同时满足各位贡献上界和原题中 k\le10^9 的限制:

s^3\le\frac{bn^2}{4}\left(\frac{10^9}{6n}\right)^2\le\frac{20\cdot10^{18}}{144}<520000^3

所以 s<520000,定长位集取 M=520005 足够。设 r 是满足 c(n-c)\le s 的最大 c,机器字长为 w=64。时间复杂度为 O(brs/w),空间复杂度为 O(bM/w+n)

又有 r\le n/2s\le10^9/(6n),所以 rs\le10^9/12。本题至多进行约 2.7\times10^7 次机器字转移。

参考代码

#include <bits/stdc++.h>
using namespace std;

using ll=long long;
using ull=unsigned long long;
const int N=1000005;
const int M=520005;
int a[N];
ull f[21][(M+63)/64];
bool get(int i,int x)
{
    return f[i][x>>6]>>(x&63)&1;
}
void shift(int i,int x,int w)
{
    int q=x>>6,r=x&63;
    for(int j=0;j+q<w;j++)
    {
        f[i][j+q]|=f[i-1][j]<<r;
        if(r&&j+q+1<w)f[i][j+q+1]|=f[i-1][j]>>(64-r);
    }
}
int main()
{
    ios::sync_with_stdio(false);
    cin.tie(nullptr);
    int n,m,k;
    cin>>n>>m>>k;
    if(k%(6*n))
    {
        cout<<-1<<'\n';
        return 0;
    }
    k/=6*n;
    int b=0;
    for(int i=m;i;i>>=1)b++;
    int h=n/2;
    ll mx=1LL*b*h*(n-h);
    if(mx<k)
    {
        cout<<-1<<'\n';
        return 0;
    }
    int w=(k>>6)+1;
    f[0][0]=1;
    for(int i=1;i<=b;i++)
    {
        for(int j=0;j<w;j++)f[i][j]=f[i-1][j];
        for(int j=1;j<=h;j++)
        {
            ll x=1LL*j*(n-j);
            if(x>k)break;
            shift(i,(int)x,w);
        }
    }
    if(!get(b,k))
    {
        cout<<-1<<'\n';
        return 0;
    }
    int c[25]={};
    int sum=k;
    for(int i=b;i>=1;i--)
    {
        for(int j=0;j<=h;j++)
        {
            ll x=1LL*j*(n-j);
            if(x>sum)break;
            if(get(i-1,sum-x))
            {
                c[i-1]=j;
                sum-=x;
                break;
            }
        }
    }
    for(int i=1;i<=c[b-1];i++)a[i]=1<<(b-1);
    for(int i=0;i<b-1;i++)for(int j=1;j<=c[i];j++)a[n-j+1]|=1<<i;
    for(int i=1;i<=n;i++)cout<<a[i]<<(i==n?'\n':' ');
    return 0;
}