Chapter 5 · Python for AI Work
Data Structures and Algorithms: Writing Fast Code
- Page 19 of 23
- 7 min read
AI programs handle a lot of data: thousands of documents, millions of embeddings, a queue of requests waiting for an API. With small data any code feels fast. With real data, the difference between a good and a bad choice of data structure (how you store things) and algorithm (the steps you take) is the difference between a second and an hour. This page gives you the essentials that every interviewer — and every production system — expects.
Big-O: how the work grows
Big-O notation describes how the number of steps grows as the input gets bigger. It ignores the exact speed of your computer and asks one question: if the data were ten times bigger, how much more work would this be?
| Big-O | Name | Steps for n = 1,000,000 | Python example |
|---|---|---|---|
O(1) | constant | 1 | d[key], x in a_set, items.append(x) |
O(log n) | logarithmic | about 20 | binary search, bisect, heapq.heappush |
O(n) | linear | 1,000,000 | a loop over every item, x in a_list, sum(items) |
O(n log n) | "n log n" | about 20,000,000 | sorted(items) |
O(n²) | quadratic | 1,000,000,000,000 | a loop inside a loop over the same data |
The last row is the one to fear. Comparing every document with every other document is O(n²): fine for 100 documents, impossible for a million.
Linear search vs binary search
To find something in a list, the simple way is to look at every item in turn. If the list is sorted, there is a much better way: look at the middle, throw away the half that cannot contain the target, and repeat.
def linear_search(items, target):
"""Check every item in turn. Returns (position, steps)."""
steps = 0
for i, item in enumerate(items):
steps += 1
if item == target:
return i, steps
return -1, steps
def binary_search(items, target):
"""Items must be sorted. Halve the search range each step."""
low, high, steps = 0, len(items) - 1, 0
while low <= high:
steps += 1
mid = (low + high) // 2
if items[mid] == target:
return mid, steps
if items[mid] < target:
low = mid + 1 # target is in the right half
else:
high = mid - 1 # target is in the left half
return -1, steps
ids = list(range(0, 2_000_000, 2)) # one million sorted even numbers
print(linear_search(ids, 1_999_998))
print(binary_search(ids, 1_999_998))(999999, 1000000)
(999999, 20)One million steps against twenty. Each step of binary search halves the range, and you can only halve a million about 20 times before one item is left — that is O(log n). You rarely write it yourself: Python's bisect module does it for you on any sorted list.
Pick the right structure: list, set or dict
A list keeps items in order, so x in a_list must check them one by one. A set or a dict uses hashing: Python turns the value into a number that says exactly where to look, so the check takes one step no matter how big it is.
import timeit
ids_list = list(range(1_000_000))
ids_set = set(ids_list)
in_list = timeit.timeit(lambda: 999_999 in ids_list, number=100)
in_set = timeit.timeit(lambda: 999_999 in ids_set, number=100)
print(f"list: {in_list * 1000:.1f} ms for 100 lookups")
print(f"set: {in_set * 1000:.3f} ms for 100 lookups")
print(f"the set was about {in_list / in_set:,.0f} times faster")list: 813.8 ms for 100 lookups
set: 0.006 ms for 100 lookups
the set was about 131,255 times fasterYour numbers will differ — they depend on your computer — but the gap will be enormous. Whenever you write if x in something inside a loop, ask whether something should be a set. Need to look up extra data by a key, such as a document's text by its id? Use a dict.
Stacks and queues
Two simple patterns appear everywhere. A stack is last in, first out — like a pile of plates, or an undo history. A plain list is a perfect stack. A queue is first in, first out — like people waiting in line, or jobs waiting for an API. Use collections.deque for that: removing from the front of a list is O(n) because every other item has to move, while deque.popleft() is O(1).
from collections import deque
# Stack: last in, first out — an undo history
history = []
history.append("typed 'Hello'")
history.append("made it bold")
history.append("deleted a line")
print("undo:", history.pop())
print("undo:", history.pop())
print("left:", history)
# Queue: first in, first out — requests waiting for a rate-limited API
waiting = deque(["summarise doc1", "summarise doc2"])
waiting.append("summarise doc3")
while waiting:
print("sending:", waiting.popleft())undo: deleted a line
undo: made it bold
left: ["typed 'Hello'"]
sending: summarise doc1
sending: summarise doc2
sending: summarise doc3Heaps: the top k without sorting everything
Retrieval in a RAG system ends with "give me the k most similar documents". Sorting all the scores is O(n log n); heapq.nlargest keeps only the best k as it goes, which is faster when k is small. A heap is also a priority queue: the item with the smallest priority number always comes out first.
import heapq
# How similar each document is to a question (from embeddings, page 13)
scores = {"refunds.md": 0.91, "shipping.md": 0.42, "warranty.md": 0.77,
"privacy.md": 0.18, "returns.md": 0.88, "careers.md": 0.05}
top3 = heapq.nlargest(3, scores.items(), key=lambda pair: pair[1])
print(top3)
# A priority queue: the smallest number comes out first
jobs = []
heapq.heappush(jobs, (2, "nightly report"))
heapq.heappush(jobs, (0, "customer is waiting"))
heapq.heappush(jobs, (1, "retry failed call"))
while jobs:
priority, job = heapq.heappop(jobs)
print(priority, job)[('refunds.md', 0.91), ('returns.md', 0.88), ('warranty.md', 0.77)]
0 customer is waiting
1 retry failed call
2 nightly reportRecursion: a function that calls itself
Some data is nested to an unknown depth: JSON from an API, folders inside folders, a tree of agent steps. A recursive function handles it naturally: it deals with the simple case directly (the base case) and calls itself on each smaller piece. Here it counts every value in a nested model response:
def count_values(data):
"""Count every plain value inside nested dicts and lists."""
if isinstance(data, dict):
return sum(count_values(v) for v in data.values())
if isinstance(data, list):
return sum(count_values(v) for v in data)
return 1 # base case: a single value
response = {
"model": "gpt-5.5",
"output": [
{"type": "text", "text": "Hello!"},
{"type": "tool_call", "name": "search", "args": {"query": "tokens", "limit": 3}},
],
"usage": {"input_tokens": 12, "output_tokens": 5},
}
print(count_values(response))9Every recursive function needs a base case that stops it. Without one, it calls itself until Python gives up with a RecursionError (by default after about 1,000 levels).
Memoisation: never compute the same thing twice
A naive recursive function can repeat the same work again and again. @cache from functools — a decorator, as on page 8 — remembers each answer the first time it is computed:
from functools import cache
calls = 0
def fib(n):
global calls
calls += 1
return n if n < 2 else fib(n - 1) + fib(n - 2)
print(fib(25), "calls:", calls)
calls = 0
@cache
def fib_fast(n):
global calls
calls += 1
return n if n < 2 else fib_fast(n - 1) + fib_fast(n - 2)
print(fib_fast(25), "calls:", calls)75025 calls: 242785
75025 calls: 26Same answer, 242,785 calls against 26. Caching is one of the most effective optimisations in AI apps too: if the same prompt or the same embedding request comes in twice, answer from the cache instead of paying for another model call.
Try it yourself
- Use
bisect.bisect_leftto find where1_999_998sits inids, and check it matches the binary search above. - You have 10,000 banned words and 1,000,000 messages. Write the check so that it is fast — which structure holds the banned words?
- Write a recursive function that finds the deepest level of nesting in a dict.