Sorting Algorithm Derived Problems

Once we master basic sorting algorithms such as bubble sort and quicksort, we will realize that sorting is not just about putting data in order.

In real programming and interviews, many problems are variants or derivative problems of sorting algorithms.

Understanding these derived problems can help us grasp the essence of sorting algorithms more deeply and improve our ability to solve complex problems.

This article will explore several classic sorting algorithm derived problems, helping you build more systematic algorithmic thinking by analyzing their core ideas, solutions, and practical applications.


What are derived problems of sorting algorithms?

Sorting algorithm derived problems refer to thoseDoes not directly require sortingBut the core ideas, algorithm flow, or data structures used to solve the problems are closely related to classic sorting algorithms.Highly relevantproblems.

Such problems typically have the following characteristics:

  • Different goals: the ultimate goal is not to output an ordered sequence
  • Common ideasUse core ideas from sorting, such as comparison, swapping, and divide-and-conquer.
  • Complexity-relatedTime complexity is often in the same order of magnitude as sorting algorithms.
  • Widely used: frequently encountered in actual development

Below, let's gain a deeper understanding through several specific problems.


Classic variant problem 1: Top K problem

Problem definition

Find from n elementsthe K largest (or smallest) elements。

Practical application scenarios:

  • E-commerce website: find the 10 best-selling products
  • Social network: find the 100 users with the most followers
  • Data analysis: find the 5 most visited pages

Comparison of solutions

Methods Time complexity Space complexity Applicable scenarios
Sort directly and take the first K O(n log n) O(1) or O(n) When K is close to n
Bubble sort K times O(n × K) O(1) K is very small (K < 10)
Min heap / Max heap O(n log K) O(K) Most commonly used, K is much smaller than n
Quickselect algorithm O(n) average, O(n²) worst case O(1) Full order is not needed; only the Top K is required.

Method 1: Quickselect based on the idea of quicksort

The quickselect algorithm is a variant of quicksort. It does not need to completely sort the entire array; it only needs to find the position of the K-th largest element.

Example

def quick_select(nums, k):
    """
Use the quickselect algorithm to find the k-th smallest element
nums: input array
k: the k-th smallest element to find (1-based)
Return: the value of the k-th smallest element
    """

    def partition(left, right, pivot_idx):
        """Partition function that divides the array into two parts: less than and greater than the pivot"""
        pivot = nums[pivot_idx]
        # Move the pivot to the end
        nums[pivot_idx], nums[right] = nums[right], nums[pivot_idx]
       
        store_idx = left
        for i in range(left, right):
            if nums[i] < pivot:
                nums[store_idx], nums[i] = nums[i], nums[store_idx]
                store_idx += 1
       
        # Move the pivot to its correct position
        nums[right], nums[store_idx] = nums[store_idx], nums[right]
        return store_idx
   
    def select(left, right, k_smallest):
        """Main selection function"""
        if left == right:
            return nums[left]
       
        # Randomly select a pivot
        pivot_idx = left + (right - left) // 2
       
        # Partition and get the final position of the pivot
        pivot_idx = partition(left, right, pivot_idx)
       
        # The pivot is exactly the k-th smallest element
        if k_smallest == pivot_idx:
            return nums[k_smallest]
        # The k-th smallest element is in the left half
        elif k_smallest < pivot_idx:
            return select(left, pivot_idx - 1, k_smallest)
        # The k-th smallest element is in the right half
        else:
            return select(pivot_idx + 1, right, k_smallest)
   
    # k-1 because 0-based indexing is used internally
    return select(0, len(nums) - 1, k - 1)


def find_top_k(nums, k, largest=True):
    """
Find the largest k elements
nums: input array
k: number of elements to find
largest: True to find the largest k, False to find the smallest k
Return: the list of the largest k elements
    """

    n = len(nums)
    if k >= n:
        return sorted(nums, reverse=largest)
   
    if largest:
        # Find the largest k elements = find the (n-k+1)-th smallest element
        kth_value = quick_select(nums, n - k + 1)
        # Collect all elements greater than or equal to kth_value
        result = [x for x in nums if x > kth_value]
        # If the count is insufficient, add equal elements
        while len(result) < k:
            result.append(kth_value)
        return result
    else:
        # Find the smallest k elements = find the k-th smallest element
        kth_value = quick_select(nums, k)
        result = [x for x in nums if x < kth_value]
        while len(result) < k:
            result.append(kth_value)
        return result


# Test data
test_data = [3, 2, 1, 5, 6, 4, 9, 8, 7, 0]
print("Original data:", test_data)
print("The largest 3 elements:", find_top_k(test_data.copy(), 3, largest=True))
print("The smallest 3 elements:", find_top_k(test_data.copy(), 3, largest=False))

Method 2: Solution based on heap sort

A heap is a classic data structure for solving the Top K problem, especially suitable for data stream scenarios.

Example

import heapq

