Skip to content

剑指 Offer 51. 数组中的逆序对

在数组中的两个数字,如果前面一个数字大于后面的数字,则这两个数字组成一个逆序对。输入一个数组,求出这个数组中的逆序对的总数。

示例 1:

plain
输入: [7,5,6,4]
输出: 5

限制:

0 <= 数组长度 <= 50000

解题思路:

直观来看,使用暴力统计法即可,即遍历数组的所有数字对并统计逆序对数量。此方法时间复杂度为 \(O(N^2)\),观察题目给定的数组长度范围 \(0 \leq N \leq 50000\) ,可知此复杂度是不能接受的。

c
class Solution {
public:
    int reversePairs(vector<int>& nums) {
        if(nums.size()<1)return 0;
        int sum = 0;
        for(int i=1;i<nums.size();++i){
            for(int j=0;j<i;++j){
                if(nums[j]>nums[i]) sum++;
            }
        }
        return sum;
    }
};

「归并排序」与「逆序对」是息息相关的。归并排序体现了 “分而治之” 的算法思想,具体为:

  • 分: 不断将数组从中点位置划分开(即二分法),将整个数组的排序问题转化为子数组的排序问题;
  • 治: 划分到子数组长度为 1 时,开始向上合并,不断将 较短排序数组 合并为 较长排序数组,直至合并至原数组时完成排序;

如下图所示,为数组\([7,3,2,6,0,1,5,4]\)的归并排序过程。

合并阶段 本质上是 合并两个排序数组 的过程,而每当遇到****左子数组当前元素 > 右子数组当前元素时,意味着 「左子数组当前元素 至 末尾元素」 与 「右子数组当前元素」 构成了若干 「逆序对」

计算逆序对有两种对称的方法:

  1. 数出每一个数前面有多少个数比它大。
  2. 数出每个数后面有多少数比它小。

这里我们采用方法一,如下图所示,为左子数组 \([2, 3, 6, 7]\)与 右子数组\([0,1,4,5]\) 的合并与逆序对统计过程。

因此,考虑在归并排序的合并阶段统计「逆序对」数量,完成归并排序时,也随之完成所有逆序对的统计。

算法流程:

merge_sort()归并排序与逆序对统计:

  1. 终止条件:\(l\ge r\)时,代表子数组长度为 1 ,此时终止划分;
  2. 递归划分: 计算数组中点 m,递归划分左子数组 merge_sort(l, m) 和右子数组 merge_sort(m + 1, r)
  3. 合并与逆序对统计:
    • 暂存数组 nums闭区间 \([i, r]\)内的元素至辅助数组 tmp
    • 循环合并: 设置双指针 i ,j分别指向左 / 右子数组的首元素;
    • \(i = m + 1\)时: 代表左子数组已合并完,因此添加右子数组当前元素 tmp[j],并执行 j = j + 1
    • 否则,当\(j=r+1\) 时: 代表右子数组已合并完,因此添加左子数组当前元素 tmp[i] ,并执行 i = i + 1
    • 否则,当 \(tmp[i] \leq tmp[j]\)时: 添加左子数组当前元素 tmp[i] ,并执行 i = i + 1
    • 否则(即 \(tmp[i] > tmp[j]\))时: 添加右子数组当前元素 tmp[j],并执行 j = j + 1此时构成m - i + 1个「逆序对」,统计添加至res
  4. 返回值:返回直至目前的逆序对总数res

reversePairs()主函数:

  1. 初始化: 辅助数组 tmp_t__m_p ,用于合并阶段暂存元素;
  2. 返回值: 执行归并排序 merge_sort() ,并返回逆序对总数即可;

如下图所示,为数组 \([7, 3, 2, 6, 0, 1, 5, 4]\)的归并排序与逆序对统计过程。

代码:

c
class Solution {
    vector<int> tmp;
public:
    int reversePairs(vector<int>& nums) {
        tmp.resize(nums.size());
        return mergeSort(nums,0,nums.size()-1);
    }

