求解将数组划分为两部分且唯一元素数量相等的方法数
优化解法:动态规划+组合计数
问题分析
你需要将长度为n的数组划分为两个等长子集,要求两个子集的唯一元素数量相等,求划分方法数。首先明确两个前提:
- 若
n为奇数,直接返回0(无法分成等长两部分)。 - 每个元素的分配方式有三种:全给左子集、全给右子集、拆分到两个子集(仅当元素出现次数≥2时可行)。
核心思路
通过**动态规划(DP)**跟踪两个关键状态:
- 左子集当前的总长度
- 左右子集唯一元素数量的差值(左唯一数 - 右唯一数)
结合组合计数计算拆分元素时的方法数,最终统计左子集长度为n/2且差值为0的状态对应的方法数。
具体步骤
1. 预处理
- 统计数组中每个元素的出现频率,得到频率字典
freq。 - 计算目标左子集长度
target_len = n // 2,总唯一元素数U = len(freq)。
2. 动态规划初始化
用字典(或二维数组)存储状态:dp[(l, d)]表示处理完部分元素后,左子集长度为l、左右唯一数差值为d的划分方法数。初始状态为dp[(0, 0)] = 1(未处理任何元素时的初始状态)。
3. 遍历每个元素的频率
对每个元素的出现次数c,基于上一轮的DP状态更新新状态:
- 情况1:全部分配到左子集
左长度变为l + c,差值变为d + 1(左唯一数+1),方法数累加dp[(l, d)] * 1(仅1种分法)。注意左长度不能超过target_len。 - 情况2:全部分配到右子集
左长度不变,差值变为d - 1(右唯一数+1),方法数累加dp[(l, d)] * 1。 - 情况3:拆分到两个子集
当c ≥ 2时,遍历左子集分配的元素数k(1 ≤ k < c):
左长度变为l + k,差值不变(左右唯一数各+1,差值抵消),方法数累加dp[(l, d)] * C(c, k)(C(c, k)是从c个元素中选k个的组合数)。
4. 计算最终结果
遍历完所有元素后,dp.get((target_len, 0), 0)就是符合要求的划分方法数。
优化点
- 空间优化:仅保留上一轮的DP状态,用两个字典交替更新,避免存储所有轮次的状态。
- 极端情况处理:若所有元素仅出现1次,直接判断
U是否等于2 * target_len:- 是:返回组合数
C(U, target_len)(选target_len个元素全给左,剩下的给右,此时左右唯一数相等)。 - 否:返回0。
- 是:返回组合数
- 组合数预计算:使用内置组合数函数(如Python的
math.comb)或预计算杨辉三角,提升计算效率。
示例验证
以数组[1,2,2,3,3,3,4,4,4,4]为例:
n=10,target_len=5,freq={1:1,2:2,3:3,4:4}。- 按上述步骤处理每个元素后,最终
dp[(5, 0)] = 8,与题目给出的结果一致。
伪代码实现
import math from collections import Counter, defaultdict def count_valid_partitions(arr): n = len(arr) if n % 2 != 0: return 0 target_len = n // 2 freq = Counter(arr) U = len(freq) # 处理所有元素仅出现1次的情况 all_single = all(c == 1 for c in freq.values()) if all_single: if U == 2 * target_len: return math.comb(U, target_len) else: return 0 # 初始化DP dp = defaultdict(int) dp[(0, 0)] = 1 for c in freq.values(): new_dp = defaultdict(int) for (l_prev, d_prev), cnt in dp.items(): # 情况1:全左 l_new = l_prev + c if l_new <= target_len: d_new = d_prev + 1 new_dp[(l_new, d_new)] += cnt # 情况2:全右 l_new = l_prev d_new = d_prev - 1 new_dp[(l_new, d_new)] += cnt # 情况3:拆分 if c >= 2: for k in range(1, c): l_new = l_prev + k if l_new <= target_len: d_new = d_prev comb = math.comb(c, k) new_dp[(l_new, d_new)] += cnt * comb dp = new_dp return dp.get((target_len, 0), 0) # 测试示例 arr = [1,2,2,3,3,3,4,4,4,4] print(count_valid_partitions(arr)) # 输出8
复杂度分析
- 时间复杂度:O(U * target_len * C_avg),其中
U是唯一元素数,target_len是n/2,C_avg是元素平均出现次数(仅当元素出现次数≥2时需要遍历k)。远优于递归的指数级复杂度。 - 空间复杂度:O(target_len * U),存储DP状态的字典大小。
内容的提问来源于stack exchange,提问作者Nikitka
相关产品推荐
相关产品推荐

