@loganthorneloe: This is a excellent explanation of JAX. Understanding how ML frameworks work internally gives you a massive advantage w…
摘要
本文详细解释了JAX的核心思想,包括函数纯度、不可变性、显式状态管理和JIT编译,帮助读者从面向对象思维转向函数式编程以优化机器学习性能。
查看缓存全文
缓存时间: 2026/06/22 21:52
This is a excellent explanation of JAX.
Understanding how ML frameworks work internally gives you a massive advantage when optimizing models for production. This video breaks down JAX in a way that clicks.
https://t.co/JsQR9qMgce
TL;DR
JAX 通过强制函数纯度(纯函数)、不可变数组和显式状态管理来实现高性能机器学习,这要求编程思维从面向对象转向函数式,从而让编译器(如 XLA)进行深度优化,并保证可重复性与可并行性。
核心思维转变:从“对象操作”到“函数应用”
JAX 是一个面向高性能机器学习的 Python 库,但它与传统编程习惯有很大不同——它要求你进入 函数式编程 的世界。在函数式范式下,你不应把事物当作可操作的对象,而是把函数应用于对象(在机器学习中,这个对象通常是数据)。更重要的是,JAX 强制推行 函数纯度(purity) 的概念。
什么是纯函数?
纯函数是指:对于相同的输入,无论在任何情况下,它总是返回相同的输出,且不改变任何外部状态。你可能认为自己写的函数已经是这样了,但仔细检查:
- 是否修改了函数外部的变量?——不纯。
- 是否读取了外部值(例如全局参数)?——不纯。
所有输入必须显式传入,所有输出必须显式返回。这听起来限制重重,但正是获得极致性能的代价。因为函数无法做任何超出其内部定义的操作,JAX 能对程序进行更深入的优化,使其运行更快、扩展更好。
不可变性:JAX 数组的“副本”哲学
传统 Python 中,你可以直接修改 NumPy 数组的值;但在 JAX 中,数组不能在原地修改。因为修改数组相当于引用函数外部的东西,破坏纯函数特性。
JAX 使用诸如 x.at[idx].set(y) 的操作来创建一个带有修改的新数组,原数组保持不变。虽然从代码表面看这似乎在制造很多临时数组,但 JAX 的编译器(XLA)会执行深层优化,经常避免实际创建中间数组,从而带来显著的速度提升。
显式状态管理:将状态“贯穿”函数
在机器学习中,我们通常用类或全局对象保存模型参数,然后直接传递。但在 JAX 中,必须将参数显式地作为函数参数传入,并在输出中返回更新后的新状态。这种模式称为“显式贯穿状态”(explicit state threading)。
例如,一个训练步骤函数需要接受当前状态(参数、优化器状态等)和数据,然后返回更新后的状态以及结果。虽然代码可能更冗长,但在复杂系统(尤其是并行系统)中,它使推理更清晰、调试更容易。基于 JAX 构建的神经网络库(如 Flax)将模型设计为无状态蓝图,参数存储在外部并显式传递。
案例:伪随机数生成(PRNG)
PRNG 是显式状态管理的绝佳例子。在传统编程中,你可能在程序开始时设置一个随机种子,之后隐式地调用随机函数。但 JAX 要求你显式传递 PRNG 键——一个用作随机状态的数组。JAX 的随机函数读取这个键,但不修改它(因为修改会破坏纯度)。因此,提供相同的键总是生成相同的随机样本,保证可靠的重复性。
当你需要不同的统计独立样本时,必须显式地将键拆分为新的子键:key, subkey = jax.random.split(key)。一般经验法则是:除非你试图生成相同的输出,否则永远不要重复使用键。这种显式管理保证了可重复性和可并行性,对可靠的科学计算和大规模研究至关重要。
JIT 编译:纯函数的又一个例子
JIT(即时编译)是 JAX 性能的关键部分。它使用 XLA(加速线性代数)将兼容 JAX 的 Python 函数编译成高度优化的机器码。JIT 的编译发生在第一次调用(且仅第一次调用)时,过程称为跟踪(tracing):JAX 用一个抽象的跟踪对象执行函数,将操作记录为 jaxpr(JAX 表达式),然后编译成高效的机器码。
从某种意义上说,你实际运行的不是 Python 代码,而是机器码。你所写的 Python 代码在 JIT 期间被跟踪,由 XLA 编译,然后执行编译后的代码。由于跟踪只针对跟踪器所走的路径进行优化,如果你在 jitted 函数内部依赖隐式状态或全局变量,函数将不会按预期工作。
动态控制流注意事项
由于 JIT 编译只运行一次(第一次调用),并且专门针对所走的路径,在 jitted 函数内部使用依赖于 JAX 数组运行时值的 if-else 语句或循环可能导致错误,因为可能只有一个分支被编译。只有运行时依赖的值才会出现问题;固定值(例如常量)没问题。如果确实需要在运行时进行动态决策,应使用一组特殊函数:
jax.lax.cond(条件)jax.lax.while_loop(循环)
这些函数会正确处理动态控制流,避免编译错误。
结语与学习资源
掌握 JAX 的关键在于接受 函数纯度 这一思维转变。一旦你理解并应用纯函数、不可变性、显式状态管理以及 JIT 编译的含义,你就能充分利用 JAX 无与伦比的优化和扩展能力。
你在理解 JAX 时遇到的最大困难是什么?欢迎在评论区分享。如果想了解如何将数据集加载到 JAX 生态系统中,可以关注 Grain 库——Yufeng Guo 很快会发布相关视频。
Source: YouTube 视频
相似文章
@snowboat84: https://x.com/snowboat84/status/2065215177029787705
本文是AI工程全景系列的中篇,详细介绍了推理优化、模型瘦身(量化、蒸馏、剪枝、MoE)和投机解码等核心技术,综述了从硬件到工程栈的最新进展。
@GitHub_Daily: 大语言模型内部是如何工作的,为什么会产生幻觉,为什么有时答非所问,想深入了解这些。 可以看下 Awesome LLM Interpretability 这份资源合集,提供一整套拆解 AI 黑盒的系统路径。 涵盖从注意力可视化、神经元分析到…
介绍了一个名为 Awesome LLM Interpretability 的资源合集,汇集了多种可解释性工具、论文和社区资源,帮助理解大语言模型的内部工作机制。
@tanzhengmc97: https://x.com/tanzhengmc97/status/2066531753762656730
用通俗易懂的语言解释了大模型的运行原理,包括词向量、Transformer注意力机制、下一个词预测训练以及涌现能力,适合初学者理解AI基础概念。
@freeman1266: 不懂数学,也能看懂大多数 AI 论文——只要理解这条链路: token → embedding → 位置编码 → attention → FFN → 残差流 → next-token prediction LLM 本质上是把 Transf…
一条中文科普推文,用直观方式解释了LLM(大语言模型)的核心链路:从token、embedding、位置编码、attention、FFN到残差流和next-token prediction,帮助非数学背景读者理解AI论文。
@snowboat84: https://x.com/snowboat84/status/2075374060637503560
本文对AI可解释性进行了系统全面的综述,介绍了其需求(调试、合规、安全)、经典方法和前沿挑战,强调了忠实的解释比看起来合理的解释更重要。