跳转至

BISHI138 【模板】多重背包

中等通过率 30.25%pypy3样例通过牛客 AC

牛客原题  源码

讲解章节背包问题

一句话

第 i 种物品有 s_i 件,多组数据。

解题思路

⚠️ 本题必须用 PyPy3 提交(牛客语言 id 25),CPython 交不过去。原因见文末。

这题考什么

多重背包的 单调队列优化,把 O(m · Σ log s_i) 压到 O(n · m)。

转移式是

f[j] = max(f[j], f[j-w] + v, f[j-2w] + 2v, ..., f[j-kw] + kv)
其中 k = min(s, j // w)

只有下标模 w 同余的位置会互相影响。于是按余数 r(0 <= r < w)分组, 组内下标写成 j = r + t·w,转移变成

g[t] = max_{t-s <= t' <= t} ( g[t'] - t'·v ) + t·v

括号里那一项只和 t' 有关,于是就是一个定长窗口最大值——单调队列的标准形态。

队列里存 (t', f[j'] - t'·v),按第二个分量单调递减:

  • 队尾:新值不小于队尾就把队尾弹掉(它永远不会再当最大值);
  • 队首:t' < t - s 说明超过了「最多取 s 件」的窗口,弹掉;
  • 取队首即为当前最优。

每个下标进队出队各一次,单组数据总计 O(n·m)。

两个必须做的预处理:

  1. 件数上限截断:s <- min(s, m // w)。背包只有 m 的容量;
  2. 体积为 0 的特判(第 15、16 个测试点专门卡这个): w = 0 的物品不占容量,直接把 s·v 全部计入答案, 否则 m // w 会除零崩溃。

为什么不用二进制拆分

二进制拆分(把 s 件拆成 1, 2, 4, ... 若干个 01 物品)是多重背包的通用写法, 配合「整段切片取 max」还能让每一轮都跑在 C 层,看上去很适合 Python。 但它把物品数放大了约 log s 倍,本题最坏情况下是 12 倍。

构造「w 全为 1、s 全为 1e6」的数据(截断后每件拆出 12 个打包件)实测:

二进制拆分 + C 层整段取 max:80.1 秒     O(m · Σ log s) ≈ 1e9
单调队列(纯 Python 层):   14.9 秒     O(n · m)       = 9e7

C 层每次操作再便宜,也架不住总操作数多一个数量级。

这里要记住的是:「把循环压进 C 层」是常数优化,替代不了降低渐进复杂度。 数据规模足够大时先看复杂度、再谈常数——这一点上 Python 与 C++ 并无分歧, 真正的分歧只在「常数相差多少倍」。

为什么必须 PyPy3

单调队列是 9e6 次纯 Python 层迭代 / 组,10 组就是 9e7 次。 CPython 实测 14.9 秒,而时限是「其他语言 10 秒」,怎么调常数都差着一截, 且这个循环有前后依赖(队列状态),没有任何办法向量化到 C 层。 PyPy3 的 JIT 能把这种紧循环编译成机器码,实测轻松通过。 这是本项目 165 题里少数几道「CPython 物理上做不到」的题之一。

坑在哪

  1. 窗口判定用 qi[head] < t - s,是「最多取 s 件」而不是 s+1 件;
  2. 入队要把当前 t 压进去再弹队首,否则 s = 0 时队列会空;
  3. w = 0 与 v = 0 的组合都要能正常跑(测试点 15/16);
  4. 队列用两个预分配的定长 list + head/tail 下标, 不要用 collections.deque——PyPy 下前者更快,也避免了对象封装。

参考实现

solutions/BISHI138.py
import sys


def main() -> None:
    data = sys.stdin.buffer.read().split()
    p = 0
    T = int(data[p]); p += 1
    out = []
    for _ in range(T):
        n = int(data[p]); m = int(data[p + 1])
        p += 2
        f = [0] * (m + 1)                    # 每组重置;不要求装满,初值全 0
        base = 0                             # 体积为 0 的物品直接全拿
        # 队列容量取 m + 2:单组内下标最多 m + 1 个,预分配两条定长 list,
        # 用 head/tail 当双指针,比 deque 少一层对象封装
        qi = [0] * (m + 2)                   # 单调队列:下标 t'
        qv = [0] * (m + 2)                   # 单调队列:键值 f[j'] - t'·v
        for _ in range(n):
            w = int(data[p]); v = int(data[p + 1]); s = int(data[p + 2])
            p += 3
            if w == 0:                       # 不占容量,能拿多少拿多少,绕开下面的 m // w 除零
                base += s * v
                continue
            if v == 0 or w > m:              # 零价值或单件就超容量,拿了也没用
                continue
            if s > m // w:                   # 件数截断:拿再多也塞不下
                s = m // w
            for r in range(w):               # 按 j mod w 分组,组内做窗口最大值
                head = 0
                tail = 0                     # 队列区间是 [head, tail),每组从空队列重开
                t = 0                        # 组内第几项,对应下标 j = r + t·w
                j = r
                while j <= m:
                    val = f[j] - t * v       # 减去 t·v 后窗口内可直接比大小,与 t 无关
                    while tail > head and qv[tail - 1] <= val:
                        tail -= 1            # 队尾比新值差,永远轮不到它
                    qi[tail] = t
                    qv[tail] = val
                    tail += 1                # 先入队再弹队首,保证 s = 0 时队列不空
                    if qi[head] < t - s:     # 超出「最多 s 件」的窗口
                        head += 1            # 窗口每步只右移一格,弹一次就够
                    f[j] = qv[head] + t * v  # 把刚才减掉的 t·v 加回来
                    t += 1
                    j += w
        out.append(f[m] + base)              # 零体积物品的收益与背包无关,最后补上
    sys.stdout.write("\n".join(map(str, out)) + "\n")


main()
[:octicons-arrow-left-16: BISHI137](BISHI137.md) [BISHI139 :octicons-arrow-right-16:](BISHI139.md)