求数组中和为指定值的不同三元组数量:代码错误排查
三元组求和计数错误排查与修复
给定随机整数数组ARR和数值X,需找出并返回数组中和为X的不同位置三元组数量。编写的Java代码在输入数组为[1, 3, 3, 3, 3, 3, 3]、指定值为9时,输出10,但预期输出为20,需排查错误并保证时间复杂度不超过O(n²)。
原代码如下:
public class Solution { public static void merge(int arr[], int lb, int mid, int ub) { int n1 = mid - lb + 1; int n2 = ub - mid; int arr1[] = new int[n1]; int arr2[] = new int[n2]; for (int i = 0; i < n1; i++) arr1[i] = arr[lb + i]; for (int j = 0; j < n2; j++) arr2[j] = arr[mid + 1 + j]; int i = 0, j = 0; int k = lb; while (i < n1 && j < n2) { if (arr1[i] <= arr2[j]) arr[k] = arr1[i++]; else arr[k] = arr2[j++]; k++; } while (i < n1) arr[k++] = arr1[i++]; while (j < n2) arr[k++] = arr2[j++]; } public static void mergeSort(int arr[], int lb, int ub) { if (lb < ub) { int mid = lb + (ub - lb) / 2; mergeSort(arr, lb, mid); mergeSort(arr, mid + 1, ub); merge(arr, lb, mid, ub); } } public static int tripletSum(int[] arr, int num) { mergeSort(arr, 0, arr.length - 1); int n = arr.length; int count = 0; for (int i = 0; i < n - 2; i++) { int sum = num - arr[i]; int j = i + 1; int k = n - 1; while (j < k) { if (arr[j] + arr[k] == sum) { count++; k--; } else if (arr[j] + arr[k] > sum) { k--; } else j++; } } return count; } }
错误原因
原代码在处理数组中存在大量重复元素的场景时,仅通过单次递增/递减指针并计数+1的方式统计组合,严重遗漏了重复元素间的有效组合。
以测试用例为例:排序后的数组为[1, 3, 3, 3, 3, 3, 3],和为9的三元组全部由三个3组成,总共有C(6,3)=20种(从6个3中选3个不同位置的元素)。但原代码在遍历每个i时,仅对j<k的情况逐个计数,最终只累加了4+3+2+1=10次,完全忽略了同一区间内重复元素的组合数。
修复方案
当arr[j] + arr[k] == sum时,需根据元素是否重复分情况计算组合数:
- 当
arr[j] == arr[k]:此时j到k之间的所有元素均相同,组合数为从m个元素中选2个的组合数(m = k-j+1),即m*(m-1)/2,一次性累加后跳出循环。 - 当
arr[j] != arr[k]:统计左侧连续等于arr[j]的元素数量countJ,右侧连续等于arr[k]的元素数量countK,组合数为countJ * countK,累加后直接跳过这些重复元素。
修改后的tripletSum方法如下:
public static int tripletSum(int[] arr, int num) { mergeSort(arr, 0, arr.length - 1); int n = arr.length; int count = 0; for (int i = 0; i < n - 2; i++) { int sum = num - arr[i]; int j = i + 1; int k = n - 1; while (j < k) { int currentSum = arr[j] + arr[k]; if (currentSum == sum) { if (arr[j] == arr[k]) { // 计算j到k区间内的组合数 int m = k - j + 1; count += m * (m - 1) / 2; break; } else { // 统计左侧相同元素数量 int countJ = 1; while (j + 1 < k && arr[j] == arr[j + 1]) { countJ++; j++; } // 统计右侧相同元素数量 int countK = 1; while (k - 1 > j && arr[k] == arr[k - 1]) { countK++; k--; } count += countJ * countK; j++; k--; } } else if (currentSum > sum) { k--; } else { j++; } } } return count; }
复杂度分析
- 排序阶段:归并排序的时间复杂度为
O(n log n)。 - 三元组统计阶段:外层循环
O(n),内层双指针遍历O(n),整体为O(n²),满足复杂度要求。
内容的提问来源于stack exchange,提问作者Arun Kumar
相关产品推荐
相关产品推荐

