Prefix sums: count possible startsLESSON 1.18 · 18 OF 23 IN CHAPTER
PART A / Data structures and algorithms
Step 20 of 252
LESSON 1.18 · 18 OF 23 IN CHAPTERWorked lesson

Prefix sums: count possible starts

“Count all contiguous transaction ranges that sum to 1 in [1,-1,1]. Negative values and repeated cumulative totals are allowed.”

There are three ranges. Explain why a window based only on whether the sum is too large cannot decide which boundary to move.

The coding-practice chapter will apply this tool to complete problems. Here, focus on the mechanism and trace how its state changes.

Working example: Count nonempty contiguous ranges summing to K. [1,-1,1], 1 → 3.

Prefix sums: count possible starts

The idea: A prior prefix equals current_prefix - K. Store its frequency. Seed {0:1} and look up before inserting.

First, what is a prefix sum?

A subarray is a contiguous stretch. A prefix sum is the total from the start through the current item. For [1, -1, 1], the running totals after each item are 1, 0, 1. If a past total was 0 and the current total is 1, the stretch between them sums to 1. Negative numbers can make a running total rise or fall, so a window that only shrinks when the total gets too large is unreliable.

counts = {0: 1}  # one empty prefix, before reading any item
prefix = 0
answer = 0
for value in [1, -1, 1]:
    prefix += value
    answer += counts.get(prefix - 1, 0)  # 1 is the desired sum
    counts[prefix] = counts.get(prefix, 0) + 1
print(answer)  # 3

This dictionary maps prefix total → number of earlier times we saw it. A set would remember only whether a total occurred, losing multiple valid starting positions. counts.get(key, 0) means “how many times, or zero if the key is absent.” The seed {0: 1} accounts for ranges starting at index 0.

Item / current total Needed earlier total Earlier count New qualifying ranges
1 / 1 0 1 [1] at positions 0–0
-1 / 0 −1 0 none
1 / 1 0 2 positions 0–2 and 2–2

Read the animation: the counter on the right is not an index map. It records the frequency of a prefix total. Notice the current total is counted only after its matches to earlier totals, excluding an empty range made by using the current state twice.

Check the mechanism

Predict each expected result, then trace the state that produces it. Explain the boundary case before opening the reference.

Cost: Running sum from each start: O(n²). Frequency map: expected O(n) time and O(n) space.

01 · Try this input

Input / starting state
[1, -1, 1], 1
Expected result
3

Qualifying ranges: positions 0–0, 0–2, 2–2

02 · Try this input

Input / starting state
[0, 0], 0
Expected result
3

Qualifying ranges: positions 0–0, 0–1, 1–1

03 · Try this input

Input / starting state
[], 0
Expected result
0

Qualifying ranges: The empty range does not count.

04 · Try this input

Input / starting state
[2, -2, 2], 0
Expected result
2

Qualifying ranges: positions 0–1 and 1–2

05 · Try this input

Input / starting state
[1, 2], 9
Expected result
0

Qualifying ranges: No qualifying range.

n is the number of input values. Trying all possible starting positions and running to every ending position examines O(n²) ranges. Keeping one running total and a dictionary uses expected O(n) lookup/update work and up to O(n) stored distinct totals. Even with an iterator, the dictionary may still grow with the stream.

Pass before moving on: Explain why a set undercounts and why sum-based shrinking fails with negative values.

Consume an iterator instead of a stored list. Which memory still grows?

After attempting: reference and explanation

Compare subarray_sum in algorithms.py (download file, source below). Use pattern notes for the invariant and contracts for complexity edge cases. Reimplement tomorrow without copying.

algorithms.py · algorithms.py
"""Reference solutions. Try the exercises in README.md before opening this file."""
from collections import Counter, OrderedDict, deque
from heapq import heappop, heappush, nlargest


def two_sum(nums, target):
    visited = {}
    for j, value in enumerate(nums):
        complement = target - value
        if complement in visited:
            return visited[complement], j
        visited.setdefault(value, j)
    return None


def longest_unique(text):
    left = best = 0
    last_seen = {}
    for right, char in enumerate(text):
        left = max(left, last_seen.get(char, -1) + 1)
        best = max(best, right - left + 1)
        last_seen[char] = right
    return best


def subarray_sum(nums, target):
    prefix_counts = {0: 1}
    prefix = result = 0
    for value in nums:
        prefix += value
        result += prefix_counts.get(prefix - target, 0)
        prefix_counts[prefix] = prefix_counts.get(prefix, 0) + 1
    return result


def lower_bound(nums, target):
    lo, hi = 0, len(nums)
    while lo < hi:
        mid = lo + (hi - lo) // 2
        if nums[mid] < target:
            lo = mid + 1
        else:
            hi = mid
    return lo


def merge_intervals(intervals):
    """Closed intervals; touching endpoints merge. Does not mutate input."""
    out = []
    for start, end in sorted(intervals):
        if start > end:
            raise ValueError('reversed interval')
        if out and start <= out[-1][1]:
            out[-1][1] = max(out[-1][1], end)
        else:
            out.append([start, end])
    return out


