記憶體上限下的演算法
範圍 — 輸入放不進 RAM 時該怎麼做:位元向量、先分桶再細化的多趟掃描、依雜湊切分檔案,以及外部合併排序——這些技巧以多掃幾趟為代價,仍能回傳精確答案。 另見:streaming_algorithms.md — 同一問題中單趟、近似的那一半(Bloom filter、count-min、reservoir sampling);bit_manipulation.md — 下文用到的位元運算子;sort.md — 這些技巧所依賴的記憶體內排序。
LeetCode 題目清單
0) 概念
面試官每次都用同樣的方式描述這類題目:「你有 40 億個數字和 10 MB 記憶體」、 「用 4 KB 在 1 到 32,000 的陣列中找重複」、「用 1 GB RAM 排序 1 TB」。問題從來不是 哪個演算法最快——而是哪種資料表示法能塞進預算,以及你願意花幾趟掃描來讓答案保持精確。
三個招式幾乎涵蓋全部:
- 縮小每個項目——用一個位元取代一個 4 位元組的 int,就是 32 倍的節省。
- 縮小問題——用一趟便宜的掃描把搜尋範圍縮到一個夠小的切片,再用第二趟好好解決。
- 切分問題——把資料切成各自放得下的獨立片段,並使用能讓相關項目待在一起的切分規則。
0-1) 先算預算 優先度 5/5 — 必備 — 幾乎每一輪面試都會出現
在提出任何方案前,先把數字說出來。這正是被評分的部分。
| 要追蹤的值 | 以 4 位元組 int 儲存 |
以每個一位元儲存 |
|---|---|---|
| 32,000 | 128 KB | 4 KB |
| 1 百萬 | 4 MB | 125 KB |
2^31(所有非負 int) |
8 GB | 256 MB |
2^32(所有 int) |
16 GB | 512 MB |
反過來看——一份預算能買到什麼:
10 MB = 10 * 2^20 bytes * 8 ~= 83.9 million bits -> 83.9M distinct flags
= 10 * 2^20 / 4 ~= 2.6 million ints -> 2.6M counters
1 GB = 2^30 bytes * 8 = 8.6 billion bits -> a bit per 32-bit value, twice over
所以「40 億個數字、1 GB」是單趟位元向量問題,而「40 億個數字、10 MB」不是——10 MB 裝不下 2^32 個位元,因此需要 §1-2 的兩趟細化。
0-2) 該用哪種技巧 優先度 4/5 — 高價值 — 這裡有缺口就會掉關
| 題目中的訊號 | 技巧 | 記憶體 | 趟數 |
|---|---|---|---|
| 在有界值域上判斷成員/重複 | 位元向量(§1-1) | 每個可能值 1 位元 | 1 |
| 整個值域的位元向量放不下 | 分桶計數,再細化(§1-2) | 每個區塊 1 個計數器 + 1 個區塊的位元 | 2 |
| 鍵必須分組(去重、計數、join、top-k) | 依雜湊切分(§1-3) | 一次 1 個桶 | 2 |
| 輸出必須有序 | 外部合併排序(§1-4) | 1 個區塊,之後每個 run 1 筆記錄 | 1 + log_k(runs) |
| 可接受近似答案 | Sketch——見 streaming_algorithms.md | 次線性 | 1 |
1) 一般形式
1-1) 位元向量——每個可能值一個位元 優先度 5/5 — 必備 — 幾乎每一輪面試都會出現
一個 int 的雜湊集合每筆要花數十個位元組——在 Java 中,一個裝箱的 Integer 加上一個
HashMap 節點約 32–48 位元組,而 Python 的 set 算上負載因子後也是同一個量級。當值
密集且有界時,就降到一個位元,並用算術來索引:pos >> 3 選出位元組,pos & 7 選出其中的位元。
# python
# GENERAL PATTERN: bit vector (bitset) over values 0 .. size-1
# IDEA: byte i of the array holds the flags for values 8i .. 8i+7
# time = O(1) per op, space = size BITS (size/8 bytes)
class BitVector(object):
def __init__(self, size):
self.bits = bytearray((size + 7) // 8) # 1 byte == 8 flags
def get(self, pos):
return (self.bits[pos >> 3] >> (pos & 7)) & 1
def set(self, pos):
self.bits[pos >> 3] |= 1 << (pos & 7)
// java
// GENERAL PATTERN: bit vector over values 0 .. size-1
// IDEA: word (pos >> 5) of an int[] holds the flags for 32 consecutive values
// time = O(1) per op, space = size BITS
class BitVector {
private final int[] words;
BitVector(int size) { words = new int[(size >> 5) + 1]; } // /32, rounded up
boolean get(int pos) { return (words[pos >> 5] & (1 << (pos & 31))) != 0; }
void set(int pos) { words[pos >> 5] |= 1 << (pos & 31); }
}
pos & 31就是pos % 32,pos >> 5就是pos / 32——兩者都精確成立,因為 32 是 2 的冪。 Java 自己的java.util.BitSet正是這樣做的;但當題目說*「用 4 KB 記憶體」*時,還是把它寫出來, 因為這段算術就是答案。
實際應用——用 4 KB 在 [1, 32000] 中找重複(CtCI 10.8):32,000 個位元是 4,000 位元組,
所以整個值域都放得下,一趟就夠。
# python
# CtCI 10.8 - print duplicates in an array of values in [1, 32000], memory ~ 4 KB
# IDEA: the value IS the index — flip its bit; a bit already set means a repeat
# time = O(n), space = 32000 bits = 4 KB
def print_duplicates(nums):
seen = BitVector(32000)
for n in nums:
if seen.get(n - 1): # values start at 1, bits start at 0
print(n)
else:
seen.set(n - 1)
這個技巧的 LeetCode 版本是陣列本身就是位元向量:LC 41(First Missing Positive)與
LC 448 透過就地將 nums[v - 1] 取負來標記「值 v 出現過」,這是一個完全不花額外記憶體的位元向量。
1-2) 分桶計數,再細化 優先度 5/5 — 必備 — 幾乎每一輪面試都會出現
當每個值一個位元仍超出預算時,就花第一趟把值計數到各區塊中。收到的值比格數少的區塊 必定缺了一個——所以第二趟只需要對那單一區塊建立位元向量。
# python
# CtCI 10.7 - find a missing non-negative int among ~4 billion, memory ~ 10 MB
# IDEA: pass 1 counts values per block of 2^20; a block holding < 2^20 values must miss one
# pass 2 builds a bit vector for THAT BLOCK ONLY (2^20 bits = 128 KB)
# time = O(n) over 2 passes, space = 2048 counters (8 KB) + 128 KB
BLOCK = 1 << 20 # 2^20 values per block
def find_missing(read_all, limit=1 << 31): # read_all() -> a fresh iterator each call
counts = [0] * -(-limit // BLOCK) # 2048 counters; ceiling covers a part-block
for v in read_all(): # ---- pass 1
counts[v // BLOCK] += 1
def slots(i): # the last block may be shorter than BLOCK
return min(BLOCK, limit - i * BLOCK)
block = next((i for i, c in enumerate(counts) if c < slots(i)), -1)
if block < 0:
return -1 # every block is full: nothing is missing
lo, width = block * BLOCK, slots(block)
seen = bytearray((width + 7) // 8) # ---- pass 2, 128 KB for a full block
for v in read_all():
if lo <= v < lo + width:
off = v - lo
seen[off >> 3] |= 1 << (off & 7)
for off in range(width):
if not (seen[off >> 3] >> (off & 7)) & 1:
return lo + off
return -1 # no gap: the input held every value
有兩件事要說出口:
- 為什麼計數就足以選出區塊。 值互不相同,所以一個大小為
2^20、收到少於2^20個值的區塊,可證明有一個空缺。不需要知道是哪一個。 - 如何決定區塊大小。 第 1 趟需要
limit / BLOCK個計數器,第 2 趟需要BLOCK個位元,兩者都必須放得下。BLOCK = 2^20對上面非負int的範圍花費 8 KB + 128 KB(2048 個計數器), 若涵蓋全部2^32個值則是 16 KB + 128 KB——兩者都遠低於 10 MB。
這個形態可推廣為先粗略計數,再放大。它是對高位元做計數排序的磁碟版本,也是在巨大檔案上 回答第 k 大查詢的方法:每個桶計數,找出包含第 k 名的桶,只重新掃描那個桶。
1-3) 依雜湊切分——讓片段彼此獨立 優先度 4/5 — 高價值 — 這裡有缺口就會掉關
計算 10 GB 檔案中的字數、URL 去重、join 兩個巨大檔案:障礙在於同一個鍵可能出現在任何地方。
選擇讓這件事不可能發生的切分規則來消除它——把每筆記錄送到桶 h(key) % B,一個鍵的每個副本
都會落在同一個桶。接著每個桶都是普通的記憶體內問題,結果直接串接即可。
# python
# GENERAL PATTERN: shard a too-big file into buckets that each fit in RAM
# IDEA: h(key) % B decides the bucket -> identical keys can never split across buckets
# time = O(n) to shard + O(n) to process, space = O(largest bucket)
import hashlib
def bucket_of(key, n_buckets):
# NOT the builtin hash(): Python salts str hashing per process, so a re-run
# would send the same key to a different file
digest = hashlib.md5(key.encode()).digest()
return int.from_bytes(digest[:4], "big") % n_buckets
def count_words(lines, n_buckets, tmpdir):
files = [open("%s/part-%d" % (tmpdir, i), "w") for i in range(n_buckets)]
for line in lines: # ---- pass 1: scatter
for word in line.split():
files[bucket_of(word, n_buckets)].write(word + "\n")
for f in files:
f.close()
for i in range(n_buckets): # ---- pass 2: gather
counts = {} # one bucket fits in RAM
with open("%s/part-%d" % (tmpdir, i)) as f:
for word in f:
word = word.rstrip("\n")
counts[word] = counts.get(word, 0) + 1
for word, c in counts.items():
yield word, c
選擇 B 時要讓最大的桶放得下,而不是平均的桶——一個偏斜的鍵(被存取十億次的 URL)
仍然只會落在單一檔案裡,而且換哪個雜湊函數都切不開它——在運算允許時,改成先行聚合(資料串流經過時就在記憶體中計數或加總)。
如果某個桶是因為太多不同的鍵擠在一起而溢出,
就用另一個雜湊函數把那個桶再切一次。這就是每個分散式引擎中
GROUP BY 做的事,也是「你會怎麼擴展它」的誠實答案:各桶彼此獨立,所以不需修改就能搬到不同機器上。
1-4) 外部合併排序 優先度 4/5 — 高價值 — 這裡有缺口就會掉關
用 1 GB RAM 排序 1 TB:把輸入切成放得下的區塊,每塊在記憶體中排序並寫出成一個已排序的 run,再用一個每個 run 持有一筆記錄的 heap 合併所有 run。
# python
# GENERAL PATTERN: external merge sort — sort more data than fits in memory
# IDEA: phase 1 turns the input into sorted runs; phase 2 k-way merges them, keeping
# only k heads (one per run) in RAM
# time = O(n log n) compares; I/O = O(n) per pass, passes = 1 + log_k(#runs)
# NOTE: ONE merge pass, so it assumes len(runs) fits the fan-in the process can hold
# open at once. More runs than that must be merged in rounds — see below.
import heapq
from contextlib import ExitStack
def external_sort(records, chunk_size, tmpdir):
runs, chunk = [], []
for rec in records: # ---- phase 1: sorted runs
chunk.append(rec)
if len(chunk) == chunk_size:
runs.append(_flush(chunk, tmpdir, len(runs)))
chunk = []
if chunk:
runs.append(_flush(chunk, tmpdir, len(runs)))
with ExitStack() as stack: # ---- phase 2: k-way merge
files = [stack.enter_context(open(p)) for p in runs] # closed on the way out,
for line in heapq.merge(*files, key=int): # even if the caller
yield int(line) # abandons the generator
def _flush(chunk, tmpdir, i):
chunk.sort()
path = "%s/run-%d" % (tmpdir, i)
with open(path, "w") as f:
f.writelines("%d\n" % x for x in chunk)
return path
- 為什麼用 heap。 靠掃描每個開頭來合併
k個 run,每筆記錄要O(k);heap 讓它變成O(log k)。這就是 LC 23(Merge k Sorted Lists)——唯一差別是每個串列都存在磁碟上。 - 為什麼
k有上限。 每個開啟的 run 都要花一個檔案描述符以及一個讀取緩衝區,所以k同時受限於RAM / buffer size與行程的描述符上限(ulimit -n,通常為 256–1024)。 太小的chunk_size會產生數千個 run,上面的單趟合併就會直接失敗——改成每輪合併k個:passes = 1 + ceil(log_k(runs))。 - 它換來什麼。 資料一旦排序,去重、分組與集合交集都各只是一次雙指標的線性掃描,不需額外記憶體。
2) LC 範例
以下是訓練同樣反射動作的記憶體內題目;面試時要點出其中的關聯。
| # | 題目 | 它訓練的記憶體思路 |
|---|---|---|
| 41 | First Missing Positive | 陣列本身就是位元向量——就地取負來標記 |
| 448 | Find All Numbers Disappeared in an Array | 同樣的就地標記,不需額外空間 |
| 268 | Missing Number | 用總和/XOR 取代儲存任何東西 |
| 287 | Find the Duplicate Number | 唯讀輸入、O(1) 空間——用 Floyd 循環偵測,而不是已見集合 |
| 23 | Merge k Sorted Lists | 外部排序的 k 路合併階段 |
| 692 / 703 | Top K Frequent Words / Kth Largest in a Stream | 有界 heap——切分後每個桶的處理步驟 |
2-1) 要說出口的四個問題 優先度 4/5 — 高價值 — 這裡有缺口就會掉關
- 值域是什麼? 有界且密集 → 位元向量。無界 → 雜湊切分。
- 以位元計的預算是多少? 換算出來,再說整個值域是否放得下。
- 我可以掃幾趟? 一趟 → sketch,答案是近似的。兩趟以上 → 先計數再細化,答案保持精確。
- 輸入能否重讀? 檔案可以。真正的串流不行,而這正是把問題推向 streaming_algorithms.md 的原因。
3) 常見陷阱
- 以項目數而非位元組數報記憶體。 「我會用一個裝 40 億個 int 的 set」就是 16 GB。 在提出資料結構前先做乘法。
- 在稀疏值域上用位元向量。 每個值一個位元只有在值密集時才划算;對散佈在 2^63 上的 1,000 個值,雜湊集合要小上好幾個數量級。
- 用不穩定的雜湊做切分。 Python 對
str的hash()每個行程都會加鹽,所以重新執行時 鍵的分散方式會不同。只要桶的壽命比行程長,就用hashlib(或 Java 的String.hashCode, 它有明確規範且穩定)。 - 假設各桶是平衡的。 依最差的桶來決定大小,並對溢出的桶重新切分。
- 分組就夠時卻去排序。 排序要花
O(n log n)次 I/O 掃描;依雜湊切分只要兩趟。只有在 輸出必須有序或需要範圍掃描時才排序。 - 忘了第二趟要重讀輸入。 如果來源是網路串流,就無法先計數再細化——選擇它之前要先說明這點。
4) 總結
| 技巧 | 精確? | 記憶體 | 趟數 | 使用時機 |
|---|---|---|---|---|
| 位元向量 | 是 | 每個值 1 位元 | 1 | 密集有界的值域 |
| 先計數,再細化 | 是 | 計數器 + 1 個區塊 | 2 | 值域太大,無法每個值一位元 |
| 依雜湊切分 | 是 | 最大的桶 | 2 | 鍵必須分組 |
| 外部合併排序 | 是 | 1 個區塊/每個 run 1 筆記錄 | 1 + log_k | 輸出必須有序 |
| Sketch(Bloom、count-min) | 否 | 次線性 | 1 | 單趟、可接受誤差 |