求满足i<j<k且和可被d整除的三元组的优化解法
优化「三元组和可被d整除」的计数算法
我们需要找出数组中所有满足i<j<k且arr[i]+arr[j]+arr[k] % d == 0的三元组数量。原O(n³)的解法对于n=1e3来说会执行约1e9次操作,显然会超时,需要优化。
核心思路转换
首先,利用模运算的性质:(a + b + c) % d == 0等价于(a%d + b%d + c%d) % d == 0。我们可以先将数组中每个元素替换为它对d取模的结果,这样问题转化为寻找三个余数r1, r2, r3,使得r1 + r2 + r3 ≡ 0 mod d,且对应的索引满足i<j<k。
解法一:O(n²)时间复杂度
思路
遍历所有i<j的数对,计算它们的余数和s = (r[i] + r[j]) % d,那么我们需要找所有k>j的元素中,余数等于(d - s) % d的数量,将这个数量累加进结果。
为了高效统计k>j的余数数量,我们可以维护一个后缀余数计数数组:
- 先统计整个数组中每个余数出现的总次数,存入
cnt数组。 - 从后往前遍历
j:- 先将
r[j]从cnt中减1(因为k必须大于j,所以当前r[j]不能被计入k的候选)。 - 再遍历所有
i < j的元素,计算需要的目标余数target = (d - (r[i] + r[j]) % d) % d,将cnt[target]加到结果中。
- 先将
代码实现
def count_triplets(arr, d): n = len(arr) remainders = [x % d for x in arr] cnt = [0] * d for r in remainders: cnt[r] += 1 count = 0 # 从后往前遍历j for j in range(n-1, 0, -1): # 先把当前j的余数从cnt中移除,因为k要大于j cnt[remainders[j]] -= 1 # 遍历所有i < j for i in range(j): s = (remainders[i] + remainders[j]) % d target = (d - s) % d count += cnt[target] return count
解法二:基于余数频率的组合计数(O(d²)时间)
思路
先统计每个余数出现的次数cnt[r],然后枚举所有可能的余数三元组组合,计算符合条件的组合数:
- 枚举所有
r1从0到d-1:- 枚举
r2从r1到d-1:- 计算
r3 = (d - (r1 + r2)) % d,确保r3 >= r2(避免重复计算)。 - 根据
r1, r2, r3的相等情况,计算对应的组合数:- 若
r1 == r2 == r3:组合数为C(cnt[r1], 3) = cnt[r1]*(cnt[r1]-1)*(cnt[r1]-2)//6 - 若
r1 == r2 != r3:组合数为C(cnt[r1], 2)*cnt[r3] = cnt[r1]*(cnt[r1]-1)//2 * cnt[r3] - 若
r1 != r2 == r3:组合数为cnt[r1] * C(cnt[r2], 2) - 若
r1 != r2 != r3:组合数为cnt[r1]*cnt[r2]*cnt[r3]
- 若
- 计算
- 枚举
代码实现
def count_triplets(arr, d): cnt = [0] * d for x in arr: cnt[x % d] += 1 count = 0 # 枚举所有r1 <= r2 <= r3的组合 for r1 in range(d): if cnt[r1] == 0: continue # 情况1: r1 == r2 == r3 if (3 * r1) % d == 0: if cnt[r1] >=3: count += cnt[r1] * (cnt[r1]-1) * (cnt[r1]-2) //6 # 情况2: r1 == r2 != r3 r3 = (d - 2*r1) % d if r3 > r1 and cnt[r3] >0: if cnt[r1] >=2: count += (cnt[r1] * (cnt[r1]-1) //2) * cnt[r3] # 情况3: r1 != r2 == r3 for r2 in range(r1+1, d): if cnt[r2] ==0: continue if (r1 + 2*r2) %d ==0: if cnt[r2] >=2: count += cnt[r1] * (cnt[r2] * (cnt[r2]-1) //2) # 情况4: r1 < r2 < r3 for r2 in range(r1+1, d): if cnt[r2] ==0: continue r3 = (d - (r1 + r2)) %d if r3 > r2 and cnt[r3] >0: count += cnt[r1] * cnt[r2] * cnt[r3] return count
复杂度对比
- 原解法:O(n³),n=1e3时约1e9次操作,超时。
- 解法一:O(n²),n=1e3时约1e6次操作,完全符合时间要求。
- 解法二:O(d²),当d较小时(比如d<=100),比解法一更快;当d较大时(比如d>1e3),解法一更优。
内容的提问来源于stack exchange,提问作者Athul Raju
相关产品推荐
相关产品推荐