def find_top_k_with_heap(nums, k, largest=True):
    """
Use heap to find the largest k elements
nums: input array
k: number of elements to find
largest: True to find the largest k, False to find the smallest k
Return: the list of the largest k elements
    """

    if largest:
        # Find the largest k elements: use a min-heap
        heap = []
        for num in nums:
            if len(heap) < k:
                heapq.heappush(heap, num)
            elif num > heap[0]:  # Current element is larger than the heap top
                heapq.heapreplace(heap, num)
        return sorted(heap, reverse=True)
    else:
        # Find the smallest k elements: use a max-heap (implemented by taking negatives)
        heap = []
        for num in nums:
            if len(heap) < k:
                heapq.heappush(heap, -num)  # Store negatives
            elif num < -heap[0]:  # Current element is smaller than the heap top
                heapq.heapreplace(heap, -num)
        return sorted([-x for x in heap])


# Test heap method
print("\nUsing heap method:")
print("The largest 3 elements:", find_top_k_with_heap(test_data, 3, largest=True))
print("The smallest 3 elements:", find_top_k_with_heap(test_data, 3, largest=False))

Classic Derived Problem 2: Counting Inversions

Problem definition

In an array, if an earlier element is greater than a later element, then these two elements form anInversionsCalculate the total number of inversions in the array.

Practical application scenarios:

  • Measure the degree of order in data
  • User preference analysis in recommendation systems
  • Modification conflict detection in version control systems

Solution: Algorithm based on merge sort

During the merge sort process, the number of inversion pairs can be naturally counted. This is a typical application of the divide-and-conquer idea.

Example

def count_inversions(nums):
    """
Count the number of inversions in an array
Use a merge sort based approach
    """

    def merge_sort_count(left, right):
        """Merge sort and count inversions"""
        if left >= right:
            return 0, [nums[left]] if left == right else []
       
        mid = (left + right) // 2
       
        # Recursively count inversions in the left and right parts
        left_count, left_arr = merge_sort_count(left, mid)
        right_count, right_arr = merge_sort_count(mid + 1, right)
       
        # Count inversions that cross the midpoint
        cross_count = 0
        merged = []
        i = j = 0
       
        while i < len(left_arr) and j < len(right_arr):
            if left_arr[i] <= right_arr[j]:
                merged.append(left_arr[i])
                i += 1
            else:
                # left_arr[i] > right_arr[j] constitutes an inversion
                merged.append(right_arr[j])
                cross_count += len(left_arr) - i  # Key: all remaining elements on the left form inversions with right_arr[j]
                j += 1
       
        # Add remaining elements
        merged.extend(left_arr[i:])
        merged.extend(right_arr[j:])
       
        total_count = left_count + right_count + cross_count
        return total_count, merged
   
    if not nums:
        return 0
   
    count, _ = merge_sort_count(0, len(nums) - 1)
    return count


# Test inversion count
test_cases = [
    ([1, 2, 3, 4, 5], 0),  # Fully ordered, inversion count is 0
    ([5, 4, 3, 2, 1], 10), # Fully reversed, inversion count is C(5,2)=10
    ([2, 4, 1, 3, 5], 3),  # Example
    ([7, 5, 6, 4], 5),     # Complex example
]

print("\nInversion Count Test:")
for nums, expected in test_cases:
    result = count_inversions(nums.copy())
    print(f"Array {nums}: computed value={result}, expected value={expected}, {'Correct' if result == expected else 'Wrong'}")

The figure below shows how to count inversion pairs during the merge sort process. The key point is:When an element in the right half is smaller than an element in the left half, all remaining elements in the left half form inversion pairs with that right-half element.。


Classic variant problem 3: K-th largest / K-th smallest element

This problem is a special case of the Top K problem, but its solutions are more diverse.

Comparison of multiple solutions

Methods Time complexity Space complexity Features
Directly take after sorting O(n log n) O(1) or O(n) Simplest and most intuitive
Quickselect O(n) average O(1) Optimal average complexity
Heap method O(n log k) O(k) Suitable for data streams
Counting sort O(n + m) O(m) Suitable for cases with a small data range

Complete implementation example

Example

def find_kth_smallest(nums, k):
    """Find the k-th smallest element (1-based)"""
    # Method 1: Direct sorting
    def method_sort():
        return sorted(nums)[k-1]
   
    # Method 2: Quickselect (implemented above)
    def method_quick_select():
        return quick_select(nums.copy(), k)
   
    # Method 3: Heap method
    def method_heap():
        # Use a max heap to find the k-th smallest
        heap = []
        for num in nums:
            if len(heap) < k:
                heapq.heappush(heap, -num)  # Max heap is implemented using negative numbers
            elif num < -heap[0]:
                heapq.heapreplace(heap, -num)
        return -heap[0]
   
    # Method 4: Counting sort (suitable for small data ranges)
    def method_counting():
        if not nums:
            return None
       
        # Find the data range
        min_val, max_val = min(nums), max(nums)
        range_size = max_val - min_val + 1
       
        # Counting
        count = [0] * range_size
        for num in nums:
            count[num - min_val] += 1
       
        # Find the k-th smallest
        accumulated = 0
        for i in range(range_size):
            accumulated += count[i]
            if accumulated >= k:
                return i + min_val
        return None
   
    # Verify all methods produce consistent results
    result1 = method_sort()
    result2 = method_quick_select()
    result3 = method_heap()
    result4 = method_counting()
   
    print(f"\n{k}-th smallest element search test (Array: {nums}):")
    print(f"Sorting method: {result1}")
    print(f"Quickselect: {result2}")
    print(f"Heap method: {result3}")
    print(f"Counting method: {result4}")
   
    # Verify consistency
    assert result1 == result2 == result3 == result4, "Different methods have inconsistent results!"
    return result1


