如何以O(N²)时间复杂度解决四元组元素和条件计数问题
问题描述
给定一个包含N个整数的数组A[],统计满足以下条件的四元组(i, j, k, l)的数量:
- i、j、k、l为不同整数,且满足1 ≤ i < j < k < l ≤ N
- 至少满足以下条件之一:
- A[i] = A[j] + A[k] + A[l]
- A[j] = A[i] + A[k] + A[l]
- A[k] = A[i] + A[j] + A[l]
- A[l] = A[i] + A[j] + A[k]
简单来说,就是找出数组中所有四元组里,至少有一个元素等于另外三个元素之和的数量。
输入输出示例
示例1
输入:
N = 5
A = [1, 1, 3, 1, 1]
输出:4
解释:符合条件的四元组为(1,2,3,4)、(1,2,3,5)、(1,3,4,5)、(2,3,4,5)
示例2
输入:
N = 2
A = [5, 2]
输出:0
错误代码分析
你编写的代码逻辑完全偏离问题要求:
A.sort() C = 0 D = dict() for i in range(N): for j in range(i + 1, N): D[A[i] + A[j]] = (i, j) for i in range(N): for j in range(i + 1, N): S = A[i] + A[j] if S in D and D[S] != (i, j): C += 1 print(C)
- 这段代码统计的是"存在两对不同索引对,它们的和相等"的数量,和问题要求的四元组条件无关。
- 字典
D会覆盖相同和的索引对,导致统计结果错误,同时完全没考虑四元组的索引顺序要求。
正确解法(O(N²)时间复杂度)
思路分析
- 先对数组排序,利用有序性优化查找,同时保证索引顺序i<j<k<l的判断更简单。
- 分四种情况处理每个位置的元素作为"目标元素"(即等于另外三个元素和的元素),分别查找符合条件的三元组。
- 使用集合存储合法四元组的索引组合,避免重复计数(比如一个四元组可能满足多个条件)。
- 对于两数之和的查找,用双指针法将时间复杂度控制在O(N²);对于单个元素的查找,用哈希表存储每个值的索引列表,快速定位。
代码实现
from collections import defaultdict def count_valid_quadruples(N, A): if N < 4: return 0 A_sorted = sorted(A) valid_quadruples = set() value_indices = defaultdict(list) for idx, num in enumerate(A_sorted): value_indices[num].append(idx) # 情况1:A[l] = A[i]+A[j]+A[k],i<j<k<l for l in range(3, N): target = A_sorted[l] for k in range(2, l): remaining = target - A_sorted[k] i, j = 0, k - 1 while i < j: current_sum = A_sorted[i] + A_sorted[j] if current_sum == remaining: # 统计所有符合的i、j对 left = i while left < j and A_sorted[left] == A_sorted[i]: left += 1 right = j while right > i and A_sorted[right] == A_sorted[j]: right -= 1 # 将所有合法四元组加入集合 for ip in range(i, left): for jp in range(right + 1, j + 1): valid_quadruples.add((ip, jp, k, l)) i = left j = right elif current_sum < remaining: i += 1 else: j -= 1 # 情况2:A[k] = A[i]+A[j]+A[l],i<j<k<l for k in range(2, N-1): target = A_sorted[k] for l in range(k+1, N): remaining = target - A_sorted[l] if remaining > 0: continue i, j = 0, k - 1 while i < j: current_sum = A_sorted[i] + A_sorted[j] if current_sum == remaining: left = i while left < j and A_sorted[left] == A_sorted[i]: left += 1 right = j while right > i and A_sorted[right] == A_sorted[j]: right -= 1 for ip in range(i, left): for jp in range(right + 1, j + 1): valid_quadruples.add((ip, jp, k, l)) i = left j = right elif current_sum < remaining: i += 1 else: j -= 1 # 情况3:A[j] = A[i]+A[k]+A[l],i<j<k<l for j in range(1, N-2): target = A_sorted[j] for k in range(j+1, N-1): for l in range(k+1, N): remaining = target - (A_sorted[k] + A_sorted[l]) if remaining not in value_indices: continue for ip in value_indices[remaining]: if ip < j: valid_quadruples.add((ip, j, k, l)) # 情况4:A[i] = A[j]+A[k]+A[l],i<j<k<l for i in range(N-3): target = A_sorted[i] for j in range(i+1, N-2): for k in range(j+1, N-1): remaining = target - (A_sorted[j] + A_sorted[k]) if remaining not in value_indices: continue for lp in value_indices[remaining]: if lp > k: valid_quadruples.add((i, j, k, lp)) return len(valid_quadruples) # 测试示例1 N = 5 A = [1, 1, 3, 1, 1] print(count_valid_quadruples(N, A)) # 输出4 # 测试示例2 N = 2 A = [5, 2] print(count_valid_quadruples(N, A)) # 输出0
代码说明
- 排序数组后,针对每个位置的元素作为目标,分别查找符合条件的三元组,确保索引顺序合法。
- 双指针法用于快速查找两数之和等于目标值的索引对,将这部分时间复杂度控制在O(N)每轮,外层循环O(N),总时间O(N²)。
- 集合存储合法四元组的索引组合,避免同一个四元组因满足多个条件被重复计数。
内容的提问来源于stack exchange,提问作者Shriram Kannuri
相关产品推荐
相关产品推荐

