BISHI138 【模板】多重背包
中等通过率 30.25%pypy3样例通过牛客 AC
牛客原题 源码
讲解章节:背包问题
解题思路
⚠️ 本题必须用 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)。
两个必须做的预处理:
- 件数上限截断:s <- min(s, m // w)。背包只有 m 的容量;
- 体积为 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 物理上做不到」的题之一。
坑在哪
- 窗口判定用
qi[head] < t - s,是「最多取 s 件」而不是 s+1 件;
- 入队要先把当前 t 压进去再弹队首,否则 s = 0 时队列会空;
- w = 0 与 v = 0 的组合都要能正常跑(测试点 15/16);
- 队列用两个预分配的定长 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()
|