Python技术迷

aesara,一个超 nice 的 Python 库!

Aesara 这玩意儿,第一眼看着像“又一个张量库”,真上手以后才发现它不是来跟 NumPy 抢活的。

它干的事更偏后厨:先把你的数学表达式变成计算图,再改写、优化、求导,最后编译成能跑的东西。官方文档对它的定义也差不多这个意思:定义、优化重写、再高效执行多维数组上的数学表达式。安装文档里还专门提了,推荐直接走 conda-forge。另一个得知道的点是,Aesara 项目在 GitHub 上已经归档了,后面 PyTensor 是它的 fork。这个别搞混。

很多人一看到“符号计算”“计算图”,脑子里就自动切到“学术”“不好用”“离业务远”。这判断我一开始也有。后来真拿它处理一点需要“表达式 + 自动求导 + 图优化”的东西,味道就出来了。

你先看一段最小代码。

import aesara
import aesara.tensor as at
import numpy as np

x = at.vector("x")
y = x ** 2 + 2 * x + 1

f = aesara.function([x], y)

data = np.array([1.0, 2.0, 3.0], dtype=np.float32)
print(f(data))
# [ 4.  9. 16.]

这段代码要是换成 NumPy,当然也能写,甚至更短。但区别不在“能不能算”,而在 Aesara 接住的不是数值,而是表达式本身。你给它的是一张图,不是一把已经炒好的菜。

这个差别在求导时特别明显。

import aesara
import aesara.tensor as at

x = at.scalar("x")
loss = x ** 4 - 3 * x ** 2 + 2 * x

grad = at.grad(loss, x)

calc = aesara.function([x], [loss, grad])

value, g = calc(2.0)
print("loss =", value)
print("grad =", g)

手写导数当然不难,4x^3 - 6x + 2 嘛。但表达式一旦长起来,或者里面套了矩阵运算、广播、条件分支,再手推一遍,基本就是给自己挖坑。Aesara 内置了符号求导,这就是它很值钱的地方。官方也把 efficient symbolic differentiation 摆在核心能力里。

我平时看这种库,先不信宣传,先看它到底适合什么场景。Aesara 不是拿来替代你日常 CRUD 代码的,它更像下面这几类活:

第一类,你要描述的是公式,不是流程。 像概率模型、优化问题、损失函数、递推关系,这种东西写成表达式比写成 if/for 更自然。

第二类,你后面一定会动求导。 参数拟合、最优化、贝叶斯建模,这种一旦上手,自动求导就不是加分项,是保命项。

第三类,你想让系统帮你改图,而不是自己抠细节。 Aesara 的文档里明确提到 rewrite system,也就是它能在图层面做改写和优化。这个事不是 NumPy 擅长的。

比如你手上有个很常见的线性回归损失:

import aesara
import aesara.tensor as at
import numpy as np

X = at.matrix("X")
y = at.vector("y")
w = at.vector("w")
b = at.scalar("b")

pred = at.dot(X, w) + b
loss = at.mean((pred - y) ** 2)

grad_w = at.grad(loss, w)
grad_b = at.grad(loss, b)

train_step = aesara.function(
    inputs=[X, y, w, b],
    outputs=[loss, grad_w, grad_b]
)

X_val = np.array([[1, 2], [2, 1], [3, 4]], dtype=np.float32)
y_val = np.array([5, 4, 9], dtype=np.float32)
w_val = np.array([0.1, 0.2], dtype=np.float32)
b_val = np.float32(0.0)

loss_val, gw, gb = train_step(X_val, y_val, w_val, b_val)
print(loss_val)
print(gw)
print(gb)

这段代码不花哨,但味道已经出来了:你关心的是损失函数怎么写,梯度怎么来,至于底层怎么组织计算图、怎么做优化、怎么执行,库替你兜一层。

再往前走一步,Aesara 还有个很多人容易忽略的点:它不只是会算,还允许你扩展自己的 Op。官方文档里就专门有“Creating a new Op”。这个能力对做一些特殊数学算子、老系统兼容、或者接一点自定义逻辑时很有用。

我给个偏现场一点的例子。比如你有个业务里常见的“分段罚分函数”,懒得在图里拼一堆 switch,你可以先包成一个自定义 Op。

import numpy as np
import aesara.tensor as at
from aesara.graph.op import Op
from aesara.graph.basic import Apply

classPiecewisePenalty(Op):
    itypes = [at.dvector]
    otypes = [at.dvector]

defperform(self, node, inputs, outputs):
        (x,) = inputs
        out = np.where(
            x < 0,
10.0,
            np.where(x < 10, x * 0.5, x * 2.0)
        )
        outputs[0][0] = out.astype(np.float64)

penalty_op = PiecewisePenalty()

x = at.dvector("x")
y = penalty_op(x)
f = aesara.function([x], y)

print(f(np.array([-2, 3, 12], dtype=np.float64)))

这代码不复杂,但思路挺重要:先把“脏活怪活”包起来,再挂到图上。真到了模型里,这种方式比把业务判断散在一堆公式里好维护得多。

当然,这库也不是没门槛。

一个坑是 编译和执行是两段心智。刚接触的人最别扭的地方,就是变量不是立即值,打印出来经常不是你想看的那个数,而是一个 symbolic variable。这个阶段容易烦。 另一个坑是 报错栈不总是友好。计算图一复杂,定位问题不像普通 Python 那么直接。 还有一个现实问题,现在你如果是新项目,往往得顺手看看 PyTensor,因为它已经明确是 Aesara 的 fork 了。继续死磕老名字,不一定划算。

所以我对 Aesara 的判断一直不是“超 nice,因为它强”,而是:

你手上的问题,恰好值得用“表达式 + 图优化 + 自动求导”这套打法时,它会显得非常顺手。 否则,它就是一把有点重的刀。

最后给个实用建议,真要上手,别一开始就啃大模型,也别先研究一堆后端优化参数。先做三步:

先拿标量和向量把 aesara.function() 跑通。 再拿一个你自己看得懂的损失函数,把 grad 打出来。 最后再碰 scan、自定义 Op、配置项这些稍微费劲的东西。文档里这些能力都有,但一口气全上,十有八九把自己绕进去。

Aesara 这种库,漂亮不在 API 有多短,漂亮在于你把公式交给它之后,后面的图、导数、执行链路,它真能接得住。