    int mergeSort(vector<int>& nums, int l, int r){
        if(l>=r) return 0;
        int mid = l+(r-l)/2;
        int sum = mergeSort(nums,l,mid)+mergeSort(nums,mid+1,r);
        for(int i = l; i<=r;++i){
            tmp[i] = nums[i];
        }
        int i = l, j = mid+1;
        for(int k=l;k<=r;++k){
            //左边数组已经没有了
            if(i>mid){
                nums[k] = tmp[j++];
            }else if(j>r){//右边数组已经没有了
                nums[k] = tmp[i++];
            }else if(tmp[i]<=tmp[j]){
                nums[k] = tmp[i++];
            }else{
                nums[k] = tmp[j++];
                sum+= mid-i+1;// 统计逆序对
            }
        }
        return sum;
    }
};

复杂度分析:

  • 时间复杂度\(O(N \log N)\) 其中 N 为数组长度;归并排序使用 \(O(N \log N)\)时间;
  • 空间复杂度\(O(N)\) 辅助数组 \(tmp\)占用 \(O(N)\)大小的额外空间;

315. Count of Smaller Numbers After Self

Given an integer array nums, return an integer array counts where counts[i] is the number of smaller elements to the right of nums[i].

Example 1:

c
Input: nums = [5,2,6,1]
Output: [2,1,1,0]
Explanation:
To the right of 5 there are 2 smaller elements (2 and 1).
To the right of 2 there is only 1 smaller element (1).
To the right of 6 there is 1 smaller element (1).
To the right of 1 there is 0 smaller element.

Example 2:

c
Input: nums = [-1]
Output: [0]

Example 3:

c
Input: nums = [-1,-1]
Output: [0,0]

Constraints:

  • 1 <= nums.length <= 10^5
  • -10^4 <= nums[i] <= 10^4

解题思路

与上一题相同,实际上也是求逆序对,只不过要求输出的结果更具体:要求计算每一个元素的右边有多少个元素比自己小。

计算逆序对有两种对称的方法:

  1. 数出每一个数前面有多少个数比它大。
  2. 数出每个数后面有多少数比它小。

根据题意,这里只能采用方法二。

如下所示,在逆序对一题中,我们数的是比0大的数字个数,也就是采用方式一。但是这里我们得数出每个数后面有多少数比它小。比如这里填入0之后应当填入2,这个时候j-mid-1个数,也就是右边一半中指针j之前的所有的数,这里就是一个0比当前从左边待填入的数小!e

问题 1:为什么引入索引数组

依然可以使用「归并排序」的「分而治之」的算法思想。接下来要解决的问题是如何知道「每一个元素的右边有多少个元素比自己小」,这一点就需要我们知道 当前归并回去的那个元素在输入数组里是哪一个元素

一种可行的办法是:把「下标」和「数值」绑在一起进行归并排序,在一些编程语言中提供了 TuplePair 这样的类可以实现,也可以自己创建一个类。

问题 2:使用下标数组(索引数组)的原因

一个更经典的做法是:由于 通过下标再回到输入数组中就可以查询到下标对应的数值,因此可以 只对下标数组进行排序

因此,解决当前问题就变成了原始数组不变,用于根据下标查询数值。对下标数组进行排序,下标数组对应的数值有序,我们的问题就得到了解决。比较的时候是比较的原始数组的值,记录结果的时候用下标数组的值。

参考代码

java:

java
import java.util.ArrayList;
import java.util.List;

public class Solution {

    public List<Integer> countSmaller(int[] nums) {
        List<Integer> result = new ArrayList<>();
        int len = nums.length;
        if (len == 0) {
            return result;
        }

        int[] temp = new int[len];
        int[] res = new int[len];

        // 索引数组,作用:归并回去的时候,方便知道是哪个下标的元素
        int[] indexes = new int[len];
        for (int i = 0; i < len; i++) {
            indexes[i] = i;
        }
        mergeAndCountSmaller(nums, 0, len - 1, indexes, temp, res);

        // 把 int[] 转换成为 List<Integer>,没有业务逻辑
        for (int i = 0; i < len; i++) {
            result.add(res[i]);
        }
        return result;
    }

