经典算法以及C++/Python临时抱佛脚指南
在AI盛行的今天,仍有一些比赛、技术面试、考试需要我们熟练地掌握各类算法,并能用一款编程语言在尽可能短的时间内手搓代码解决问题。这项技能从我们学习编程的第一天就开始培养,又在实际工作中逐渐生疏——我们常常遇到曾经已经解决的问题再度变成不会的难题的情况,对STL的熟练度也往往会成为成败的分水岭。笔者最近正好需要在短时间内准备一场比赛,为了不坑队友,决定花一点时间梳理一下常用的知识。
本文不是按照由难到易或者由易到难的顺序组织的,而是按照笔者多年来对容易遗忘程度进行排序的,排在前面的是那些如果背不出基本只能放弃的特定算法,接着是一些需要用到STL关键特性的解决方案(如果不会STL,手搓将耗费过多时间),然后是经典的各类算法知识点。希望读者也能从中获益,在临时抱佛脚的时候会感激这篇文章的存在。
线段树 - Segment Tree
首先讲一道最基础的需要用线段树解决的题目: LeetCode 307 Range Sum Query。这道题几乎是线段树的定义和解释——对一个数组保持下面两个特性:(1)支持更新其中的任何一个元素;(2)能快速计算[left, right]之间元素的和。
线段树的特性是建树复杂度O(n), 单点修改O(log n),区间求和O(log n)。如果不会线段树的实现,区间求和这部分很难临场想出好办法,最笨的办法需要O(n)的复杂度计算每次求和,导致超时错误。
假设我们有nums=[1, 3, 5, 7]这个数组,线段树的结构大概如下所示。叶子节点都只有一个元素,每个树节点存的是它子节点元素的和。
1
2
3
4
5
[0,3] = 16
/ \
[0,1] = 4 [2,3] = 12
/ \ / \
[0,0]=1 [1,1]=3 [2,2]=5 [3,3]=7
如果元素个数不是2的次方,比如5个元素,nums=[1, 3, 5, 7, 4],这棵树会长成下面这样:
1
2
3
4
5
6
7
[0,4] = 20
/ \
[0,2] = 9 [3,4] = 11
/ \ / \
[0,1] = 4 [2,2] = 5 [3,3] = 7 [4,4] = 4
/ \
[0,0] = 1 [1,1] = 3
针对每个节点下标si,访问子节点的方法是
1
2
left_child = 2 * si + 1;
right_child = 2 * si + 2;
这棵树需要多少个节点?比较宽松快捷的方式是直接使用4n(节点数不可能超过4n),更节省空间的算法是2 * pow(2, ceil(log2(n))) - 1。这个计算公式的原理是这样的:
- 一颗满二叉树的叶子节点有L个,那么它总共有2L-1个节点。
- 把一个长度为n的数组作为完全二叉树的叶子节点,要塞下,这个完全二叉树有多少叶子节点?这个L需要是2的次方,并且至少有n那么大
这棵线段树中的每个节点都表示了一个线段[start, end],每个节点存储的信息就是这条线段的长度,叶子节点就是每个数组元素的值。
容易忘的是参数个数,可以这么记:
- query需要(1)st - segment tree,一个用来存放线段树的数组,长度用前面的方法计算,简记
4n;(2) ss - segment start,当前线段开始点;(3) se - segmemt end,当前线段结束点;(4) qs - query start,查询的起始点;(5) qe - query end,查询的终止点;(6) si - segment index,当前访问的线段节点下标; - build的时候没有query,所以少两个参数qs, qe;
- update需要下标和差值,用一个额外接口函数包装一下。
建树、查询、更新的过程都是用二分查找的方式自顶向下遍历二叉树:
- 建树的退出情况是当
ss==se的时候存数组本身的元素值返回,否则就递归取两个子树的值之和,这个值取完都要更新到st[si]中。 - 查询时当
[qs,qe]包含了[ss,se],说明当前这条线段被包含在查询中,直接返回当前节点的值,如果[qs,qe]完全在[ss, se]之外,说明当前这条线段不应被考虑直接返回0,否则用二分的方式继续找子树,求和就返回两个子树的和,求最大/最小值就返回两个子树的最大/最小值。 - 更新的时候,如果下标
i不在[ss,se]之内,说明与当前线段无关,直接返回;否则就更新当前节点的值加上st[i] +=diff,然后递归更新子树。
下面是这道题的完整解法。
1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
22
23
24
25
26
27
28
29
30
31
32
33
34
35
36
37
38
39
40
41
42
43
44
45
46
47
48
49
50
51
52
53
54
55
56
57
58
59
60
61
62
63
64
65
66
/**
* Your NumArray object will be instantiated and called as such:
* NumArray* obj = new NumArray(nums);
* obj->update(index,val);
* int param_2 = obj->sumRange(left,right);
*//**
* Your NumArray object will be instantiated and called as such:
* NumArray* obj = new NumArray(nums);
* obj->update(index,val);
* int param_2 = obj->sumRange(left,right);
*/
class NumArray {
public:
int* segment_tree;
int size;
vector<int> nums;
NumArray(vector<int>& nums) {
this->size = 2*(int)pow(2, ceil(log2(nums.size())))-1;
this->segment_tree = new int[this->size];
this->nums = nums;
constructSTUtil(segment_tree, 0, nums.size()-1, 0);
}
void update(int index, int val) {
int diff = val - nums[index];
nums[index] = val;
updateValueUtil(segment_tree, 0, nums.size()-1, index, diff, 0);
}
int sumRange(int left, int right) {
return getSumUtil(segment_tree, 0, nums.size()-1, left, right, 0);
}
int getSumUtil(int* st, int ss, int se, int qs, int qe, int si) {
if(qs<=ss && qe>=se){
return st[si];
}
if(se<qs || ss > qe){
return 0;
}
int mid = (ss+se)/2;
return getSumUtil(st, ss, mid, qs, qe, 2*si+1) + getSumUtil(st, mid+1, se, qs, qe, 2*si+2);
}
void updateValueUtil(int *st, int ss, int se, int i, int diff, int si) {
if(i<ss || i>se) {
return;
}
st[si] = st[si] + diff;
if(se!=ss) {
int mid = (ss+se)/2;
updateValueUtil(st, ss, mid, i, diff, 2*si+1);
updateValueUtil(st, mid+1, se, i, diff, 2*si+2);
}
}
int constructSTUtil(int* st, int ss, int se, int si) {
if(ss==se){
st[si] = nums[ss];
return nums[ss];
}
int mid = (ss+se)/2;
st[si] = constructSTUtil(st, ss, mid, si*2+1) + constructSTUtil(st, mid+1, se, si*2+2);
return st[si];
}
};
Python实现:
1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
22
23
24
25
26
27
28
29
30
31
32
33
34
35
36
37
38
39
40
41
42
43
44
45
46
47
48
49
50
51
52
53
54
55
56
57
58
59
60
61
62
63
64
65
66
67
68
69
70
71
72
73
74
75
76
77
78
79
80
81
82
83
84
85
86
87
88
89
90
91
92
93
94
95
96
97
98
99
100
101
102
103
104
105
106
107
108
109
110
111
112
113
114
115
116
117
118
119
120
121
122
123
124
125
126
127
128
129
130
131
132
133
134
135
136
137
138
139
140
from typing import List
class NumArray:
def __init__(self, nums: List[int]):
self.nums = nums.copy()
# 找到不小于 n 的最小 2 的幂
leaf_count = 1 << (len(nums) - 1).bit_length()
# 具有 leaf_count 个叶子节点的满二叉树共有 2L - 1 个节点
tree_size = 2 * leaf_count - 1
self.segment_tree = [0] * tree_size
self._build(0, len(nums) - 1, 0)
def update(self, index: int, val: int) -> None:
diff = val - self.nums[index]
self.nums[index] = val
self._update(
segment_start=0,
segment_end=len(self.nums) - 1,
index=index,
diff=diff,
tree_index=0
)
def sumRange(self, left: int, right: int) -> int:
return self._query(
segment_start=0,
segment_end=len(self.nums) - 1,
query_start=left,
query_end=right,
tree_index=0
)
def _build(
self,
segment_start: int,
segment_end: int,
tree_index: int
) -> int:
# 叶子节点
if segment_start == segment_end:
self.segment_tree[tree_index] = self.nums[segment_start]
return self.nums[segment_start]
mid = (segment_start + segment_end) // 2
left_sum = self._build(
segment_start,
mid,
tree_index * 2 + 1
)
right_sum = self._build(
mid + 1,
segment_end,
tree_index * 2 + 2
)
self.segment_tree[tree_index] = left_sum + right_sum
return self.segment_tree[tree_index]
def _update(
self,
segment_start: int,
segment_end: int,
index: int,
diff: int,
tree_index: int
) -> None:
# index 不在当前节点表示的区间中
if index < segment_start or index > segment_end:
return
# 当前节点包含 index,因此节点和需要加上 diff
self.segment_tree[tree_index] += diff
# 如果不是叶子节点,继续递归更新子节点
if segment_start != segment_end:
mid = (segment_start + segment_end) // 2
if index <= mid:
self._update(
segment_start,
mid,
index,
diff,
tree_index * 2 + 1
)
else:
self._update(
mid + 1,
segment_end,
index,
diff,
tree_index * 2 + 2
)
def _query(
self,
segment_start: int,
segment_end: int,
query_start: int,
query_end: int,
tree_index: int
) -> int:
# 情况一:当前节点的区间完全被查询区间包含
if query_start <= segment_start and segment_end <= query_end:
return self.segment_tree[tree_index]
# 情况二:当前节点的区间与查询区间完全没有交集
if segment_end < query_start or segment_start > query_end:
return 0
# 情况三:部分重叠,分别查询左右子树
mid = (segment_start + segment_end) // 2
left_sum = self._query(
segment_start,
mid,
query_start,
query_end,
tree_index * 2 + 1
)
right_sum = self._query(
mid + 1,
segment_end,
query_start,
query_end,
tree_index * 2 + 2
)
return left_sum + right_sum
前缀树 Prefix Tree - Trie
LeetCode 208. Implement Trie。Prefix Tree主要用于高效地存储和查询大量字符串,特别适合处理与前缀有关的问题——可以判断一个完整单词是否存在,判断是否存在某个前缀,根据前缀寻找所有匹配的单词,在大量字符串中进行词典匹配。对于长度为L的字符串,插入、查询完整字符串、查询前缀、删除字符串都可以在O(L)的时间复杂度内完成。它的核心思想是让具有相同前缀的字符串共享同一段路径。
每个树节点包含一个_isEnd的标记位用于判断当前节点是否是一个词的结尾,还需要维护一个长度为所有种类字符总长度的字符集的children数组,这里因为只有英文字母,可以确定长度是26。
insert, search, startsWtih三个操作都是递归的,判断当前下标是否满足到达词末尾或到达prefix字符串末尾,如果尚未到达,从children[cur_index]继续进行下一步的操作,cur_index是字符串当前下标的字符对应到字符集的下标。退出的时候都是index==word.size()的时候,因为从一开始就是对root->children进行操作而不是root本身,只有到index到word.size()的时候才结束,不是word.size()-1的时候结束。
这个children也可以用哈希表实现,这样没有字符集的限制,递归的过程也可以不用函数调用,而用循环的方式更快完成,我们在Python实现中采用这种方式。
完整的C++实现如下所示:
1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
22
23
24
25
26
27
28
29
30
31
32
33
34
35
36
37
38
39
40
41
42
43
44
45
46
47
48
49
50
51
52
53
54
55
56
57
58
59
60
61
62
63
64
65
66
67
68
69
70
71
72
73
74
75
76
77
78
79
80
81
82
83
84
85
86
87
88
class TrieNode {
std::vector<TrieNode*> children;
bool _isEnd;
public:
TrieNode() {
this->children = std::vector<TrieNode*>(26, nullptr);
this->_isEnd = false;
}
~TrieNode() {
for(int i=0;i<children.size();i++) {
if(children[i] != nullptr) {
delete children[i];
}
}
}
void insert(const string& word, int index) {
if(index == word.size()){
this->_isEnd = true;
return;
}
int cur_index = word[index] - 'a';
if(children[cur_index] == nullptr) {
children[cur_index] = new TrieNode();
}
children[cur_index]->insert(word, index+1);
}
bool search(const string& word, int index) {
if(index == word.size()) {
if(_isEnd){
return true;
}else{
return false;
}
}
int cur_index = word[index] - 'a';
if(children[cur_index] == nullptr){return false;}
else{
return children[cur_index]->search(word, index+1);
}
}
bool startsWith(string prefix, int index) {
if(index == prefix.size()) {
return true;
}
int cur_index = prefix[index] - 'a';
if(children[cur_index] == nullptr){return false;}
else{
return children[cur_index]->startsWith(prefix, index+1);
}
}
};
class Trie {
TrieNode* root;
public:
Trie() {
root = new TrieNode();
}
~Trie() {
delete root;
}
void insert(string word) {
root->insert(word, 0);
}
bool search(string word) {
return root->search(word, 0);
}
bool startsWith(string prefix) {
return root->startsWith(prefix, 0);
}
};
/**
* Your Trie object will be instantiated and called as such:
* Trie* obj = new Trie();
* obj->insert(word);
* bool param_2 = obj->search(word);
* bool param_3 = obj->startsWith(prefix);
*/
Python实现如下:
1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
22
23
24
25
26
27
28
29
30
31
32
33
34
35
36
37
38
39
40
41
42
43
44
45
46
47
48
49
50
class TrieNode:
# Initialize your data structure here.
def __init__(self):
self.word=False
self.children={}
class Trie:
def __init__(self):
self.root = TrieNode()
# @param {string} word
# @return {void}
# Inserts a word into the trie.
def insert(self, word):
node=self.root
for i in word:
if i not in node.children:
node.children[i]=TrieNode()
node=node.children[i]
node.word=True
# @param {string} word
# @return {boolean}
# Returns if the word is in the trie.
def search(self, word):
node=self.root
for i in word:
if i not in node.children:
return False
node=node.children[i]
return node.word
# @param {string} prefix
# @return {boolean}
# Returns if there is any word in the trie
# that starts with the given prefix.
def startsWith(self, prefix):
node=self.root
for i in prefix:
if i not in node.children:
return False
node=node.children[i]
return True
# Your Trie object will be instantiated and called as such:
# trie = Trie()
# trie.insert("somestring")
# trie.search("key")
接下来是一个使用Trie的案例LeetCode 211. Design Add and Search Words Data Structure。
我们在C++中用哈希表实现一下,增强对前缀树的理解。由于有通配符的存在,search难以用循环实现,对通配符需要对所有children进行枚举判断是否存在一条路径满足条件。
1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
22
23
24
25
26
27
28
29
30
31
32
33
34
35
36
37
38
39
40
41
42
43
44
45
46
47
48
49
50
51
52
53
54
55
56
57
58
59
60
61
62
63
64
65
66
67
struct Node {
std::unordered_map<char, Node*> children;
bool is_end;
Node(){
this->is_end = false;
}
};
class WordDictionary {
private:
Node* root;
public:
WordDictionary() {
root = new Node();
}
void addWord(string word) {
Node* p = root;
int index;
for(int i=0;i<word.size();i++) {
char cur = word[i];
if(p->children.find(cur) == p->children.end()) {
p->children[cur] = new Node();
}
p = p->children[cur];
if(i == word.size() - 1) {
p->is_end = true;
}
}
}
bool search(string word) {
return searchImpl(word, 0, root);
}
bool searchImpl(const string& word, int index, Node* p) {
if(index == word.size()) {
if(p->is_end) {
return true;
} else {
return false;
}
}
char cur = word[index];
if(cur == '.') {
bool found = false;
for(auto it = p->children.begin(); it!=p->children.end(); it++) {
found = found || searchImpl(word, index+1, it->second);
}
return found;
} else {
auto iter = p->children.find(cur);
if(iter == p->children.end()) {
return false;
} else {
return searchImpl(word, index+1, iter->second);
}
}
}
};
/**
* Your WordDictionary object will be instantiated and called as such:
* WordDictionary* obj = new WordDictionary();
* obj->addWord(word);
* bool param_2 = obj->search(word);
*/
最短路径算法和优先队列
最短路径算法其实应该属于普通的经典算法,但因为里面同时需用到优先队列,如果不及时复习很容易写不出来,我把它单独放在前面讲解。
Dijkstra 算法
LeetCode 743. Network Delay Time是一道经典的边权非负的最短路径问题,求的是从一个给定点出发发射信号,到所有节点都能收到所需的最短时间。这类问题可以用Dijkstra算法解决。
Dijkstra需要优先队列的原因是每次需要从当前尚未处理的节点中,找到起点距离最小的节点,如果每次遍历所有节点取最小值,复杂度为 $O(V^2)$,使用优先队列后,复杂度为 $O((V+E) \log V)$,对于一个连通图,边的数量大于等于V-1,所以这个复杂度可以简记为$O(E \log V)$.
Dijkstra的核心性质是:当前距离最小的节点出堆后,其最短路径已经确定,不会再被后续路径缩短,这一点必须要求边权非负。
Dijkstra找到到某一个点的最短路径和找到到所有点的最短路径所需的流程是一样的。
该算法的通用实现方式如下:
1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
22
23
24
25
26
27
28
29
30
31
32
33
34
35
36
37
38
39
40
41
42
43
44
45
using State = pair<long long, int>;
// {从起点到当前节点的距离, 当前节点}
std::vector<long long> dijkstra(
int start,
const std::vector<std::vector<std::pair<int, int>>>& graph
) {
int n = graph.size();
const long long INF = std::numeric_limits<long long>::max();
std::vector<long long> distance(n, INF);
distance[start] = 0;
std::priority_queue<
State,
std::vector<State>,
std::greater<State>
> min_heap;
min_heap.push({0, start});
while (!min_heap.empty()) {
auto [current_distance, node] = min_heap.top();
min_heap.pop();
// 堆中可能存在同一个节点的旧距离
if (current_distance > distance[node]) {
continue;
}
for (auto [next_node, weight] : graph[node]) {
long long new_distance =
current_distance + weight;
if (new_distance < distance[next_node]) {
distance[next_node] = new_distance;
min_heap.push({
new_distance,
next_node
});
}
}
}
return distance;
}
下面是将这个算法用于解决LeetCode 743的完整实现:
1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
22
23
24
25
26
27
28
29
30
31
32
33
34
35
36
37
38
39
40
41
42
43
44
45
46
47
48
49
50
51
52
53
54
class Solution {
public:
const int INF = std::numeric_limits<int>::max();
std::vector<int> dijkstra(
int start,
const std::vector<std::vector<std::pair<int, int>>>& graph
) {
int n = graph.size();
std::vector<int> distance(n, INF);
distance[start] = 0;
std::priority_queue<
std::pair<int, int>,
std::vector<std::pair<int, int>>,
std::greater<std::pair<int,int>>
> min_heap;
min_heap.push({0, start});
while(!min_heap.empty()) {
auto [current_distance, node] = min_heap.top();
min_heap.pop();
if(current_distance > distance[node]) {
continue;
}
for (auto [next_node, weight] : graph[node]) {
int new_distance = current_distance + weight;
if (new_distance < distance[next_node]) {
distance[next_node] = new_distance;
min_heap.push({new_distance, next_node});
}
}
}
return distance;
}
int networkDelayTime(vector<vector<int>>& times, int n, int k) {
std::vector<std::vector<std::pair<int, int>>> graph(n);
for(const auto& edge : times) {
graph[edge[0] - 1].push_back({edge[1] - 1, edge[2]});
}
auto distances = dijkstra(k - 1, graph);
int maximum_distance = 0;
for (auto distance : distances) {
if (distance == INF) {return -1;}
maximum_distance = std::max(maximum_distance, distance);
}
return maximum_distance;
}
};
如果不仅需要记录最短路径长度,还需要具体的这条路径是什么,需要往里加一个parent来记录。
1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
22
23
24
25
26
27
28
29
30
31
32
33
34
35
36
37
38
39
40
41
42
43
44
45
46
47
48
49
50
51
52
53
54
55
56
57
58
59
60
61
62
63
64
65
66
67
68
69
70
71
72
73
74
75
76
77
78
79
class Solution {
public:
using Edge = pair<int, int>;
// {next_node, weight}
using State = pair<int, int>;
// {distance, node}
pair<int, vector<int>> shortestPath(
int start,
int target,
const vector<vector<Edge>>& graph
) {
int n = static_cast<int>(graph.size());
const int INF = numeric_limits<int>::max();
vector<int> distance(n, INF);
vector<int> parent(n, -1);
priority_queue<
State,
vector<State>,
greater<State>
> min_heap;
distance[start] = 0;
min_heap.push({0, start});
while (!min_heap.empty()) {
auto [current_distance, node] = min_heap.top();
min_heap.pop();
if (current_distance > distance[node]) {
continue;
}
// target 第一次以有效状态出堆,
// 它的最短距离已经确定。
if (node == target) {
break;
}
for (const auto& [next_node, weight] : graph[node]) {
int new_distance =
current_distance + weight;
if (new_distance < distance[next_node]) {
distance[next_node] = new_distance;
// 记录 next_node 是从 node 到达的
parent[next_node] = node;
min_heap.push({
new_distance,
next_node
});
}
}
}
if (distance[target] == INF) {
return {-1, {}};
}
vector<int> path;
// 从终点沿 parent 倒推到起点
for (int node = target;
node != -1;
node = parent[node]) {
path.push_back(node);
}
// 当前是 target -> ... -> start,需要反转
reverse(path.begin(), path.end());
return {distance[target], path};
}
};
如果涉及到多条路径还需要求解路径的时候用DFS找出所有路径。
1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
22
23
24
25
26
27
28
29
30
31
32
33
34
35
36
37
38
39
40
41
42
43
44
45
46
47
48
49
// parents的定义变更如下
vector<vector<int>> parents(n);
// 更新的部分需要在距离相等的时候,将当前节点添加为父节点
if (new_distance < distance[next_node]) {
distance[next_node] = new_distance;
parents[next_node].clear();
parents[next_node].push_back(node);
min_heap.push({
new_distance,
next_node
});
} else if (new_distance == distance[next_node]) {
parents[next_node].push_back(node);
}
// 最后计算路径的时候需要用dfs向前回溯
void buildPaths(
int node,
int start,
const vector<vector<int>>& parents,
vector<int>& current_path,
vector<vector<int>>& result
) {
current_path.push_back(node);
if (node == start) {
vector<int> path(
current_path.rbegin(),
current_path.rend()
);
result.push_back(path);
} else {
for (int previous : parents[node]) {
buildPaths(
previous,
start,
parents,
current_path,
result
);
}
}
current_path.pop_back();
}
一般Dijkstra无法解决的最短路径问题
787. Cheapest Flights Within K Stops,除了要求边权和小,还有轮数限制,用一般的Dijkstra算法就无法解决,因为可能存在这样的情况:
1
2
3
4
到达 A:
路径一:价格 100,用 3 条边
路径二:价格 150,只用了 1 条边
用Dijkstra会选到价格100的路径,但它的边条数超过了限制,不满足要求。这种有额外限制的最短路径问题一般需要使用Bellman-Ford动态规划。
1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
22
23
24
25
26
27
28
29
30
31
32
33
34
35
36
37
38
39
40
41
42
class Solution {
public:
int findCheapestPrice(
int n,
vector<vector<int>>& flights,
int src,
int dst,
int k
) {
const int INF = numeric_limits<int>::max();
vector<int> distance(n, INF);
distance[src] = 0;
// 最多 K 个中转站,即最多使用 K + 1 条边
for (int edges = 0; edges <= k; ++edges) {
// 必须复制上一轮结果
vector<int> next_distance = distance;
for (const auto& flight : flights) {
int from = flight[0];
int to = flight[1];
int price = flight[2];
if (distance[from] == INF) {
continue;
}
next_distance[to] = min(
next_distance[to],
distance[from] + price
);
}
distance = std::move(next_distance);
}
return distance[dst] == INF
? -1
: distance[dst];
}
};
如果一定要用Dijkstra算法,需要把边作为一个状态存入优先队列,上题的另一种解法如下:
1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
22
23
24
25
26
27
28
29
30
31
32
33
34
35
36
37
38
39
40
41
42
43
44
45
46
47
48
49
50
51
52
53
54
55
56
57
58
59
60
61
62
63
64
65
66
67
68
69
70
71
72
73
74
75
76
77
78
79
80
81
82
83
class Solution {
public:
struct State {
int cost;
int node;
int edges;
bool operator>(const State& other) const {
return cost > other.cost;
}
};
int findCheapestPrice(
int n,
vector<vector<int>>& flights,
int src,
int dst,
int k
) {
vector<vector<pair<int, int>>> graph(n);
for (const auto& flight : flights) {
int from = flight[0];
int to = flight[1];
int price = flight[2];
graph[from].push_back({to, price});
}
int max_edges = k + 1;
const int INF = numeric_limits<int>::max();
// distance[node][edges]:
// 恰好使用 edges 条边到达 node 的最低价格
vector<vector<int>> distance(
n,
vector<int>(max_edges + 1, INF)
);
priority_queue<
State,
vector<State>,
greater<State>
> min_heap;
distance[src][0] = 0;
min_heap.push({0, src, 0});
while (!min_heap.empty()) {
auto [cost, node, edges] = min_heap.top();
min_heap.pop();
if (cost > distance[node][edges]) {
continue;
}
if (node == dst) {
return cost;
}
if (edges == max_edges) {
continue;
}
for (const auto& [next_node, price] : graph[node]) {
int new_cost = cost + price;
int new_edges = edges + 1;
if (new_cost < distance[next_node][new_edges]) {
distance[next_node][new_edges] = new_cost;
min_heap.push({
new_cost,
next_node,
new_edges
});
}
}
}
return -1;
}
};
滑动窗口 - Sliding Window
滑动窗口是双指针技巧的延伸,专门处理”数组/字符串的连续子区间”问题。核心是用左右两个指针维护一个窗口 [l, r],根据约束条件动态调整窗口大小,把 O(n²) 的枚举优化到 O(n)。如果临场想不起来,每次只能 O(n) 重算窗口状态,很容易超时。
滑动窗口分为两类:
- 定长窗口:窗口大小固定为 K(如”长度为 K 的子数组最大和”、”字符串的排列”)。
- 变长窗口:窗口大小由约束条件动态决定,找到满足约束的”最长”或”最短”子区间。
变长窗口的两套模板
把滑动窗口想象成两只手在数组上滑动,每一步都在做三件事:
- 右手扩张:把
s[r]加入窗口(窗口变成[l, r])。 - 左手收缩:在需要的时候,把
s[l]从窗口移除,l右移(窗口收缩)。 - 更新答案:当窗口恰好满足约束时,记下当前的长度。
窗口里还需要维护一个计数表 window[c]:进入窗口时 +1,离开窗口时 -1,这两个动作必须严格对称,缺一就会出错。
因为 l 和 r 都只往右走,最坏各走 n 步,所以总复杂度一定是 O(n)——这就是滑动窗口不会超时的根本原因。
找最长:窗口不满足约束就一直收缩
约束示例:”窗口内字符不能重复”。只要 window[c] > 1 就说明重复了,需要收缩。
1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
int l = 0; // 窗口左端点(闭区间,包含 s[l])
int ans = 0; // 记录目前为止的最长长度
std::unordered_map<char, int> window; // window[c] = 字符 c 当前在窗口内出现的次数
for (int r = 0; r < n; ++r) {
// ① 右手扩张:把 s[r] 加入窗口
char c = s[r];
window[c]++;
// ② 左手收缩:只要窗口"不满足约束",就一直把 s[l] 移出窗口
while (window[c] > 1) { // 以"窗口内不能有重复字符"为例
char out = s[l];
window[out]--; // s[l] 离开窗口,计数 -1
l++; // 窗口左端点右移
}
// ③ 此时窗口一定满足约束,记下当前窗口长度
ans = std::max(ans, r - l + 1);
}
1
2
3
4
5
6
7
8
9
# TODO: 自己用 Python 实现"找最长"模板
# 提示:
# l = 0, ans = 0, window = {}
# for r, c in enumerate(s):
# window[c] = window.get(c, 0) + 1
# while window[c] > 1: # 以"不能有重复字符"为例
# window[s[l]] -= 1
# l += 1
# ans = max(ans, r - l + 1)
找最短:窗口满足约束就一直收缩
约束示例:”窗口内必须包含 t 的所有字符”。这种约束没法用一个计数搞定,需要一个 need 表,再用一个 formed 记录”已经凑齐数量的字符种类数”。
1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
22
23
24
25
26
27
std::unordered_map<char, int> need, window;
int required = (int)need.size(); // t 中一共有多少种字符需要凑齐
int formed = 0; // 当前窗口里,已经"恰好凑够 need 数量"的字符种类数
int l = 0;
int ans = INT_MAX;
for (int r = 0; r < n; ++r) {
// ① 右手扩张:把 s[r] 加入窗口
char c = s[r];
window[c]++;
// 如果这个字符是 need 里的,并且加入后"刚好达到" need 的要求,
// 就把它计入"已满足"
if (window[c] == need[c]) ++formed;
// ② 左手收缩:只要窗口"已经满足约束",就一直往左压,尝试找更短的
while (formed == required && l <= r) {
ans = std::min(ans, r - l + 1);
char out = s[l];
window[out]--;
// 移走 out 之后,如果它的数量变得不够了,就把 formed 减 1
if (window[out] < need[out]) --formed;
l++;
}
}
1
2
3
4
5
6
7
8
9
10
11
12
13
# TODO: 自己用 Python 实现"找最短"模板
# 提示:
# need = Counter(t); required = len(need); formed = 0
# window = {}; l = 0; ans = float('inf')
# for r, c in enumerate(s):
# window[c] = window.get(c, 0) + 1
# if window[c] == need[c]: formed += 1
# while formed == required:
# ans = min(ans, r - l + 1)
# out = s[l]
# window[out] -= 1
# if window[out] < need[out]: formed -= 1
# l += 1
注意 formed == required 才表示窗口里所有需要的字符种类都凑齐了(数量也刚好够)。一旦凑齐,就不断尝试从左边压,看最短能压到多长。
用 s = "abcabcbb" 走一遍”找最长”
约束:窗口内字符不能重复。
| r | s[r] | 操作 | window | l | 窗口 | ans |
|---|---|---|---|---|---|---|
| 0 | a | 加入 | {a:1} | 0 | “a” | 1 |
| 1 | b | 加入 | {a:1, b:1} | 0 | “ab” | 2 |
| 2 | c | 加入 | {a:1, b:1, c:1} | 0 | “abc” | 3 |
| 3 | a | 加入 → 窗口内 a 重复,收缩 1 次 | {a:1, b:1, c:1} | 1 | “bca” | 3 |
| 4 | b | 加入 → 窗口内 b 重复,收缩 1 次 | {a:1, b:1, c:1} | 2 | “cab” | 3 |
| 5 | c | 加入 → 窗口内 c 重复,收缩 1 次 | {a:1, b:1, c:1} | 3 | “abc” | 3 |
| 6 | b | 加入 → 窗口内 b 重复,连续收缩到 l=5 | {a:0, b:1, c:1} | 5 | “cb” | 3 |
| 7 | b | 加入 → 窗口内 b 重复,连续收缩到 l=7 | {b:1} | 7 | “b” | 3 |
最终 ans = 3,对应的最长无重复子串是 abc / bca / cab。
记忆要点:
- 定长窗口先初始化前 K 个元素,再循环
n - K次,每次”出左入右”。 - 变长窗口的
while条件是核心——”找最长”用”不满足就收缩”,”找最短”用”满足就收缩”。 - 窗口状态的加入和移除要严格对称,缺一就会出错。
- 字符类问题优先用
int cnt[128]数组,比unordered_map快得多。 l和r都只往右走,总复杂度一定是O(n)——这是滑动窗口不会超时的根本原因。
例题 1:最长无重复子串
LC 3. Longest Substring Without Repeating Characters。变长窗口的入门题:找到不含重复字符的最长子串。
1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
class Solution {
public:
int lengthOfLongestSubstring(string s) {
std::vector<int> last(128, -1);
int ans = 0;
int l = 0;
for (int r = 0; r < (int)s.size(); ++r) {
unsigned char c = s[r];
if (last[c] >= l) {
l = last[c] + 1;
}
last[c] = r;
ans = std::max(ans, r - l + 1);
}
return ans;
}
};
1
2
3
4
5
6
7
8
9
10
11
class Solution:
def lengthOfLongestSubstring(self, s: str) -> int:
last = {}
ans = 0
l = 0
for r, c in enumerate(s):
if c in last and last[c] >= l:
l = last[c] + 1
last[c] = r
ans = max(ans, r - l + 1)
return ans
这题有一个常数优化的写法:不用哈希表,而是用 int last[128] 数组记录字符 c 上次出现的下标。如果上次出现的下标在窗口内,就把 l 跳到它的下一个位置。这种写法比纯哈希表更快。
例题 2:最小覆盖子串
LC 76. Minimum Window Substring。变长窗口找最短的经典题:找到包含 t 所有字符的 s 的最短子串。
1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
22
23
24
25
26
27
28
29
30
31
32
class Solution {
public:
string minWindow(string s, string t) {
std::vector<int> need(128, 0);
for (char c : t) need[(unsigned char)c]++;
int required = 0;
for (int c : need) if (c > 0) ++required;
std::vector<int> window(128, 0);
int formed = 0;
int l = 0;
int min_len = INT_MAX;
int min_l = 0;
for (int r = 0; r < (int)s.size(); ++r) {
unsigned char c = s[r];
if (++window[c] == need[c]) ++formed;
while (formed == required && l <= r) {
if (r - l + 1 < min_len) {
min_len = r - l + 1;
min_l = l;
}
unsigned char cl = s[l];
if (--window[cl] < need[cl]) --formed;
++l;
}
}
return min_len == INT_MAX ? "" : s.substr(min_l, min_len);
}
};
1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
22
23
24
25
26
27
28
from collections import Counter
class Solution:
def minWindow(self, s: str, t: str) -> str:
need = Counter(t)
required = len(need)
formed = 0
window = {}
l = 0
min_len = float('inf')
min_l = 0
for r, c in enumerate(s):
window[c] = window.get(c, 0) + 1
if c in need and window[c] == need[c]:
formed += 1
while formed == required and l <= r:
if r - l + 1 < min_len:
min_len = r - l + 1
min_l = l
cl = s[l]
window[cl] -= 1
if cl in need and window[cl] < need[cl]:
formed -= 1
l += 1
return "" if min_len == float('inf') else s[min_l:min_l + min_len]
记忆要点:
formed表示窗口内已经”满足 need 中字符数量要求”的字符种类数。- 当
formed == required时窗口内包含了所有需要的字符,可以尝试收缩。 - 收缩过程中可能破坏约束,每次
window[c]--后都要判断window[c] < need[c]。
例题 3:字符串的排列
LC 567. Permutation in String。定长窗口:判断 s2 是否包含 s1 的某个排列(即 s2 中是否存在长度为 |s1| 的窗口,字符计数和 s1 完全一致)。
1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
class Solution {
public:
bool checkInclusion(string s1, string s2) {
int n = s1.size();
if (n > (int)s2.size()) return false;
std::vector<int> cnt1(26, 0), cnt2(26, 0);
for (int i = 0; i < n; ++i) {
cnt1[s1[i] - 'a']++;
cnt2[s2[i] - 'a']++;
}
if (cnt1 == cnt2) return true;
for (int r = n; r < (int)s2.size(); ++r) {
cnt2[s2[r] - 'a']++;
cnt2[s2[r - n] - 'a']--;
if (cnt1 == cnt2) return true;
}
return false;
}
};
1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
class Solution:
def checkInclusion(self, s1: str, s2: str) -> bool:
n = len(s1)
if n > len(s2):
return False
cnt1 = [0] * 26
cnt2 = [0] * 26
for i in range(n):
cnt1[ord(s1[i]) - ord('a')] += 1
cnt2[ord(s2[i]) - ord('a')] += 1
if cnt1 == cnt2:
return True
for r in range(n, len(s2)):
cnt2[ord(s2[r]) - ord('a')] += 1
cnt2[ord(s2[r - n]) - ord('a')] -= 1
if cnt1 == cnt2:
return True
return False
记忆要点:定长窗口的循环只走 n - 1 次(从 r = n 到末尾),每次循环把 r 处字符入窗,把 r - n 处字符出窗。
BFS - 广度优先搜索
BFS 适合”层序遍历”和”无权图最短路”两类问题,核心是用队列维护”待访问节点”,按”由近及远”的顺序访问。和 DFS 相比,BFS 的特点是可以方便地按层处理(每层节点共享同一个距离),代价是要维护一个队列和 visited 集合。
如果临场写不出 BFS 模板,主要是因为漏了 visited 标记,导致队列无限增长。
通用模板
1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
std::queue<State> q;
std::unordered_set<State> visited;
q.push(start);
visited.insert(start);
while (!q.empty()) {
State cur = q.front();
q.pop();
// 处理 cur
for (State next : neighbors(cur)) {
if (!visited.count(next)) {
visited.insert(next);
q.push(next);
}
}
}
层序 BFS(一次处理一层):
1
2
3
4
5
6
7
8
9
while (!q.empty()) {
int size = q.size(); // 一次性取出一整层
for (int i = 0; i < size; ++i) {
State cur = q.front();
q.pop();
// 处理 cur,把邻居 push 进队列
}
// 一层结束,更新层数
}
多源 BFS(从多个起点同时出发):把所有起点一次性 push 进队列即可。
记忆要点:
visited不能漏! 没有 visited 的话,环或者重复边会让队列无限增长。- 层序 BFS 通过
int size = q.size()一次性取出一整层,这是层序遍历和最短路的关键。 - 多源 BFS 的初始化就是把所有起点都 push 到队列里,跑法和单源完全一样。
- BFS 求无权图最短路时,节点”第一次出队”的距离就是最短距离。
例题 1:二叉树的层序遍历
LC 102. Binary Tree Level Order Traversal。BFS 的基础应用:按层输出二叉树节点。
1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
22
23
24
class Solution {
public:
vector<vector<int>> levelOrder(TreeNode* root) {
vector<vector<int>> ans;
if (!root) return ans;
queue<TreeNode*> q;
q.push(root);
while (!q.empty()) {
int size = q.size();
vector<int> level;
for (int i = 0; i < size; ++i) {
TreeNode* cur = q.front();
q.pop();
level.push_back(cur->val);
if (cur->left) q.push(cur->left);
if (cur->right) q.push(cur->right);
}
ans.push_back(level);
}
return ans;
}
};
1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
22
23
24
25
26
from collections import deque
from typing import Optional, List
class TreeNode:
def __init__(self, val=0, left=None, right=None):
self.val = val
self.left = left
self.right = right
class Solution:
def levelOrder(self, root: Optional[TreeNode]) -> List[List[int]]:
ans = []
if not root:
return ans
q = deque([root])
while q:
level = []
for _ in range(len(q)):
cur = q.popleft()
level.append(cur.val)
if cur.left:
q.append(cur.left)
if cur.right:
q.append(cur.right)
ans.append(level)
return ans
记忆要点:层序遍历模板就是 size = q.size()(C++)/len(q)(Python)取整层大小,循环结束就代表一层走完。
例题 2:岛屿数量
LC 200. Number of Islands。网格 BFS 的经典题:统计 ‘1’ 相连的连通分量数。
1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
22
23
24
25
26
27
28
29
30
31
32
33
34
35
36
class Solution {
public:
int numIslands(vector<vector<char>>& grid) {
int m = grid.size(), n = grid[0].size();
int ans = 0;
std::queue<std::pair<int, int>> q;
std::vector<std::pair<int, int>> dirs;
dirs.push_back({-1, 0});
dirs.push_back({1, 0});
dirs.push_back({0, -1});
dirs.push_back({0, 1});
for (int i = 0; i < m; ++i) {
for (int j = 0; j < n; ++j) {
if (grid[i][j] != '1') continue;
++ans;
q.push({i, j});
grid[i][j] = '0'; // 直接把访问过的格子置 0 当 visited
while (!q.empty()) {
auto [x, y] = q.front();
q.pop();
for (auto [dx, dy] : dirs) {
int nx = x + dx, ny = y + dy;
if (nx >= 0 && nx < m && ny >= 0 && ny < n
&& grid[nx][ny] == '1') {
grid[nx][ny] = '0';
q.push({nx, ny});
}
}
}
}
}
return ans;
}
};
1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
22
23
24
25
26
from collections import deque
from typing import List
class Solution:
def numIslands(self, grid: List[List[str]]) -> int:
m, n = len(grid), len(grid[0])
ans = 0
q = deque()
dirs = [(-1, 0), (1, 0), (0, -1), (0, 1)]
for i in range(m):
for j in range(n):
if grid[i][j] != '1':
continue
ans += 1
q.append((i, j))
grid[i][j] = '0'
while q:
x, y = q.popleft()
for dx, dy in dirs:
nx, ny = x + dx, y + dy
if 0 <= nx < m and 0 <= ny < n and grid[nx][ny] == '1':
grid[nx][ny] = '0'
q.append((nx, ny))
return ans
记忆要点:在网格上做 BFS 时,直接把访问过的格子改成 ‘0’ 就能省掉一个 visited 数组,这种”原地修改做 visited”的技巧在网格题里非常常用。
例题 3:腐烂的橘子
LC 994. Rotting Oranges。多源 BFS 的入门题:每分钟腐烂橘子会感染上下左右相邻的新鲜橘子,问几分钟所有橘子都腐烂(或返回 -1 表示不可能)。
1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
22
23
24
25
26
27
28
29
30
31
32
33
34
35
36
37
38
39
40
41
42
43
44
class Solution {
public:
int orangesRotting(vector<vector<int>>& grid) {
int m = grid.size(), n = grid[0].size();
std::queue<std::pair<int, int>> q;
int fresh = 0;
for (int i = 0; i < m; ++i) {
for (int j = 0; j < n; ++j) {
if (grid[i][j] == 2) q.push({i, j});
else if (grid[i][j] == 1) ++fresh;
}
}
if (fresh == 0) return 0;
int minutes = 0;
std::vector<std::pair<int, int>> dirs;
dirs.push_back({-1, 0});
dirs.push_back({1, 0});
dirs.push_back({0, -1});
dirs.push_back({0, 1});
while (!q.empty()) {
int size = q.size();
for (int i = 0; i < size; ++i) {
auto [x, y] = q.front();
q.pop();
for (auto [dx, dy] : dirs) {
int nx = x + dx, ny = y + dy;
if (nx >= 0 && nx < m && ny >= 0 && ny < n
&& grid[nx][ny] == 1) {
grid[nx][ny] = 2;
q.push({nx, ny});
--fresh;
}
}
}
++minutes;
}
return fresh == 0 ? minutes - 1 : -1;
}
};
1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
22
23
24
25
26
27
28
29
30
31
32
33
34
from collections import deque
from typing import List
class Solution:
def orangesRotting(self, grid: List[List[int]]) -> int:
m, n = len(grid), len(grid[0])
q = deque()
fresh = 0
for i in range(m):
for j in range(n):
if grid[i][j] == 2:
q.append((i, j))
elif grid[i][j] == 1:
fresh += 1
if fresh == 0:
return 0
minutes = 0
dirs = [(-1, 0), (1, 0), (0, -1), (0, 1)]
while q:
for _ in range(len(q)):
x, y = q.popleft()
for dx, dy in dirs:
nx, ny = x + dx, y + dy
if 0 <= nx < m and 0 <= ny < n and grid[nx][ny] == 1:
grid[nx][ny] = 2
q.append((nx, ny))
fresh -= 1
minutes += 1
return minutes - 1 if fresh == 0 else -1
记忆要点:
- 多源 BFS 的初始化:把所有腐烂的橘子都 push 进队列。
- 每一轮 BFS 表示”经过了一分钟”,所以
minutes在每轮结束后加 1。 - 最后需要返回
minutes - 1,因为最后一轮”没有新的橘子被感染”也计入了一次循环(哨兵轮),如果 fresh 不为 0 则说明存在无法腐烂的橘子,返回 -1。
DFS - 深度优先搜索
DFS 和 BFS 是图遍历的两大基础,但风格迥异——DFS 倾向于”一条路走到底再回溯”,常用于树/网格的递归处理。和 BFS 相比,DFS 写起来更简洁(天然递归),缺点是不容易控制层数。
临场写不出 DFS 通常是因为递归三要素没记牢:终止条件、状态修改、状态撤销。尤其是带回溯的 DFS,”撤销”一步漏了就会得到错误答案。
递归 DFS 模板
1
2
3
4
5
6
7
8
9
10
11
12
void dfs(State cur, ...) {
if (终止条件) {
// 处理答案或返回结果
return;
}
for (State next : neighbors(cur)) {
if (跳过条件) continue;
// 修改状态(如果是回溯)
dfs(next, ...);
// 撤销修改(如果是回溯)
}
}
迭代 DFS 模板(用显式栈)
1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
std::stack<State> stk;
stk.push(start);
visited.insert(start);
while (!stk.empty()) {
State cur = stk.top();
stk.pop();
// 处理 cur
for (State next : neighbors(cur)) {
if (!visited.count(next)) {
visited.insert(next);
stk.push(next);
}
}
}
记忆要点:
- 递归 DFS 一定要有终止条件(base case),否则会栈溢出。
- “修改状态 + 递归 + 撤销修改” 三个动作必须严格配对——这就是回溯的核心。
- 网格 DFS 通常要把访问过的格子做标记(原地改成 ‘0’ 或 ‘#’),避免重复访问。
- 迭代 DFS 的访问顺序和递归 DFS 略有不同(栈是 LIFO),但都能遍历所有节点。
例题 1:岛屿最大面积
LC 695. Max Area of Island。网格 DFS 的基础题:求 ‘1’ 连通区域的最大面积。
1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
22
23
24
class Solution {
public:
int maxAreaOfIsland(vector<vector<int>>& grid) {
int m = grid.size(), n = grid[0].size();
int ans = 0;
for (int i = 0; i < m; ++i) {
for (int j = 0; j < n; ++j) {
if (grid[i][j] == 1) {
ans = std::max(ans, dfs(grid, i, j));
}
}
}
return ans;
}
int dfs(vector<vector<int>>& grid, int i, int j) {
int m = grid.size(), n = grid[0].size();
if (i < 0 || i >= m || j < 0 || j >= n || grid[i][j] != 1) return 0;
grid[i][j] = 0; // 标记访问,避免回环
return 1 + dfs(grid, i + 1, j) + dfs(grid, i - 1, j)
+ dfs(grid, i, j + 1) + dfs(grid, i, j - 1);
}
};
1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
from typing import List
class Solution:
def maxAreaOfIsland(self, grid: List[List[int]]) -> int:
m, n = len(grid), len(grid[0])
def dfs(i: int, j: int) -> int:
if i < 0 or i >= m or j < 0 or j >= n or grid[i][j] != 1:
return 0
grid[i][j] = 0
return 1 + dfs(i + 1, j) + dfs(i - 1, j) + dfs(i, j + 1) + dfs(i, j - 1)
ans = 0
for i in range(m):
for j in range(n):
if grid[i][j] == 1:
ans = max(ans, dfs(i, j))
return ans
记忆要点:进入一个格子后立即标记为 0,这样既能从相邻格子 DFS 进来时立刻退出(边界检查不通过),又能避免对同一格子重复计算。
例题 2:单词搜索
LC 79. Word Search。网格 DFS + 回溯的经典题:判断 word 是否能从 board 某格出发,沿着上下左右连续匹配。
1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
22
23
24
25
26
27
class Solution {
public:
bool exist(vector<vector<char>>& board, string word) {
int m = board.size(), n = board[0].size();
for (int i = 0; i < m; ++i) {
for (int j = 0; j < n; ++j) {
if (dfs(board, word, 0, i, j)) return true;
}
}
return false;
}
bool dfs(vector<vector<char>>& board, const string& word, int k, int i, int j) {
int m = board.size(), n = board[0].size();
if (i < 0 || i >= m || j < 0 || j >= n || board[i][j] != word[k]) return false;
if (k == (int)word.size() - 1) return true; // 匹配到最后一个字符
char tmp = board[i][j];
board[i][j] = '#'; // 标记"访问中"
bool found = dfs(board, word, k + 1, i + 1, j)
|| dfs(board, word, k + 1, i - 1, j)
|| dfs(board, word, k + 1, i, j + 1)
|| dfs(board, word, k + 1, i, j - 1);
board[i][j] = tmp; // 回溯:撤销标记
return found;
}
};
1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
22
23
from typing import List
class Solution:
def exist(self, board: List[List[str]], word: str) -> bool:
m, n = len(board), len(board[0])
def dfs(k: int, i: int, j: int) -> bool:
if i < 0 or i >= m or j < 0 or j >= n or board[i][j] != word[k]:
return False
if k == len(word) - 1:
return True
tmp = board[i][j]
board[i][j] = '#'
found = (dfs(k + 1, i + 1, j) or dfs(k + 1, i - 1, j) or
dfs(k + 1, i, j + 1) or dfs(k + 1, i, j - 1))
board[i][j] = tmp
return found
for i in range(m):
for j in range(n):
if dfs(0, i, j):
return True
return False
记忆要点:
- 在网格 DFS + 回溯问题中,用
'#'或者其他特殊字符标记”访问中”,递归返回后还原原字符。 - 比起维护一个 visited 数组,”标记 + 还原”节省了内存分配,而且能在同一格子上多起点搜索。
- 终止条件的顺序:先判断”匹配失败”(边界/字符不匹配),再判断”匹配成功”(已匹配到 word 的最后一个字符),如果颠倒会因为字符已经被覆盖而出错。
例题 3:路径总和
LC 112. Path Sum。树 DFS 的基础题:判断是否存在从根到叶的路径使得节点值之和等于 targetSum。
1
2
3
4
5
6
7
8
9
10
11
12
class Solution {
public:
bool hasPathSum(TreeNode* root, int targetSum) {
if (!root) return false;
if (!root->left && !root->right) {
return root->val == targetSum;
}
int remain = targetSum - root->val;
return hasPathSum(root->left, remain)
|| hasPathSum(root->right, remain);
}
};
1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
from typing import Optional
class TreeNode:
def __init__(self, val=0, left=None, right=None):
self.val = val
self.left = left
self.right = right
class Solution:
def hasPathSum(self, root: Optional[TreeNode], targetSum: int) -> bool:
if not root:
return False
if not root.left and not root.right:
return root.val == targetSum
remain = targetSum - root.val
return (self.hasPathSum(root.left, remain) or
self.hasPathSum(root.right, remain))
记忆要点:树 DFS 的关键是把”自顶向下”的累加转化为”自底向上”的递归——把当前节点的值减去后传给子节点,叶子节点判断剩余是否为 0。注意要先判断空节点(!root),再判断叶子节点,否则 root->val 会因为 root 为空而崩溃。
单调栈 - Monotonic Stack
单调栈用于解决”下一个更大/更小元素”类型的问题。它通过维护一个单调(递增或递减)的栈,在 O(n) 时间复杂度内完成所有元素的”下一个更大元素”查询。如果临场用暴力 O(n²) 求解,会卡数据范围比较大的题。
典型应用场景:
- 下一个更大元素 / 下一个更小元素
- 柱状图最大矩形
- 每日温度(下一个更暖的日子)
- 循环数组的下一个更大元素
核心模板(下一个更大元素)
1
2
3
4
5
6
7
8
9
10
11
12
13
14
std::vector<int> nextGreater(const std::vector<int>& nums) {
int n = nums.size();
std::vector<int> ans(n, -1);
std::stack<int> stk; // 存下标,栈中元素对应的 nums[stk] 单调递增
for (int i = 0; i < n; ++i) {
while (!stk.empty() && nums[i] > nums[stk.top()]) {
ans[stk.top()] = nums[i];
stk.pop();
}
stk.push(i);
}
return ans;
}
记忆要点:
- 栈中存下标而不是元素本身——这样既能比较值,又能得到位置关系。
- 单调性的判断是针对
nums的,不是栈本身的”高度”——栈只是辅助结构,nums[stk] 保持单调。 - 找”下一个更大元素”用单调递增栈(栈底到栈顶对应 nums[stk] 单调递增);找”下一个更小元素”用单调递减栈。
- 处理循环数组的方法:把数组”复制一份”接在后面(
i % n取模),或者用取模 + 数组长度翻倍的循环。 - 处理”还没找到答案的元素”:循环结束后,栈中剩余的下标对应的答案就是
-1(如果题目默认填 -1)。
例题 1:每日温度
LC 739. Daily Temperatures。单调栈入门题:求每天之后第一个更暖的日子距离几天。
1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
class Solution {
public:
vector<int> dailyTemperatures(vector<int>& temperatures) {
int n = temperatures.size();
std::vector<int> ans(n, 0);
std::stack<int> stk; // 单调递增栈,存下标
for (int i = 0; i < n; ++i) {
while (!stk.empty() && temperatures[i] > temperatures[stk.top()]) {
int prev = stk.top();
stk.pop();
ans[prev] = i - prev;
}
stk.push(i);
}
return ans;
}
};
1
2
3
4
5
6
7
8
9
10
11
12
13
14
from typing import List
class Solution:
def dailyTemperatures(self, temperatures: List[int]) -> List[int]:
n = len(temperatures)
ans = [0] * n
stk = [] # 单调递增栈,存下标
for i in range(n):
while stk and temperatures[i] > temperatures[stk[-1]]:
prev = stk.pop()
ans[prev] = i - prev
stk.append(i)
return ans
记忆要点:栈中存的是”还没找到下一个更大元素”的下标。每当遇到一个更暖的日子,就把栈中所有比它冷的下标都”解决”掉——栈顶对应的答案就是 i - prev(天数差)。
例题 2:下一个更大元素 I
LC 496. Next Greater Element I。这题的变形是 nums1 是 nums2 的子集,只需要返回 nums1 中每个元素在 nums2 中的”下一个更大元素”。
1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
22
23
24
25
26
class Solution {
public:
vector<int> nextGreaterElement(vector<int>& nums1, vector<int>& nums2) {
std::unordered_map<int, int> next;
std::stack<int> stk;
for (int num : nums2) {
while (!stk.empty() && num > stk.top()) {
next[stk.top()] = num;
stk.pop();
}
stk.push(num);
}
while (!stk.empty()) {
next[stk.top()] = -1;
stk.pop();
}
std::vector<int> ans;
ans.reserve(nums1.size());
for (int num : nums1) {
ans.push_back(next[num]);
}
return ans;
}
};
1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
from typing import List
class Solution:
def nextGreaterElement(self, nums1: List[int], nums2: List[int]) -> List[int]:
next_greater = {}
stk = []
for num in nums2:
while stk and num > stk[-1]:
smaller = stk.pop()
next_greater[smaller] = num
stk.append(num)
while stk:
next_greater[stk.pop()] = -1
return [next_greater[num] for num in nums1]
记忆要点:这题栈中存的是元素本身(不是下标),因为我们只需要输出”值”而不需要位置。处理完 nums2 后,栈中剩余元素的”下一个更大元素”就是 -1,要在循环结束后补齐。
例题 3:柱状图中最大的矩形
LC 84. Largest Rectangle in Histogram。单调栈的进阶应用:给定柱状图高度,求最大矩形面积。
1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
class Solution {
public:
int largestRectangleArea(vector<int>& heights) {
int n = heights.size();
std::stack<int> stk; // 单调递增栈,存下标
int ans = 0;
// 多走一轮 i == n,用 cur_h = 0 把栈里所有柱子都弹出来
for (int i = 0; i <= n; ++i) {
int cur_h = (i == n) ? 0 : heights[i];
while (!stk.empty() && cur_h < heights[stk.top()]) {
int h = heights[stk.top()];
stk.pop();
int w = stk.empty() ? i : i - stk.top() - 1;
ans = std::max(ans, h * w);
}
stk.push(i);
}
return ans;
}
};
1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
from typing import List
class Solution:
def largestRectangleArea(self, heights: List[int]) -> int:
n = len(heights)
stk = [] # 单调递增栈,存下标
ans = 0
for i in range(n + 1):
cur_h = 0 if i == n else heights[i]
while stk and cur_h < heights[stk[-1]]:
h = heights[stk.pop()]
w = i if not stk else i - stk[-1] - 1
ans = max(ans, h * w)
stk.append(i)
return ans
记忆要点:
- 哨兵技巧:循环
n + 1次,最后一次用cur_h = 0把栈里所有柱子都弹出来,避免在结尾额外处理栈。 - 弹出柱子
h = heights[top]时,宽度是i - stk.top() - 1(栈顶的新栈顶是左边界)。如果栈已经空了,宽度就是i。 - 这种”用 0 高度清空栈”的技巧可以写出非常简洁的代码,是这题的标志性写法。
前缀和 / 差分数组
前缀和是一种”预处理换查询”的技巧:把”区间和”的多次查询从 O(n) 优化到 O(1),代价是预处理 O(n)。差分数组是前缀和的逆运算,用来高效处理”区间加”操作。
这两个技巧的模板都很短,但临场如果一时想不起来,多次 O(n) 区间求和照样会导致 TLE。
一维前缀和
1
2
3
4
5
6
7
// 构造:pre[i] = nums[0] + nums[1] + ... + nums[i - 1]
std::vector<int> pre(n + 1, 0);
for (int i = 0; i < n; ++i) {
pre[i + 1] = pre[i] + nums[i];
}
// 查询区间 [l, r] 的元素和(包含 l 和 r)
int sum = pre[r + 1] - pre[l];
一维差分数组
1
2
3
4
5
6
7
8
// 给区间 [l, r](包含 l 和 r)每个元素加上 val
diff[l] += val;
diff[r + 1] -= val;
// 最后对 diff 求前缀和就能还原数组
for (int i = 1; i < n; ++i) {
diff[i] += diff[i - 1];
}
二维前缀和(了解即可)
1
2
3
4
5
6
7
8
9
// 构造
for (int i = 0; i < m; ++i)
for (int j = 0; j < n; ++j)
pre[i + 1][j + 1] = nums[i][j] + pre[i][j + 1]
+ pre[i + 1][j] - pre[i][j];
// 查询矩形 (x1, y1) 到 (x2, y2) 的元素和(包含边界)
int sum = pre[x2 + 1][y2 + 1] - pre[x1][y2 + 1]
- pre[x2 + 1][y1] + pre[x1][y1];
记忆要点:
- 前缀和数组长度是
n + 1,这样pre[0] = 0可以优雅地处理l = 0的情况,避免特判。 - 区间
[l, r](包含两端)的和是pre[r + 1] - pre[l],因为pre[r + 1]包含了 nums[0..r],减去pre[l](包含 nums[0..l-1])得到 nums[l..r]。 - 差分数组的核心:
diff[i] = nums[i] - nums[i - 1]。对区间[l, r]加val,等价于diff[l] += val和diff[r + 1] -= val。 - 差分数组处理”多个区间加,最后求每个点的最终值”的问题特别合适,时间复杂度 O(n + k),其中 k 是操作数。
例题 1:区域和检索 - 不可变
LC 303. Range Sum Query - Immutable。前缀和的入门题:构造时算前缀和,查询时 O(1) 返回区间和。
1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
class NumArray {
public:
std::vector<int> pre;
NumArray(vector<int>& nums) {
int n = nums.size();
pre.resize(n + 1, 0);
for (int i = 0; i < n; ++i) {
pre[i + 1] = pre[i] + nums[i];
}
}
int sumRange(int left, int right) {
return pre[right + 1] - pre[left];
}
};
1
2
3
4
5
6
7
8
9
10
from typing import List
class NumArray:
def __init__(self, nums: List[int]):
self.pre = [0]
for num in nums:
self.pre.append(self.pre[-1] + num)
def sumRange(self, left: int, right: int) -> int:
return self.pre[right + 1] - self.pre[left]
记忆要点:构造函数里算 pre[i + 1] = pre[i] + nums[i],查询时直接 pre[right + 1] - pre[left]。如果忘了把 pre 数组长度设为 n + 1,left = 0 的查询就需要特判。
例题 2:航班预订统计
LC 1109. Corporate Flight Bookings。差分数组的入门题:每个 bookings[i] = [l, r, v] 表示航班 l 到 r 预订了 v 个座位,返回每个航班的总预订数。
1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
class Solution {
public:
vector<int> corpFlightBookings(vector<vector<int>>& bookings, int n) {
std::vector<int> diff(n + 1, 0);
for (const auto& booking : bookings) {
int l = booking[0] - 1; // 转成 0-indexed
int r = booking[1] - 1;
int v = booking[2];
diff[l] += v;
diff[r + 1] -= v;
}
std::vector<int> ans(n);
ans[0] = diff[0];
for (int i = 1; i < n; ++i) {
ans[i] = ans[i - 1] + diff[i];
}
return ans;
}
};
1
2
3
4
5
6
7
8
9
10
11
12
13
14
from typing import List
class Solution:
def corpFlightBookings(self, bookings: List[List[int]], n: int) -> List[int]:
diff = [0] * (n + 1)
for l, r, v in bookings:
diff[l - 1] += v
diff[r] -= v
ans = []
cur = 0
for i in range(n):
cur += diff[i]
ans.append(cur)
return ans
记忆要点:差分数组长度是 n + 1,最后一个位置用作 r + 1 的”越界保护”,最后不必取出来。最后扫一遍 diff 累加得到的就是每个航班的总预订数。
例题 3:拼车
LC 1094. Car Pooling。差分数组 + 容量判断:判断从起点到终点接送所有乘客时车上是否超过 capacity(乘客在同一站下就马上能上)。
1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
class Solution {
public:
bool carPooling(vector<vector<int>>& trips, int capacity) {
std::vector<int> diff(1001, 0); // 题目限制 0 <= from < to <= 1000
for (const auto& trip : trips) {
int passengers = trip[0];
int from = trip[1];
int to = trip[2];
diff[from] += passengers;
diff[to] -= passengers;
}
int cur = 0;
for (int i = 0; i <= 1000; ++i) {
cur += diff[i];
if (cur > capacity) return false;
}
return true;
}
};
1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
from typing import List
class Solution:
def carPooling(self, trips: List[List[int]], capacity: int) -> bool:
diff = [0] * 1001
for passengers, from_, to in trips:
diff[from_] += passengers
diff[to] -= passengers
cur = 0
for i in range(1001):
cur += diff[i]
if cur > capacity:
return False
return True
记忆要点:差分数组配合”扫描求前缀和”可以高效判断某个时刻是否超限。这题和上一题的模式完全一样——多个区间”加上乘客”(差分),最后扫一遍看每个位置的累积值(求前缀和),累加过程中一旦超过 capacity 就返回 false。
经典算法和数据结构
哈希表 - Hash Map
哈希表可以用来在O(1)复杂度内判断过去是否遇到过相同值。需要熟练掌握,如果你忽视这种最简单的数据结构的话,你会发现多年不用,可能写不出来!
Two Sum
LeetCode 1. Two Sum:给定一个数组和目标值,如何通过一遍遍历找出可以求和得到目标值的二元组下标?
这题简单的一点在于只需要找到一个解,因为题设只有一个符合条件的解,到后面Three Sum,我们会进一步解决多个解的问题。
关键在于掌握C++中unordered_map的使用方法(如何判断表中是否已有某个元素,如何往表中插入一个新值)。哈希表中的key是target减去当前元素的值(因为下次遇到这个key的值时,就说明找到了),value是当前元素的下标。
1
2
3
4
5
6
7
8
9
10
11
12
13
14
class Solution {
public:
vector<int> twoSum(vector<int>& nums, int target) {
unordered_map<int, int> map_;
vector<int> result;
for(int i=0;i<nums.size();i++) {
if(map_.find(nums[i]) != map_.end()) {
return vector<int>{map_[nums[i]], i};
}
map_[target - nums[i]] = i;
}
return vector<int>{};
}
};
Python实现:
1
2
3
4
5
6
7
8
class Solution:
def twoSum(self, nums: List[int], target: int) -> List[int]:
storage = {}
for index, num in enumerate(nums):
if target - num in storage:
return [storage[target - num], index]
storage[num] = index
return []
Three Sum
LeetCode 15. 3Sum。这题需要找出数组中所有不重复的三元组,求和等于0。比前面TwoSum复杂的在于解不止一个,而且不能出现重复。
这个问题的一个思路是,对数组中每个值$value_i$都可以往后用TwoSum把target设为$-value_i$来求得三元组,为了不出现重复的,可以对数组进行排序,并且在TwoSum内使用std::set来避免插入重复元素(这里每次插入的复杂度是O(log n),排序的时间复杂度是$O(nlog n)$,前面这个操作是$O(n^2 log n)$,所以最后总的时间复杂度是$O(n^2 log n)$。这不是最优解,时间复杂度最优可以到$O(n^2)$,不过这是一个最容易想到的解决方案。
1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
22
23
24
25
26
27
28
29
30
31
32
33
34
35
36
37
38
39
40
41
42
43
44
45
46
47
48
49
50
51
class Solution {
public:
std::set<std::pair<int, int>> twoSum(
const std::vector<int>& nums,
int start,
int target
) {
std::unordered_map<int, int> needed;
std::set<std::pair<int, int>> result;
for (int i = start; i < static_cast<int>(nums.size()); ++i) {
auto it = needed.find(nums[i]);
if (it != needed.end()) {
result.insert({it->second, nums[i]});
}
// If a later number equals target - nums[i],
// combine it with the current nums[i].
needed[target - nums[i]] = nums[i];
}
return result;
}
std::vector<std::vector<int>> threeSum(std::vector<int>& nums) {
std::vector<std::vector<int>> final_result;
std::sort(nums.begin(), nums.end());
int n = static_cast<int>(nums.size());
for (int i = 0; i + 2 < n; ++i) {
// Avoid generating the same first element again.
if (i > 0 && nums[i] == nums[i - 1]) {
continue;
}
auto results = twoSum(nums, i + 1, -nums[i]);
for (const auto& [first, second] : results) {
final_result.push_back({
nums[i],
first,
second
});
}
}
return final_result;
}
};
上面这种方法还是会导致产生重复元组,是利用了std::set的方式进行清除,一个更好的方式是用TwoPointers的思想,从源头上避免重复,这个方案的时间复杂度是$O(n^2)$。这里的关键思路是在从小到大排序的情况下,求和想要值变大,一定是左边的pointer右移,要想值变小,一定是右边的pointer往左移。
1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
22
23
24
25
26
27
28
29
30
31
32
33
34
35
36
37
38
39
40
41
42
43
44
45
46
47
48
49
50
51
52
53
54
55
56
57
58
class Solution {
public:
std::vector<std::vector<int>> threeSum(std::vector<int>& nums) {
std::vector<std::vector<int>> result;
std::sort(nums.begin(), nums.end());
int n = static_cast<int>(nums.size());
for (int i = 0; i + 2 < n; ++i) {
// Once nums[i] is positive, the remaining values are also
// positive, so their sum cannot be zero.
if (nums[i] > 0) {
break;
}
// Skip duplicate choices for the first number.
if (i > 0 && nums[i] == nums[i - 1]) {
continue;
}
int left = i + 1;
int right = n - 1;
while (left < right) {
int sum = nums[i] + nums[left] + nums[right];
if (sum < 0) {
++left;
} else if (sum > 0) {
--right;
} else {
result.push_back({
nums[i],
nums[left],
nums[right]
});
++left;
--right;
// Skip duplicate second numbers.
while (left < right &&
nums[left] == nums[left - 1]) {
++left;
}
// Skip duplicate third numbers.
while (left < right &&
nums[right] == nums[right + 1]) {
--right;
}
}
}
}
return result;
}
};
红黑树/平衡树 - std::set
一道经典的问题是选择最匹配的内存部署虚拟机:一批物理机记录于数组capacities中,capacities[i]表示编号为i的物理机的初始内存大小;同时给出一批虚拟机部署请求requests,requests[i]表示某虚拟机的所需内存。请按如下规则依次处理每个虚拟机部署请求,并返回每个虚拟机部署所在物理机的编号(或者-1):
- 如果所有物理机的可用内存不足,则部署失败,返回-1
- 否则,在满足虚拟机所需内存的所有物理机中,选择可用内存最小的;若仍有多台,选择其中编号最小的。
这题最重要的是了解平衡树的特性在C++中如何使用,可以使用std::set来完成。std::set是通过平衡树实现的,通常是红黑树(如果不会用要现场手搓红黑树,想想就酸爽)。
需要记住的知识点主要有这两点:(1)了解C++ STL中自定义排序的写法,最方便的是使用lambda函数;(2) std::set的lower_bound函数的用法来进行O(log n)的搜索找到“第一个容量大于等于请求”。这个lower_bound的特性是优先队列没有的,所以这题需要平衡树来解决,而不是使用std::priority_queue。
1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
22
23
24
25
26
27
28
29
30
31
32
33
34
struct Machine {
int capacity;
int id;
Machine(int capacity, int id):capacity(capacity), id(id) {}
Machine():capacity(0), id(-1){}
};
class Solution {
vector<int> DispatchRequests(const vector<int>& capacities, const vector<int>& requests) {
auto cmp = [](const Machine& a, const Machine& b) {
if(a.capacity == b.capacity) {return a.id < b.id;}
return a.capacity < b.capacity;
};
std::set<Machine, decltype(cmp)> container(cmp); // C++ 20之后可以不传递cmp参数到构造函数中
for (size_t i=0;i<capacities.size(); i++) {
container.insert(Machine{capacities[i], i});
}
std::vector<int> result;
result.reserve(requests.size()); // 减少push_back重新分配内存的时间开销
for(auto request: requests) {
auto it = container.lower_bound(Machine{request, -1});
if(it != container.end()) {
Machine tmp = *it;
result.push_back(tmp.id);
container.erase(it);
tmp.capacity -= request;
container.insert(tmp);
} else {
result.push_back(-1);
}
}
return result;
}
};
Python的实现方式 —— Python没有内置的红黑树实现,需要一个三方库sortedconatiners中的SortedList。
1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
22
23
24
25
26
27
28
29
30
31
32
33
34
from typing import List
from sortedcontainers import SortedList
class Solution:
def dispatch_requests(
self,
capacities: List[int],
requests: List[int]
) -> List[int]:
# 每个元素为 (剩余容量, 物理机编号)
machines = SortedList(
(capacity, machine_id)
for machine_id, capacity in enumerate(capacities)
)
result = []
for request in requests:
# 寻找第一个不小于 (request, -1) 的元素
index = machines.bisect_left((request, -1))
if index == len(machines):
result.append(-1)
continue
capacity, machine_id = machines.pop(index)
result.append(machine_id)
# 更新物理机的剩余容量
machines.add((capacity - request, machine_id))
return result
Two Pointers
Container With Most Water
LeetCode 11. Container With Most Water 是经典的Two Pointers入门题
主要思想是两边的柱子只有更矮的那根往中间走,才可能让水面变高,从而围出更大的面积。具体实现如下:
1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
class Solution {
public:
int maxArea(vector<int>& height) {
int left = 0, right = height.size() - 1;
int cur_max = std::min(height[left], height[right]) * (right - left);
while(left < right) {
if(height[left] < height[right]) {
left++;
} else {
right--;
}
cur_max = std::max(cur_max, std::min(height[left], height[right]) * (right - left));
}
return cur_max;
}
};
前面提到的Three Sum问题是另外一个经典的TwoPointers可以解决的场景。
二分查找 - Binary Search
二分查找是看起来简单、实则最容易写错的算法之一。ACM 赛场上”边界条件没写好导致死循环或漏解”几乎是最常见的 bug 之一——核心难点就在 mid 的计算方式和区间的开闭。本节整理三种最常考的二分模板。
标准库的三个函数
C++ STL 的 <algorithm> 提供了三个核心函数(要求区间已经排好序):
std::binary_search(begin, end, value):判断value是否在区间内,返回bool,时间复杂度O(log n)。std::lower_bound(begin, end, value):返回指向第一个大于等于value的元素的迭代器;如果不存在,返回end。std::upper_bound(begin, end, value):返回指向第一个严格大于value的元素的迭代器;如果不存在,返回end。
手写二分模板
最容易记错的是 mid 的计算和边界收缩,建议背下面两个版本之一:
版本一:闭区间 [l, r],寻找 target(target 存在时返回任一下标,不存在返回 -1)
1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
int binarySearch(const std::vector<int>& nums, int target) {
int l = 0;
int r = static_cast<int>(nums.size()) - 1;
while (l <= r) {
int mid = l + (r - l) / 2; // 防止 (l+r) 整数相加溢出
if (nums[mid] == target) {
return mid;
} else if (nums[mid] < target) {
l = mid + 1;
} else {
r = mid - 1;
}
}
return -1;
}
版本二:在答案上二分——找满足谓词 predicate(x) 的最小/最大 x,关键是 predicate 在定义域上单调(一段 false 后跟一段 true,或反过来)。
1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
22
23
24
25
26
27
// 寻找最小值 x 使得 predicate(x) == true
// 假设 predicate 在 [lo, hi] 上是 false ... false true ... true
int lowerBound(int lo, int hi) {
while (lo < hi) {
int mid = lo + (hi - lo) / 2;
if (predicate(mid)) {
hi = mid; // mid 满足条件,答案在 [lo, mid]
} else {
lo = mid + 1; // mid 不满足条件,答案在 [mid+1, hi]
}
}
return lo;
}
// 寻找最大值 x 使得 predicate(x) == true
// 假设 predicate 在 [lo, hi] 上是 true ... true false ... false
int upperBound(int lo, int hi) {
while (lo < hi) {
int mid = lo + (hi - lo + 1) / 2; // 向上取整,否则会死循环
if (predicate(mid)) {
lo = mid; // mid 满足条件,答案在 [mid, hi]
} else {
hi = mid - 1; // mid 不满足条件,答案在 [lo, mid-1]
}
}
return lo;
}
记忆要点:
mid = lo + (hi - lo) / 2而不是(lo + hi) / 2是为了避免整数相加溢出。- 找最大满足条件的时候,
mid要向上取整(lo + (hi - lo + 1) / 2),否则lo = mid不会前进,会死循环。 - 在答案上二分的核心套路是:把”求最优”转化为”判定”——给定一个候选答案,O(n) 或 O(n log n) 判断它是否可行,然后二分搜索最优解。
例题 1:寻找目标元素的第一个和最后一个位置
LeetCode 34. Find First and Last Position of Element in Sorted Array。这题是 std::lower_bound 和 std::upper_bound 的最佳应用:找到 target 第一次出现的位置,就是 lower_bound(target);找到最后一次出现的位置,就是 upper_bound(target) - 1。如果两者相等,说明 target 不存在。
C++ 实现:
1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
22
23
class Solution {
public:
std::vector<int> searchRange(
std::vector<int>& nums,
int target
) {
auto lower = std::lower_bound(
nums.begin(), nums.end(), target
);
auto upper = std::upper_bound(
nums.begin(), nums.end(), target
);
if (lower == upper) {
return {-1, -1};
}
return {
static_cast<int>(lower - nums.begin()),
static_cast<int>(upper - nums.begin() - 1)
};
}
};
Python 实现(bisect_left 和 bisect_right 分别对应 lower_bound 和 upper_bound,两者的差就是 target 的出现次数):
1
2
3
4
5
6
7
8
9
10
11
from typing import List
from bisect import bisect_left, bisect_right
class Solution:
def searchRange(self, nums: List[int], target: int) -> List[int]:
lower = bisect.bisect_left(nums, target)
upper = bisect.bisect_right(nums, target)
if lower == upper:
return [-1, -1]
return [lower, upper -1]
例题 2:在答案上二分
LeetCode 410. Split Array Largest Sum。这题需要把一个数组分成 m 段,让最大的子段和最小。答案是”最大子段和”的最小值,典型的在答案上二分的题目。
我们可以把问题转化为一个判定问题:给定一个最大子段和 limit,能不能把数组分成不超过 m 段,使得每段的和都不超过 limit?这个判定函数是 O(n) 的贪心扫描,而 limit 的范围是 [max(nums), sum(nums)]。对 limit 做二分,总复杂度为 O(n log(sum))。
C++ 实现:
1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
22
23
24
25
26
27
28
29
30
31
32
33
34
35
36
37
38
39
40
41
42
43
44
45
46
47
48
class Solution {
public:
// 给定 limit,能否把 nums 分成不超过 m 段使得每段和 <= limit
bool canSplit(
const std::vector<int>& nums,
int m,
long long limit
) {
int pieces = 1;
long long current = 0;
for (int num : nums) {
if (current + num <= limit) {
current += num;
} else {
pieces++;
current = num;
if (pieces > m) {
return false;
}
}
}
return true;
}
int splitArray(
std::vector<int>& nums,
int m
) {
long long lo = *std::max_element(
nums.begin(), nums.end()
);
long long hi = std::accumulate(
nums.begin(), nums.end(), 0LL
);
// 寻找最小值 x 使得 canSplit(..., x) == true
// canSplit 在 x 上单调:x 越大越容易满足
while (lo < hi) {
long long mid = lo + (hi - lo) / 2;
if (canSplit(nums, m, mid)) {
hi = mid;
} else {
lo = mid + 1;
}
}
return static_cast<int>(lo);
}
};
Python 实现(请自行完成——提示:用 max(nums) 作为下界,sum(nums) 作为上界,谓词函数写成贪心的 can_split,然后套用”在答案上二分”的模板):
1
2
3
4
5
6
7
8
9
from typing import List
class Solution:
def splitArray(self, nums: List[int], m: int) -> int:
# TODO: 请实现"在答案上二分"的 splitArray
# 1) 写一个 can_split(limit) 函数
# 2) 对 [max(nums), sum(nums)] 二分
pass
字符串处理 - String Algorithms
这一章梳理四个最容易临场写不出来的字符串算法:字符串哈希、KMP、Z 函数、Manacher。它们都有非常固定的”模板代码”——背下来就能解决一大类问题,但每次都自己推一遍,几乎一定会写错或写超时。
字符串哈希 - Rolling Hash
字符串哈希(也叫滚动哈希 / Rabin-Karp)的核心思想是:把字符串看作一个 $B$ 进制的”大整数”,让两个字符串的比较从 O(L) 退化到 O(1)。具体地:
\[H(s) = s_0 \cdot B^{L-1} + s_1 \cdot B^{L-2} + \dots + s_{L-1} \cdot B^{0}\]预处理出前缀哈希和 $B$ 的幂次后,任意子串 $s[l..r]$ 的哈希值都可以 O(1) 求出。配合 uint64_t 自然溢出(等价于 mod $2^{64}$),碰撞概率极低;想要绝对安全,可以用双模($10^9+7$ 和 $10^9+9$)配 pair。
模板
1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
22
class StringHash {
public:
static const uint64_t B = 1315423911ULL; // 大奇数基底
std::vector<uint64_t> h; // h[i] = s[0..i-1] 的哈希
std::vector<uint64_t> p; // p[i] = B^i
StringHash(const std::string& s) {
int n = (int)s.size();
h.assign(n + 1, 0);
p.assign(n + 1, 0);
p[0] = 1;
for (int i = 0; i < n; ++i) {
h[i + 1] = h[i] * B + (unsigned char)s[i];
p[i + 1] = p[i] * B;
}
}
// s[l..r](闭区间)的哈希值,O(1)
uint64_t get(int l, int r) const {
return h[r + 1] - h[l] * p[r - l + 1];
}
};
1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
class StringHash:
B = 1315423911
MASK = (1 << 64) - 1 # Python 整数无上限,需手动截断
def __init__(self, s: str):
n = len(s)
self.h = [0] * (n + 1)
self.p = [0] * (n + 1)
self.p[0] = 1
for i, ch in enumerate(s):
self.h[i + 1] = ((self.h[i] * self.B) + ord(ch)) & self.MASK
self.p[i + 1] = (self.p[i] * self.B) & self.MASK
def get(self, l: int, r: int) -> int:
# s[l..r](闭区间)的哈希值
return (self.h[r + 1] - self.h[l] * self.p[r - l + 1]) & self.MASK
记忆要点:
- 基底选择:131、13131、137 等大奇数都是常见基底,配
uint64_t自然溢出时碰撞概率极低。 - 前缀哈希数组长度是 n+1,这样
get(0, r)不需要特判——和前缀和数组的 trick 一样。 - 子串哈希公式:
h[r+1] - h[l] * pow(B, r-l+1)。这一行最容易写错——r-l+1必须严格匹配窗口长度。 - 滚动更新:把窗口右端字符加入、左端字符移出,得到新的窗口哈希,公式为
h_new = (h - s[left] * p) * B + s[right](其中p = B^(window_size))。 - 自然溢出 vs 双模:自然溢出实现最简单但理论上可能碰撞;想要绝对安全用
pair<uint64_t, uint64_t>配两个大质数($10^9+7$ 和 $10^9+9$)。
例题:重复的 DNA 序列
LC 187. Repeated DNA Sequences。给定长度为 $n$ 的 DNA 序列(只含 A/C/G/T),找出所有出现超过一次的长度为 10 的子串。
1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
22
23
24
25
26
27
28
29
30
31
32
33
34
class Solution {
public:
vector<string> findRepeatedDnaSequences(string s) {
int n = (int)s.size();
if (n <= 10) return {};
const uint64_t B = 1315423911ULL;
uint64_t h = 0, p = 1;
// 初始化:把前 10 个字符的哈希算出来
for (int i = 0; i < 10; ++i) {
h = h * B + (unsigned char)s[i];
if (i > 0) p *= B;
}
// 哈希 → 出现的起点位置列表
unordered_map<uint64_t, vector<int>> pos;
pos[h].push_back(0);
// 滑动窗口:每个窗口哈希对应起点 i - 9
for (int i = 10; i < n; ++i) {
h = (h - (uint64_t)(unsigned char)s[i - 10] * p) * B
+ (unsigned char)s[i];
pos[h].push_back(i - 9);
}
vector<string> ans;
for (auto& kv : pos) {
if (kv.second.size() > 1) {
ans.push_back(s.substr(kv.second[0], 10));
}
}
return ans;
}
};
1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
22
23
24
25
from typing import List
class Solution:
B = 1315423911
MASK = (1 << 64) - 1
def findRepeatedDnaSequences(self, s: str) -> List[str]:
n = len(s)
if n <= 10:
return []
h = 0
p = 1
for i in range(10):
h = ((h * self.B) + ord(s[i])) & self.MASK
if i > 0:
p = (p * self.B) & self.MASK
pos = {h: [0]}
for i in range(10, n):
h = ((h - ord(s[i - 10]) * p) * self.B + ord(s[i])) & self.MASK
pos.setdefault(h, []).append(i - 9)
ans = []
for v in pos.values():
if len(v) > 1:
ans.append(s[v[0]:v[0] + 10])
return ans
记忆要点:
- 滚动窗口哈希的核心是
h_new = (h - s[left] * p) * B + s[right],其中p = B^(window_size)是预计算好的。 - 这道题用
unordered_map<uint64_t, vector<int>>把哈希映射回起点位置,是为了从哈希反推回原字符串;如果只关心”出现过几次”,可以直接用unordered_map<uint64_t, int>计数。 - 在 ACGT 这种小字符集上,也可以用 2-bit 编码
((h << 2) | code) & mask把 20 个 bit 塞进一个 int,完全避免哈希碰撞,但这就局限到固定窗口大小了。
KMP 算法
KMP(Knuth-Morris-Pratt)能在 $O(n+m)$ 内完成”模式串 P 在文本串 S 中的所有匹配”。核心是预先算出 next[i]:p[0..i] 中”最长相等前后缀”的长度。这样匹配失败时,模式串不用从头来过,而是直接跳到 next[i-1] 继续尝试。
next 数组的构造
next[i] 表示 p[0..i] 的最长相等真前后缀的长度(即 prefix = suffix,且都严格短于 p[0..i+1])。
1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
// next[i] = p[0..i] 中最长相等真前后缀的长度
std::vector<int> buildNext(const std::string& p) {
int m = (int)p.size();
std::vector<int> nxt(m, 0);
// nxt[0] = 0(单个字符没有真前后缀)
for (int i = 1; i < m; ++i) {
int j = nxt[i - 1];
while (j > 0 && p[j] != p[i]) {
j = nxt[j - 1]; // 回退
}
if (p[j] == p[i]) ++j;
nxt[i] = j;
}
return nxt;
}
记忆要点:
j = nxt[i-1]是关键——”尝试用p[0..i-1]已经有的最长相等前后缀来扩展”。- 如果
p[j] != p[i],就把j退回到nxt[j-1],这是一个自我递归的过程(”前缀的前缀还是前缀”)。 - 最后如果
p[j] == p[i],把j加 1;否则保持 0。
匹配过程
1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
std::vector<int> kmpSearch(const std::string& s, const std::string& p) {
std::vector<int> nxt = buildNext(p);
int n = (int)s.size(), m = (int)p.size();
std::vector<int> ans;
int j = 0;
for (int i = 0; i < n; ++i) {
while (j > 0 && p[j] != s[i]) {
j = nxt[j - 1];
}
if (p[j] == s[i]) ++j;
if (j == m) {
ans.push_back(i - m + 1); // 找到一个匹配,起点是 i - m + 1
j = nxt[j - 1]; // 继续寻找下一个匹配
}
}
return ans;
}
记忆要点:
- 匹配失败时
j = nxt[j-1],匹配成功时j++——这两个动作必须严格对称。 - 找到一个匹配后要
j = nxt[j-1]而不是j = 0,因为可能存在重叠匹配(例如p = "aaa",在s = "aaaa"中匹配位置是 0, 1, 2)。
例题:找出字符串中第一个匹配项的下标
LC 28. Find the Index of the First Occurrence in a String。返回 needle 在 haystack 中第一次出现的下标;如果不存在,返回 -1。
1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
22
23
24
25
26
27
class Solution {
public:
int strStr(string haystack, string needle) {
int n = (int)haystack.size(), m = (int)needle.size();
if (m > n) return -1;
std::vector<int> nxt(m, 0);
for (int i = 1; i < m; ++i) {
int j = nxt[i - 1];
while (j > 0 && needle[j] != needle[i]) {
j = nxt[j - 1];
}
if (needle[j] == needle[i]) ++j;
nxt[i] = j;
}
int j = 0;
for (int i = 0; i < n; ++i) {
while (j > 0 && needle[j] != haystack[i]) {
j = nxt[j - 1];
}
if (needle[j] == haystack[i]) ++j;
if (j == m) return i - m + 1;
}
return -1;
}
};
1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
22
class Solution:
def strStr(self, haystack: str, needle: str) -> int:
n, m = len(haystack), len(needle)
if m > n:
return -1
nxt = [0] * m
for i in range(1, m):
j = nxt[i - 1]
while j > 0 and needle[j] != needle[i]:
j = nxt[j - 1]
if needle[j] == needle[i]:
j += 1
nxt[i] = j
j = 0
for i in range(n):
while j > 0 and needle[j] != haystack[i]:
j = nxt[j - 1]
if needle[j] == haystack[i]:
j += 1
if j == m:
return i - m + 1
return -1
记忆要点:构造 next 和匹配 s 用的循环结构几乎完全一样——都是”匹配失败就退到 next[j-1],匹配成功 j++“。背下一种就等于背下两种。
Z 函数 - Z-Algorithm
Z 函数 $z[i]$ 表示 $s$ 和它的后缀 $s[i..n-1]$ 的最长公共前缀长度。预处理 $O(n)$ 后,可以用 Z 函数做模式匹配。
模板
1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
std::vector<int> zFunction(const std::string& s) {
int n = (int)s.size();
std::vector<int> z(n, 0);
int l = 0, r = 0; // 当前已知的最右 Z-box [l, r]
for (int i = 1; i < n; ++i) {
if (i <= r) {
z[i] = std::min(r - i + 1, z[i - l]);
}
while (i + z[i] < n && s[z[i]] == s[i + z[i]]) {
++z[i];
}
if (i + z[i] - 1 > r) {
l = i;
r = i + z[i] - 1;
}
}
return z;
}
记忆要点:
- Z-box 维护:维护一个最右的 Z-box
[l, r]。如果当前i在 box 内,可以直接”借”z[i-l]的值(但不超过r-i+1)。 - 扩展:从初始猜测开始,尝试向右扩展直到不匹配为止。
- 更新 box:如果扩展后的
i + z[i] - 1超过了r,更新l = i, r = i + z[i] - 1。
模式匹配:把模式串拼到文本前面
Z 函数最常见的应用是模式匹配:把 p + "#" + s 拼成一个新串 concat,对 concat 求 Z 数组。如果 z[i] == |p|,说明 p 从 s[i - |p| - 1] 开始匹配成功。
1
2
3
4
5
6
7
8
9
10
11
12
13
// 找出 pattern p 在 s 中的所有匹配位置
std::vector<int> findMatches(const std::string& s, const std::string& p) {
std::string concat = p + "#" + s;
auto z = zFunction(concat);
int m = (int)p.size();
std::vector<int> ans;
for (int i = m + 1; i < (int)concat.size(); ++i) {
if (z[i] == m) {
ans.push_back(i - m - 1);
}
}
return ans;
}
记忆要点:Z 函数和 KMP 都能做模式匹配,KMP 更短一些(代码量小一半),但 Z 函数更容易扩展到”找所有公共前缀长度”等更复杂的问题。
Manacher 算法
Manacher 能在 $O(n)$ 内求出字符串的最长回文子串。核心 trick 是在字符之间(和两端)插入特殊字符 #,把”奇数长度”和”偶数长度”的回文统一成”奇数长度”。
模板
1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
22
23
24
25
26
27
28
29
30
31
32
33
34
35
// 返回 s 的最长回文子串
std::string manacher(const std::string& s) {
// 插入 #:例如 "aba" -> "^#a#b#a#$"
std::string t = "^#";
for (char c : s) {
t += c;
t += "#";
}
t += "$";
int n = (int)t.size();
std::vector<int> p(n, 0); // p[i] = 以 i 为中心的回文半径
int c = 0, r = 0; // 当前最右回文中心和右边界
for (int i = 1; i < n - 1; ++i) {
int mirror = 2 * c - i;
if (i < r) {
p[i] = std::min(r - i, p[mirror]);
}
while (t[i + p[i] + 1] == t[i - p[i] - 1]) {
++p[i];
}
if (i + p[i] > r) {
c = i;
r = i + p[i];
}
}
// 找到最大半径
int max_i = 0;
for (int i = 1; i < n - 1; ++i) {
if (p[i] > p[max_i]) max_i = i;
}
int start = (max_i - p[max_i]) / 2;
return s.substr(start, p[max_i]);
}
记忆要点:
- 插入
#和边界哨兵^/$:边界哨兵让while扩展自动终止,不用特判越界。 - p[i] 的含义:以 i 为中心的回文”半径”,长度是
p[i](含中心),对应原串的回文长度就是p[i]。 - mirror 公式
mirror = 2*c - i:这是以 c 为中心的对称点。 - 更新 c, r:每一步都可能更新最右回文边界,这正是 Manacher 能维持 $O(n)$ 的关键。
例题:最长回文子串
LC 5. Longest Palindromic Substring。直接套用上面的 Manacher 实现即可。
1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
22
23
24
25
26
27
28
29
30
31
32
33
34
35
36
class Solution {
public:
string longestPalindrome(string s) {
std::string t = "^#";
for (char c : s) {
t += c;
t += "#";
}
t += "$";
int n = (int)t.size();
std::vector<int> p(n, 0);
int c = 0, r = 0;
for (int i = 1; i < n - 1; ++i) {
int mirror = 2 * c - i;
if (i < r) {
p[i] = std::min(r - i, p[mirror]);
}
while (t[i + p[i] + 1] == t[i - p[i] - 1]) {
++p[i];
}
if (i + p[i] > r) {
c = i;
r = i + p[i];
}
}
int max_i = 0;
for (int i = 1; i < n - 1; ++i) {
if (p[i] > p[max_i]) max_i = i;
}
int start = (max_i - p[max_i]) / 2;
return s.substr(start, p[max_i]);
}
};
1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
class Solution:
def longestPalindrome(self, s: str) -> str:
t = "^#" + "#".join(s) + "#$"
n = len(t)
p = [0] * n
c = r = 0
for i in range(1, n - 1):
mirror = 2 * c - i
if i < r:
p[i] = min(r - i, p[mirror])
while t[i + p[i] + 1] == t[i - p[i] - 1]:
p[i] += 1
if i + p[i] > r:
c = i
r = i + p[i]
max_i = max(range(1, n - 1), key=lambda i: p[i])
start = (max_i - p[max_i]) // 2
return s[start:start + p[max_i]]
记忆要点:
- Manacher 比”中心扩展”(每个位置向两侧扩展)快很多——后者最坏 $O(n^2)$,前者保证 $O(n)$。
- 输出时回文长度 =
p[max_i],起点 =(max_i - p[max_i]) / 2。 - 如果不需要”重建”最长回文子串,只需要长度,可以直接返回
p[max_i]。
并查集 - Union-Find / DSU
并查集(Disjoint Set Union)是 ACM 比赛中最容易写不出的数据结构之一——核心代码只有十几行,但如果没有背下”路径压缩 + 按秩合并”两个优化,临场大概率会写错或者写出退化到 O(n) 的版本。它主要用于处理元素的分组关系和判断两个元素是否属于同一组——典型场景包括:图中的连通分量、岛屿问题、生成树相关(Kruskal)、冗余边检测等。
核心数据结构
并查集维护一个森林,每个集合用一棵树表示,根节点是这个集合的”代表”。需要三个关键的状态:
parent[i]:节点i的父节点(根节点的parent指向自身)。rank[i]或size[i]:以i为根的树的深度/大小,用于按秩合并。- 任意一个”非根”节点都可以通过
find操作回到根。
三个基础操作
find(x):找到x所在集合的根,同时把路径上所有节点直接挂到根上(路径压缩)。union(x, y):合并x和y所在的集合。先find出各自的根,再把秩/大小更小的根挂到更大的根下面(按秩合并)。connected(x, y):等价于find(x) == find(y)。
记忆要点:
find必须带路径压缩,否则树可能退化成链,时间复杂度退化到O(n)。写法是先递归找根,再把parent[x]指向根——return parent[x] = find(parent[x])。union必须按秩合并:两棵树深度不同时,把深度小的树根指向深度大的树根;如果深度相同,新根的深度加 1。这样能保证树的高度是O(log n)。- 路径压缩 + 按秩合并后,单次操作均摊复杂度为
O(α(n)),其中α是反阿克曼函数,实际使用中可以视为常数。
模板代码
1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
22
23
24
25
26
27
28
29
30
31
32
33
34
35
36
37
38
39
40
class DSU {
public:
std::vector<int> parent;
std::vector<int> rank_;
DSU(int n) {
parent.resize(n);
rank_.assign(n, 0);
for (int i = 0; i < n; ++i) {
parent[i] = i;
}
}
int find(int x) {
if (parent[x] != x) {
parent[x] = find(parent[x]); // 路径压缩
}
return parent[x];
}
void unite(int x, int y) {
int rx = find(x);
int ry = find(y);
if (rx == ry) {
return;
}
// 按秩合并
if (rank_[rx] < rank_[ry]) {
std::swap(rx, ry);
}
parent[ry] = rx;
if (rank_[rx] == rank_[ry]) {
rank_[rx]++;
}
}
bool connected(int x, int y) {
return find(x) == find(y);
}
};
例题 1:省份数量
LeetCode 547. Number of Provinces。这题是并查集的最直接应用:给定城市之间的邻接矩阵,求连通分量的数量。遍历邻接矩阵,把每对相连的城市 unite 起来,最后统计有多少个节点的 parent[i] == i(即根的数量)。
C++ 实现:
1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
22
23
24
25
26
27
28
29
30
31
32
33
34
35
36
37
38
39
40
41
42
43
44
45
46
47
48
49
class DSU {
public:
std::vector<int> parent;
std::vector<int> rank_;
DSU(int n) : parent(n), rank_(n, 0) {
for (int i = 0; i < n; ++i) {
parent[i] = i;
}
}
int find(int x) {
if (parent[x] != x) {
parent[x] = find(parent[x]);
}
return parent[x];
}
void unite(int x, int y) {
int rx = find(x);
int ry = find(y);
if (rx == ry) return;
if (rank_[rx] < rank_[ry]) std::swap(rx, ry);
parent[ry] = rx;
if (rank_[rx] == rank_[ry]) rank_[rx]++;
}
};
class Solution {
public:
int findCircleNum(std::vector<std::vector<int>>& isConnected) {
int n = isConnected.size();
DSU dsu(n);
for (int i = 0; i < n; ++i) {
for (int j = i + 1; j < n; ++j) {
if (isConnected[i][j] == 1) {
dsu.unite(i, j);
}
}
}
int provinces = 0;
for (int i = 0; i < n; ++i) {
if (dsu.find(i) == i) {
provinces++;
}
}
return provinces;
}
};
Python 实现(请自行完成——提示:可以用 list 当 parent,或者直接用 Python 自带的字典记录父节点;记得在 find 时做路径压缩):
1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
from typing import List
class DSU:
def __init__(self, n: int):
# TODO: 初始化 parent 和 rank_
pass
def find(self, x: int) -> int:
# TODO: 带路径压缩的 find
pass
def unite(self, x: int, y: int) -> None:
# TODO: 按秩合并
pass
class Solution:
def findCircleNum(self, isConnected: List[List[int]]) -> int:
# TODO: 调用 DSU,统计根节点数量
pass
例题 2:冗余连接
LeetCode 684. Redundant Connection。这题给一棵树加上一条边后形成了带环的图,要求找到这条多余的边。思路:依次尝试加入每条边,加入前如果发现两个节点已经连通(即 find(x) == find(y)),说明这条边会形成环,它就是要找的冗余边。
C++ 实现:
1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
22
23
24
25
26
27
28
29
30
31
32
33
34
35
36
37
38
39
40
class Solution {
public:
std::vector<int> parent;
std::vector<int> rank_;
int find(int x) {
if (parent[x] != x) {
parent[x] = find(parent[x]);
}
return parent[x];
}
bool unite(int x, int y) {
int rx = find(x);
int ry = find(y);
if (rx == ry) {
return false; // 已在同一集合
}
if (rank_[rx] < rank_[ry]) std::swap(rx, ry);
parent[ry] = rx;
if (rank_[rx] == rank_[ry]) rank_[rx]++;
return true;
}
std::vector<int> findRedundantConnection(
std::vector<std::vector<int>>& edges
) {
int n = edges.size();
parent.resize(n + 1);
rank_.assign(n + 1, 0);
for (int i = 0; i <= n; ++i) parent[i] = i;
for (const auto& edge : edges) {
if (!unite(edge[0], edge[1])) {
return edge; // 这条边会造成环
}
}
return {};
}
};
Python 实现(请自行完成——提示:unite 在加入一条边之前如果发现两个节点已连通,就返回这条边本身):
1
2
3
4
5
6
7
8
9
from typing import List
class Solution:
def findRedundantConnection(self, edges: List[List[int]]) -> List[int]:
# TODO:
# 1) 实现一个轻量的 DSU(parent list + find + unite)
# 2) 遍历 edges,第一次让 unite 返回 False 时返回当前边
pass
动态规划 - Dynamic Programming
动态规划(DP)解决”原问题的解可由子问题的解推出”的问题。它不是某个具体算法,而是一种”用空间换时间”的范式——把子问题的答案存下来,下次需要时直接查表,避免重复计算。
DP 的四要素:
- 状态定义:
dp[i](或dp[i][j])表示什么? - 转移方程:
dp[i]怎么由之前的dp推出? - 初始化:哪些基础情况要直接给出?
- 遍历顺序:保证计算
dp[i]时,所需的小问题已经算好。
最容易出错的是”状态定义”——状态没想清楚,转移方程就写不出来。临场判断能不能用 DP 的标志是:问题有重叠子问题 + 最优子结构。
一维 DP
一维 DP 的状态只用一个下标描述,转移通常从前面(或后面)的状态推出。最经典的入门题是斐波那契 / 爬楼梯。
例题 1:爬楼梯
LC 70. Climbing Stairs。一次爬 1 或 2 阶,爬到第 n 阶有几种方法。
1
2
3
4
5
6
7
8
9
10
11
12
13
class Solution {
public:
int climbStairs(int n) {
if (n <= 2) return n;
int a = 1, b = 2;
for (int i = 3; i <= n; ++i) {
int c = a + b;
a = b;
b = c;
}
return b;
}
};
1
2
3
4
5
6
7
8
class Solution:
def climbStairs(self, n: int) -> int:
if n <= 2:
return n
a, b = 1, 2
for _ in range(3, n + 1):
a, b = b, a + b
return b
记忆要点:dp[i] = dp[i-1] + dp[i-2] 滚动数组后只需要 a, b 两个变量。dp[1] = 1, dp[2] = 2 是基础情况,循环从 i = 3 开始。
例题 2:打家劫舍
LC 198. House Robber。每间房有一定金额,相邻的两间不能同时偷,求最大金额。
状态:dp[i] = 偷到第 i 间房为止的最大金额(不一定偷第 i 间)。转移:dp[i] = max(dp[i-1], dp[i-2] + nums[i])。
1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
class Solution {
public:
int rob(vector<int>& nums) {
int n = (int)nums.size();
if (n == 0) return 0;
if (n == 1) return nums[0];
int a = nums[0];
int b = std::max(nums[0], nums[1]);
for (int i = 2; i < n; ++i) {
int c = std::max(b, a + nums[i]);
a = b;
b = c;
}
return b;
}
};
1
2
3
4
5
6
7
8
9
10
11
class Solution:
def rob(self, nums: List[int]) -> int:
n = len(nums)
if n == 0:
return 0
if n == 1:
return nums[0]
a, b = nums[0], max(nums[0], nums[1])
for i in range(2, n):
a, b = b, max(b, a + nums[i])
return b
记忆要点:滚动数组的关键是 a, b = b, max(b, a + nums[i])——Python 元组赋值保证先算右边再赋值,不会出现”一边更新一边引用”的问题。
0/1 背包
0/1 背包问题:$N$ 件物品,第 $i$ 件重 $w[i]$、价值 $v[i]$,背包容量 $W$。每件物品最多选一次,求最大价值。
模板(压缩到一维):dp[j] = 容量为 $j$ 的背包能装的最大价值。倒序遍历 $j$ 从 $W$ 到 $w[i]$,转移 dp[j] = max(dp[j], dp[j - w[i]] + v[i])。
例题:分割等和子集
LC 416. Partition Equal Subset Sum。判断数组能否分成两个子集,使它们的和相等。
变形:能否从数组中选出若干个数,使得它们的和恰好为 total / 2。这就是 0/1 背包(物品价值 = 物品重量 = nums[i],背包容量 = total/2)。
1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
class Solution {
public:
bool canPartition(vector<int>& nums) {
int total = std::accumulate(nums.begin(), nums.end(), 0);
if (total % 2 != 0) return false;
int W = total / 2;
std::vector<int> dp(W + 1, 0);
for (int num : nums) {
for (int j = W; j >= num; --j) {
dp[j] = std::max(dp[j], dp[j - num] + num);
}
}
return dp[W] == W;
}
};
1
2
3
4
5
6
7
8
9
10
11
12
13
from typing import List
class Solution:
def canPartition(self, nums: List[int]) -> bool:
total = sum(nums)
if total % 2 != 0:
return False
W = total // 2
dp = [0] * (W + 1)
for num in nums:
for j in range(W, num - 1, -1):
dp[j] = max(dp[j], dp[j - num] + num)
return dp[W] == W
记忆要点:
- 0/1 背包内层 j 必须倒序遍历:因为每件物品只能用一次,正序遍历会让
dp[j-num]已经是”用了当前物品”的值,导致重复使用。 - 完全背包内层 j 必须正序遍历:因为每件物品可以用无限次,正序遍历能让同一个物品被多次使用。
- 状态压缩到一维:原本是
dp[i][j](前 i 件物品、容量 j),压缩后只需dp[j],因为第 i 件的处理逻辑只依赖第 i-1 件的dp值。
完全背包
完全背包问题:每件物品可以选无限次。
例题:零钱兑换
LC 322. Coin Change。给定不同面额的硬币和总金额,求凑出该金额所需的最少硬币数;不能则返回 -1。
1
2
3
4
5
6
7
8
9
10
11
12
13
14
class Solution {
public:
int coinChange(vector<int>& coins, int amount) {
const int INF = amount + 1;
std::vector<int> dp(amount + 1, INF);
dp[0] = 0;
for (int coin : coins) {
for (int j = coin; j <= amount; ++j) {
dp[j] = std::min(dp[j], dp[j - coin] + 1);
}
}
return dp[amount] == INF ? -1 : dp[amount];
}
};
1
2
3
4
5
6
7
8
9
class Solution:
def coinChange(self, coins: List[int], amount: int) -> int:
INF = amount + 1
dp = [INF] * (amount + 1)
dp[0] = 0
for coin in coins:
for j in range(coin, amount + 1):
dp[j] = min(dp[j], dp[j - coin] + 1)
return -1 if dp[amount] == INF else dp[amount]
记忆要点:
- 完全背包”最值”问题:内层 j 正序遍历,每件物品可以用无限次。
- 完全背包”组合数”问题:要先遍历物品、再遍历容量(
for coin: for j);如果是”排列数”,要先遍历容量、再遍历物品。 - 这题
dp[j]初始化为INF = amount + 1(凑不出比 amount+1 还多硬币),最后判断dp[amount] == INF即可。
二维 DP
二维 DP 的状态用两个下标描述,转移需要看”当前格子”和”周围的格子”。最常见的是网格路径和字符串编辑距离。
例题 1:不同路径
LC 62. Unique Paths。机器人从网格左上角到右下角,每次只能向右或向下走,求路径数。
1
2
3
4
5
6
7
8
9
10
11
12
class Solution {
public:
int uniquePaths(int m, int n) {
std::vector<int> dp(n, 1);
for (int i = 1; i < m; ++i) {
for (int j = 1; j < n; ++j) {
dp[j] += dp[j - 1];
}
}
return dp[n - 1];
}
};
1
2
3
4
5
6
7
class Solution:
def uniquePaths(self, m: int, n: int) -> int:
dp = [1] * n
for i in range(1, m):
for j in range(1, n):
dp[j] += dp[j - 1]
return dp[n - 1]
记忆要点:二维 DP 经常可以压缩到一维——dp[j] 不断被更新,最终一行结束时 dp[n-1] 就是答案。
例题 2:最长公共子序列
LC 1143. Longest Common Subsequence。求两个字符串的最长公共子序列长度。
1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
class Solution {
public:
int longestCommonSubsequence(string s, string t) {
int m = (int)s.size(), n = (int)t.size();
std::vector<std::vector<int>> dp(m + 1,
std::vector<int>(n + 1, 0));
for (int i = 1; i <= m; ++i) {
for (int j = 1; j <= n; ++j) {
if (s[i - 1] == t[j - 1]) {
dp[i][j] = dp[i - 1][j - 1] + 1;
} else {
dp[i][j] = std::max(
dp[i - 1][j],
dp[i][j - 1]
);
}
}
}
return dp[m][n];
}
};
1
2
3
4
5
6
7
8
9
10
11
class Solution:
def longestCommonSubsequence(self, s: str, t: str) -> int:
m, n = len(s), len(t)
dp = [[0] * (n + 1) for _ in range(m + 1)]
for i in range(1, m + 1):
for j in range(1, n + 1):
if s[i - 1] == t[j - 1]:
dp[i][j] = dp[i - 1][j - 1] + 1
else:
dp[i][j] = max(dp[i - 1][j], dp[i][j - 1])
return dp[m][n]
记忆要点:
- 字符串类 DP 的状态几乎都是
dp[i][j]=s[0..i-1]和t[0..j-1]的某种”度量”,边界dp[0][j] = dp[i][0] = 0。 - “字符相等” → 取左上角 + 1,”字符不等” → 取左/上中的较大值。
例题 3:编辑距离
LC 72. Edit Distance。给定两个单词 word1 和 word2,返回将 word1 转换为 word2 所使用的最少操作数(插入、删除、替换,每次一个字符)。
1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
22
23
24
class Solution {
public:
int minDistance(string word1, string word2) {
int m = (int)word1.size(), n = (int)word2.size();
std::vector<std::vector<int>> dp(m + 1,
std::vector<int>(n + 1, 0));
for (int i = 0; i <= m; ++i) dp[i][0] = i;
for (int j = 0; j <= n; ++j) dp[0][j] = j;
for (int i = 1; i <= m; ++i) {
for (int j = 1; j <= n; ++j) {
if (word1[i - 1] == word2[j - 1]) {
dp[i][j] = dp[i - 1][j - 1];
} else {
dp[i][j] = std::min({
dp[i - 1][j] + 1, // 删除 word1[i-1]
dp[i][j - 1] + 1, // 插入 word2[j-1]
dp[i - 1][j - 1] + 1 // 替换
});
}
}
}
return dp[m][n];
}
};
1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
class Solution:
def minDistance(self, word1: str, word2: str) -> int:
m, n = len(word1), len(word2)
dp = [[0] * (n + 1) for _ in range(m + 1)]
for i in range(m + 1):
dp[i][0] = i
for j in range(n + 1):
dp[0][j] = j
for i in range(1, m + 1):
for j in range(1, n + 1):
if word1[i - 1] == word2[j - 1]:
dp[i][j] = dp[i - 1][j - 1]
else:
dp[i][j] = min(
dp[i - 1][j] + 1, # 删除
dp[i][j - 1] + 1, # 插入
dp[i - 1][j - 1] + 1 # 替换
)
return dp[m][n]
记忆要点:
- 初始化:
dp[i][0] = i(删 i 个字符),dp[0][j] = j(插 j 个字符)。 - 三个操作的转移方向:删除看上、插入看左、替换看左上。
- “字符相等”时直接取
dp[i-1][j-1],因为不需要任何操作。
子序列 DP - LIS
例题:最长上升子序列
LC 300. Longest Increasing Subsequence。求数组中最长严格递增子序列的长度。
1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
class Solution {
public:
int lengthOfLIS(vector<int>& nums) {
int n = (int)nums.size();
std::vector<int> dp(n, 1);
int ans = 1;
for (int i = 1; i < n; ++i) {
for (int j = 0; j < i; ++j) {
if (nums[j] < nums[i]) {
dp[i] = std::max(dp[i], dp[j] + 1);
}
}
ans = std::max(ans, dp[i]);
}
return ans;
}
};
1
2
3
4
5
6
7
8
9
10
11
class Solution:
def lengthOfLIS(self, nums: List[int]) -> int:
n = len(nums)
dp = [1] * n
ans = 1
for i in range(1, n):
for j in range(i):
if nums[j] < nums[i]:
dp[i] = max(dp[i], dp[j] + 1)
ans = max(ans, dp[i])
return ans
记忆要点:标准 $O(n^2)$ LIS 模板——dp[i] 表示以 nums[i] 结尾的最长递增子序列长度。如果需要 $O(n \log n)$,用”贪心 + 二分”维护一个 tail 数组(bisect_left 在 Python 中)。
状态机 DP - 股票系列
股票系列是一类经典的状态机 DP:每天有”持有股票”和”不持有股票”两种状态,转移由”买/卖/不动”三种动作构成。
例题 1:买卖股票的最佳时机
LC 121. Best Time to Buy and Sell Stock。只能买卖一次,求最大利润。
状态:dp_no_stock = 不持有股票的最大利润;dp_have_stock = 持有股票的最大利润。
1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
class Solution {
public:
int maxProfit(vector<int>& prices) {
int n = (int)prices.size();
int dp_no_stock = 0;
int dp_have_stock = -prices[0];
for (int i = 1; i < n; ++i) {
int new_no_stock = std::max(dp_no_stock,
dp_have_stock + prices[i]);
int new_have_stock = std::max(dp_have_stock,
dp_no_stock - prices[i]);
dp_no_stock = new_no_stock;
dp_have_stock = new_have_stock;
}
return dp_no_stock;
}
};
1
2
3
4
5
6
7
8
9
10
11
class Solution:
def maxProfit(self, prices: List[int]) -> int:
n = len(prices)
dp_no_stock = 0
dp_have_stock = -prices[0]
for i in range(1, n):
new_no_stock = max(dp_no_stock, dp_have_stock + prices[i])
new_have_stock = max(dp_have_stock, dp_no_stock - prices[i])
dp_no_stock = new_no_stock
dp_have_stock = new_have_stock
return dp_no_stock
记忆要点:
- 这题也可以用”维护至今最低价”的贪心 O(n) 解法——但 DP 写法更容易推广到多次交易 / 冷冻期 / 手续费等变种。
- 状态机 DP 的关键是:先算
new_no_stock和new_have_stock,再赋值——不能边算边覆盖,否则会出现”同一天既买又卖”的错误。
例题 2:含冷冻期的股票买卖
LC 309. Best Time to Buy and Sell Stock with Cooldown。卖出后第二天不能买入。
1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
class Solution {
public:
int maxProfit(vector<int>& prices) {
int n = (int)prices.size();
if (n == 0) return 0;
std::vector<int> hold(n, 0), free(n, 0);
hold[0] = -prices[0];
free[0] = 0;
for (int i = 1; i < n; ++i) {
int prev_free = (i >= 2) ? free[i - 2] : 0;
hold[i] = std::max(hold[i - 1], prev_free - prices[i]);
free[i] = std::max(free[i - 1], hold[i - 1] + prices[i]);
}
return free[n - 1];
}
};
1
2
3
4
5
6
7
8
9
10
11
12
class Solution:
def maxProfit(self, prices: List[int]) -> int:
n = len(prices)
if n == 0:
return 0
hold = [-prices[0]] + [0] * (n - 1)
free = [0] * n
for i in range(1, n):
prev_free = free[i - 2] if i >= 2 else 0
hold[i] = max(hold[i - 1], prev_free - prices[i])
free[i] = max(free[i - 1], hold[i - 1] + prices[i])
return free[n - 1]
记忆要点:
hold[i] = max(hold[i-1], prev_free - prices[i]):不操作 / 从”两天前的 free”买入。free[i] = max(free[i-1], hold[i-1] + prices[i]):不操作 / 卖出昨天的持股。- “两天前”是关键——冷冻期让”昨天刚卖出”的状态不能立刻买入。
区间 DP
区间 DP 的状态是 dp[i][j] = 区间 [i, j] 上的最优解,转移通常枚举分割点 k:dp[i][j] = min(dp[i][k] + dp[k+1][j]) + cost(i, j)。
例题:最长回文子序列
LC 516. Longest Palindromic Subsequence。求字符串的最长回文子序列长度。
1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
class Solution {
public:
int longestPalindromeSubseq(string s) {
int n = (int)s.size();
std::vector<std::vector<int>> dp(n, std::vector<int>(n, 0));
for (int i = 0; i < n; ++i) dp[i][i] = 1;
for (int len = 2; len <= n; ++len) {
for (int i = 0; i + len <= n; ++i) {
int j = i + len - 1;
if (s[i] == s[j]) {
dp[i][j] = (len == 2) ? 2 : dp[i + 1][j - 1] + 2;
} else {
dp[i][j] = std::max(dp[i + 1][j], dp[i][j - 1]);
}
}
}
return dp[0][n - 1];
}
};
1
2
3
4
5
6
7
8
9
10
11
12
13
14
class Solution:
def longestPalindromeSubseq(self, s: str) -> int:
n = len(s)
dp = [[0] * n for _ in range(n)]
for i in range(n):
dp[i][i] = 1
for length in range(2, n + 1):
for i in range(n - length + 1):
j = i + length - 1
if s[i] == s[j]:
dp[i][j] = 2 if length == 2 else dp[i + 1][j - 1] + 2
else:
dp[i][j] = max(dp[i + 1][j], dp[i][j - 1])
return dp[0][n - 1]
记忆要点:
- 区间 DP 通常按”区间长度”遍历(
len = 2, 3, ..., n),保证dp[i+1][j-1]已经算好。 - “两端字符相等” → 取内层 + 2(长度为 2 时直接是 2,因为内层是空串)。
- “两端字符不等” → 取去掉任一端的最大值。
记忆要点(DP 全章):
- 状态定义先行:写不出转移方程,先想想
dp[i](或dp[i][j])到底表示什么。 - 滚动数组的边界:滚动时要注意
dp[i-2]是不是已经被覆盖(冷冻期那题就要特判i >= 2)。 - 背包内层 j 的方向:0/1 背包倒序、完全背包正序——最容易记错的一条。
- 二维 DP 可以压成一维:当转移只依赖”上一行”或”左、上、左上”时,可以滚动数组到一维。
- 区间 DP 按长度遍历:从小区间推到大区间,避免顺序问题。
- 状态机 DP 先算 new 再赋值:避免”同一天既买又卖”导致状态错乱。