# Test data
test_nums = [3, 2, 3, 1, 2, 4, 5, 5, 6]
for k in [1, 3, 5, 7]:
    find_kth_smallest(test_nums.copy(), k)

Classic Derived Problem 4: Interval Merging

Problem definition

Given a set of intervals, merge all overlapping intervals.

Practical application scenarios:

  • Time period merging in calendar applications
  • Task time scheduling in project management
  • Merging network IP address segments

Solution: Sorting + linear scan

Example

def merge_intervals(intervals):
    """
Merge overlapping intervals
intervals: a list of lists, each sublist represents an interval [start, end]
Returns: the merged list of intervals
    """

    if not intervals:
        return []
   
    # Key step 1: Sort intervals by start point
    intervals.sort(key=lambda x: x[0])
   
    merged = []
    # Key step 2: Linear scan and merge
    current_start, current_end = intervals[0]
   
    for interval in intervals[1:]:
        start, end = interval
       
        if start <= current_end:  # Overlap
            # Merge intervals: take the largest end point
            current_end = max(current_end, end)
        else:  # No overlap
            # Save the current interval
            merged.append([current_start, current_end])
            # Start a new interval
            current_start, current_end = start, end
   
    # Add the last interval
    merged.append([current_start, current_end])
   
    return merged


# Test interval merging
test_intervals = [
    [[1, 3], [2, 6], [8, 10], [15, 18]],
    [[1, 4], [4, 5]],
    [[1, 4], [0, 4]],
    [[1, 4], [2, 3]],  # Fully contained
    [[1, 4], [0, 2], [3, 5]],  # Complex case
]

print("\nInterval merge test:")
for i, intervals in enumerate(test_intervals, 1):
    result = merge_intervals(intervals.copy())
    print(f"Test case {i}:")
    print(f" Input: {intervals}")
    print(f" Output: {result}")
    print()

Algorithm process analysis


Practice exercises

Now, let's reinforce what we have learned through a few exercises:

Exercise 1: Sort Colors (Dutch national flag problem)

Given an array containing n elements that are red, white, or blue, sort them in place so that elements of the same color are adjacent and are arranged in the order red, white, and blue.

Example

def sort_colors(nums):
    """
Sort the color array in place
0: red, 1: white, 2: blue
Use the three-pointer method (similar to the partition idea of quicksort)
    """

    # Initialize three pointers
    left = 0  # Right boundary of the red region
    current = 0  # Current element being checked
    right = len(nums) - 1  # Left boundary of the blue region
   
    while current <= right:
        if nums[current] == 0:  # Red
            nums[left], nums[current] = nums[current], nums[left]
            left += 1
            current += 1
        elif nums[current] == 1:  # White
            current += 1
        else:  # Blue
            nums[current], nums[right] = nums[right], nums[current]
            right -= 1
            # Note: do not increment current here, because the swapped-in element has not been checked yet
   
    return nums


# Test
colors = [2, 0, 2, 1, 1, 0]
print("Before sorting colors:", colors)
print("After sorting colors:", sort_colors(colors.copy()))

Exercise 2: K-th largest element in an array

Find the k-th largest element in an unsorted array.

Example

def find_kth_largest(nums, k):
    """
Find the kth largest element
Use the quickselect algorithm
    """

    # The kth largest = the (n-k+1)th smallest
    k_smallest = len(nums) - k + 1
   
    def quick_select(arr, left, right, k_small):
        if left == right:
            return arr[left]
       
        # Choose a pivot
        pivot_idx = (left + right) // 2
        pivot = arr[pivot_idx]
       
        # Partition
        arr[pivot_idx], arr[right] = arr[right], arr[pivot_idx]
        store_idx = left
       
        for i in range(left, right):
            if arr[i] < pivot:
                arr[store_idx], arr[i] = arr[i], arr[store_idx]
                store_idx += 1
       
        arr[right], arr[store_idx] = arr[store_idx], arr[right]
       
        # Recursive selection
        if k_small == store_idx:
            return arr[store_idx]
        elif k_small < store_idx:
            return quick_select(arr, left, store_idx - 1, k_small)
        else:
            return quick_select(arr, store_idx + 1, right, k_small)
   
    return quick_select(nums.copy(), 0, len(nums) - 1, k_smallest - 1)


# Test
test_array = [3, 2, 1, 5, 6, 4]
k = 2
print(f"\nThe {k}-th largest element in the array {test_array} is: {find_kth_largest(test_array, k)})
other extensions