    /**
     * 针对数组 nums 指定的区间 [left, right] 进行归并排序,在排序的过程中完成统计任务
     *
     * @param nums
     * @param left
     * @param right
     */
    private void mergeAndCountSmaller(int[] nums, int left, int right, int[] indexes, int[] temp, int[] res) {
        if (left == right) {
            return;
        }
        int mid = left + (right - left) / 2;
        mergeAndCountSmaller(nums, left, mid, indexes, temp, res);
        mergeAndCountSmaller(nums, mid + 1, right, indexes, temp, res);

        // 归并排序的优化,如果索引数组有序,则不存在逆序关系,没有必要合并
        if (nums[indexes[mid]] <= nums[indexes[mid + 1]]) {
            return;
        }
        mergeOfTwoSortedArrAndCountSmaller(nums, left, mid, right, indexes, temp, res);
    }

    /**
     * [left, mid] 是排好序的,[mid + 1, right] 是排好序的
     *
     * @param nums
     * @param left
     * @param mid
     * @param right
     * @param indexes
     * @param temp
     * @param res
     */
    private void mergeOfTwoSortedArrAndCountSmaller(int[] nums, int left, int mid, int right, int[] indexes, int[] temp, int[] res) {
        for (int i = left; i <= right; i++) {
            temp[i] = indexes[i];
        }

        int i = left;
        int j = mid + 1;
        for (int k = left; k <= right; k++) {
            if (i > mid) {
                indexes[k] = temp[j];
                j++;
            } else if (j > right) {
                indexes[k] = temp[i];
                i++;
                res[indexes[k]] += (right - mid);
            } else if (nums[temp[i]] <= nums[temp[j]]) {
                // 注意:这里是 <= ,保证稳定性
                indexes[k] = temp[i];
                i++;
                res[indexes[k]] += (j - mid - 1);
            } else {
                indexes[k] = temp[j];
                j++;
            }
        }
    }

    public static void main(String[] args) {
        int[] nums = new int[]{5, 2, 6, 1};
        Solution solution = new Solution();
        List<Integer> countSmaller = solution.countSmaller(nums);
        System.out.println(countSmaller);
    }
}

c++代码

注意这里不能在递归函数里使用vector<int> tmp(indexs)每次都新建临时数组,会超时!!!直接建一个全局的。

c
class Solution {
    vector<int> res;
    vector<int> tmp;
    void mergeSort(vector<int>& nums,vector<int>& indexs,int left, int right){
        if(left>=right) return;
        int mid = left + (right - left)/2;
        mergeSort(nums,indexs,left,mid);
        mergeSort(nums,indexs,mid+1,right);
        for(int i = left;i<=right;++i){
            tmp[i] = indexs[i];
        }
        int l = left,r = mid+1;
        for(int i = left;i<=right;++i){
            if(r>right){//右边没了
                indexs[i] = tmp[l];
                res[tmp[l]] += right - mid;
                l++;
            }else if(l>mid){ //左边没了
                indexs[i] = tmp[r++];
            }else if(nums[tmp[l]]<=nums[tmp[r]]){// 注意:这里是 <= ,保证稳定性
                 indexs[i] = tmp[l];
                 res[tmp[l]] += r - mid - 1;
                 l++;
            }else{
                indexs[i] = tmp[r++];
            }
        }
    }
public:
    vector<int> countSmaller(vector<int>& nums) {
        vector<int> indexs(nums.size());
        res.resize(nums.size());
        tmp.resize(nums.size());
        for(int i = 0; i < nums.size(); ++i){
            indexs[i] = i;
        }
        mergeSort(nums,indexs,0,nums.size()-1);
        return res;
    }
};

复杂度分析:

  • 时间复杂度:\(O(N \log N)\),数组的元素个数是 \(N\),递归执行分治法,时间复杂度是对数级别的,因此时间复杂度是 \(O(N \log N)\)
  • 空间复杂度:\(O(N)\),需要 3 个数组,一个索引数组,一个临时数组用于索引数组的归并,还有一个结果数组,它们的长度都是 \(N\),故空间复杂度是 \(O(N)\)

用心记录,持续成长