@loganthorneloe: This is a excellent explanation of JAX. Understanding how ML frameworks work internally gives you a massive advantage w…

X AI KOLs Timeline 工具

摘要

本文详细解释了JAX的核心思想,包括函数纯度、不可变性、显式状态管理和JIT编译,帮助读者从面向对象思维转向函数式编程以优化机器学习性能。

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
查看原文
查看缓存全文

缓存时间: 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 视频

相似文章

@GitHub_Daily: 大语言模型内部是如何工作的,为什么会产生幻觉,为什么有时答非所问,想深入了解这些。 可以看下 Awesome LLM Interpretability 这份资源合集,提供一整套拆解 AI 黑盒的系统路径。 涵盖从注意力可视化、神经元分析到…

X AI KOLs Timeline

介绍了一个名为 Awesome LLM Interpretability 的资源合集,汇集了多种可解释性工具、论文和社区资源,帮助理解大语言模型的内部工作机制。