优化元素和固定的排列生成器:迭代器转换与结果总数预计算需求
Great question! Let's break down how to solve both of your requirements step by step.
1. Convert to an Iterator (Streaming Results)
Your current recursive approach builds all results upfront and stores them in a list, which is slow and memory-heavy for large inputs. Switching to a generator (iterator) with backtracking will let you stream results one at a time, saving memory and improving responsiveness.
Here's an optimized implementation using yield to create a generator:
def permutations_with_sum(sample_size, desired_sum, set_of_numbers): # Deduplicate and sort numbers to optimize pruning sorted_unique_nums = sorted(set(set_of_numbers)) def backtrack(current_pos, current_arr, remaining_sum): # If we've filled all positions, check if we hit the target sum if current_pos == sample_size: if remaining_sum == 0: yield current_arr.copy() # Return a copy to avoid mutation issues return # Iterate through valid numbers (prune early since sorted) for num in sorted_unique_nums: if num > remaining_sum: break # No need to check larger numbers, since list is sorted current_arr[current_pos] = num # Recurse to fill the next position, with updated remaining sum yield from backtrack(current_pos + 1, current_arr, remaining_sum - num) # Initialize the array to reuse memory during backtracking yield from backtrack(0, [0] * sample_size, desired_sum)
How it works:
- Backtracking with Pruning: We sort the number set so we can stop iterating once a number exceeds the remaining sum (since larger numbers will also be too big).
- Streaming with Yield: Instead of collecting all results in a list, we
yieldeach valid permutation as it's found. This means you can iterate through results without loading everything into memory. - Memory Efficiency: We reuse a single array during backtracking, only copying it when we yield a valid result.
Usage Example:
# Iterate through results one at a time for perm in permutations_with_sum(4, 10, [1, 2, 3]): print(perm)
This will output the same permutations as your original code, but in a sorted order (due to the sorted number set).
2. Precalculate Total Number of Results
To find how many valid permutations exist without generating them, we can use dynamic programming (DP). This avoids the overhead of building all permutations and is much faster for large inputs.
DP Approach
We define dp[i][s] as the number of permutations of i elements that sum to s. The recurrence relation is:
dp[i][s] = sum(dp[i-1][s - num])for everynumin your number set wheres >= num- Base case:
dp[0][0] = 1(one way to make sum 0 with 0 elements)
Here's the implementation:
def count_permutations_with_sum(sample_size, desired_sum, set_of_numbers): nums = sorted(set(set_of_numbers)) # Initialize DP table: rows = number of elements, columns = possible sums dp = [[0] * (desired_sum + 1) for _ in range(sample_size + 1)] dp[0][0] = 1 # Base case for i in range(1, sample_size + 1): for s in range(desired_sum + 1): for num in nums: if s >= num: dp[i][s] += dp[i-1][s - num] return dp[sample_size][desired_sum]
Optimized Space Version
Since we only need the previous row (i-1) to compute the current row (i), we can reduce space complexity from O(n*s) to O(s):
def count_permutations_with_sum_optimized(sample_size, desired_sum, set_of_numbers): nums = sorted(set(set_of_numbers)) dp = [0] * (desired_sum + 1) dp[0] = 1 # Base case for _ in range(sample_size): new_dp = [0] * (desired_sum + 1) for s in range(desired_sum + 1): for num in nums: if s >= num: new_dp[s] += dp[s - num] dp = new_dp return dp[desired_sum]
Usage Example:
total = count_permutations_with_sum(4, 10, [1, 2, 3]) print(total) # Output: 10 (matches your original result count)
内容的提问来源于stack exchange,提问作者mullin

