寻求Codejam 2015无限煎饼屋问题O(D√M+M)算法的Python实现思路
Great question! Let's walk through how to implement the O(D√M + M) optimized approach step by step, with clear Python code examples and explanations.
Core Idea Recap
The key insight is that the sequence ceil(p/x) for a single pancake stack p only changes O(√p) times as x increases. We use a difference array to efficiently track these changes across all stacks, then compute the prefix sum to get sum(ceil(p_i/x)) for all x in O(M) time.
Step-by-Step Implementation
1. Initialize Variables & Difference Array
First, we'll read input values, find the maximum stack size M, and set up a difference array to track changes in sum(ceil(p_i/x)):
def solve_optimized(): import sys input = sys.stdin.read().split() ptr = 0 T = int(input[ptr]) ptr +=1 for case in range(1, T+1): D = int(input[ptr]) # number of diners with non-empty plates ptr +=1 diners = list(map(int, input[ptr:ptr+D])) ptr +=D M = max(diners) # Difference array: delta[x] tracks the change in sum_ceil at position x delta = [0]*(M +2) # +2 to avoid index out of bounds for x_max+1
2. Populate the Difference Array
For each pancake stack p, we split into two parts to cover all possible x values efficiently:
- Part 1: Handle small
qvalues (≤√p), which correspond to largexvalues (>√p). We update the difference array for the range ofxwhereceil(p/x) = q. - Part 2: Handle small
xvalues (≤√p), which correspond to largeqvalues (>√p). We directly add theqvalue to the difference array for eachx.
for p in diners: sqrt_p = int(p**0.5) # Part 1: q <= sqrt(p) → x >= sqrt(p) for q in range(1, sqrt_p +1): if q ==1: # ceil(p/x)=1 when x >=p x_min = p x_max = M else: # Calculate x range where ceil(p/x)=q using integer arithmetic x_min = (p + q -1) // q # ceil(p/q) x_max = (p -1) // (q -1) # floor( (p-1)/(q-1) ) # Clamp to our max x (M) since we don't care about x>M x_min = min(x_min, M) x_max = min(x_max, M) if x_min > x_max: continue # Update difference array for this interval delta[x_min] += q if x_max +1 <= M: delta[x_max +1] -= q # Part 2: x <= sqrt(p) → q=ceil(p/x) > sqrt(p) for x in range(1, sqrt_p +1): q = (p + x -1) // x # Equivalent to ceil(p/x) delta[x] += q delta[x+1] -= q
3. Compute Prefix Sum to Get sum_ceil
Now we calculate the prefix sum of the difference array to get sum_ceil[x] = sum(ceil(p_i/x)) for each x:
# Compute prefix sum to get sum_ceil[x] sum_ceil = [0]*(M+1) current_sum =0 for x in range(1, M+1): current_sum += delta[x] sum_ceil[x] = current_sum
4. Find the Minimum Time
Finally, we compute the candidate time for each x (time = x + sum(ceil(p_i/x)-1) → which simplifies to x + sum_ceil[x] - D) and track the minimum:
min_time = M # Initial value is the max stack size for x in range(1, M+1): total_time = x + (sum_ceil[x] - D) if total_time < min_time: min_time = total_time print(f'Case #{case}: {min_time}')
Key Notes
- Integer Arithmetic: We use
(p + k -1) // kto computeceil(p/k)without floating-point errors, which is crucial for correctness. - Difference Array Efficiency: This data structure lets us perform range updates in O(1) time per interval, which is why we can handle all stacks in O(D√M) time.
- Boundary Handling: We clamp
x_minandx_maxtoMbecause we don't need to considerxvalues larger than the biggest stack (since the initial candidate time is alreadyM).
This implementation will be significantly faster than the original O(D*M) approach for large values of M (e.g., when M is 1e5 or higher).
内容的提问来源于stack exchange,提问作者Eugene Yarmash

