从零构建 LLM:深入理解自动微分
用代码从头实现自动微分机制,深入理解现代深度学习框架的核心原理,适合想掌握 LLM 本质的开发者。
用代码从头实现自动微分机制,深入理解现代深度学习框架的核心原理,适合想掌握 LLM 本质的开发者。
from typing import Any, Optional, List
import networkx as nx
从零构建 LLM:自动微分
我正在完全从零开始构建一个具备各种现代特性的语言模型:从原生 Python 一路做到功能完备的编程助手。借鉴(厚颜无耻地照搬)电子游戏的做法,我为自己认为实现一个功能完备的语言模型所需的一切构建了一棵科技树。如果你觉得其中缺少了什么,请告诉我:
在继续构建旋转位置编码(Rotary Positional Encodings)等现代特性之前,我们首先需要弄清楚如何使用计算机进行微分。支撑整个深度学习领域的反向传播算法,需要具备计算神经网络输出相对于其输入的导数的能力。在本文中,我们将从零开始,构建一个(不得不承认功能非常有限的)自动微分库,它能够对任意标量值函数进行微分。
这个算法将成为我们深度学习库的核心。最终,这个库将包含训练语言模型所需的一切。
如果连可供微分的数值都没有,我们就无法进行任何微分。我们还希望添加一些标准 float 类型不具备的额外功能,因此需要创建自己的数值类型。我们就把它称为 Tensor。
class Tensor:
"""
Just a number (for now)
"""
value: float
def __init__(self, value: float):
self.value = value
def __repr__(self) -> str:
"""
Create a printable string representation of this
object
This function gets called when you pass a Tensor to print
Without this function:
>>> print(Tensor(5))
<__main__.Tensor at 0x104fd1950>
With this function:
>>> print(Tensor(5))
Tensor(5)
"""
return f"Tensor({self.value})"
# try it out
Tensor(5)
Tensor(5)
接下来,我们需要实现一些想要执行的简单运算:加法、减法和乘法。
def _add(a: Tensor, b: Tensor):
"""
Add two tensors
"""
return Tensor(a.value + b.value)
def _sub(a: Tensor, b: Tensor):
"""
Subtract tensor b from tensor a
"""
return Tensor(a.value - b.value)
def _mul(a: Tensor, b: Tensor):
"""
Multiply two tensors
"""
return Tensor(a.value * b.value)
我们可以像下面这样使用这些运算:
def test(got: Any, want: Any):
"""
Check that two objects are equal to each other
"""
indicator = "✅" if want == got else "❌"
print(f"{indicator} - Want: {want}, Got: {got}")
a = Tensor(3)
b = Tensor(4)
test(_add(a, b).value, 7)
test(_sub(a, b).value, -1)
test(_mul(a, b).value, 12)
✅ - Want: 7, Got: 7
✅ - Want: -1, Got: -1
✅ - Want: 12, Got: 12
直接一头扎进矩阵微分听起来太难了,所以我们先从简单一些的内容开始:标量微分。我能想到的最简单的标量导数,就是计算一个张量相对于自身的导数:[\frac{dx}{dx} = 1]
一个更有意思的例子,是计算两个张量相加所得结果的导数(注意,由于我们的函数有多个输入,因此这里使用偏导数):[f(x, y) = x + y] [\frac{\partial f}{\partial x} = 1] [\frac{\partial f}{\partial y} = 1]
对于乘法和减法,我们也可以做类似的处理。
现在,我们已经从数学上推导出了这些导数,下一步就是将它们转换成代码。在上表中,当我们通过某种运算组合两个张量来创建一个新张量时,导数始终只取决于输入和所执行的运算,不存在任何“隐藏状态”。
这意味着,我们唯一需要存储的信息就是某个运算的输入,以及一个用于计算相对于各个输入的导数的函数。有了这些信息,我们应该就能计算任意二元函数相对于其输入的导数。存储这些信息的一个合适位置,就是该运算所产生的张量。
我们将为 Tensor 添加一些新属性:args 和 local_derivatives。如果某个张量是一个运算的输出,那么 args 将存储该运算的参数,而 local_derivatives 将存储相对于每个输入的导数。我们把它称为 local_derivatives,是为了避免在开始嵌套函数时产生混淆。
计算出导数(根据 args 和 local_derivatives)之后,我们还需要将它存储起来。事实证明,最简洁的做法是把它放在输出所对应的求导变量张量中。我们将这个属性称为 derivative。
class Tensor:
"""
A number that can be differentiated
"""
# If the tensor was made by an operation, the operation arguments
# are stored in args
args: tuple["Tensor"] = ()
# If the tensor was made by an operation, the derivatives wrt
# operation inputs are stored in derivatives
local_derivatives: tuple["Tensor"] = ()
# The derivative we have calculated
derivative: Optional["Tensor"] = None
def __init__(self, value: float):
self.value = value
def __repr__(self) -> str:
"""
Create a printable string representation of this
object
This function gets called when you pass a Tensor to print
Without this function:
>>> print(Tensor(5))
<__main__.Tensor at 0x104fd1950>
With this function:
>>> print(Tensor(5))
Tensor(5)
"""
return f"Tensor({self.value})"
例如,如果有:
a = Tensor(3)
b = Tensor(4)
output = _mul(a, b)
那么 output.args 和 output.local_derivatives 应该被设置为:
output.args == (Tensor(3), Tensor(4))
output.derivatives == (
b, # derivative of output wrt a is b
a, # derivative of output wrt b is a
)
实际计算出导数后,output 相对于 a 的导数将存储在 a.derivative 中,并且应该等于 b(在本例中为 4)。
当下面这些测试通过时,我们就能确定所有操作都是正确的:
a = Tensor(3)
b = Tensor(4)
output = _mul(a, b)
# TODO: differentiate here
test(got=output.args, want=(a, b))
test(got=output.local_derivatives, want=(b, a))
test(got=a.derivative, want=b)
test(got=b.derivative, want=a)
❌ - Want: (Tensor(3), Tensor(4)), Got: ()
❌ - Want: (Tensor(4), Tensor(3)), Got: ()
❌ - Want: Tensor(4), Got: None
❌ - Want: Tensor(3), Got: None
首先,我们为 Tensor 添加一个函数,让它真正计算函数各个参数对应的导数。Pytorch 将这个函数称为 backward,所以我们也采用相同的名称。
class Tensor:
"""
A number that can be differentiated
"""
# If the tensor was made by an operation, the operation arguments
# are stored in args
args: tuple["Tensor"] = ()
# If the tensor was made by an operation, the derivatives wrt
# operation inputs are stored in
local_derivatives: tuple["Tensor"] = ()
# The derivative we have calculated
derivative: Optional["Tensor"] = None
# optionally give this tensor a name
name: Optional[str] = None
# Later, we'll want to record the path we followed to get
# to this tensor and some operations we did along the way
# don't worry about these for now
paths: List["Tensor"] = None
chains: List["Tensor"] = None
def __init__(self, value: float):
self.value = value
def backward(self):
if self.args is None or self.local_derivatives is None:
raise ValueError(
"Cannot differentiate a Tensor that is not a function of other Tensors"
)
for arg, derivative in zip(self.args, self.local_derivatives):
arg.derivative = derivative
def __repr__(self) -> str:
"""
Create a printable string representation of this
object
This function gets called when you pass a Tensor to print
Without this function:
>>> print(Tensor(5))
<__main__.Tensor at 0x104fd1950>
With this function:
>>> print(Tensor(5))
Tensor(5)
"""
return f"Tensor({self.value})"
只有同时在运算的输出张量中存储参数和导数,这段代码才能正常工作。
def _add(a: Tensor, b: Tensor):
"""
Add two tensors
"""
result = Tensor(a.value + b.value)
result.local_derivatives = (Tensor(1), Tensor(1))
result.args = (a, b)
return result
def _sub(a: Tensor, b: Tensor):
"""
Subtract tensor b from a
"""
result = Tensor(a.value - b.value)
result.local_derivatives = (Tensor(1), Tensor(-1))
result.args = (a, b)
return result
def _mul(a: Tensor, b: Tensor):
"""
Multiply two tensors
"""
result = Tensor(a.value * b.value)
result.local_derivatives = (b, a)
result.args = (a, b)
return result
让我们重新运行测试,看看它是否有效。
a = Tensor(3)
b = Tensor(4)
output = _mul(a, b)
output.backward()
test(got=output.args, want=(a, b))
test(got=output.local_derivatives, want=(b, a))
test(a.derivative, b)
test(b.derivative, a)
✅ - Want: (Tensor(3), Tensor(4)), Got: (Tensor(3), Tensor(4))
✅ - Want: (Tensor(4), Tensor(3)), Got: (Tensor(4), Tensor(3))
✅ - Want: Tensor(4), Got: Tensor(4)
✅ - Want: Tensor(3), Got: Tensor(3)
到目前为止一切顺利,让我们尝试嵌套操作。
a = Tensor(3)
b = Tensor(4)
output_1 = _mul(a, b)
# z = a + (a * b)
output_2 = _add(a, output_1)
output_2.backward()
# should get
# dz/db = 0 + a = a
test(b.derivative, a)
❌ - Want: Tensor(3), Got: None
出了问题。
我们应该得到 a 作为 b 的导数,但我们得到了 None。查看 .backward() 函数,问题很清楚:我们还没有考虑嵌套函数。要使这个例子生效,我们需要弄清楚如何通过多个函数计算导数,而不仅仅是一个函数。
要计算嵌套函数的导数,我们可以使用微积分中的一条法则:链式法则。
对于由嵌套函数 $f$ 和 $g$ 生成的变量 $z$,使得 $$z = f(g(x))$$
那么 $z$ 关于 $x$ 的导数为:$$\frac{\partial z}{\partial x} = \frac{\partial f(u)}{\partial u} \frac{\partial g(x)}{\partial x}$$
这里,$u$ 是一个虚拟变量。$\frac{\partial f(u)}{\partial u}$ 表示 $f$ 关于其输入的导数。
设 $$f(x) = g(x)^2$$ 那么我们可以定义 $u=g(x)$ 并用 $u$ 重新表述 $f$ $$f(u) = u^2 \implies \frac{\partial f(u)}{\partial u} = 2u = 2 g(x)$$
链式法则对于多变量函数也按预期工作。当对一个变量求导时,我们可以将其他变量视为常数并正常求导 $$z = f(g(x), h(y))$$
$$\frac{\partial z}{\partial x} = \frac{\partial f(u)}{\partial u} \frac{\partial g(x)}{\partial x}$$ $$\frac{\partial z}{\partial y} = \frac{\partial f(u)}{\partial u} \frac{\partial h(y)}{\partial y}$$
如果我们有不同的函数取相同的输入,我们分别对它们求导,然后将它们加在一起
$$z = f(g(x), h(x))$$
我们得到 $$\frac{\partial z}{\partial x} = \frac{\partial f(u)}{\partial u}\frac{\partial g(x)}{\partial x} + \frac{\partial f(u)}{\partial u}\frac{\partial h(x)}{\partial x}$$
如果我们将 3 个函数链在一起,我们仍然只是将每个函数的导数相乘:
$$\frac{\partial z}{\partial x} = \frac{\partial f(u)}{\partial u} \frac{\partial g(x)}{\partial x} = \frac{\partial f(u)}{\partial u} \frac{\partial g(u)}{\partial u}\frac{\partial h(x)}{\partial x}$$
这推广到任何数量的嵌套
$$z = f_1(f_2(....f_{n-1}(f_n(x))...)) \implies \frac{\partial z}{\partial x} = \frac{\partial f_1(u)}{\partial u}\frac{\partial f_2(u)}{\partial u}...\frac{\partial f_{n-1}(u)}{\partial u}\frac{\partial f_{n}(x)}{\partial x}$$
正如你可能注意到的,数学开始变得相当密集。当我们开始使用神经网络时,我们很容易深入数百或数千个函数,所以为了掌握情况,我们需要一个不同的策略。幸运的是,有一个:将其转换为图。
我们可以从一些规则开始:
例如,这是 $z = mx$ 的图示。
仅此而已!我们将使用的所有方程都可以使用这些简单的规则用图形表示。要尝试一下,让我们为更复杂的公式绘制图表。
这是一种称为图(也称为网络)的结构的示例。计算机科学中的许多问题如果能用图表表示就会容易得多,这也不例外。
这些图的真正威力在于它们也可以帮助我们计算导数。取 $$y = mx + p = \texttt{add}(p, \texttt{mul}(m ,x))$$
从之前开始,我们可以通过对每个操作相对于其输入求导并将结果相乘来找到其导数。在这种情况下,我们得到: $$\frac{\partial y}{\partial p} = \frac{\partial \texttt{add}(u_1, u_2)}{\partial u_1} = 1$$ $$\frac{\partial y}{\partial m} = \frac{\partial \texttt{add}(u_1, u_2)}{\partial u_2}\frac{\partial \texttt{mul}(u_1, u_2)}{\partial u_2} = 1 \times x = x$$ $$\frac{\partial y}{\partial x} = \frac{\partial \texttt{add}(u_1, u_2)}{\partial u_2}\frac{\partial \texttt{mul}(u_1, u_2)}{\partial u_1} = 1 \times m = m$$
我们也可以像这样绘制图表。
如果你想象从 $y$ 走到每个输入,你可能会注意到你经过的边和上面的方程之间的相似性。如果你从 $y$ 走到 $x$,你会经过 a→c→d。类似地,如果你从 $y$ 走到 $m$,你会经过 a→d→e。注意两条路径都经过 c,这是来自 add 的边,对应于输入 $u_2$。此外,两个方程都包含术语 $\frac{\partial \texttt{add}(u_1, u_2)}{\partial u_2}$。
如果我按如下方式重命名边:
我们可以看到从 $y$ 到 $x$ 走,我们经过 $1$、$\frac{\partial \texttt{add}(u_1, u_2)}{\partial u_2}$ 和 $\frac{\partial \texttt{mul}(u_1, u_2)}{\partial u_1}$。如果我们将这些相乘,我们会得到正好是 $\frac{\partial \texttt{add}(u_1, u_2)}{\partial u_2}\frac{\partial \texttt{mul}(u_1, u_2)}{\partial u_1} = \frac{\partial y}{\partial x}$!
事实证明这条规则具有普遍适用性:
如果我们有某个操作 $\texttt{op}(u_1, u_2, ..., u_n)$,我们应该用 $\frac{\partial \texttt{op}(u_1, u_2, ..., u_n)}{\partial u_i}$ 标记对应于输入 $u_i$ 的边。
然后,如果我们想找到输出节点关于任何输入的导数,
输出变量关于输入变量之一的导数可以通过从输出遍历图到输入并将路径上每条边的导数相乘来找到
为了涵盖每个边界情况,有一些额外的细节
如果图包含从输出到输入的多条路径,则导数是每条路径的乘积之和
这来自我们之前看到的情况,当我们有不同的函数具有相同的输入时,我们必须将它们的导数链加在一起。
如果边不是任何函数的输入,则其导数为 1
这涵盖从最终操作到输出的边。你可以认为边的导数为 $\frac{\partial y}{\partial y}=1$
仅此而已!让我们用 $z = (x + c)x$ 尝试一下:
这里,我不是为每个导数写公式,而是继续计算它们在我们代入输入参数时的实际值。我们不仅要找出导数的公式,还要在代入输入参数时计算其值。
剩下的就是沿着每条路径将局部导数相乘在一起。我们会将沿单条路径的导数乘积称为一条链(以链式法则命名)。
我们可以通过绿色路径和红色路径从 $z$ 到 $x$。按照这些路径,我们得到: $$\text{red path} = 1 \times (x + c) = x + c$$
沿着绿色路径我们得到: $$\text{green path} = 1 \times x \times 1 = x$$
将这些加在一起,我们得到 $(x+c) + x = 2x + c$
如果我们从代数上计算导数:
$$\frac{\partial z}{\partial x} = \frac{\partial}{\partial x}((x+c)x) = \frac{\partial}{\partial x}(x^2 + cx) = \frac{\partial x^2}{\partial x} + c\frac{\partial x}{\partial x} = 2x + c$$
我们可以看到它似乎有效!计算 $\frac{\partial z}{\partial c}$ 留给读者作为练习(我一直想说这个)。
总结一下,我们已经发明了以下算法来计算变量关于其输入的导数:
现在我们已经有了图表和文字形式的算法,让我们将其转换为代码。
令人惊讶的是,我们实际上已经将我们的函数转换为图。如果你还记得,当我们从操作生成张量时,我们在输出张量中记录操作的输入(在 .args 中)。我们还在 .local_derivatives 中存储了为每个输入计算导数的函数,这意味着我们知道指向给定节点的每条边的目标和导数。这意味着我们已经完成了步骤 1 和 2。
下一个挑战是找出从我们想要求导的张量,到创建它的输入张量之间的所有路径。由于我们的操作都不存在自引用(输出永远不会被重新作为输入),并且所有边都有方向,因此我们的运算图是一个有向无环图,也就是 DAG。图中不存在环这一性质意味着,我们可以非常容易地使用广度优先搜索(也可以使用深度优先搜索,但正如我们将在第 2 部分看到的那样,BFS 更便于进行某些优化)找到通往每个参数的所有路径。
为了试验一下,让我们重新创建之前那个巨大的图。首先,可以根据输入计算 \(L\):
y = Tensor(1) m = Tensor(2) x = Tensor(3) c = Tensor(4)
left = _sub(y, _add(_mul(m, x), c)) right = _sub(y, _add(_mul(m, x), c))
L = _mul(left, right)
y.name = "y" m.name = "m" x.name = "x" c.name = "c" L.name = "L"
然后使用广度优先搜索完成三件事:
找出从 \(L\) 到参数的所有路径
我们还没有实现一种简单的方法来检查两个张量是否相同,因此需要比较它们的哈希值。
edges = []
stack = [(L, [L])]
nodes = [] edges = [] while stack: node, current_path = stack.pop() # Record nodes we haven't seen before if hash(node) not in [hash(n) for n in nodes]: nodes.append(node)
# because it wasn't created by an operation) then
# record the path taken to get here
if not node.args:
if node.paths is None:
node.paths = []
node.paths.append(current_path)
continue
for arg in node.args: stack.append((arg, current_path + [arg])) # Record every new edge edges.append((hash(node), hash(arg)))
现在我们已经获得了所有边和节点,也就掌握了计算图的完整信息。接下来使用 networkx 将它绘制出来:
labels = {} for i, node in enumerate(nodes): if node.name is None: labels[hash(node)] = str(i) else: labels[hash(node)] = node.name
graph = nx.DiGraph() graph.add_edges_from(edges) pos = nx.nx_agraph.pygraphviz_layout(graph, prog="dot") nx.draw(graph, pos=pos, labels=labels)
如果稍微眯起眼睛看,就会发现它和我们之前创建的图很像!让我们看看算法找到的从 \(L\) 到 \(x\) 的路径。
for path in x.paths: steps = [] for step in path: steps.append(labels[hash(step)]) print("->".join(steps))
L->1->2->4->x L->8->9->10->x
这些路径看起来是正确的!现在只需稍微修改一下算法,让它记录每条路径上的导数链。
y = Tensor(1) m = Tensor(2) x = Tensor(3) c = Tensor(4)
left = _sub(y, _add(_mul(m, x), c)) right = _sub(y, _add(_mul(m, x), c))
L = _mul(left, right)
y.name = "y" m.name = "m" x.name = "x" c.name = "c" L.name = "L"
stack = [(L, [L], [])]
nodes = [] edges = [] while stack: node, current_path, current_chain = stack.pop() # Record nodes we haven't seen before if hash(node) not in [hash(n) for n in nodes]: nodes.append(node)
# because it wasn't created by an operation) then
# record the path taken to get here
if not node.args:
if node.paths is None:
node.paths = []
if node.chains is None:
node.chains = []
node.paths.append(current_path)
node.chains.append(current_chain)
continue
for arg, op in zip(node.args, node.local_derivatives): next_node = arg next_path = current_path + [arg] next_chain = current_chain + [op]
stack.append((arg, next_path, next_chain))
edges.append((hash(node), hash(arg)))
让我们检查导数是否被正确记录。
print(f"Number of chains: {len(x.chains)}") for chain in x.chains: print(chain)
Number of chains: 2 [Tensor(-9), Tensor(-1), Tensor(1), Tensor(2)] [Tensor(-9), Tensor(-1), Tensor(1), Tensor(2)]
到目前为止看起来很合理。正如预期的那样,我们有两条完全相同的路径,每条路径包含 4 个导数(路径中的每条边对应一个导数)。
接下来,将每条路径上的导数相乘,再把各条路径的结果相加,看看能否得到正确答案。
根据我的计算(以及 Wolfram Alpha),\(L\) 关于 \(x\) 的导数是:\[\frac{\partial L}{\partial x} = 2m (c + mx - y)\] 将张量的值代入,得到 \[2\times2 (4 + (2\times3) - 1) = 36\]
total_derivative = Tensor(0) for chain in x.chains: chain_total = Tensor(1) for step in chain: chain_total = _mul(chain_total, step) total_derivative = _add(total_derivative, chain_total)
total_derivative
Tensor(36)
答案正确!看来我们的算法有效。剩下要做的,就是把所有部分组合起来。
整合所有内容
在构思算法时,我们记录了节点、边和路径,这让绘图与调试变得更加容易。既然现在已经知道算法有效,就可以删除这些内容,稍微简化一下实现。
def backward(root_node: Tensor) -> None: stack = [(root_node, [])]
while stack: node, current_derivative = stack.pop()
# because it wasn't created by an operation) then
# record the path taken to get here
if not node.args:
if node.chains is None:
node.chains = []
node.chain.append(current_derivative)
continue
for arg, op in zip(node.args, node.local_derivatives): stack.append((arg, current_derivative + [op]))
此外,(目前)也没有必要存储导数,再单独进行计算。相反,我们可以在遍历过程中直接将导数相乘,从而避免大量重复计算。
def backward(root_node: Tensor) -> None: stack = [(root_node, Tensor(1))]
while stack: node, current_derivative = stack.pop()
# b