这道题的核心是贪心排序,Python实现的关键在于自定义排序规则。
核心解题思路
每个片段形如 "111...000..."(nums1[i] 个 1 后跟 nums0[i] 个 0)。排序规则:
1. 纯 1 片段优先:nums0[i] == 0 的片段全由 1 组成
2. 1 多的靠前:1 的数量越多,高位 1 越多
3. 1 相同时,0 少的靠前
Python3 代码实现
```python
class Solution:
def maxValue(self, nums1: List[int], nums0: List[int]) -> int:
MOD = 10**9 + 7
n = len(nums1)
# 1. 创建片段列表 (ones, zeros)
fragments = list(zip(nums1, nums0))
# 2. 自定义排序(使用 functools.cmp_to_key)
from functools import cmp_to_key
def compare(a, b):
ones_a, zeros_a = a
ones_b, zeros_b = b
# 纯 1 片段优先
if zeros_a == 0 and zeros_b == 0:
return ones_b - ones_a # 1 多的在前
if zeros_a == 0:
return -1
if zeros_b == 0:
return 1
# 都含 0:1 多的在前
if ones_a != ones_b:
return ones_b - ones_a
# 1 相同:0 少的在前
return zeros_a - zeros_b
fragments.sort(key=cmp_to_key(compare))
# 3. 拼接并计算结果
ans = 0
for ones, zeros in fragments:
# 追加 ones 个 '1'
for _ in range(ones):
ans = (ans * 2 + 1) % MOD
# 追加 zeros 个 '0'
for _ in range(zeros):
ans = (ans * 2) % MOD
return ans
```
更简洁的实现(使用 key 参数)
如果不想用 cmp_to_key,可以构造一个复合的排序键:
```python
class Solution:
def maxValue(self, nums1: List[int], nums0: List[int]) -> int:
MOD = 10**9 + 7
fragments = list(zip(nums1, nums0))
# 排序键:
# 1. 是否纯 1(纯 1 优先)
# 2. 1 的数量(降序)
# 3. 0 的数量(升序)
fragments.sort(key=lambda x: (
x[1] != 0, # False(纯1)排在 True(含0)前面
-x[0], # 1 多的在前
x[1] # 0 少的在前
))
ans = 0
for ones, zeros in fragments:
# 批量处理优化(避免逐位循环)
# 先处理 1
ans = (ans * pow(2, ones, MOD) + (pow(2, ones, MOD) - 1)) % MOD
# 再处理 0
ans = (ans * pow(2, zeros, MOD)) % MOD
return ans
```
批量优化版本(处理大片段)
当 nums1[i] 或 nums0[i] 很大时,逐位循环可能较慢,可以用幂运算批量处理:
```python
class Solution:
def maxValue(self, nums1: List[int], nums0: List[int]) -> int:
MOD = 10**9 + 7
fragments = list(zip(nums1, nums0))
# 排序(同上)
fragments.sort(key=lambda x: (x[1] != 0, -x[0], x[1]))
ans = 0
for ones, zeros in fragments:
# 追加 ones 个 1:ans = ans * 2^ones + (2^ones - 1)
if ones > 0:
pow2_ones = pow(2, ones, MOD)
ans = (ans * pow2_ones + (pow2_ones - 1)) % MOD
# 追加 zeros 个 0:ans = ans * 2^zeros
if zeros > 0:
ans = (ans * pow(2, zeros, MOD)) % MOD
return ans
```
测试用例
```python
# 测试
sol = Solution()
# 示例 1
print(sol.maxValue([1, 1], [1, 1])) # 输出:6("10"+"10" = "1010" = 10,但 "11"+"00" 不存在)
# 实际:排序后 [1,1] 和 [1,1] 顺序不影响,1010 = 10
# 示例 2
print(sol.maxValue([2, 1], [0, 1]))
# fragments: [2,0] = "11", [1,1] = "10"
# 排序后:["11", "10"] -> "1110" = 14
# 输出:14
# 示例 3
print(sol.maxValue([1, 2, 1], [2, 0, 1]))
# fragments: [1,2]="100", [2,0]="11", [1,1]="10"
# 排序:["11", "10", "100"] -> "1110100" = 116
# 输出:116
```
排序规则证明(简洁版)
比较两个片段 A 和 B,我们需要判断 A+B 和 B+A 哪个更大:
· 如果 A 全是 1,A+B 前缀是 1,B+A 前缀是 B 的第一个字符(可能是 0),所以 A 应在前
· 如果都有 0,比较 1 的数量,多的在前(因为高位 1 越多越大)
· 如果 1 数量相同,0 少的在前(因为 0 越早出现,字典序越小)
复杂度分析
· 时间复杂度:O(n log n + L),其中 L 是总长度(sum(nums1) + sum(nums0))
· 空间复杂度:O(n),用于存储片段列表
关键注意事项
1. 取模运算:结果要对 10^9+7 取模
2. 批量处理:使用 pow(2, k, MOD) 可以快速处理连续的 1 或 0
3. 排序稳定性:Python 的 sort 是稳定的,但建议明确定义所有比较规则
如果还有疑问,欢迎继续追问!