基于基准值(pivot)实现Python列表元素分区的技术问询
Hey folks, let's tackle this list partitioning problem in Python 3.x. The goal is to rearrange the list so that all elements ≤ our pivot value come before it, and everything greater comes after. Your existing code has a good start, but let's fix the logic to handle this correctly (especially since there are duplicate pivot values in your example list!).
Problem with the Original Code Snippet
Your current approach uses b.index(pivot) in each loop iteration, which has two key issues:
- If there are multiple instances of the pivot (like two 4s in your list),
index()only returns the first occurrence's index—this will break the partitioning logic. - You haven't defined what to do when
a[i] > pivot, leaving the code incomplete.
Solution 1: Simple Two-List Approach (Easy to Read)
This method uses two separate lists to collect elements that belong on each side of the pivot, then combines them. It's straightforward and handles duplicates perfectly:
a = [1,4,3,7,4,7,6,3,7,8,9,9,2,5] print("Original List") print(a) pivot = 4 # Feel free to change this to any value from the list # Split elements into two groups less_or_equal = [] greater = [] for num in a: if num <= pivot: less_or_equal.append(num) else: greater.append(num) # Combine the groups to get the partitioned list partitioned_list = less_or_equal + greater print("\nPartitioned List (pivot =", pivot, ")") print(partitioned_list)
Output:
Original List [1, 4, 3, 7, 4, 7, 6, 3, 7, 8, 9, 9, 2, 5] Partitioned List (pivot = 4 ) [1, 4, 3, 4, 3, 2, 7, 7, 6, 7, 8, 9, 9, 5]
Solution 2: In-Place Partitioning (Space-Efficient)
If you want to avoid using extra lists and modify a copy of the original list directly, you can use a two-pointer technique to rearrange elements in place:
a = [1,4,3,7,4,7,6,3,7,8,9,9,2,5] print("Original List") print(a) pivot = 4 # or select any number from the list b = list(a) # Work on a copy to keep the original list intact # Pointer to track where to place the next element <= pivot current_position = 0 # Iterate through each element for i in range(len(b)): if b[i] <= pivot: # Swap the current element with the element at current_position b[current_position], b[i] = b[i], b[current_position] current_position += 1 print("\nPartitioned List (pivot =", pivot, ")") print(b)
How This Works:
- We start with
current_positionat 0 (the start of the list). - For every element ≤ pivot, we swap it with the element at
current_position, then increment the pointer. - By the end of the loop, all elements before
current_positionare ≤ pivot, and everything after is > pivot.
Both solutions run in O(n) time (we only traverse the list once) and get the job done cleanly.
内容的提问来源于stack exchange,提问作者jjacobson