def top_k_frequent(nums, k):
    if k < 0:
        raise ValueError('k must be nonnegative')
    counts = Counter(nums)
    # Higher count wins; smaller value breaks ties deterministically.
    return nlargest(k, counts, key=lambda value: (counts[value], -value))


def islands(grid):
    """Rectangular list of lists of '0'/'1'; input remains unchanged."""
    if not grid:
        return 0
    cols = len(grid[0])
    if any(len(row) != cols or any(x not in ('0', '1') for x in row) for row in grid):
        raise ValueError('expected rectangular binary grid')
    seen = set()
    count = 0
    for r, row in enumerate(grid):
        for c, value in enumerate(row):
            if value != '1' or (r, c) in seen:
                continue
            count += 1
            seen.add((r, c))
            queue = deque([(r, c)])
            while queue:
                x, y = queue.popleft()
                for a, b in ((x-1, y), (x+1, y), (x, y-1), (x, y+1)):
                    if 0 <= a < len(grid) and 0 <= b < cols and grid[a][b] == '1' and (a, b) not in seen:
                        seen.add((a, b))  # Mark on enqueue, not dequeue.
                        queue.append((a, b))
    return count


def course_order(n, prerequisites):
    """Each pair is (course, prerequisite); [] means a cycle or no courses."""
    if n < 0:
        raise ValueError('negative course count')
    adjacency = [set() for _ in range(n)]
    degree = [0] * n
    for course, prerequisite in prerequisites:
        if not 0 <= course < n or not 0 <= prerequisite < n:
            raise ValueError('course outside graph')
        if course not in adjacency[prerequisite]:
            adjacency[prerequisite].add(course)
            degree[course] += 1
    queue = deque(i for i, d in enumerate(degree) if d == 0)
    order = []
    while queue:
        node = queue.popleft()
        order.append(node)
        for neighbor in adjacency[node]:
            degree[neighbor] -= 1
            if degree[neighbor] == 0:
                queue.append(neighbor)
    return order if len(order) == n else []


def shortest_paths(n, edges, source):
    """Directed nonnegative weighted graph; unreachable distances are infinity."""
    if not 0 <= source < n:
        raise ValueError('invalid source')
    graph = [[] for _ in range(n)]
    for u, v, weight in edges:
        if not 0 <= u < n or not 0 <= v < n or weight < 0:
            raise ValueError('invalid edge')
        graph[u].append((v, weight))
    distance = [float('inf')] * n
    distance[source] = 0
    heap = [(0, source)]
    while heap:
        cost, node = heappop(heap)
        if cost != distance[node]:
            continue
        for neighbor, weight in graph[node]:
            candidate = cost + weight
            if candidate < distance[neighbor]:
                distance[neighbor] = candidate
                heappush(heap, (candidate, neighbor))
    return distance


class LRU:
    """Capacity in entries; None denotes a miss. Not thread-safe."""
    def __init__(self, capacity):
        if capacity < 0:
            raise ValueError('negative capacity')
        self.capacity = capacity
        self.data = OrderedDict()

    def get(self, key):
        if key not in self.data:
            return None
        self.data.move_to_end(key)
        return self.data[key]

    def put(self, key, value):
        if self.capacity == 0:
            return
        self.data[key] = value
        self.data.move_to_end(key)
        if len(self.data) > self.capacity:
            self.data.popitem(last=False)


def min_coins(coins, amount):
    if amount < 0 or any(c <= 0 for c in coins):
        raise ValueError('nonnegative amount and positive coins required')
    dp = [0] + [amount + 1] * amount
    for total in range(1, amount + 1):
        for coin in coins:
            if coin <= total:
                dp[total] = min(dp[total], dp[total-coin] + 1)
    return -1 if dp[amount] > amount else dp[amount]


def daily_temperatures(temperatures):
    answer = [0] * len(temperatures)
    stack = []
    for i, temp in enumerate(temperatures):
        while stack and temperatures[stack[-1]] < temp:
            previous = stack.pop()
            answer[previous] = i - previous
        stack.append(i)
    return answer


def word_exists(board, word):
    """4-neighbor word search; do not reuse a cell or mutate the board."""
    if not word:
        return True
    if not board:
        return False
    cols = len(board[0])
    if any(len(row) != cols for row in board):
        raise ValueError('ragged board')
    def visit(r, c, index, used):
        if not (0 <= r < len(board) and 0 <= c < cols) or (r, c) in used or board[r][c] != word[index]:
            return False
        if index == len(word) - 1:
            return True
        used.add((r, c))
        found = any(visit(a, b, index+1, used) for a, b in ((r+1,c),(r-1,c),(r,c+1),(r,c-1)))
        used.remove((r, c))
        return found
    return any(visit(r, c, 0, set()) for r in range(len(board)) for c in range(cols))


class Trie:
    END = object()

    def __init__(self):
        self.root = {}

    def insert(self, word):
        node = self.root
        for char in word:
            node = node.setdefault(char, {})
        node[self.END] = True

    def contains(self, word):
        node = self.root
        for char in word:
            if char not in node:
                return False
            node = node[char]
        return self.END in node