LeetCode热题100【Python】
ACM模式
import sys
def solve():
data = sys.stdin.read().strip().split()
# 所有数据都在data列表中
# 根据题目要求解析...
# 例如:第一个数是n,后面跟着n个数
if data:
n = int(data[0])
nums = list(map(int, data[1:n+1]))
# 处理...
if __name__ == "__main__":
solve()
哈希
1. 两数之和
输入:nums = [2,7,11,15], target = 9
输出:[0,1]
解释:因为 nums[0] + nums[1] == 9 ,返回 [0, 1] 。
思路:遍历数组,用哈希表存储{数值: 下标}。对于当前数num,检查target - num是否已在哈希表中,如果在就直接返回两个下标,否则将当前数加入哈希表。时间复杂度O(n)。
def twoSum(nums, target):
seen = {}
for i, num in enumerate(nums):
complement = target - num
if complement in seen:
return [seen[complement], i]
seen[num] = i
return []
49. 字母异位词分组
输入: strs = [“eat”, “tea”, “tan”, “ate”, “nat”, “bat”]
输出: [[“bat”],[“nat”,“tan”],[“ate”,“eat”,“tea”]]
思路:异位词排序后会得到相同的字符串。遍历所有单词,将每个单词排序后作为key,存入哈希表(value是异位词列表)。最后返回哈希表的所有value即可。时间复杂度O(n·klogk),k为单词最大长度。
def groupAnagrams(strs):
from collections import defaultdict
groups = defaultdict(list)
for s in strs:
# 将字符串排序后作为key
key = ''.join(sorted(s))
groups[key].append(s)
return list(groups.values())
def groupAnagrams(strs):
from collections import defaultdict
groups = defaultdict(list)
for s in strs:
# 用26个字母的计数作为key
count = [0] * 26
for c in s:
count[ord(c) - ord('a')] += 1
groups[tuple(count)].append(s)
return list(groups.values())
128. 最长连续序列
输入:nums = [100,4,200,1,3,2]
输出:4
解释:最长数字连续序列是 [1, 2, 3, 4]。它的长度为 4。
思路:先用set去重。遍历每个数,只从序列的起点开始计算(即num-1不在set中时)。然后不断找num+1、num+2…统计当前连续序列的长度,更新最大值。这样每个数只会访问一次,时间复杂度O(n)。
def longestConsecutive(nums):
if not nums:
return 0
num_set = set(nums)
max_len = 0
for num in num_set:
# 只从序列的起点开始计算
if num - 1 not in num_set:
curr_num = num
curr_len = 1
while curr_num + 1 in num_set:
curr_num += 1
curr_len += 1
max_len = max(max_len, curr_len)
return max_len
双指针
两个升序数组合并
算法复杂度
时间复杂度:O (m + n) (只遍历一遍)
空间复杂度:O (m + n) (存储结果)
这是最优解法,没有比这更快的了。
def merge_two_sorted_arrays(nums1, nums2):
i = j = 0 # 双指针
res = [] # 结果数组
# 同时遍历两个数组,谁小取谁
while i < len(nums1) and j < len(nums2):
if nums1[i] < nums2[j]:
res.append(nums1[i])
i += 1
else:
res.append(nums2[j])
j += 1
# 处理剩余元素(只会剩下一个数组)
res += nums1[i:]
res += nums2[j:]
return res
283. 移动零
输入: nums = [0,1,0,3,12]
输出: [1,3,12,0,0]
思路:快慢指针。j指向第一个0的位置(即下一个非零元素应该放的位置)。遍历数组,遇到非零元素就与j位置交换,然后j后移。这样所有非零元素都被移到前面,零自然被挤到后面。
def moveZeroes(nums):
# 双指针:j指向第一个0的位置
j = 0
for i in range(len(nums)):
if nums[i] != 0:
nums[i], nums[j] = nums[j], nums[i]
j += 1
return nums
11. 盛最多水的容器
输入:[1,8,6,2,5,4,8,3,7]
输出:49
解释:在此情况下,容器能够容纳水(表示为蓝色部分)的最大值为 49。min(8,7)*(9-2)=7*7=49
思路:左右指针从两端向中间移动。面积由较短的边决定,所以每次移动较短的边才有可能找到更大的面积(因为宽度在减小,只有高度可能变大)。时间复杂度O(n)。
def maxArea(height):
left, right = 0, len(height) - 1
max_water = 0
while left < right:
# 计算当前面积
h = min(height[left], height[right])
w = right - left
max_water = max(max_water, h * w)
# 移动较短的边
if height[left] < height[right]:
left += 1
else:
right -= 1
return max_water
15. 三数之和
输入:nums = [-1,0,1,2,-1,-4]
输出:[[-1,-1,2],[-1,0,1]]
解释:
nums[0] + nums[1] + nums[2] = (-1) + 0 + 1 = 0 。
nums[1] + nums[2] + nums[4] = 0 + 1 + (-1) = 0 。
nums[0] + nums[3] + nums[4] = (-1) + 2 + (-1) = 0 。
不同的三元组是 [-1,0,1] 和 [-1,-1,2] 。
注意,输出的顺序和三元组的顺序并不重要。
思路:排序+双指针。固定第一个数nums[i],然后用双指针在i+1到末尾之间找两数之和等于-nums[i]。关键点是去重:①第一个数去重 ②找到解后对左右指针去重。时间复杂度O(n²)。
def threeSum(nums):
nums.sort()
res = []
n = len(nums)
for i in range(n - 2):
# 去重:跳过相同的第一个数
if i > 0 and nums[i] == nums[i-1]:
continue
left, right = i + 1, n - 1
target = -nums[i]
while left < right:
curr_sum = nums[left] + nums[right]
if curr_sum == target:
res.append([nums[i], nums[left], nums[right]])
# 去重
while left < right and nums[left] == nums[left + 1]:
left += 1
while left < right and nums[right] == nums[right - 1]:
right -= 1
left += 1
right -= 1
elif curr_sum < target:
left += 1
else:
right -= 1
return res
42. 接雨水
输入:height = [0,1,0,2,1,0,1,3,2,1,2,1]
输出:6
解释:在这种情况下,可以接 6 个单位的雨水(蓝色部分表示雨水)。
思路:双指针记录左右两边的最大高度。对于每个位置,能接的雨水量取决于它左右两边最大高度的较小值减去当前高度。通过比较height[left]和height[right],可以确定哪边的最大高度是确定的,从而计算该位置的雨水量并移动指针。
def trap(height):
if not height:
return 0
left, right = 0, len(height) - 1
left_max = right_max = water = 0
while left < right:
if height[left] < height[right]:
if height[left] >= left_max:
left_max = height[left]
else:
water += left_max - height[left]
left += 1
else:
if height[right] >= right_max:
right_max = height[right]
else:
water += right_max - height[right]
right -= 1
return water
滑动窗口
3. 无重复字符的最长子串
输入: s = “abcabcbb”
输出: 3
解释: 因为无重复字符的最长子串是 “abc”,所以其长度为 3。注意 “bca” 和 “cab” 也是正确答案。
思路:滑动窗口+哈希表。用left指向窗口左边界,right遍历字符串。used哈希表记录每个字符最近出现的位置。当遇到重复字符且该字符在窗口内(即used[c] >= left)时,将left移动到重复字符的下一个位置。每次更新最长长度。时间复杂度O(n)。
def lengthOfLongestSubstring(s):
left = max_len = 0
used = {}
for right, c in enumerate(s):
if c in used and used[c] >= left:
left = used[c] + 1
used[c] = right
max_len = max(max_len, right - left + 1)
return max_len
438. 找到字符串中所有的字母异位词
输入: s = “cbaebabacd”, p = “abc”
输出: [0,6]
解释:
起始索引等于 0 的子串是 “cba”, 它是 “abc” 的异位词。
起始索引等于 6 的子串是 “bac”, 它是 “abc” 的异位词。
思路:固定大小的滑动窗口+字符计数数组。用两个长度为26的数组分别统计p和窗口内字符出现次数。先初始化第一个窗口,然后每次移动窗口时:加入一个新字符,移除一个旧字符。比较两个计数数组是否相等,相等则记录起始下标。时间复杂度O(n)。
def findAnagrams(s, p):
if len(p) > len(s):
return []
p_count = [0] * 26
s_count = [0] * 26
# 初始化
for i in range(len(p)):
p_count[ord(p[i]) - ord('a')] += 1
s_count[ord(s[i]) - ord('a')] += 1
res = []
if s_count == p_count:
res.append(0)
for i in range(len(p), len(s)):
# 滑动窗口
s_count[ord(s[i]) - ord('a')] += 1
s_count[ord(s[i - len(p)]) - ord('a')] -= 1
if s_count == p_count:
res.append(i - len(p) + 1)
return res
思路:滑动窗口维护与p等长的子串,用Counter统计字符频率,当窗口计数与p计数相等时记录起始位置。
def findAnagrams(s, p):
from collections import Counter
len_p, len_s = len(p), len(s)
if len_p > len_s:
return []
# 统计p的字符频率
p_count = Counter(p)
window_count = Counter()
res = []
for i in range(len_s):
# 加入当前字符
window_count[s[i]] += 1
# 窗口大小超过p的长度时,移除左边字符
if i >= len_p:
if window_count[s[i - len_p]] == 1:
del window_count[s[i - len_p]]
else:
window_count[s[i - len_p]] -= 1
# 比较窗口和p的字符频率
if window_count == p_count:
res.append(i - len_p + 1)
return res
子串
14. 最长公共前缀
输入:strs = [“flower”,“flow”,“flight”]
输出:“fl”
def longestCommonPrefix(strs):
if not strs:
return ""
# 取第一个字符串作为基准
prefix = strs[0]
# 依次与每个字符串比较
for s in strs[1:]:
# 不断缩短prefix直到是s的前缀
while s.find(prefix) != 0:
prefix = prefix[:-1]
if not prefix:
return ""
return prefix
560. 和为K的子数组
输入:nums = [1,2,3], k = 3
输出:2 len([[1,2],[3]])
思路:前缀和+哈希表。遍历数组计算前缀和curr_sum,用哈希表记录每个前缀和出现的次数。对于当前curr_sum,查找之前有多少个前缀和等于curr_sum - k,这些位置到当前索引的子数组和就是k。注意初始化{0:1}表示空前缀。
def subarraySum(nums, k):
# 前缀和 + 哈希表
prefix_sum = {0: 1} # 前缀和 -> 出现次数
curr_sum = count = 0
for num in nums:
curr_sum += num
# 查找是否存在前缀和 = curr_sum - k
if curr_sum - k in prefix_sum:
count += prefix_sum[curr_sum - k]
# 更新当前前缀和的出现次数
prefix_sum[curr_sum] = prefix_sum.get(curr_sum, 0) + 1
return count
239. 滑动窗口最大值
输入:nums = [1,3,-1,-3,5,3,6,7], k = 3
输出:[3,3,5,5,6,7]
解释:
以 nums = [1,3,-1,-3,5,3,6,7], k = 3 为例:
| i | n | 队列变化(存下标) | 说明 | 窗口 | 最大值 |
|---|---|---|---|---|---|
| 0 | 1 | [0] | 队列空,直接加入 | - | - |
| 1 | 3 | [1] | 3>1,弹出0,加入1 | - | - |
| 2 | -1 | [1,2] | -1<3,直接加入 | [1,3,-1] | 3 |
| 3 | -3 | [1,2,3] | -3<-1,直接加入 | [3,-1,-3] | 3 |
| 4 | 5 | [4] | 5>所有,清空后加入 | [-1,-3,5] | 5 |
| 5 | 3 | [4,5] | 3<5,直接加入 | [-3,5,3] | 5 |
| 6 | 6 | [6] | 6>所有,清空后加入 | [5,3,6] | 6 |
| 7 | 7 | [7] | 7>6,弹出6加入7 | [3,6,7] | 7 |
本题难点: 如何在每次窗口滑动后,将 “获取窗口内最大值” 的时间复杂度从 O(k) 降低至 O(1) 。
为什么用双端队列?
队尾:需要频繁弹出较小的元素(pop())
队首:需要移除滑出窗口的元素(popleft())
为什么存下标不存值?
需要判断元素是否还在窗口内(通过下标差)
通过下标可以随时获取对应的值
思路:单调队列。维护一个双端队列,始终保持队首是当前窗口最大值的下标。遍历数组时:①从队尾移除所有比当前元素小的值 ②将当前元素下标入队 ③如果队首已滑出窗口则移除 ④当窗口形成后,队首就是最大值。
def maxSlidingWindow(nums, k):
dq = collections.deque()
res = []
for i, n in enumerate(nums):
while dq and nums[dq[-1]] < n:
dq.pop()
dq.append(i)
if dq[0] == i - k:
dq.popleft()
if i >= k - 1:
res.append(nums[dq[0]])
return res
76. 最小覆盖字串
输入:s = “ADOBECODEBANC”, t = “ABC”
输出:“BANC”
解释:最小覆盖子串 “BANC” 包含来自字符串 t 的 ‘A’、‘B’ 和 ‘C’。
思路:滑动窗口+计数。用need字典记录t中字符的需求量,missing表示还缺多少个字符。右指针扩展窗口直到包含所有t中字符,然后收缩左指针移除多余字符,记录最小窗口。之后左指针右移一位打破平衡,继续寻找下一个满足条件的窗口。
def minWindow(s, t):
from collections import Counter
if not s or not t or len(s) < len(t):
return ""
need = Counter(t) # 需要凑齐的字符及数量
missing = len(t) # 还缺多少个字符
left = start = end = 0
for right, char in enumerate(s, 1):
# right从1开始,方便计算长度
# 如果当前字符是需要的,missing减少
if need[char] > 0:
missing -= 1
need[char] -= 1
# 当窗口包含所有需要的字符时
if missing == 0:
# 收缩左边界,移除多余的字符
while left < right and need[s[left]] < 0:
need[s[left]] += 1
left += 1
# 更新最小窗口
if end == 0 or right - left < end - start:
start, end = left, right
# 移动左边界,打破满足条件的状态,继续寻找下一个窗口
need[s[left]] += 1
missing += 1
left += 1
return s[start:end]
普通数组
53. 最大子数组和
输入:nums = [-2,1,-3,4,-1,2,1,-5,4]
输出:6
解释:连续子数组 [4,-1,2,1] 的和最大,为 6 。
思路:遍历数组时,对于每个位置,计算以当前元素结尾的最大子数组和:
- curr_sum表示以当前元素结尾的子数组最大和
- 要么只取当前元素(重新开始),要么加上之前的连续子数组
- 状态转移方程:curr_sum = max(num, curr_sum + num)
用max_sum记录遍历过程中出现的最大值,最后返回即可。
def maxSubArray(nums):
# 动态规划,Kadane算法
curr_sum = max_sum = nums[0]
for num in nums[1:]:
curr_sum = max(num, curr_sum + num)
max_sum = max(max_sum, curr_sum)
return max_sum
56. 合并区间
输入:intervals = [[1,3],[2,6],[8,10],[15,18]]
输出:[[1,6],[8,10],[15,18]]
解释:区间 [1,3] 和 [2,6] 重叠, 将它们合并为 [1,6].
思路:排序+贪心。先按区间起点排序,然后遍历:如果当前区间起点 ≤ 上一个合并区间的终点,说明重叠,更新上一个区间的终点为两者最大值;否则不重叠,直接加入结果。
def merge(intervals):
if not intervals:
return []
# 按区间起点排序
intervals.sort(key=lambda x: x[0])
merged = [intervals[0]]
for interval in intervals[1:]:
# 如果有重叠,合并
if interval[0] <= merged[-1][1]:
merged[-1][1] = max(merged[-1][1], interval[1])
else:
merged.append(interval)
return merged
189. 轮转数组
输入: nums = [1,2,3,4,5,6,7], k = 3
输出: [5,6,7,1,2,3,4]
思路:三次翻转。先整体翻转,再翻转前k个,最后翻转剩余部分。例如[1,2,3,4,5,6,7], k=3:整体→[7,6,5,4,3,2,1],前3个→[5,6,7,4,3,2,1],剩余→[5,6,7,1,2,3,4]。时间复杂度O(n),空间O(1)。
def rotate(nums, k):
"""
Do not return anything, modify nums in-place instead.
"""
n = len(nums)
k %= n # 处理k大于n的情况
# 方法1:三次翻转(最优)
def reverse(start, end):
while start < end:
nums[start], nums[end] = nums[end], nums[start]
start += 1
end -= 1
reverse(0, n - 1) # 整体翻转
reverse(0, k - 1) # 翻转前k个
reverse(k, n - 1) # 翻转剩余部分
238. 除了自身之外数组的乘积
输入: nums = [1,2,3,4]
输出: [24,12,8,6]
思路:左右乘积列表。第一遍遍历计算每个位置左边所有数的乘积存入result;第二遍从右向左遍历,用right_product记录右边乘积,乘到result对应位置。这样每个位置的result就是左边乘积×右边乘积,且不使用除法。
def productExceptSelf(nums):
n = len(nums)
result = [1] * n
# 计算左边乘积
left_product = 1
for i in range(n):
result[i] = left_product
left_product *= nums[i]
# 乘上右边乘积
right_product = 1
for i in range(n - 1, -1, -1):
result[i] *= right_product
right_product *= nums[i]
return result
41. 缺失的第一个正数
输入:nums = [3,4,-1,1]
输出:2
解释:1 在数组中,但 2 没有。
思路:原地哈希(索引标记)。长度为n的数组,答案一定在[1, n+1]范围内。第一遍将所有不在[1,n]的数标记为n+1;第二遍遍历,将出现过的正数对应的索引位置的值标记为负数;第三遍找到第一个正数索引,其索引+1就是答案。利用数组本身作为哈希表,空间O(1)。
这个算法的精妙之处在于利用数组本身作为哈希表,通过正负号来标记某个数字是否出现过,既节省了空间,又保持了O(n)的时间复杂度。
让我们一步步分析:
初始状态 nums = [3, 4, -1, 1], n = 4
第一步:预处理
将所有不在 [1, 4] 范围内的数标记为 n+1 = 5
nums = [3, 4, 5, 1]
第二步:标记出现过的数
遍历数组,用绝对值作为索引,将对应位置标记为负数
i=0: val = |3| = 3,3 在 [1,4] 范围内,索引 3-1 = 2 的值是 5 > 0,将 nums[2] 变成 -5,nums = [3, 4, -5, 1]
i=1: val = |4| = 4,4 在 [1,4] 范围内,索引 4-1 = 3 的值是 1 > 0,将 nums[3] 变成 -1,nums = [3, 4, -5, -1]
i=2: val = |-5| = 5,5 不在 [1,4] 范围内,跳过
i=3: val = |-1| = 1,1 在 [1,4] 范围内,索引 1-1 = 0 的值是 3 > 0,将 nums[0] 变成 -3,nums = [-3, 4, -5, -1]
第三步:查找第一个正数
遍历数组找第一个 > 0 的数:
索引0: -3 < 0
索引1: 4 > 0 ✅ 找到第一个正数
第一个正数的索引是 1,所以返回 1 + 1 = 2
最终结果:缺失的最小正数是 2
def firstMissingPositive(nums):
n = len(nums)
# 将不在[1, n]范围内的数标记为n+1
for i in range(n):
if nums[i] <= 0 or nums[i] > n:
nums[i] = n + 1
# 将出现的正数对应的索引位置标记为负数
for i in range(n):
val = abs(nums[i])
if 1 <= val <= n:
if nums[val - 1] > 0:
nums[val - 1] = -nums[val - 1]
# 第一个正数的索引+1就是缺失的最小正数
for i in range(n):
if nums[i] > 0:
return i + 1
return n + 1
矩阵
矩阵乘法
def matrix_mult(a, b):
# 获取矩阵行列数
m = len(a) # a 的行数
n = len(b) # b 的行数 = a 的列数
p = len(b[0]) # b 的列数
# 初始化结果矩阵(全 0)
result = [[0 for _ in range(p)] for _ in range(m)]
# 三重循环计算
for i in range(m):
for j in range(p):
for k in range(n):
result[i][j] += a[i][k] * b[k][j]
return result
73. 矩阵置零
输入:matrix = [[0,1,2,0],[3,4,5,2],[1,3,1,5]]
输出:[[0,0,0,0],[0,4,5,0],[0,3,1,0]]
思路:用第一行和第一列作为标记位。先遍历矩阵,若某个元素为0,则将其所在行首和列首置0,并用两个布尔变量记录第一行/列本身是否含0。然后根据标记将对应行列置零,最后处理第一行/列。
def setZeroes(matrix):
m, n = len(matrix), len(matrix[0])
first_row = first_col = False
# 标记
for i in range(m):
for j in range(n):
if matrix[i][j] == 0:
if i == 0: first_row = True
if j == 0: first_col = True
matrix[i][0] = matrix[0][j] = 0
# 置零(除第一行第一列)
for i in range(1, m):
for j in range(1, n):
if matrix[i][0] == 0 or matrix[0][j] == 0:
matrix[i][j] = 0
# 处理第一行
if first_row:
for j in range(n):
matrix[0][j] = 0
# 处理第一列
if first_col:
for i in range(m):
matrix[i][0] = 0
54. 螺旋矩阵
输入:matrix = [[1,2,3,4],[5,6,7,8],[9,10,11,12]]
输出:[1,2,3,4,8,12,11,10,9,5,6,7]
思路:模拟上下左右四个边界。初始化top、bottom、left、right,按右→下→左→上顺序遍历,每走完一条边就收缩对应边界。注意在向左和向上前要检查边界是否合法。
def spiralOrder(matrix):
if not matrix or not matrix[0]:
return []
result = []
top, bottom = 0, len(matrix) - 1
left, right = 0, len(matrix[0]) - 1
while top <= bottom and left <= right:
# 从左到右
for j in range(left, right + 1):
result.append(matrix[top][j])
top += 1
# 从上到下
for i in range(top, bottom + 1):
result.append(matrix[i][right])
right -= 1
# 从右到左(需要检查是否还有行)
if top <= bottom:
for j in range(right, left - 1, -1):
result.append(matrix[bottom][j])
bottom -= 1
# 从下到上(需要检查是否还有列)
if left <= right:
for i in range(bottom, top - 1, -1):
result.append(matrix[i][left])
left += 1
return result
48. 旋转图像
输入:matrix = [[1,2,3],[4,5,6],[7,8,9]]
输出:[[7,4,1],[8,5,2],[9,6,3]]
思路:两种方法。①转置+每行反转:先沿主对角线转置,再水平翻转每一行。②直接旋转四个点:按层处理,每层循环交换四个对应位置的元素。
def rotate(matrix):
"""
Do not return anything, modify matrix in-place instead.
"""
n = len(matrix)
# 方法1:转置 + 每行反转
# 转置
for i in range(n):
for j in range(i + 1, n):
matrix[i][j], matrix[j][i] = matrix[j][i], matrix[i][j]
# 每行反转
for i in range(n):
matrix[i].reverse()
方法2:直接旋转四个点:
def rotate(matrix):
n = len(matrix)
for i in range(n // 2):
for j in range(i, n - i - 1):
# 保存左上角
temp = matrix[i][j]
# 左下 -> 左上
matrix[i][j] = matrix[n - 1 - j][i]
# 右下 -> 左下
matrix[n - 1 - j][i] = matrix[n - 1 - i][n - 1 - j]
# 右上 -> 右下
matrix[n - 1 - i][n - 1 - j] = matrix[j][n - 1 - i]
# 左上 -> 右上
matrix[j][n - 1 - i] = temp
240. 搜索二维矩阵 II
编写一个高效的算法来搜索 m x n 矩阵 matrix 中的一个目标值 target 。该矩阵具有以下特性:
- 每行的元素从左到右升序排列。
- 每列的元素从上到下升序排列。
输入:matrix = [[1,4,7,11,15],[2,5,8,12,19],[3,6,9,16,22],[10,13,14,17,24],[18,21,23,26,30]], target = 5
输出:true
思路:从右上角开始搜索。若当前值大于target,列左移(这一列下方都更大);若小于target,行下移(这一行左边都更小)。利用矩阵行列分别递增的特性,每次排除一行或一列。
def searchMatrix(matrix, target):
if not matrix or not matrix[0]:
return False
m, n = len(matrix), len(matrix[0])
# 从右上角开始
i, j = 0, n - 1
while i < m and j >= 0:
if matrix[i][j] == target:
return True
elif matrix[i][j] > target:
j -= 1 # 列左移
else:
i += 1 # 行下移
return False
链表
class ListNode:
def __init__(self, val=0, next=None):
self.val = val
self.next = next
| 题目 | 时间复杂度 | 空间复杂度 | 关键技巧 |
|---|---|---|---|
| 相交链表 | O(m+n) | O(1) | 双指针遍历 |
| 反转链表 | O(n) | O(1) | 迭代/递归 |
| 回文链表 | O(n) | O(1) | 快慢指针+反转 |
| 环形链表 | O(n) | O(1) | 快慢指针 |
| 环形链表 II | O(n) | O(1) | 快慢指针+数学 |
| 合并有序链表 | O(n+m) | O(1) | 虚拟头节点 |
| 两数相加 | O(max(m,n)) | O(1) | 进位处理 |
| 删除倒数第N个 | O(n) | O(1) | 快慢指针+虚拟头 |
| 两两交换 | O(n) | O(1) | 迭代/递归 |
| K个一组翻转 | O(n) | O(1) | 分组翻转 |
| 随机链表复制 | O(n) | O(1)/O(n) | 节点穿插/哈希表 |
| 排序链表 | O(n log n) | O(log n) | 归并排序 |
| 合并K个链表 | O(n log k) | O(k) | 堆/分治 |
| LRU缓存 | O(1) | O(capacity) | 哈希表+双向链表 |
链表问题通用技巧:
使用虚拟头节点(dummy)简化边界处理
快慢指针解决环、中点、倒数第N个等问题
160. 相交链表
def getIntersectionNode(headA, headB):
if not headA or not headB:
return None
pa, pb = headA, headB
while pa != pb:
pa = pa.next if pa else headB
pb = pb.next if pb else headA
return pa
206. 反转链表
def reverseList(head):
prev = None
curr = head
while curr:
next_temp = curr.next
curr.next = prev
prev = curr
curr = next_temp
return prev
递归版:
def reverseList(head):
if not head or not head.next:
return head
new_head = reverseList(head.next)
head.next.next = head
head.next = None
return new_head
234. 回文链表
def isPalindrome(head):
if not head or not head.next:
return True
# 找到中点
slow = fast = head
while fast and fast.next:
slow = slow.next
fast = fast.next.next
# 反转后半部分
prev = None
while slow:
next_temp = slow.next
slow.next = prev
prev = slow
slow = next_temp
# 比较前后半部分
left, right = head, prev
while right:
if left.val != right.val:
return False
left = left.next
right = right.next
return True
141. 环形链表
def hasCycle(head):
if not head or not head.next:
return False
slow = fast = head
while fast and fast.next:
slow = slow.next
fast = fast.next.next
if slow == fast:
return True
return False
142. 环形链表 II
def detectCycle(head):
if not head or not head.next:
return None
# 判断是否有环
slow = fast = head
has_cycle = False
while fast and fast.next:
slow = slow.next
fast = fast.next.next
if slow == fast:
has_cycle = True
break
if not has_cycle:
return None
# 找环的入口
slow = head
while slow != fast:
slow = slow.next
fast = fast.next
return slow
21. 合并两个有序链表
def mergeTwoLists(l1, l2):
dummy = ListNode(0)
curr = dummy
while l1 and l2:
if l1.val <= l2.val:
curr.next = l1
l1 = l1.next
else:
curr.next = l2
l2 = l2.next
curr = curr.next
curr.next = l1 or l2
return dummy.next
递归版:
def mergeTwoLists(l1, l2):
if not l1 or not l2:
return l1 or l2
if l1.val <= l2.val:
l1.next = mergeTwoLists(l1.next, l2)
return l1
else:
l2.next = mergeTwoLists(l1, l2.next)
return l2
2. 两数相加
def addTwoNumbers(l1, l2):
dummy = ListNode(0)
curr = dummy
carry = 0
while l1 or l2 or carry:
val1 = l1.val if l1 else 0
val2 = l2.val if l2 else 0
total = val1 + val2 + carry
carry = total // 10
curr.next = ListNode(total % 10)
curr = curr.next
if l1: l1 = l1.next
if l2: l2 = l2.next
return dummy.next
19. 删除链表的倒数第N个结点
def removeNthFromEnd(head, n):
dummy = ListNode(0)
dummy.next = head
fast = slow = dummy
# fast先走n+1步
for i in range(n + 1):
fast = fast.next
# 同时移动,fast到结尾时slow指向待删节点的前一个
while fast:
fast = fast.next
slow = slow.next
slow.next = slow.next.next
return dummy.next
24. 两两交换链表中的节点
def swapPairs(head):
dummy = ListNode(0)
dummy.next = head
prev = dummy
while prev.next and prev.next.next:
# 获取要交换的两个节点
first = prev.next
second = first.next
# 交换
first.next = second.next
second.next = first
prev.next = second
# 移动prev
prev = first
return dummy.next
递归版:
def swapPairs(head):
if not head or not head.next:
return head
new_head = head.next
head.next = swapPairs(new_head.next)
new_head.next = head
return new_head
25. K个一组翻转链表
def reverseKGroup(head, k):
dummy = ListNode(0)
dummy.next = head
prev = dummy
while True:
# 检查剩余节点是否够k个
tail = prev
for i in range(k):
tail = tail.next
if not tail:
return dummy.next
# 记录下一组的起始点
next_group = tail.next
# 翻转当前k个节点
curr = prev.next
prev_next = prev.next
while curr != next_group:
temp = curr.next
curr.next = prev.next
prev.next = curr
curr = temp
prev_next.next = next_group
prev = prev_next
138. 随机链表的复制
def copyRandomList(head):
if not head:
return None
# 第一步:在每个节点后面复制一个新节点
curr = head
while curr:
new_node = Node(curr.val)
new_node.next = curr.next
curr.next = new_node
curr = new_node.next
# 第二步:复制random指针
curr = head
while curr:
if curr.random:
curr.next.random = curr.random.next
curr = curr.next.next
# 第三步:分离两个链表
dummy = Node(0)
curr_new = dummy
curr = head
while curr:
curr_new.next = curr.next
curr.next = curr.next.next
curr_new = curr_new.next
curr = curr.next
return dummy.next
哈希表法:
def copyRandomList(head):
if not head:
return None
# 创建原节点到新节点的映射
node_map = {}
curr = head
while curr:
node_map[curr] = Node(curr.val)
curr = curr.next
# 连接next和random指针
curr = head
while curr:
if curr.next:
node_map[curr].next = node_map[curr.next]
if curr.random:
node_map[curr].random = node_map[curr.random]
curr = curr.next
return node_map[head]
148. 排序链表
def sortList(head):
if not head or not head.next:
return head
# 找到中点
slow, fast = head, head.next
while fast and fast.next:
slow = slow.next
fast = fast.next.next
# 分割链表
mid = slow.next
slow.next = None
# 递归排序
left = sortList(head)
right = sortList(mid)
# 合并
return merge(left, right)
def merge(l1, l2):
dummy = ListNode(0)
curr = dummy
while l1 and l2:
if l1.val <= l2.val:
curr.next = l1
l1 = l1.next
else:
curr.next = l2
l2 = l2.next
curr = curr.next
curr.next = l1 or l2
return dummy.next
23. 合并K个升序链表
def mergeKLists(lists):
if not lists:
return None
import heapq
# 使用最小堆
dummy = ListNode(0)
curr = dummy
heap = []
# 将每个链表的头节点加入堆
for i, node in enumerate(lists):
if node:
heapq.heappush(heap, (node.val, i, node))
while heap:
val, i, node = heapq.heappop(heap)
curr.next = node
curr = curr.next
if node.next:
heapq.heappush(heap, (node.next.val, i, node.next))
return dummy.next
146. LRU缓存
class LRUCache:
def __init__(self, capacity):
self.capacity = capacity
self.cache = {} # key -> node
self.head = Node(0, 0) # 虚拟头节点
self.tail = Node(0, 0) # 虚拟尾节点
self.head.next = self.tail
self.tail.prev = self.head
def get(self, key):
if key in self.cache:
node = self.cache[key]
self._remove(node)
self._add(node)
return node.val
return -1
def put(self, key, value):
if key in self.cache:
self._remove(self.cache[key])
node = Node(key, value)
self.cache[key] = node
self._add(node)
if len(self.cache) > self.capacity:
lru = self.head.next
self._remove(lru)
del self.cache[lru.key]
def _remove(self, node):
node.prev.next = node.next
node.next.prev = node.prev
def _add(self, node):
node.prev = self.tail.prev
node.next = self.tail
self.tail.prev.next = node
self.tail.prev = node
class Node:
def __init__(self, key, val):
self.key = key
self.val = val
self.prev = None
self.next = None
二叉树
class TreeNode:
def __init__(self, val=0, left=None, right=None):
self.val = val
self.left = left
self.right = right
94. 二叉树的中序遍历
# 递归
def inorderTraversal(root):
def inorder(node):
if not node:
return
inorder(node.left)
res.append(node.val)
inorder(node.right)
res = []
inorder(root)
return res
# 迭代(栈)
def inorderTraversal(root):
res, stack = [], []
curr = root
while curr or stack:
while curr:
stack.append(curr)
curr = curr.left
curr = stack.pop()
res.append(curr.val)
curr = curr.right
return res
104. 二叉树的最大深度
# 递归
def maxDepth(root):
if not root:
return 0
return max(maxDepth(root.left), maxDepth(root.right)) + 1
# 迭代(BFS)
def maxDepth(root):
if not root:
return 0
depth = 0
queue = [root]
while queue:
depth += 1
level_size = len(queue)
for _ in range(level_size):
node = queue.pop(0)
if node.left:
queue.append(node.left)
if node.right:
queue.append(node.right)
return depth
226. 翻转二叉树
# 递归
def invertTree(root):
if not root:
return None
# 交换左右子树
root.left, root.right = root.right, root.left
# 递归翻转
invertTree(root.left)
invertTree(root.right)
return root
# 迭代
def invertTree(root):
if not root:
return None
queue = [root]
while queue:
node = queue.pop(0)
# 交换
node.left, node.right = node.right, node.left
if node.left:
queue.append(node.left)
if node.right:
queue.append(node.right)
return root
101. 对称二叉树
def isSymmetric(root):
if not root:
return True
def is_mirror(left, right):
if not left and not right:
return True
if not left or not right:
return False
return (left.val == right.val and
is_mirror(left.left, right.right) and
is_mirror(left.right, right.left))
return is_mirror(root.left, root.right)
# 迭代
def isSymmetric(root):
if not root:
return True
queue = [(root.left, root.right)]
while queue:
left, right = queue.pop(0)
if not left and not right:
continue
if not left or not right:
return False
if left.val != right.val:
return False
queue.append((left.left, right.right))
queue.append((left.right, right.left))
return True
543. 二叉树的直径
def diameterOfBinaryTree(root):
def depth(node):
nonlocal diameter
if not node:
return 0
left_depth = depth(node.left)
right_depth = depth(node.right)
# 更新直径(经过当前节点的最长路径)
diameter = max(diameter, left_depth + right_depth)
return max(left_depth, right_depth) + 1
diameter = 0
depth(root)
return diameter
108. 将有序数组转换为二叉搜索树
def sortedArrayToBST(nums):
def build(left, right):
if left > right:
return None
# 选择中间元素作为根节点
mid = (left + right) // 2
root = TreeNode(nums[mid])
# 递归构建左右子树
root.left = build(left, mid - 1)
root.right = build(mid + 1, right)
return root
return build(0, len(nums) - 1)
图论
200. 岛屿面积
遍历网格中的每个格子,遇到’1’(陆地)时:岛屿计数+1
DFS实现:递归地向上下左右四个方向探索,遇到边界或水就返回,将访问过的陆地置’0’避免重复计数
时间复杂度:O(m×n),每个格子最多访问一次
def numIslands(grid):
if not grid:
return 0
def dfs(i, j):
if i < 0 or i >= m or j < 0 or j >= n or grid[i][j] != '1':
return
grid[i][j] = '0' # 标记已访问
dfs(i+1, j)
dfs(i-1, j)
dfs(i, j+1)
dfs(i, j-1)
m, n = len(grid), len(grid[0])
count = 0
for i in range(m):
for j in range(n):
if grid[i][j] == '1':
count += 1
dfs(i, j)
return count
回溯
46. 全排列
输入:nums = [1,2,3]
输出:[[1,2,3],[1,3,2],[2,1,3],[2,3,1],[3,1,2],[3,2,1]]
思路:交换法回溯。固定第first位,将后面每个元素交换到first位置,递归排列剩余部分,回溯时换回。无需额外空间。
# 交换法(空间优化)
def permute(nums):
def backtrack(first):
# 当 first 指向最后一个位置时,说明得到一个完整排列
if first == len(nums):
res.append(nums[:]) # nums[:] 是当前排列的副本
return
# 将 first 到末尾的每个元素轮流放到 first 位置
for i in range(first, len(nums)):
# 1. 做选择:将 nums[i] 交换到 first 位置
nums[first], nums[i] = nums[i], nums[first]
# 2. 递归:固定 first 位置,排列剩余元素
backtrack(first + 1)
# 3. 撤销选择:恢复原状,以便尝试下一个元素
nums[first], nums[i] = nums[i], nums[first]
res = []
backtrack(0)
return res
78. 子集
输入:nums = [1,2,3]
输出:[[],[1],[2],[1,2],[3],[1,3],[2,3],[1,2,3]]
思路:回溯遍历所有组合。每个元素有选或不选两种可能,用start控制只能向后选避免重复。迭代法更巧妙:每遇到新元素,就将其加入已有所有子集中。
def subsets(nums):
def backtrack(start, path):
res.append(path[:])
for i in range(start, len(nums)):
path.append(nums[i])
backtrack(i + 1, path)
path.pop()
res = []
backtrack(0, [])
return res
# 迭代法
def subsets(nums):
res = [[]]
for num in nums:
res += [curr + [num] for curr in res]
return res
17. 电话号码的字母组合
输入:digits = “23”
输出:[“ad”,“ae”,“af”,“bd”,“be”,“bf”,“cd”,“ce”,“cf”]
思路:多叉树回溯。按顺序处理每个数字,枚举其对应字母,递归拼接下一数字的字母,长度达标时记录结果。
def letterCombinations(digits):
if not digits:
return []
phone = {
'2': 'abc', '3': 'def', '4': 'ghi', '5': 'jkl',
'6': 'mno', '7': 'pqrs', '8': 'tuv', '9': 'wxyz'
}
def backtrack(index, path):
if index == len(digits):
res.append(''.join(path))
return
for char in phone[digits[index]]:
path.append(char)
backtrack(index + 1, path)
path.pop()
res = []
backtrack(0, [])
return res
39. 组合总和
输入:candidates = [2,3,6,7], target = 7
输出:[[2,2,3],[7]]
解释:
2 和 3 可以形成一组候选,2 + 2 + 3 = 7 。注意 2 可以使用多次。
7 也是一个候选, 7 = 7 。
仅有这两种组合。
思路:排序+回溯剪枝。从start开始选数避免重复组合,因为可重复使用同一元素,所以递归时传i而非i+1,剩余值小于0时剪枝。
def combinationSum(candidates, target):
def backtrack(start, path, remaining):
if remaining == 0:
res.append(path[:])
return
if remaining < 0:
return
for i in range(start, len(candidates)):
path.append(candidates[i])
backtrack(i, path, remaining - candidates[i])
# 可以重复使用,所以传i
path.pop()
res = []
candidates.sort() # 可选,用于剪枝
backtrack(0, [], target)
return res
22. 括号生成
输入:n = 3
输出:[“((()))”,“(()())”,“(())()”,“()(())”,“()()()”]
思路:约束条件回溯。维护已用左括号left和右括号right数量,保证right ≤ left ≤ n,左右都用完时记录结果。
def generateParenthesis(n):
def backtrack(left, right, path):
if left == n and right == n:
res.append(''.join(path))
return
if left < n:
path.append('(')
backtrack(left + 1, right, path)
path.pop()
if right < left:
path.append(')')
backtrack(left, right + 1, path)
path.pop()
res = []
backtrack(0, 0, [])
return res
79. 单词搜索
输入:board = [[‘A’,‘B’,‘C’,‘E’],[‘S’,‘F’,‘C’,‘S’],[‘A’,‘D’,‘E’,‘E’]], word = “ABCCED”
输出:true
思路:DFS回溯+原地标记。从每个格子出发尝试匹配单词,匹配成功则向四个方向递归。用临时字符#标记已访问路径,回溯时恢复。
# 优化版(不使用额外visited数组)
def exist(board, word):
def dfs(i, j, index):
if index == len(word):
return True
if (i < 0 or i >= m or j < 0 or j >= n or
board[i][j] != word[index]):
return False
# 临时标记
temp = board[i][j]
board[i][j] = '#'
found = (dfs(i+1, j, index+1) or
dfs(i-1, j, index+1) or
dfs(i, j+1, index+1) or
dfs(i, j-1, index+1))
board[i][j] = temp
return found
m, n = len(board), len(board[0])
for i in range(m):
for j in range(n):
if dfs(i, j, 0):
return True
return False
131. 分割回文串
请你将 s 分割成一些 子串,使每个子串都是 回文串 。返回 s 所有可能的分割方案。
输入:s = “aab”
输出:[[“a”,“a”,“b”],[“aa”,“b”]]
思路:分割点回溯。从start开始尝试不同长度的子串,若是回文则加入路径并递归剩余部分,回溯时移除。
def partition(s):
def is_palindrome(sub):
return sub == sub[::-1]
def backtrack(start, path):
if start == len(s):
res.append(path[:])
return
for end in range(start, len(s)):
if is_palindrome(s[start:end+1]):
path.append(s[start:end+1])
backtrack(end + 1, path)
path.pop()
res = []
backtrack(0, [])
return res
51. N皇后
按照国际象棋的规则,皇后可以攻击与之处在同一行或同一列或同一斜线上的棋子。n 皇后问题 研究的是如何将 n 个皇后放置在 n×n 的棋盘上,并且使皇后彼此之间不能相互攻击。
给你一个整数 n ,返回所有不同的 n 皇后问题 的解决方案。
每一种解法包含一个不同的 n 皇后问题 的棋子放置方案,该方案中 ‘Q’ 和 ‘.’ 分别代表了皇后和空位。
输入:n = 4
输出:[[“.Q…”,“…Q”,“Q…”,“…Q.”],[“…Q.”,“Q…”,“…Q”,“.Q…”]]
思路:逐行放置+列/对角线约束。用三个数组记录已占用的列、主对角线(row-col恒定)、副对角线(row+col恒定),每行尝试所有列,满足条件则递归下一行。
# 简洁版(用列表代替集合)
def solveNQueens(n):
def backtrack(row):
if row == n:
board = []
for r in range(n):
line = ['.'] * n
line[queens[r]] = 'Q'
board.append(''.join(line))
res.append(board)
return
for col in range(n):
if (col not in cols and
row - col not in diag1 and
row + col not in diag2):
queens[row] = col
cols.append(col)
diag1.append(row - col)
diag2.append(row + col)
backtrack(row + 1)
cols.pop()
diag1.pop()
diag2.pop()
res = []
queens = [-1] * n
cols, diag1, diag2 = [], [], []
backtrack(0)
return res
二分查找
def binary_search(nums, target):
left, right = 0, len(nums) - 1 # 定义target在左闭右闭的区间里
while left <= right: # 当 left == right时,区间[left, right]依然有效
mid = left + (right - left) // 2 # 防止溢出,等同于 (left + right) // 2
if nums[mid] == target:
return mid # 找到目标值,返回下标
elif nums[mid] < target:
left = mid + 1 # target在右半部分,缩小区间为[mid + 1, right]
else:
right = mid - 1 # target在左半部分,缩小区间为[left, mid - 1]
return -1 # 未找到目标值
# 示例
nums = [1, 3, 5, 7, 9, 11, 13]
target = 7
index = binary_search(nums, target)
print(f"目标值 {target} 的索引是: {index}") # 输出: 目标值 7 的索引是: 3
35. 搜索插入位置
输入: nums = [1,3,5,6], target = 5
输出: 2
输入: nums = [1,3,5,6], target = 2
输出: 1
思路:标准二分查找。找到目标值返回索引,没找到时返回left,此时left指向第一个大于等于target的位置,即插入位置。
def searchInsert(nums, target):
left, right = 0, len(nums) - 1
while left <= right:
mid = (left + right) // 2
if nums[mid] == target:
return mid
elif nums[mid] < target:
left = mid + 1
else:
right = mid - 1
return left # 插入位置
74. 搜索二维矩阵
输入:matrix = [[1,3,5,7],[10,11,16,20],[23,30,34,60]], target = 3
输出:true
思路:将二维矩阵展开为一维数组进行二分查找。通过mid // n和mid % n将一维索引映射回二维坐标。
def searchMatrix(matrix, target):
if not matrix or not matrix[0]:
return False
m, n = len(matrix), len(matrix[0])
left, right = 0, m * n - 1
while left <= right:
mid = (left + right) // 2
# 将一维索引转换为二维坐标
row = mid // n
col = mid % n
val = matrix[row][col]
if val == target:
return True
elif val < target:
left = mid + 1
else:
right = mid - 1
return False
34. 在排序数组中查找元素的第一个和最后一个位置
输入:nums = [5,7,7,8,8,10], target = 8
输出:[3,4]
思路:两次二分查找。找第一个位置时,找到目标后继续向左搜索(right = mid - 1);找最后一个位置时,找到目标后继续向右搜索(left = mid + 1)。
def searchRange(nums, target):
def binary_search(is_first):
left, right = 0, len(nums) - 1
pos = -1
while left <= right:
mid = (left + right) // 2
if nums[mid] == target:
pos = mid
if is_first:
right = mid - 1 # 继续向左找
else:
left = mid + 1 # 继续向右找
elif nums[mid] < target:
left = mid + 1
else:
right = mid - 1
return pos
return [binary_search(True), binary_search(False)]
33. 搜索旋转排序数组
输入:nums = [4,5,6,7,0,1,2], target = 0
输出:4
思路:二分查找时,总有一半是有序的。先判断哪半有序,再判断target是否在该半内,以此缩小搜索范围。
def search(nums, target):
left, right = 0, len(nums) - 1
while left <= right:
mid = (left + right) // 2
if nums[mid] == target:
return mid
# 判断哪边是有序的
if nums[left] <= nums[mid]: # 左半部分有序
if nums[left] <= target < nums[mid]:
right = mid - 1
else:
left = mid + 1
else: # 右半部分有序
if nums[mid] < target <= nums[right]:
left = mid + 1
else:
right = mid - 1
return -1
153. 寻找旋转排序数组中的最小值
输入:nums = [3,4,5,1,2]
输出:1
解释:原数组为 [1,2,3,4,5] ,旋转 3 次得到输入数组。
思路:比较mid与right的值。若nums[mid] > nums[right],说明最小值在右半部分;否则在左半部分(包括mid)。
def findMin(nums):
left, right = 0, len(nums) - 1
while left < right:
mid = (left + right) // 2
if nums[mid] > nums[right]:
left = mid + 1
else:
right = mid
return nums[left]
4. 寻找两个正序数组的中位数
输入:nums1 = [1,3], nums2 = [2]
输出:2.00000
解释:合并数组 = [1,2,3] ,中位数 2
思路:在较短的数组上二分查找切分点i,使两个数组左半部分元素个数之和为(m+n+1)//2。通过比较A[i-1]与B[j]、B[j-1]与A[i]确定切分是否合适,然后计算中位数。
def findMedianSortedArrays(nums1, nums2):
A, B = nums1, nums2
if len(A) > len(B):
A, B = B, A
m, n = len(A), len(B)
total = m + n
half = total // 2
l, r = 0, m
while l <= r:
i = (l + r) // 2
j = half - i
A_left = A[i-1] if i > 0 else float('-inf')
A_right = A[i] if i < m else float('inf')
B_left = B[j-1] if j > 0 else float('-inf')
B_right = B[j] if j < n else float('inf')
if A_left <= B_right and B_left <= A_right:
if total % 2 == 1:
return min(A_right, B_right)
return (max(A_left, B_left) + min(A_right, B_right)) / 2
elif A_left > B_right:
r = i - 1
else:
l = i + 1
return -1
栈
20. 有效的括号
输入:s = “()[]{}”
输出:true
思路:遇到左括号入栈,遇到右括号时检查栈顶是否是对应的左括号,若是则弹出,否则无效。最后栈应为空。
def isValid(s):
stack = []
pairs = {'(': ')', '{': '}', '[': ']'}
for char in s:
if char in pairs: # 左括号
stack.append(char)
else: # 右括号
if not stack or pairs[stack.pop()] != char:
return False
return len(stack) == 0
最长有效括号
输入:s = “)()())”
输出:4
解释:最长有效括号子串是 “()()”
思路:栈里存放的是括号的索引,栈底始终保存着"最后一个无法匹配的右括号"的索引,当遇到右括号时,弹出栈顶,然后用当前索引减去新的栈顶得到有效长度。
def longestValidParentheses(s: str) -> int:
stack = [-1] # 栈底放-1作为哨兵
max_len = 0
for i, char in enumerate(s):
if char == '(':
stack.append(i) # 左括号入栈
else: # 右括号
stack.pop() # 弹出匹配的左括号
if not stack: # 如果栈空了
stack.append(i) # 把当前位置作为新的基准
else:
max_len = max(max_len, i - stack[-1])
return max_len
155. 最小栈
输入:
[“MinStack”,“push”,“push”,“push”,“getMin”,“pop”,“top”,“getMin”]
[[],[-2],[0],[-3],[],[],[],[]]
输出:
[null,null,null,null,-3,null,0,-2]
解释:
MinStack minStack = new MinStack();
minStack.push(-2);
minStack.push(0);
minStack.push(-3);
minStack.getMin(); --> 返回 -3.
minStack.pop();
minStack.top(); --> 返回 0.
minStack.getMin(); --> 返回 -2.
思路:用辅助栈同步存储当前最小值。push时,若新值≤辅助栈顶则同时入辅助栈;pop时,若弹出的值等于辅助栈顶则辅助栈也弹出。
class MinStack:
def __init__(self):
self.stack = []
self.min_stack = [] # 辅助栈,存储最小值
def push(self, val):
self.stack.append(val)
# 如果min_stack为空或val小于等于当前最小值,则入栈
if not self.min_stack or val <= self.min_stack[-1]:
self.min_stack.append(val)
def pop(self):
if self.stack:
val = self.stack.pop()
if val == self.min_stack[-1]:
self.min_stack.pop()
def top(self):
return self.stack[-1] if self.stack else None
def getMin(self):
return self.min_stack[-1] if self.min_stack else None
394. 字符串解码
输入:s = “3[a]2[bc]”
输出:“aaabcbc”
思路:遇到数字时累积,遇到’[‘时将当前字符串和数字入栈并重置,遇到’]'时出栈拼接字符串(前字符串 + 当前字符串×数字)。注意处理多位数。
def decodeString(s):
stack = []
curr_num = 0
curr_str = ''
for char in s:
if char.isdigit():
curr_num = curr_num * 10 + int(char)
elif char == '[':
# 将当前数字和字符串入栈
stack.append((curr_str, curr_num))
curr_str = ''
curr_num = 0
elif char == ']':
# 出栈,拼接字符串
prev_str, num = stack.pop()
curr_str = prev_str + curr_str * num
else:
curr_str += char
return curr_str
# 递归解法
def decodeString(s):
def decode(i):
res = ''
num = 0
while i < len(s):
if s[i].isdigit():
num = num * 10 + int(s[i])
elif s[i] == '[':
i, inner = decode(i + 1)
res += inner * num
num = 0
elif s[i] == ']':
return i, res
else:
res += s[i]
i += 1
return i, res
return decode(0)[1]
739. 每日温度
给定一个整数数组 temperatures ,表示每天的温度,返回一个数组 answer ,其中 answer[i] 是指对于第 i 天,下一个更高温度出现在几天后。如果气温在这之后都不会升高,请在该位置用 0 来代替。
输入: temperatures = [73,74,75,71,69,72,76,73]
输出: [1,1,4,2,1,1,0,0]
思路:单调递减栈(存索引)。遍历温度,当当前温度 > 栈顶温度时,说明找到了 warmer day,计算天数差并弹出栈顶。
def dailyTemperatures(temperatures):
n = len(temperatures)
res = [0] * n
stack = [] # 存储索引,保持递减
for i in range(n):
# 当前温度比栈顶温度高,说明找到了 warmer day
while stack and temperatures[i] > temperatures[stack[-1]]:
prev_idx = stack.pop()
res[prev_idx] = i - prev_idx
stack.append(i)
return res
84. 柱状图中最大的矩形
输入:heights = [2,1,5,6,2,3]
输出:10
解释:最大的矩形为图中红色区域,面积为 10
思路:单调递增栈。当遇到比栈顶低的高度时,以栈顶高度为矩形高,左右扩展计算宽度(当前索引i与新的栈顶索引之间的距离-1),最后加入一个高度0强制清空栈。
对于每个柱子作为矩形的高,我们需要找到左边第一个比它矮和右边第一个比它矮的位置,这两个位置之间的宽度就是该高度能形成的最大宽度。
def largestRectangleArea(heights):
stack = [] # 存储索引,保持递增
max_area = 0
n = len(heights)
for i in range(n + 1):
# 当遍历完或当前高度小于栈顶高度时
curr_height = heights[i] if i < n else 0
while stack and curr_height < heights[stack[-1]]:
height = heights[stack.pop()]
# 如果栈为空,宽度为i;否则宽度为i - stack[-1] - 1
width = i if not stack else i - stack[-1] - 1
max_area = max(max_area, height * width)
if i < n:
stack.append(i)
return max_area
堆
215. 数组中的第K个最大元素
输入: [3,2,1,5,6,4], k = 2
输出: 5
思路:维护一个大小为k的最小堆。遍历数组,当堆未满时直接入堆;堆满后,若当前元素大于堆顶,则替换堆顶。最终堆顶就是第k大的元素。也可用最大堆,pop k次得到结果。
import heapq
def findKthLargest(nums, k):
# 方法1:最小堆,维护大小为k的堆
heap = nums[:k]
heapq.heapify(heap) # 构建最小堆
for num in nums[k:]:
if num > heap[0]: # 如果当前元素大于堆顶
heapq.heapreplace(heap, num) # 替换堆顶
return heap[0] # 堆顶就是第k大的元素
# 方法2:最大堆(Python只有最小堆,通过取负数实现)
def findKthLargest(nums, k):
heap = []
for num in nums:
heapq.heappush(heap, -num) # 取负数构建最大堆
for _ in range(k):
result = -heapq.heappop(heap)
return result
347. 前K个高频元素
输入:nums = [1,1,1,2,2,3], k = 2
输出:[1,2]
思路:先用Counter统计频率,然后维护一个大小为k的最小堆,按频率排序。遍历频率表,将(频率, 元素)入堆,堆满后弹出最小频率,最终堆中留下的就是前k个高频元素。
import heapq
from collections import Counter
# 方法1:最小堆
def topKFrequent(nums, k):
# 统计频率
count = Counter(nums)
heap = []
for num, freq in count.items():
heapq.heappush(heap, (freq, num))
if len(heap) > k:
heapq.heappop(heap) # 弹出频率最小的
return [num for freq, num in heap]
# 方法2:最大堆
def topKFrequent(nums, k):
count = Counter(nums)
heap = [(-freq, num) for num, freq in count.items()]
heapq.heapify(heap)
result = []
for _ in range(k):
result.append(heapq.heappop(heap)[1])
return result
295. 数据流的中位数
输入
[“MedianFinder”, “addNum”, “addNum”, “findMedian”, “addNum”, “findMedian”]
[[], [1], [2], [], [3], []]
输出
[null, null, null, 1.5, null, 2.0]
解释
MedianFinder medianFinder = new MedianFinder();
medianFinder.addNum(1); // arr = [1]
medianFinder.addNum(2); // arr = [1, 2]
medianFinder.findMedian(); // 返回 1.5 ((1 + 2) / 2)
medianFinder.addNum(3); // arr[1, 2, 3]
medianFinder.findMedian(); // return 2.0
思路:用两个堆维护数据流:small为最大堆(存较小的一半),large为最小堆(存较大的一半)。添加元素时保持两堆大小平衡(small最多比large多一个),中位数就是small堆顶(奇数个)或两堆顶平均值(偶数个)。
class MedianFinder:
def __init__(self):
self.small = [] # 最大堆(存较小的一半)
self.large = [] # 最小堆(存较大的一半)
def addNum(self, num):
# 统一加入large,再把large的最小值移到small
heapq.heappush(self.small, -heapq.heappushpop(self.large, num))
# 保持small >= large(最多多一个)
if len(self.small) > len(self.large) + 1:
heapq.heappush(self.large, -heapq.heappop(self.small))
def findMedian(self):
if len(self.small) > len(self.large):
return -self.small[0]
return (-self.small[0] + self.large[0]) / 2
贪心算法
121. 买卖股票的最佳时机
输入:[7,1,5,3,6,4]
输出:5
解释:在第 2 天(股票价格 = 1)的时候买入,在第 5 天(股票价格 = 6)的时候卖出,最大利润 = 6-1 = 5 。
思路:遍历价格,记录历史最低点,每天计算如果当天卖出能获得的最大利润(当天价格-历史最低价),更新全局最大利润。
def maxProfit(prices):
min_price = float('inf')
max_profit = 0
for price in prices:
min_price = min(min_price, price)
max_profit = max(max_profit, price - min_price)
return max_profit
55. 跳跃游戏
给你一个非负整数数组 nums ,你最初位于数组的 第一个下标 。数组中的每个元素代表你在该位置可以跳跃的最大长度。
输入:nums = [2,3,1,1,4]
输出:true
解释:可以先跳 1 步,从下标 0 到达下标 1, 然后再从下标 1 跳 3 步到达最后一个下标。
思路:维护当前能到达的最远位置。遍历数组,如果当前位置不可达(i > max_reach)则失败;否则更新最远可达位置。若最远位置已覆盖末尾,提前返回true。
def canJump(nums):
max_reach = 0 # 当前能到达的最远位置
for i in range(len(nums)):
# 如果当前位置不可达
if i > max_reach:
return False
# 更新最远可达位置
max_reach = max(max_reach, i + nums[i])
# 提前结束
if max_reach >= len(nums) - 1:
return True
return True
45. 跳跃游戏 II
输入: nums = [2,3,1,1,4]
输出: 2
解释: 跳到最后一个位置的最小跳跃数是 2。
从下标为 0 跳到下标为 1 的位置,跳 1 步,然后跳 3 步到达数组的最后一个位置。
思路:贪心+BFS思想。记录当前跳跃能到达的边界(cur_end)和所有选择中能到达的最远位置(farthest)。遍历时更新farthest,当到达边界时必须增加一次跳跃,并将边界更新为farthest。
def jump(nums):
if len(nums) <= 1:
return 0
jumps = 0
cur_end = 0 # 当前跳跃能到达的边界
farthest = 0 # 所有选择中能到达的最远位置
for i in range(len(nums) - 1):
# 更新最远可达位置
farthest = max(farthest, i + nums[i])
# 到达当前跳跃的边界,必须进行下一次跳跃
if i == cur_end:
jumps += 1
cur_end = farthest
# 如果已经可以到达末尾,提前结束
if cur_end >= len(nums) - 1:
break
return jumps
763. 划分字母区间
给你一个字符串 s 。我们要把这个字符串划分为尽可能多的片段,同一字母最多出现在一个片段中。
输入:s = “ababcbacadefegdehijhklij”
输出:[9,7,8]
解释:
划分结果为 “ababcbaca”、“defegde”、“hijhklij” 。
每个字母最多出现在一个片段中。
像 “ababcbacadefegde”, “hijhklij” 这样的划分是错误的,因为划分的片段数较少。
思路:先记录每个字母最后出现的位置。遍历字符串,不断更新当前片段的结束位置为当前字母的最后出现位置。当遍历到结束位置时,说明当前片段可以切割,记录长度并开始新片段。
def partitionLabels(s):
# 记录每个字母最后出现的位置
last_pos = {}
for i, char in enumerate(s):
last_pos[char] = i
result = []
start = end = 0
for i, char in enumerate(s):
# 更新当前片段的结束位置
end = max(end, last_pos[char])
# 如果当前位置是片段的结束位置
if i == end:
result.append(end - start + 1)
start = i + 1
return result
动态规划
DP解题五步曲:
- 确定状态(dp数组的含义)
- 确定状态转移方程
- 确定初始条件
- 确定遍历顺序
- 返回结果
空间优化技巧:
一维DP:只依赖前几个状态时用滚动数组
二维DP:只依赖上一行时可用一维数组
70. 爬楼梯
输入:n = 2
输出:2
解释:有两种方法可以爬到楼顶。
- 1 阶 + 1 阶
- 2 阶
思路:斐波那契数列DP。爬到第i阶可以从i-1跨1步或从i-2跨2步到达,所以dp[i] = dp[i-1] + dp[i-2]。可优化为滚动变量省空间。
def climbStairs(n):
if n <= 2:
return n
# dp[i] = 爬到第i阶的方法数
dp = [0] * (n + 1)
dp[1] = 1
dp[2] = 2
for i in range(3, n + 1):
dp[i] = dp[i-1] + dp[i-2]
return dp[n]
118. 杨辉三角
输入: numRows = 5
输出: [[1],[1,1],[1,2,1],[1,3,3,1],[1,4,6,4,1]]
思路:逐行构建。每行首尾为1,中间元素row[j]等于上一行的[j-1]和[j]之和。利用上一行数据计算当前行。
def generate(numRows):
triangle = []
for i in range(numRows):
row = [1] * (i + 1) # 每行首尾都是1
# 计算中间的值
for j in range(1, i):
row[j] = triangle[i-1][j-1] + triangle[i-1][j]
triangle.append(row)
return triangle
198. 打家劫舍
你是一个专业的小偷,计划偷窃沿街的房屋。每间房内都藏有一定的现金,影响你偷窃的唯一制约因素就是相邻的房屋装有相互连通的防盗系统,如果两间相邻的房屋在同一晚上被小偷闯入,系统会自动报警。
输入:[1,2,3,1]
输出:4
解释:偷窃 1 号房屋 (金额 = 1) ,然后偷窃 3 号房屋 (金额 = 3)。
偷窃到的最高金额 = 1 + 3 = 4 。
思路:线性DP。对于第i家,有两种选择:偷(则i-1不能偷,总金额=dp[i-2]+nums[i])或不偷(总金额=dp[i-1])。取最大值,dp[i] = max(dp[i-1], dp[i-2] + nums[i])。
def rob(nums):
if not nums:
return 0
if len(nums) == 1:
return nums[0]
n = len(nums)
dp = [0] * n
dp[0] = nums[0]
dp[1] = max(nums[0], nums[1])
for i in range(2, n):
# 偷第i家:dp[i-2] + nums[i]
# 不偷第i家:dp[i-1]
dp[i] = max(dp[i-1], dp[i-2] + nums[i])
return dp[n-1]
279. 完全平方数
输入:n = 12
输出:3
解释:12 = 4 + 4 + 4
输入:n = 13
输出:2
解释:13 = 4 + 9
思路:完全背包DP。dp[i]表示组成i的最少平方数个数。遍历所有平方数,对于每个i,尝试减去一个平方数并取最小值:dp[i] = min(dp[i], dp[i-square] + 1)。
def numSquares(n):
# dp[i] = 组成i的最少完全平方数个数
dp = [float('inf')] * (n + 1)
dp[0] = 0
# 预计算所有平方数
squares = [i * i for i in range(1, int(n ** 0.5) + 1)]
for i in range(1, n + 1):
for square in squares:
if square > i:
break
dp[i] = min(dp[i], dp[i - square] + 1)
return dp[n]
322. 零钱兑换
给你一个整数数组 coins ,表示不同面额的硬币;以及一个整数 amount ,表示总金额。计算并返回可以凑成总金额所需的 最少的硬币个数 。如果没有任何一种硬币组合能组成总金额,返回 -1 。你可以认为每种硬币的数量是无限的。
输入:coins = [1, 2, 5], amount = 11
输出:3
解释:11 = 5 + 5 + 1
思路:完全背包DP。dp[i]表示组成金额i的最少硬币数。遍历每种硬币,对于每个金额i,尝试使用该硬币:dp[i] = min(dp[i], dp[i-coin] + 1)。最后检查是否可达。
def coinChange(coins, amount):
# dp[i] = 组成金额i所需的最少硬币数
dp = [float('inf')] * (amount + 1)
dp[0] = 0
for i in range(1, amount + 1):
for coin in coins:
if coin <= i:
dp[i] = min(dp[i], dp[i - coin] + 1)
return dp[amount] if dp[amount] != float('inf') else -1
多维动态规划
# 二维DP通用模板
def solve_2d_dp(grid):
m, n = len(grid), len(grid[0])
dp = [[0] * n for _ in range(m)]
# 初始化边界
dp[0][0] = grid[0][0]
for i in range(1, m):
dp[i][0] = dp[i-1][0] + grid[i][0] # 根据题目调整
for j in range(1, n):
dp[0][j] = dp[0][j-1] + grid[0][j] # 根据题目调整
# 状态转移
for i in range(1, m):
for j in range(1, n):
dp[i][j] = min(dp[i-1][j], dp[i][j-1]) + grid[i][j] # 根据题目调整
return dp[m-1][n-1]
62. 不同路径
一个机器人位于一个 m x n 网格的左上角 (起始点在下图中标记为 “Start” )。
机器人每次只能向下或者向右移动一步。机器人试图达到网格的右下角(在下图中标记为 “Finish” )。
问总共有多少条不同的路径?
def uniquePaths(m, n):
# dp[i][j] = 到达(i,j)的不同路径数
dp = [[1] * n for _ in range(m)]
for i in range(1, m):
for j in range(1, n):
dp[i][j] = dp[i-1][j] + dp[i][j-1]
return dp[m-1][n-1]
# 空间优化(一维数组)
def uniquePaths(m, n):
dp = [1] * n
for i in range(1, m):
for j in range(1, n):
dp[j] += dp[j-1]
return dp[n-1]
64. 最小路径和
输入:grid = [[1,3,1],[1,5,1],[4,2,1]]
输出:7
解释:因为路径 1→3→1→1→1 的总和最小。
思路:动态规划,每个位置的最小路径和等于从上方或左方来的较小路径和加上当前位置的值。
def minPathSum(grid):
if not grid or not grid[0]:
return 0
m, n = len(grid), len(grid[0])
dp = [[0] * n for _ in range(m)]
dp[0][0] = grid[0][0]
# 初始化第一行
for j in range(1, n):
dp[0][j] = dp[0][j-1] + grid[0][j]
# 初始化第一列
for i in range(1, m):
dp[i][0] = dp[i-1][0] + grid[i][0]
# 动态规划
for i in range(1, m):
for j in range(1, n):
dp[i][j] = min(dp[i-1][j], dp[i][j-1]) + grid[i][j]
return dp[m-1][n-1]
5. 最长回文子串
输入:s = “babad”
输出:“bab”
解释:“aba” 同样是符合题意的答案。
思路:遍历每个字符(及两字符之间)作为回文中心,向两边扩展直到不能构成回文,记录最长回文子串的起止位置。
奇数长度回文:left = right(同一个中心)
偶数长度回文:left = i, right = i + 1(两个相邻字符作为中心)
# 中心扩展法(更优)
def longestPalindrome(s):
if not s or len(s) < 2:
return s
def expand_around_center(left, right):
while left >= 0 and right < len(s) and s[left] == s[right]:
left -= 1
right += 1
return left + 1, right - 1
start, end = 0, 0
for i in range(len(s)):
# 奇数长度回文
l1, r1 = expand_around_center(i, i)
# 偶数长度回文
l2, r2 = expand_around_center(i, i + 1)
if r1 - l1 > end - start:
start, end = l1, r1
if r2 - l2 > end - start:
start, end = l2, r2
return s[start:end + 1]
1143. 最长公共子序列
def longestCommonSubsequence(text1, text2):
m, n = len(text1), len(text2)
dp = [[0] * (n + 1) for _ in range(m + 1)]
for i in range(1, m + 1):
for j in range(1, n + 1):
if text1[i-1] == text2[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]
# 空间优化(一维数组)
def longestCommonSubsequence(text1, text2):
m, n = len(text1), len(text2)
dp = [0] * (n + 1)
for i in range(1, m + 1):
prev = 0 # dp[i-1][j-1]
for j in range(1, n + 1):
temp = dp[j] # 保存dp[i-1][j]
if text1[i-1] == text2[j-1]:
dp[j] = prev + 1
else:
dp[j] = max(dp[j], dp[j-1])
prev = temp
return dp[n]
72. 编辑距离
def minDistance(word1, word2):
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], # 删除
dp[i][j-1], # 插入
dp[i-1][j-1] # 替换
) + 1
return dp[m][n]
# 空间优化(一维数组)
def minDistance(word1, word2):
m, n = len(word1), len(word2)
dp = [0] * (n + 1)
# 初始化第一行
for j in range(n + 1):
dp[j] = j
for i in range(1, m + 1):
prev = dp[0] # dp[i-1][0]
dp[0] = i # dp[i][0]
for j in range(1, n + 1):
temp = dp[j] # 保存dp[i-1][j]
if word1[i-1] == word2[j-1]:
dp[j] = prev
else:
dp[j] = min(
temp, # 删除
dp[j-1], # 插入
prev # 替换
) + 1
prev = temp
return dp[n]
技巧
136. 只出现一次的数字
输入:nums = [2,2,1]
输出:1
def singleNumber(nums):
# 异或运算:相同为0,不同为1,任何数与0异或等于本身
result = 0
for num in nums:
result ^= num
return result
# 一行版本
def singleNumber(nums):
from functools import reduce
return reduce(lambda x, y: x ^ y, nums)
# 通用解法(其他数出现k次)
def singleNumberK(nums, k):
# 统计每个位上1的个数
result = 0
for i in range(32):
count = 0
for num in nums:
count += (num >> i) & 1
if count % k != 0:
result |= (1 << i)
# 处理负数
if result >= 2**31:
result -= 2**32
return result
169. 多数元素
输入:nums = [3,2,3]
输出:3
def majorityElement(nums):
# 摩尔投票法
candidate = None
count = 0
for num in nums:
if count == 0:
candidate = num
count += 1 if num == candidate else -1
return candidate
# 排序法
def majorityElement(nums):
nums.sort()
return nums[len(nums) // 2]
# 哈希表法
def majorityElement(nums):
from collections import Counter
return Counter(nums).most_common(1)[0][0]
75. 颜色分类
输入:nums = [2,0,2,1,1,0]
输出:[0,0,1,1,2,2]
def sortColors(nums):
"""
荷兰国旗问题:三指针
0: [0, p0)
1: [p0, i)
2: (p2, n-1]
"""
p0, p2 = 0, len(nums) - 1
i = 0
while i <= p2:
if nums[i] == 0:
nums[i], nums[p0] = nums[p0], nums[i]
p0 += 1
i += 1
elif nums[i] == 2:
nums[i], nums[p2] = nums[p2], nums[i]
p2 -= 1
# 不移动i,因为交换过来的数还需要检查
else:
i += 1
return nums
# 计数排序法
def sortColors(nums):
count = [0] * 3
for num in nums:
count[num] += 1
index = 0
for color in range(3):
for _ in range(count[color]):
nums[index] = color
index += 1
return nums
31. 下一个排列
整数数组的 下一个排列 是指其整数的下一个字典序更大的排列。
例如,arr = [1,2,3] 的下一个排列是 [1,3,2] 。
类似地,arr = [2,3,1] 的下一个排列是 [3,1,2] 。
而 arr = [3,2,1] 的下一个排列是 [1,2,3] ,因为 [3,2,1] 不存在一个字典序更大的排列。
def nextPermutation(nums):
"""
1. 从右向左找到第一个升序对 (nums[i] < nums[i+1])
2. 从右向左找到第一个大于nums[i]的数 nums[j]
3. 交换 nums[i] 和 nums[j]
4. 反转 i+1 到末尾(使其升序)
"""
n = len(nums)
# 找到第一个升序对
i = n - 2
while i >= 0 and nums[i] >= nums[i+1]:
i -= 1
if i >= 0:
# 找到第一个大于nums[i]的数
j = n - 1
while j >= 0 and nums[j] <= nums[i]:
j -= 1
# 交换
nums[i], nums[j] = nums[j], nums[i]
# 反转 i+1 到末尾
left, right = i + 1, n - 1
while left < right:
nums[left], nums[right] = nums[right], nums[left]
left += 1
right -= 1
return nums
287. 寻找重复数
输入:nums = [1,3,4,2,2]
输出:2
def findDuplicate(nums):
"""
快慢指针法(Floyd判圈算法)
将数组视为链表:nums[i] 指向 nums[nums[i]]
"""
# 找到相遇点
slow = nums[0]
fast = nums[0]
while True:
slow = nums[slow]
fast = nums[nums[fast]]
if slow == fast:
break
# 找到环的入口(即重复数)
slow = nums[0]
while slow != fast:
slow = nums[slow]
fast = nums[fast]
return slow
# 二分查找法
def findDuplicate(nums):
left, right = 1, len(nums) - 1
while left < right:
mid = (left + right) // 2
count = sum(1 for num in nums if num <= mid)
if count > mid:
right = mid
else:
left = mid + 1
return left
# 位运算法
def findDuplicate(nums):
n = len(nums) - 1
result = 0
# 逐位确定
for bit in range(32):
x = y = 0
for i in range(n + 1):
if nums[i] & (1 << bit):
x += 1
if i >= 1 and (i & (1 << bit)):
y += 1
if x > y:
result |= (1 << bit)
return result
非常规题
热门激活函数
| 激活函数 | 梯度特性 | 输出范围 | 主要应用 | 计算复杂度 |
|---|---|---|---|---|
| Sigmoid | 两端饱和,易梯度消失(1)饱和区:当x绝对值很大时,梯度趋近于0(梯度消失)(2)最大梯度:x=0时,梯度=0.25 (3)非零中心:输出恒正,可能导致梯度更新呈Z字形 | (0, 1) | 二分类输出层、LSTM/GRU中的门控机制、需要概率输出的场景 | 中(指数运算) |
| Swish | 无界正区间,平滑 (1)无界性:正区间梯度有界但非零 (2)平滑性:处处可导,无拐点 (3)自门控:x < 0时也有小梯度,避免死亡神经元 | (-0.28, ∞) | 现代CNN、Transformer、需要平滑激活的场景 | 中(指数运算) |
| GLU | 梯度可调节 (1)梯度调节:sigmoid(gate)控制梯度流动 (2)双路径:值和门控的梯度相互影响 (3)选择性:可以选择性通过信息 | (-∞, ∞) | 语言模型(BERT、GPT)、序列建模 | 高(双线性+门控) |
| SwiGLU | 平滑门控 (1)平滑门控:Swish比Sigmoid梯度更优 (2)无界激活:正区间无上限 (3)负区间保留:保留小负梯度 | (-∞, ∞) | LLaMA系列模型、PaLM等LLM、需要强表达能力的FFN层 | 高(双线性+Swish) |
| ReLU | 正区间常数梯度 (1)正区间:梯度恒为1(无梯度消失)(2)负区间:梯度为0(神经元死亡风险)(3)计算效率:最简单的梯度计算 | [0, ∞) | CNN、通用隐藏层、计算资源受限的场景 | 低(max操作) |
| Leaky ReLU | 避免死亡ReLU (2) 负区间:梯度恒为α(避免死亡) (3)平衡性:在稀疏性和信息保留间平衡 | (-∞, ∞) | 深度GAN(生成对抗网络)、缓解神经元死亡、需要负区间信息的任务 | 低(max+乘法) |
| Softmax | 概率分布 (1)概率约束:输出和为1 (2)竞争机制:梯度相互影响 (3)雅可比矩阵:非对角元为负 | (0, 1) | 多分类输出层、注意力机制、概率分布建模 | 中(指数+归一化) |
Sigmoid:
σ
(
x
)
=
1
/
(
1
+
e
−
x
)
σ(x)= 1/(1+e^{−x})
σ(x)=1/(1+e−x)return 1 / (1 + np.exp(-x))
导数:
σ
(
x
)
(
1
−
σ
(
x
)
)
σ(x)(1-σ(x))
σ(x)(1−σ(x))s = sigmoid(x), s * (1 - s)
Swish:
x
⋅
σ
(
x
)
x⋅σ(x)
x⋅σ(x)x * sigmoid(x)
导数:s + x * s * (1 - s)
GLU (Gated Linear Unit):
(
x
W
+
b
)
⊗
σ
(
x
V
+
c
)
(xW+b)⊗σ(xV+c)
(xW+b)⊗σ(xV+c)
- ⊗ \otimes ⊗ 表示逐元素相乘
- σ \sigma σ 是 sigmoid 函数
- x W + b xW + b xW+b 是线性变换(值路径)
- x V + c xV + c xV+c 是门控路径
SwiGLU (Swish Gated Linear Unit): SwiGLU 是 GLU (Gated Linear Unit) 和 Swish 的结合,公式为: Swish ( x W + b ) ⊗ ( x V + c ) \text{Swish}(xW+b)⊗(xV+c) Swish(xW+b)⊗(xV+c)
def sigmoid(x):
return 1 / (1 + np.exp(-x))
def swish(x, beta=1.0):
return x * sigmoid(beta * x)
def glu(x, w_gate, w_value):
gate = np.dot(x, w_gate)
value = np.dot(x, w_value)
return sigmoid(gate) * value
def swiglu(x, w1, w2, b1, b2, beta=1.0):
hidden1 = np.dot(x, w1) + b1
hidden2 = np.dot(x, w2) + b2
return swish(hidden1) * hidden2
ReLU (Rectified Linear Unit):
m
a
x
(
0
,
x
)
max(0,x)
max(0,x) np.maximum(0, x)
导数:
1
if x>0, else
0
1~\text{if x>0, else}~0
1 if x>0, else 0 np.where(x > 0, 1.0, 0.0)
Leaky ReLU:
m
a
x
(
α
x
,
x
)
max(αx,x)
max(αx,x),通常
α
=
0.01
\alpha=0.01
α=0.01 np.maximum(0, alpha*x)
导数:
1
if x>0, else
α
1~\text{if x>0, else}~\alpha
1 if x>0, else α np.where(x > 0, 1.0, alpha)
Softmax:
Softmax
(
x
i
)
=
e
x
i
∑
j
=
1
n
e
x
j
\text{Softmax}(x_i)=\frac{e^{x_i}}{∑_{j=1}^{n}e^{x_j}}
Softmax(xi)=∑j=1nexjexi
def softmax(x):
x_max = np.max(x, axis=-1, keepdims=True)
exp_x = np.exp(x - x_max)
return exp_x / np.sum(exp_x, axis=-1, keepdims=True)
热门损失函数
交叉熵损失(Cross-Entropy Loss)衡量两个概率分布之间的差异,常用于分类任务。
- “熵”最早起源于热力学,后来被香农引入信息论,用来衡量信息的不确定性或信息的平均量。衡量的是单个分布的不确定性。
- “交叉”(Cross)表示两个分布之间的相互作用
- 在机器学习中,我们最小化交叉熵,就是让预测分布尽量接近真实分布
二分类:
L
=
−
[
y
l
o
g
(
p
)
+
(
1
−
y
)
l
o
g
(
1
−
p
)
]
L=−[ylog(p)+(1−y)log(1−p)]
L=−[ylog(p)+(1−y)log(1−p)]
多分类:
L
=
−
∑
i
=
1
N
y
i
l
o
g
(
y
^
i
)
L=-\sum_{i=1}^{N}y_i log(\hat{y}_i)
L=−∑i=1Nyilog(y^i)
KL散度 (Kullback-Leibler Divergence)
D
K
L
(
P
∣
∣
Q
)
=
∑
x
P
(
x
)
l
o
g
(
P
(
x
)
Q
(
x
)
)
D_{KL}(P||Q)=\sum_{x}P(x)log(\frac{P(x)}{Q(x)})
DKL(P∣∣Q)=∑xP(x)log(Q(x)P(x))
KL散度可以分解为交叉熵减去熵:
D
K
L
(
P
∣
∣
Q
)
=
H
(
P
,
Q
)
−
H
(
P
)
D_{KL}(P||Q)=H(P, Q) - H(P)
DKL(P∣∣Q)=H(P,Q)−H(P)
其中:
- H ( P , Q ) H(P, Q) H(P,Q) 是交叉熵
- H ( P ) = − ∑ P ( x ) log P ( x ) H(P) = -\sum P(x)\log P(x) H(P)=−∑P(x)logP(x) 是分布P的熵
负对数似然 (NLL)
N
L
L
=
−
∑
i
=
1
n
l
o
g
(
p
i
,
c
i
)
NLL=−\sum_{i=1}^{n}log(p_{i,c_i})
NLL=−∑i=1nlog(pi,ci)
其中
p
i
,
c
i
p_{i,c_i}
pi,ci 是第i个样本真实类别
c
i
c_i
ci的预测概率
def cross_entropy(y_pred, y_true):
y_pred_soft = softmax(y_pred)
loss = -np.sum(y_true * np.log(y_pred_soft))
return loss
def NLLloss(y_pred, y_true):
return -np.sum(y_true * np.log(y_pred))
def KL(y_pred, y_true):
return np.sum(y_true * np.log(y_true / y_pred))
CTC损失 (Connectionist Temporal Classification)
CTC通过以下三点创新,巧妙地绕过了序列对齐的难题:
- 空白标签:引入特殊标签(-),处理沉默间隔或不确定区域
- 路径积分:考虑所有可能的对齐路径,而不依赖单一正确对齐
- 动态规划优化:使用前向-后向算法高效计算所有路径的概率和,避免枚举所有可能路径
计算步骤:
-
前向算法:计算前向变量 α t ( s ) \alpha_{t}(s) αt(s),表示在时间步 t t t达到序列位置 s s s的概率。
- 初始化( t = 1 t=1 t=1):只有序列开头的 Blank 或者第一个真实标签才有可能作为起点。 α 1 ( 1 ) = y − , 1 , α 1 ( 2 ) = y L 2 ′ , 1 , α 1 ( s ) = 0 对于 \alpha_1(1)=y_{-,1}, \alpha_1(2)=y_{L_{2}^{'},1}, \alpha_1(s)=0对于 α1(1)=y−,1,α1(2)=yL2′,1,α1(s)=0对于
- 递推:对于 t>1,
α
t
(
s
)
\alpha_t(s)
αt(s) 的值可以从上一个时间步 t−1 的几种合法状态转移过来。转移规则基于 CTC 的核心约束:
- 允许在同标签上停留 α t − 1 ( s ) \alpha_{t-1}(s) αt−1(s)
- 允许在 Blank 和标签之间跳转 α t − 1 ( s − 1 ) \alpha_{t-1}(s-1) αt−1(s−1)
- 不允许跳过标签
- 特殊跳跃,跳过空白符 α t − 1 ( s − 2 ) \alpha_{t-1}(s-2) αt−1(s−2),通常是当 L s ′ L_{s}^{'} Ls′不是blank,并且 L s ′ L_{s}^{'} Ls′和 L s − 2 ′ L_{s-2}^{'} Ls−2′不同
- 公式合并:

-
后向算法:后向算法与前者对称。 β t ( s ) \beta_{t}(s) βt(s)表示在时间步 t t t,且当前已经位于扩充后序列 L ′ L^{'} L′的第 s s s个字符从t到T这一段路径能够产生剩余后缀 L s ′ L_{s}^{'} Ls′的概率。
- 初始化( t = T t=T t=T):最后一个时间步只能对应最后两个位置:最后的blank或最后一个标签。 β T ( ∣ L ′ ∣ ) = 1 , β T ( ∣ L ′ ∣ − 1 ) = 1 \beta_{T}(|L^{'}|)=1,\beta_{T}(|L^{'}|-1)=1 βT(∣L′∣)=1,βT(∣L′∣−1)=1
- 递推:从T往前推,转移逻辑与前向类似,但是方向相反。

-
计算损失:当我们有了所有 α 和 β 后,我们就可以计算在时间步 t 经过特定字符 L s ′ L_{s}^{'} Ls′的所有路径概率和。最终,整个标签序列 L 的总概率 p(L∣X) 就是所有路径的概率之和。这个值可以通过在最终时间步 T 收集所有合法终点的 α 值得到:给定输入语音特征X,模型输出正确标签序列 L 的概率:

各种组件
CNN
当你拿到一个CNN架构时,逐层计算的方式如下:
确定输入:
H
i
n
,
W
i
n
,
C
i
n
H_{in}, W_{in}, C_{in}
Hin,Win,Cin
确定该层参数:
K
,
S
,
P
,
C
o
u
t
K, S, P, C_{out}
K,S,P,Cout
计算新尺寸:
H
o
u
t
=
[
(
H
i
n
+
2
P
−
K
)
/
S
]
+
1
H_{out} = [(H_{in}+2P-K)/S] + 1
Hout=[(Hin+2P−K)/S]+1
W
o
u
t
=
[
(
W
i
n
+
2
P
−
K
)
/
S
]
+
1
W_{out} = [(W_{in}+2P-K)/S] + 1
Wout=[(Win+2P−K)/S]+1
C
o
u
t
=
卷积核数量(每个卷积核运算后生成一张特征图)
C_{out} = 卷积核数量(每个卷积核运算后生成一张特征图)
Cout=卷积核数量(每个卷积核运算后生成一张特征图)
重复步骤:将计算出的作为下一层的输入
在处理图像输入时,影响输出特征图尺寸的主要有四个参数,我习惯用记忆法“KSPD”来概括:
K:卷积核大小(Kernel Size)
S:步长(Stride) 每次滑动多少个像素
P:填充(Padding) 在输入边缘补多少圈0
D:空洞率(Dilation) 卷积核点之间的间隔,默认为1(标准卷积)。
情况 A:尺寸不变(Same Padding / Half Padding)
W
o
u
t
=
W
i
n
W_{out}=W_{in}
Wout=Win
解得
P
=
(
K
−
1
)
/
2
P=(K-1)/2
P=(K−1)/2
情况 B:尺寸减半
W
o
u
t
=
W
i
n
/
2
W_{out}=W_{in}/2
Wout=Win/2, 通常配合步长S=2 实现。
CNN 降采样
import torch.nn as nn
import torch
class FeatureCNN(nn.Module):
def __init__(self):
super(FeatureCNN, self).__init__()
self.conv1 = nn.Conv1d(1, 16, kernel_size=3, stride=2, padding=1)
self.conv2 = nn.Conv1d(16, 32, kernel_size=3, stride=5, padding=1)
self.conv3 = nn.Conv1d(32, 64, kernel_size=3, stride=5, padding=1)
def forward(self, x):
x = self.conv1(x)
x = self.conv2(x)
x = self.conv3(x)
return x
def test_feature_cnn():
model = FeatureCNN()
input_tensor = torch.randn(1, 1, 16000) # batch_size=1, channels=1, length=FS
output = model(input_tensor)
print(f"Input shape: {input_tensor.shape}")
print(f"Output shape: {output.shape}")
test_feature_cnn()
Normalization
BatchNorm中的running_mean 和 running_var 是统计量,不需要梯度更新,适合图片任务
训练时:
- 每个批次的分布可能不同,需要根据当前数据归一化
- 同时用滑动平均方式估计全局分布(为推理做准备)
推理时:
- 推理时可能 batch size 很小,甚至为 1,用当前批次统计会很不稳定
- 推理样本应该基于训练集的整体分布进行归一化
LayerNorm不需要很大的Batch,只是对每个特征的均值做归一化,适合语言建模
两种归一化的主要区别:
- LayerNorm:减去均值并除以标准差,然后进行仿射变换(缩放+平移)
- RMSNorm:只除以均方根(RMS),只进行缩放变换(没有平移),计算量更小
这两种归一化方法在Transformer架构中都很常用,特别是RMSNorm因为计算效率更高,在一些现代大语
import torch
import torch.nn as nn
class BatchNorm(nn.Module):
"""
自定义的 Batch Normalization 层,适用于输入形状为 [batch_size, seq_len, num_features] 的数据。
在训练和推理阶段分别使用批统计量和运行统计量进行归一化。
"""
def __init__(self, num_features, eps=1e-5, rho=0.1):
"""
参数:
num_features: 特征维度(即输入张量的最后一个维度大小)
eps: 防止除零的小常数
rho: 动量参数,用于更新 running_mean 和 running_var
"""
super(BatchNorm, self).__init__()
# 可学习的缩放因子和偏移量
self.gamma = nn.Parameter(torch.ones(num_features))
self.beta = nn.Parameter(torch.zeros(num_features))
# 用于数值稳定的小常数
self.eps = eps
# 动量参数,控制运行均值和方差的更新速度
self.rho = rho
# 注册为缓冲区(buffer),不会被优化器更新,但会随模型保存和加载
self.register_buffer("running_mean", torch.zeros(num_features))
self.register_buffer("running_var", torch.ones(num_features))
def forward(self, x):
"""
前向传播:
x: 输入张量,形状为 [batch_size, seq_len, num_features]
返回:
归一化后的张量,形状与输入相同
"""
# 训练模式:使用当前批次的均值和方差
if self.training:
# 在 batch 和 seq_len 维度上计算均值和方差(每个特征独立)
# keepdim=True 保持维度,便于后续广播
mean = x.mean(dim=(0, 1), keepdim=True) # 形状: [1, 1, num_features]
var = x.var(dim=(0, 1), keepdim=True, unbiased=False) # 无偏估计设为 False,使用有偏方差
# 更新运行统计量(使用动量平滑)
self.running_mean = (1 - self.rho) * self.running_mean \
+ self.rho * mean.squeeze()
self.running_var = (1 - self.rho) * self.running_var \
+ self.rho * var.squeeze()
else:
# 推理模式:使用累计的运行均值和方差
# 将形状从 [num_features] 扩展为 [1, 1, num_features] 以匹配输入
mean = self.running_mean.view(1, 1, -1)
var = self.running_var.view(1, 1, -1)
# 归一化: (x - mean) / sqrt(var + eps)
x_normalized = (x - mean) / torch.sqrt(var + self.eps)
# 缩放和平移: gamma * x_normalized + beta
return self.gamma * x_normalized + self.beta
class LayerNorm(nn.Module):
"""
Layer Normalization 层归一化
公式: y = gamma * (x - mean) / sqrt(var + eps) + beta
"""
def __init__(self, d_model, eps=1e-5):
"""
Args:
d_model: 特征维度
eps: 防止除零的小常数
"""
super(LayerNorm, self).__init__()
# 可学习的缩放参数 gamma,初始化为1
self.gamma = nn.Parameter(torch.ones(d_model))
# 可学习的偏置参数 beta,初始化为0
self.beta = nn.Parameter(torch.zeros(d_model))
self.eps = eps
def forward(self, x):
"""
Args:
x: 输入张量,形状为 [batch_size, seq_len, d_model] 或 [batch_size, d_model]
Returns:
归一化后的张量,形状与输入相同
"""
# 计算最后一个维度的均值,保持维度以便广播
mean = x.mean(dim=-1, keepdim=True)
# 计算最后一个维度的方差,保持维度
var = x.var(dim=-1, keepdim=True, unbiased=False) # unbiased=False 使用有偏方差估计
# 归一化: (x - mean) / sqrt(var + eps)
x_norm = (x - mean) / torch.sqrt(var + self.eps)
# 应用可学习的参数: gamma * x_norm + beta
return self.gamma * x_norm + self.beta
class RMSNorm(nn.Module):
"""
Root Mean Square Layer Normalization
论文: https://arxiv.org/abs/1910.07467
公式: y = gamma * x / sqrt(mean(x^2) + eps)
"""
def __init__(self, d_model, eps=1e-5):
"""
Args:
d_model: 特征维度
eps: 防止除零的小常数
"""
super(RMSNorm, self).__init__()
# 可学习的缩放参数 gamma,初始化为1
self.gamma = nn.Parameter(torch.ones(d_model))
self.eps = eps
def forward(self, x):
"""
Args:
x: 输入张量,形状为 [batch_size, seq_len, d_model] 或 [batch_size, d_model]
Returns:
归一化后的张量,形状与输入相同
"""
# 计算RMS: sqrt(mean(x^2) + eps)
# torch.mean(x ** 2, dim=-1, keepdim=True) 计算最后一个维度的平方均值
rms = torch.sqrt(torch.mean(x ** 2, dim=-1, keepdim=True) + self.eps)
# 归一化: x / rms
x_norm = x / rms
# 应用可学习的缩放参数: gamma * x_norm
# RMSNorm没有beta参数(偏置项)
return self.gamma * x_norm
Transformer
Transformer的核心结构是一个完全基于自注意力机制、摒弃了循环与卷积的Encoder-Decoder框架。在编码器端,输入序列会先叠加位置编码以注入顺序信息,随后经过多个相同的层;每一层主要由两个子模块构成:第一个是多头自注意力,它通过计算输入中所有位置两两之间的Q、K、V点积来捕捉全局依赖关系,第二个是前馈神经网络,负责对每个位置独立进行非线性变换,并且每个子模块前后都使用了残差连接与层归一化。解码器端结构与之类似,但在第一个自注意力层中使用了掩码,防止当前位置看到未来的输出,同时额外插入了一个交叉注意力层,用于将解码器的查询与编码器输出的键和值进行交互,从而让解码过程能够聚焦于输入序列的相关部分。最终,解码器顶层的输出通过线性层和Softmax生成目标序列的概率分布。

SelfAttention
import torch
import torch.nn.functional as F
# 补全单头注意力计算代码
class SelfAttention(torch.nn.Module):
def __init__(self, input_dim, hidden_dim):
super(SelfAttention, self).__init__()
self.query_linear = torch.nn.Conv1d(input_dim, hidden_dim, kernel_size=1)
self.key_linear = torch.nn.Conv1d(input_dim, hidden_dim, kernel_size=1)
self.value_linear = torch.nn.Conv1d(input_dim, hidden_dim, kernel_size=1)
def forward(self, x, mask):
# x: (b, c, n)
# mask: (b, n, n)
Q = self.query_linear(x).permute(0, 2, 1) # (b, n, hidden_dim)
K = self.key_linear(x).permute(0, 2, 1)
V = self.value_linear(x).permute(0, 2, 1)
d_k = Q.size(-1)
score = torch.matmul(Q, K.transpose(-2, -1)) / torch.sqrt(torch.tensor(d_k))
score = score.masked_fill(mask, float('-inf')
attention = F.softmax(score, dim=-1)
output = torch.matmul(attention, V)
return output
if __name__ == '__main__':
# 测试SelfAttention
b, c, n = 2, 3, 5
x = torch.rand(b, c, n)
mask = torch.randint(0, 2, (b, n, n)).bool()
sa = SelfAttention(c, c)
y = sa(x, mask)
print(y.shape)
MHA
import torch
from torch import nn
import torch.functional as F
import math
class multi_head_attention(nn.Module):
def __init__(self, d_model, n_head):
super(multi_head_attention, self).__init__()
self.n_head = n_head
self.d_model = d_model
self.w_q = nn.Linear(d_model, d_model)
self.w_k = nn.Linear(d_model, d_model)
self.w_v = nn.Linear(d_model, d_model)
self.w_combine = nn.Linear(d_model, d_model)
self.softmax = nn.Softmax(dim=-1)
def forward(self, q, k, v):
batch, time, dimension = q.shape
n_d = self.d_model // self.n_head
q, k, v = self.w_q(q), self.w_k(k), self.w_v(v)
q = q.view(batch, time, self.n_head, n_d).permute(0, 2, 1, 3)
k = k.view(batch, time, self.n_head, n_d).permute(0, 2, 1, 3)
v = v.view(batch, time, self.n_head, n_d).permute(0, 2, 1, 3)
score = q @ k.transpose(2, 3) / math.sqrt(n_d)
mask = torch.tril(torch.ones(time, time, dtype=bool))
score = score.masked_fill(mask == 0, float("-inf"))
score = self.softmax(score) @ v
score = score.permute(0, 2, 1, 3).contiguous().view(batch, time, dimension)
output = self.w_combine(score)
return output
if __name__ == '__main__':
d_model = 512
n_head = 8
X = torch.randn(128, 64, 512) # B, T, D
attention = multi_head_attention(d_model, n_head)
output = attention(X, X, X)
print(output, output.shape)
TransformerEncoderBlock
import torch
import torch.nn as nn
import torch.nn.functional as F
import math
class TransformerEncoderBlock(nn.Module):
"""Transformer编码器块 - 面试简洁版"""
def __init__(self, d_model, nhead, dim_feedforward=2048, dropout=0.1):
super().__init__()
# 1. 多头自注意力
self.self_attn = nn.MultiheadAttention(d_model, nhead, dropout=dropout)
# 2. 前馈网络 (FFN)
self.linear1 = nn.Linear(d_model, dim_feedforward)
self.linear2 = nn.Linear(dim_feedforward, d_model)
# 3. 层归一化
self.norm1 = nn.LayerNorm(d_model)
self.norm2 = nn.LayerNorm(d_model)
# 4. Dropout
self.dropout = nn.Dropout(dropout)
def forward(self, x, mask=None):
"""
x: [seq_len, batch_size, d_model] 或 [batch_size, seq_len, d_model]
mask: 注意力掩码
"""
# 残差连接 + 多头注意力 + 层归一化
attn_out, _ = self.self_attn(x, x, x, attn_mask=mask)
x = self.norm1(x + self.dropout(attn_out))
# 残差连接 + 前馈网络 + 层归一化
ff_out = self.linear2(F.relu(self.linear1(x)))
x = self.norm2(x + self.dropout(ff_out))
return x
# 完整的Transformer编码器
class CompleteTransformerEncoder(nn.Module):
def __init__(self, ...):
super().__init__()
self.embedding = nn.Embedding(vocab_size, d_model)
self.pos_encoder = PositionalEncoding(d_model, dropout)
self.encoder_blocks = nn.ModuleList([
TransformerEncoderBlock(d_model, nhead, dim_feedforward, dropout)
for _ in range(num_layers)
])
def forward(self, x):
x = self.embedding(x) * math.sqrt(self.d_model)
x = self.pos_encoder(x) # 位置编码在这里添加!
for block in self.encoder_blocks:
x = block(x)
return x
Position Embedding
Absolute Position Embedding
Rotary Position Embedding (RoPE)
大模型的代码
后训练
对话模型
- 基础模型 (Base Model):知识渊博但不会对话。
- 经过监督微调 (SFT):学会基本的对话和指令遵循。
- 经过偏好对齐 (DPO/RLHF):回答更符合人类偏好,更安全、更有帮助。
- 经过持续学习和 specialization (OSFT/LoRA):能够在不遗忘旧知识的前提下,高效地学习新领域知识。
ASR模型
- 领域与说话人自适应(Fine-tuning / Adaptation)让模型更贴合实际应用场景:领域自适应医疗、法律、客服、会议等|解决术语(OOV)问题|说话人自适应|特定口音 / 发音习惯|噪声环境适配|车载、电话、远场【方法:全量微调 / LoRA|数据重采样(domain balancing)】
- 语言模型融合(LM Fusion)ASR 的关键后处理步骤之一:外接语言模型(LM)提升文本合理性|常见方式:shallow fusion(解码时加权)deep fusion / cold fusion(训练中融合)【例如:声学模型输出:recognize speach LM 修正为:recognize speech】
- 解码优化(Decoding Optimization)直接影响最终输出:beam search 调参(beam size、长度惩罚)|CTC prefix beam search(若使用 CTC)|时间对齐优化(timestamp refinement)|适用于模型如 Whisper 或 RNN-T / CTC 系列。
- 后处理与文本规范化(Post-processing / TN)把“可读文本”变成“正确文本”:标点恢复(punctuation restoration)|大小写恢复(truecasing)|数字规范化|“twenty twenty six” → “2026”|专有名词修正【通常使用:seq2seq 模型或规则 + LM】
- 错误纠正(Error Correction)进一步降低 WER:拼写纠错(spell correction)语义级纠错(context-aware correction)N-best 重排序(rescoring)【例如:“ice cream” vs “I scream”】
- 数据增强与再训练(Data Augmentation Loop)提升鲁棒性:SpecAugment(频谱遮挡)加噪声(SNR 不同级别)速度扰动(speed perturbation)混响(RIR)然后:重新微调模型(闭环优化)
- 对齐与时间戳优化(Alignment / Timestamp)提升时间信息质量:word-level timestampforced | alignment|VAD(语音活动检测)优化切分【用于:字幕生成|语音检索】
- 多语言与口音增强|针对复杂语言场景:增加低资源语言数据|code-switch(中英混说)方言适配
- 推理优化(Inference Optimization)部署阶段关键:
模型压缩:quantization(INT8 / FP16)|流式识别(streaming ASR)|chunk-based 推理(低延迟)|ONNX / TensorRT 加速 - 稳定性与工程处理线上系统必须处理:长音频切分(segmentation)|异常音频 fallback|重试机制|VAD + ASR pipeline 协同
TTS模型
- 领域/说话人微调(Fine-tuning)在通用模型基础上做定向优化:说话人微调(voice cloning)|少量数据适配特定音色|领域微调|客服、播报、有声书等特定风格|多语言/方言适配|补充长尾发音和表达【常见方法:全参数微调|LoRA / Adapter(更轻量)】
- 对齐与时长修正(Alignment Refinement)
很多 TTS(尤其是类似 Tacotron 2 或 VITS 系列)依赖文本-语音对齐:修复:漏字 / 重复读发音跳跃【方法:强制对齐(forced alignment)duration predictor 再训练monotonic alignment search 优化】 - 音质增强(Audio Quality Enhancement)
生成的语音通常需要后处理:降噪 / 去伪影|带宽扩展(Bandwidth Extension)|超分辨率(Super-resolution)|响度标准化(loudness normalization)【常见技术:GAN-based vocoder refinement(如 HiFi-GAN)|diffusion-based enhancer】 - Prosody(韵律)与可控性优化|让语音更自然,而不是“读字”:语速(speed)音高(pitch)情感(emotion)停顿(pause)【方法:引入 prosody embedding style token(GST)prompt-based 控制(类似大模型)】
- 评估与对齐人类偏好(Evaluation / RLHF-like)类似大模型的对齐过程:主观评测:MOS(Mean Opinion Score)自动指标:CER / WER(用 ASR 反评估)偏好优化:ranking loss|reward model(少见但在增加)
- 文本前处理(Text Normalization, TN)部署前非常关键:数字 → 读法(2026 → “二零二六”)单位 → 规范化(kg, km)缩写展开(Dr., etc.)【通常包括:G2P(grapheme-to-phoneme)多语言 phoneme 统一】
- 推理优化(Inference Optimization)面向实际部署:模型压缩:quantization(INT8 / FP16)|pruning|加速:ONNX / TensorRT|streaming inference(低延迟)|batch / chunk 推理
- 稳定性与异常处理|解决线上问题:防止:卡顿爆音NaN 音频【fallback 策略:切换到安全模型|重试机制】
- 多模态/大模型融合(新趋势)最新 TTS(如 GPT-style TTS)会加入:instruction tuning(文本指令控制语音)语音 prompt(few-shot voice cloning)LLM + codec token 统一建模
数学公式推导
经典算法
强化学习代码
小于n的最大整数
用数组 A 中的数字(可重复使用)拼出一个小于 n 的尽可能大的整数。n 是一个正整数,可以有不同的位数。
- n=23121, A={2,4,9} → 22999
- n=23121, A={9} → 9999
- n=23333, A={2,3} → 23332
- n=22222, A={2} → 222
def find_max_less_than_n(n, A):
s = str(n)
A.sort()
def get_max_less_than_digit(d):
for x in A[::-1]:
if x < d:
return x
return None
res = []
# 标记是否已经找到更小的位
smaller = False
for i, ch in enumerate(s):
d = int(ch)
if smaller:
res.append(str(A[-1]))
continue
# 找 <= d 的最大数字
for x in A[::-1]:
if x < d:
res.append(str(x))
smaller = True
break
elif x == d:
res.append(str(x))
break
else: # 没找到 <= d 的数字
# 回溯
j = i - 1
while j >= 0:
prev = int(res[j])
smaller_num = get_max_less_than_digit(prev)
if smaller_num is not None:
res[j] = str(smaller_num)
res = res[:j+1]
smaller = True
break
j -= 1
res.pop()
if j < 0: # 回溯失败,降位数
return int(str(A[-1]) * (len(s) - 1))
return int(''.join(res))
WER 计算
WER = (S + D + I) / N × 100%
- S(Substitutions):替换错误(识别词替换了正确词)
- D(Deletions):删除错误(漏识别)
- I(Insertions):插入错误(多识别)
- N:参考文本(正确文本)的总词数
WER(Word Error Rate,词错率)和 CER(Character Error Rate,字错率/字符错率)的核心区别在于计算的最小单位不同:
WER 基于“词”(英文按空格分隔的单词)
CER 基于“字符”(英文是字母,中文是单个汉字)
直观举例(中文)
参考答案:今天天气真好
识别结果:今天天起真好
| 指标 | 计算过程 | 结果 |
|---|---|---|
| WER | 分词后参考:今天 / 天气 / 真好(3个词)。识别结果:今天 / 天起 / 真好。“天气”错成“天起”算1个词错。 | 1/3 ≈ 33.3% |
| CER | 参考有7个字。识别结果把“气”换成“起”(1个替换)。7个字里错1个。 | 1/7 ≈ 14.3% |
同一个错误,CER 比 WER 数值低(因为字多分母大)。如果每错一个字就算整个词错,WER 会更敏感。
直观举例(英文)
参考答案:cat
识别结果:bat
| 指标 | 计算过程 | 结果 |
|---|---|---|
| CER | 3个字母,第1个字母 c→b(1替换) | 1/3 ≈ 33.3% |
| WER | 1个词,整个词错(1替换) | 1/1 = 100% |
英文中 CER 很严(字母错就算),但实际意义有时不大;WER 更符合语义层面评估。
def wer_simple(reference, hypothesis):
"""简化版:只计算WER值"""
ref_words = reference.split()
hyp_words = hypothesis.split()
distance = edit_distance(ref_words, hyp_words)
return distance / len(ref_words) * 100
def edit_distance(ref, hyp):
"""计算词级别的编辑距离"""
m, n = len(ref), len(hyp)
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 ref[i-1] == hyp[j-1]:
dp[i][j] = dp[i-1][j-1]
else:
dp[i][j] = min(
dp[i-1][j-1] + 1, # 替换
dp[i-1][j] + 1, # 删除
dp[i][j-1] + 1 # 插入
)
return dp[m][n]
更多推荐



所有评论(0)