关于一类特殊范围下的 0-1 背包优化算法

· · 题解

这是一个 0-1 背包模板题,但物品数 n 高达 10^6,背包容量 k 高达 10^5,且需要求出每个容量下的最大价值。

朴素背包 DP 的复杂度为 O(nk),显然不可行。注意到本题的每个物品大小至多为 300,考虑从这里入手。

由于物品大小的范围很小,于是对其进行分组,相同大小的物品分为一组。然后对于同一组的物品,按价值从大到小排序,然后求前缀和。设 a_{i,j} 表示大小为 i 的物品中前 j 大的价值和,可以发现,函数 g(j)=a_{i,j} 是一个凹函数。

这样就转换为了分组背包了,设 dp_{i,j} 表示拿前 i 组、背包容量为 j 时的最大价值。

若第 i 组有 s_i 个物品,转移方程如下:

dp_{i,j}=\max_{0\le k\le s_i}\{dp_{i-1,j-i\times k}+a_{i,k}\}

注意到第一维可以去掉,设 f_i 表示上一轮背包容量为 i 的最大价值,dp_i 表示当前这一轮背包容量为 i 的最大价值,设当前处于第 t 轮,则有转移方程:

dp_{i}=\max_{0\le k\le s_t}\{f_{i-t\times k}+a_{t,k}\}

考虑去掉 k 的限制,由于 f_jdp_i 的贡献为 f_j+a_{t,\min(\lfloor \frac{i-j}{t}\rfloor,s_t)},于是可以得出:

dp_{i}=\max_{0\le j\le i}\{f_j+a_{t,\min(\lfloor \frac{i-j}{t}\rfloor,s_t)}\}

i,j 按模 t 意义下分组,对每组分别 DP,设 x=i\bmod t,则第 i 个数的实际下标 id(i)=it-t+x\frac{id(i)-id(j)}{t}=i-j,得转移方程:

dp_{id(i)}=\max_{0\le j\le i}\{f_{id(j)}+a_{t,\min(i-j,s_t)}\}

观察我们最终得到的式子,不难发现转移实际上是 fg 的 max-plus 卷积,而由于先前我们已经得出 g 是一个凹函数,由 max-plus 卷积的性质,当参与卷积的两函数之一为凹函数时,最优决策下标单调不减,因此我们惊喜地发现该转移具有决策单调性!

补充:之所以要对 i,j 按模 t 意义下分组,是为了把下取整符号去掉,这样 i-j\le s_tg(j)=a_{t,i-j}f(id(j)) 一一对应,满足 i-j 单调递增。要是有下取整存在的话,i-j 只满足单调不降,g 函数就不具有凹性了,可以从函数图像上理解:

g(j)=a_{t,\min(\lfloor \frac{i-j}{t}\rfloor,s_t)},其图像可能长这样:

这显然不是一个凹函数,而当 g(j)=a_{t,\min(i-j,s_t)} 时,图像就可能长这样:

这很明显就是一个凹函数,即使 i-j>s_t 的部分 g(j)=a_{t,s_t},其也依旧具有凹性。

于是套个决策单调性分治优化的板子就写完了,这么做的时间复杂度为 O(sk\log k+n\log n),其中 s=300

参考代码如下:

#include<bits/stdc++.h>
#define cin_fast ios::sync_with_stdio(false) , cin.tie(0) , cout.tie(0)
//#define int long long 
#define in(a) a = read()
#define PII pair<int , int>
using namespace std;
typedef long long ll;
const int N = 1e6 + 5 , mod = 998244353;
const int inf = 0x3f3f3f3f;
const long long INF = 0x3f3f3f3f3f3f3f3f; 
inline int read() {
    int x = 0;
    char ch = getchar();
    bool f = 0;
    while('9' < ch || ch < '0') f |= ch == '-' , ch = getchar();
    while('0' <= ch && ch <= '9') x = (x << 3) + (x << 1) + ch - '0' , ch = getchar();
    return f ? -x : x;
}
int t , x;
ll f[N] , dp[N];
vector<ll>a[N];
ll w(int i , int j) {
    if(i == j) return f[j * t - t + x];
    return f[j * t - t + x] + a[t][min(i - j , (int)a[t].size()) - 1];
}
void solve(int l , int r , int optl , int optr) {
    int mid = (l + r) >> 1 , id = mid * t - t + x , opt = optl;
    for(int i = optl ; i <= min(optr , mid) ; i ++) {
        if(w(mid , i) >= dp[id]) dp[id] = w(mid , i) , opt = i; 
    }
    if(l < mid) solve(l , mid - 1 , optl , opt);
    if(mid < r) solve(mid + 1 , r , opt , optr);
}
signed main() { 
    //cin_fast;
    int n , k;
    in(n) , in(k);
    for(int i = 1 ; i <= n ; i ++) {
        int v , w;
        in(v) , in(w);
        a[v].emplace_back(w);
    }
    for(int i = 1 ; i <= 300 ; i ++) {
        sort(a[i].begin() , a[i].end() , greater<int>());
        for(int j = 1 ; j < a[i].size() ; j ++) a[i][j] += a[i][j - 1];
    }
    for(t = 1 ; t <= min(300 , k) ; t ++) {
        if(a[t].empty()) continue;
        for(x = 0 ; x < t ; x ++) {
            int cnt = (k - x) / t + 1;
            solve(1 , cnt , 1 , cnt);
        }
        for(int j = 0 ; j <= k ; j ++) f[j] = dp[j];
    }
    for(int i = 1 ; i <= k ; i ++) cout << f[i] << ' ';
    return 0;
}