【智能纪元厅】
2018年,JAX 高性能计算正式问世。由Google主导开发,面向跨平台平台用户。
可微分编程的先锋,让GPU加速与自动求导完美融合。
技术特色:高性能计算、自动微分、开源。
影响力评估:技术维度 9/10,商业维度 6/10,文化维度 5/10,用户维度 5/10。
作为智能纪元厅的经典代表,JAX 高性能计算在软件发展史上留下了深刻的印记。
可微分编程的先锋,让GPU加速与自动求导完美融合。
【智能纪元厅】
2018年,JAX 高性能计算正式问世。由Google主导开发,面向跨平台平台用户。
可微分编程的先锋,让GPU加速与自动求导完美融合。
技术特色:高性能计算、自动微分、开源。
影响力评估:技术维度 9/10,商业维度 6/10,文化维度 5/10,用户维度 5/10。
作为智能纪元厅的经典代表,JAX 高性能计算在软件发展史上留下了深刻的印记。
2018年的深度学习世界,正处在一个微妙的十字路口。TensorFlow和PyTorch的“框架战争”如火如荼,研究人员在动态图的灵活性与静态图的性能之间反复权衡。GPU的算力虽在飞速增长,但绝大多数代码仍像在沙滩上写字——每次运行都要重新解释Python,而硬件的真正潜力被一层层抽象所掩盖。就在这个时刻,一个来自Google的小团队,带着一个看似叛逆的想法,悄然发布了JAX。它不叫“框架”,而自称“库”;它不提供高级API,却承诺让NumPy代码在GPU上飞驰,并附赠自动微分。这个当时只有几千行核心代码的项目,像一颗种子,将在未来五年里长成改写高性能机器学习计算规则的参天大树。
要理解JAX的诞生,必须先回到它孕育的技术土壤。2017年,Google的XLA(Accelerated Linear Algebra)编译器已经能够将TensorFlow的计算图编译成针对GPU和TPU优化的机器码,但它的使用体验并不理想——用户需要显式构建静态图,调试如同在黑箱中摸索。与此同时,DeepMind的研究团队正饱受性能瓶颈的折磨。他们需要一种能够灵活表达复杂计算(如强化学习中的环境模拟、物理引擎中的微分方程求解),同时又能自动利用硬件加速的工具。当时的主流方案要么牺牲灵活性换取性能(如TensorFlow的静态图),要么牺牲性能换取灵活性(如PyTorch的动态图)。有没有可能两者兼得?
答案藏在Google内部一个代号为“JAX”的实验性项目中。核心开发者包括Matt Johnson、Roy Frostig、Dougal Maclaurin和Chris Leary等人,他们大多来自Google Brain和XLA团队。这些人有一个共同特点:痴迷于函数式编程和编译器技术。Matt Johnson曾参与过NumPy的早期开发,对Python数值计算的底层机制了如指掌;Roy Frostig和Dougal Maclaurin则在自动微分领域深耕多年,写过名为“Autograd”的库——一个能够对NumPy代码进行自动微分的工具。Autograd在学术界小有名气,但它只能运行在CPU上,且无法利用GPU加速。团队意识到,如果将Autograd的自动微分能力与XLA的编译优化结合起来,就能创造出一个全新的工具:它看起来像NumPy,用起来像Python函数,但底层却能将计算编译成高效的硬件指令。
JAX的核心设计哲学可以用四个词概括:函数式、可组合、可编译、可微分。这听起来像一句口号,但每一个词背后都是深思熟虑的技术决策。函数式意味着JAX中的计算被建模为纯函数——没有副作用,没有可变状态。这在Python中几乎是“反直觉”的,因为NumPy本身支持原地修改数组。但正是这种约束,让JAX能够安全地执行JIT编译、自动并行化和自动微分。如果你写了一个修改数组的函数,JAX会直接报错,迫使你改用函数式写法。这种“不近人情”的设计,在初期劝退了大量开发者,却为后来的高性能计算铺平了道路。
JAX的另一大创新是它的转换系统。它提供四个核心转换函数:grad(自动微分)、jit(即时编译)、vmap(向量化映射)和pmap(并行映射)。这些转换可以任意组合,像乐高积木一样搭出复杂的计算流程。例如,你可以先定义一个计算损失函数的Python函数,用grad得到它的梯度函数,再用jit将其编译成GPU内核,然后用vmap自动批量化处理,最后用pmap分布到多台TPU上。整个过程只需要几行代码,且完全在Python层面完成,没有图构建、没有会话执行、没有上下文管理器。这种“一切皆函数,一切皆可转换”的理念,让JAX在表达力上达到了前所未有的高度。
2018年6月,JAX在GitHub上首次公开亮相。最初版本只有寥寥数千行代码,文档也相当简陋。但它的核心能力已经清晰可见:你可以用几乎与NumPy完全相同的API编写代码,然后通过一行@jit装饰器,让函数在GPU上以接近硬件极限的速度运行。这种体验对于习惯了Python缓慢循环的研究人员来说,无异于发现新大陆。DeepMind是最早的“吃螃蟹者”。他们的强化学习团队需要频繁运行环境模拟,而传统的Python模拟器速度极慢。用JAX重写后,模拟速度提升了数百倍,且自动微分能力让他们能够直接对模拟过程本身进行梯度计算,这在之前几乎是不可能的。
真正让JAX从“小众工具”走向“行业标准”的,是2020年之后大模型训练的爆发。当GPT-3、PaLM等超大规模模型出现时,传统框架的局限性暴露无遗。TensorFlow的静态图在分布式训练中效率不错,但调试和扩展极为困难;PyTorch的动态图灵活,但在多机多卡场景下需要大量手动优化。JAX的函数式设计和编译优化,恰好命中了大模型训练的核心需求:计算图的可重排性、内存的精细管理、跨设备的高效通信。特别是JAX的“检查点”机制——它允许用户在前向传播中丢弃中间结果,在反向传播时重新计算,从而在显存和计算量之间做出灵活权衡。这一特性对于训练数百亿参数的模型至关重要。
2021年,Google发布了PaLM(Pathways Language Model),一个拥有5400亿参数的大模型。PaLM的训练使用了6144块TPU v4芯片,而JAX正是其底层的计算引擎。据公开资料,PaLM的训练效率达到了硬件理论峰值的57.8%,这是一个令人震惊的数字。相比之下,同期使用其他框架的大模型训练效率通常在30%-40%之间。JAX的编译优化能力在这里发挥了关键作用:它能够自动融合计算核、消除冗余内存拷贝、智能调度通信与计算。更重要的是,JAX的函数式纯计算特性,使得分布式训练中的梯度同步和模型并行变得异常简洁——研究人员只需要用pmap将计算函数映射到多个设备上,剩下的通信细节由框架自动处理。
JAX的崛起并非一帆风顺。它面临的最大挑战是生态的贫瘠。TensorFlow和PyTorch拥有庞大的社区、丰富的预训练模型和成熟的工具链,而JAX就像一个刚出生的婴儿,除了核心库外几乎一无所有。为了解决这个问题,Google和DeepMind联手推动了几个关键的生态项目。Flax和Haiku是专门为JAX设计的高级神经网络库,提供了类似PyTorch的模块化接口;Optax提供了优化器集合;Chex提供了测试工具;Orbax提供了模型保存和加载功能。这些库共同构成了JAX的“扩展宇宙”,让研究人员能够在不牺牲性能的前提下,享受到类似主流框架的开发体验。
另一个有趣的轶事是JAX社区的文化。由于JAX本身极度推崇函数式编程,它的用户群体中聚集了大量对FP(Functional Programming)有热情的程序员。在JAX的GitHub issue和Discord频道里,你经常能看到关于“纯函数与副作用”、“类型系统与编译器优化”的热烈讨论。这种技术品味上的“精英主义”,让JAX在初期显得有些高冷,但也吸引了一批顶尖的研究者和工程师。他们愿意忍受文档不完善、调试困难、生态不成熟等问题,只因为JAX赋予了他们前所未有的计算表达力。
从市场影响来看,JAX虽然没有像PyTorch那样成为“全民框架”,但在高性能计算和前沿研究领域,它的地位已经无可撼动。根据2023年的一项调查,在NeurIPS、ICML等顶级机器学习会议上,使用JAX的论文比例从2019年的不到1%增长到了2023年的约15%。在强化学习、物理模拟、计算生物学等对计算效率极度敏感的领域,JAX的渗透率甚至更高。DeepMind几乎所有的核心项目——从AlphaFold到AlphaGo的后续版本——都基于JAX构建。Google Research的多个大模型训练基础设施,也全面转向了JAX。
JAX的文化遗产,或许比它的商业成功更为深远。它重新定义了“深度学习框架”的边界:不再是一个包含所有工具的大一统平台,而是一个专注于编译器优化和自动微分的“计算核心”。这种“小而精”的理念,影响了后续多个项目,如PyTorch的TorchDynamo和TensorFlow的TFXLA。更重要的是,JAX证明了函数式编程在机器学习领域的巨大潜力。它让“计算”本身成为可以被转换、组合、优化的第一类对象,这为未来的自动机器学习、可微分编程甚至神经符号系统打开了新的可能性。
站在2025年回望,JAX的故事远未结束。随着硬件架构的多样化(GPU、TPU、NPU、量子处理器)和计算需求的指数级增长,JAX所代表的“编译优先、函数式核心”的路径,可能会成为高性能计算的主流范式。它或许永远不会成为“大众情人”,但它为那些追求极限性能的研究者提供了一个近乎完美的工具——一个让你能够用Python的优雅,写出接近C++性能的代码,同时还能自动求导的工具。这,就是JAX在科技史上留下的独特印记。
对技术发展和工程实践的推动程度
对商业模式和市场格局的影响深度
在科技文化和社会层面的持久影响力
用户群体的广度和普及程度