Jax
PulseAugur coverage of Jax — every cluster mentioning Jax across labs, papers, and developer communities, ranked by signal.
- used by Fortran 90%
- used by Flax 90%
- used by Pallas 90%
- used by reinforcement learning 90%
- developed alphaXiv 90%
- used by NumPy 70%
- used by Orbax Distributed Checkpointing With Jax 70%
- used by vLLM 70%
- used by NVIDIA H100 70%
- used by Hugging Face Transformers 70%
- used by AI accelerator 70%
- instance of AI accelerator 70%
10 天有情绪数据
-
BluffJAX 套件为不完美信息博弈中的强化学习带来新挑战
研究人员开发了 BluffJAX,这是一个用 JAX 实现的对抗不完美信息博弈的开源套件。该套件专为高模拟吞吐量和 GPU 并行化而设计,提供了德州扑克和库恩扑克等既定基准的规范实现,以及诸如 Bluff、Stud Poker 和 Kemps 等研究较少的博弈。目标是通过基准测试性能和为各种算法提供基线结果,为博弈论方法中的强化学习研究带来新挑战,并促进比较。
-
新的无模拟方法能更快地学习种群动力学
研究人员开发了一种名为 Double-Stitch 的新无模拟方法来学习种群动力学。该技术利用 Wasserstein 拉格朗日残差,从不成对的快照中重建和外推概率分布(如细胞或流体)的演变。与先前在每个训练步骤都需要运行数值求解器的基于模拟的方法不同,Double-Stitch 通过沿学习路径惩罚运动方程残差来显著加快训练速度。该方法在包括合成、单细胞和海洋涡旋数据在内的各种数据集上,均表现出与现有方法相当或更优的性能。
-
新的CutBCE方法优化推荐系统训练,减少内存使用并加速计算
研究人员开发了一种名为CutBCE的新颖方法,用于优化大型词汇推荐系统中的二元交叉熵(BCE)损失。该方法解决了在海量商品目录上训练模型时出现的内存限制和内存溢出错误。CutBCE利用JAX和Pallas通过硬件加速,实现精确的融合重构和自定义向量-雅可比积(Vector-Jacobian Product)的片上计算,以避免在大带宽内存(High Bandwidth Memory)中存储大型张量。在TPU v5e/v6e和8芯片TPU…
-
新基准评估人工智能在解谜方面的推理能力,结果喜忧参半
研究人员开发了 PuzzleJAX,这是一个 GPU 加速引擎和领域特定语言,用于评估人工智能的推理和学习能力。该引擎支持动态编译游戏,能够处理各种直观但对人工智能模型具有挑战性的任务。另外,Nonobench 发布了一个开放基准,评估了 49 种大型语言模型在数图谜题上的表现,结果显示随着谜题复杂度的增加,解题率显著下降。
-
Google's Gemma 4 models repacked for enhanced performance on single TPU v5e
一份技术指南详细介绍了如何重新打包Google的量化感知训练(QAT)Gemma 4模型,以提高在单个Google Cloud TPU v5e芯片上的性能。重新打包后的模型,特别是12B参数版本,与bfloat16版本相比,在吞吐量和准确性方面具有竞争力,其中12B模型成为能装入该芯片的最大模型。这种优化允许在单个TPU v5e上部署各种大小的Gemma 4模型,从E2B到26B,为高效部署这些模型提供了一种实用的方法。
-
新的 JAx 方法通过对齐预测来加速扩散模型训练
研究人员推出了一种新颖的扩散模型预测监督方法 JAx(Just Align x),该方法可以对齐不同噪声水平下的干净图像预测。与表示对齐不同,JAx 专注于改进预测目标本身,从而实现更稳定和加速的训练。该方法在 ImageNet 256x256 的各种 JiT 配置上,在 Fréchet inception distance (FID) 和收敛速度方面均表现出一致的改进,且无需架构更改或外部编码器。
-
Google Research 的 Kauldron:用于模块化的 JAX 训练库
本教程提供了 Google Research 的 Kauldron 指南,这是一个专为研究速度和模块化设计的 JAX 训练库。它详细介绍了 Kauldron 的核心机制:konfig 用于将实验转换为 JSON 可序列化字典,kontext 用于通过字符串路径连接组件以避免直接导入,以及一个带有命名轴的运行时形状检查器。该指南演示了如何实现自定义损失和指标,在没有加速器的情况下对合成数据进行模型训练,以及运行实验扫描。
-
新的JAX框架vidax优化了Cloud TPUs的视频生成
研究人员开发了vidax,一个使用JAX和Flax构建的新开源框架,旨在优化Cloud TPU Pod上的视频生成模型。该引擎包含一个用于PyTorch权重的零拷贝转换器,能够高效地推理各种时空模型(如Diffusion Transformers和3D VAEs),而无需在执行期间依赖PyTorch。Vidax统一了张量和序列并行,集成了TPU特定内核,并实现了权重卸载以处理高分辨率,同时提供了TPU v4-8硬件的基准测试。
-
新研究探索协作无人机群的高级训练 · 2 篇论文
两篇新研究论文探讨了训练和部署协作无人机群的高级方法。第一篇论文介绍了 AeroWeaver,一种具身智能体工具,可将大型语言模型的决策连接到可执行的无人机技能,从而在没有中央控制代理的情况下实现分布式协调和自适应学习。第二篇论文提出了一种混合保真度训练方案,使用低保真度模拟器,并通过简短的高保真度校准飞行的残差学习进行校正,显著降低了计算成本并提高了各种团队规模下的性能。
-
新方法通过不确定性量化校准海洋模型
研究人员开发了一种使用基于仿真的推理(SBI)来校准单柱海洋模型的新方法。该方法通过量化参数估计相关的不确定性,解决了先前方法的局限性,这在逆问题适定性差时至关重要。该研究将SBI应用于基于JAX的`tunax`海洋模型,以校准其k-epsilon闭合的系数,并利用分块主成分分析来压缩模拟器输出,使推理变得可行。
-
新的JAX库加速AI临时协作研究
研究人员开发了JaxAHT,一个基于JAX的新开源库,旨在加速和标准化人工智能中临时协作(AHT)的研究过程。该库旨在克服先前阻碍AHT进展的计算成本和标准化基准的缺乏。JaxAHT提供了一个统一的框架,用于生成队友、训练自我代理,并评估它们在面对未见过伙伴时的表现,与PyTorch实现相比,速度显著提升。该库还包括一套跨越Level-Based Foraging、Overcooked和Hanabi等不同领域的评估队友,并用于对不同的…
-
Sakana AI 提出层局部训练方法以训练 1000 层网络
Sakana AI 的研究人员开发了一种名为增强拉格朗日预测编码 (PC-ALM) 的新颖训练方法,它提供了一种传统的反向传播的层局部替代方法。这种新方法允许训练非常深(多达 1000 层)的神经网络,同时保持接近反向传播的性能。该方法已在 MNIST 和 Fashion-MNIST 等小型图像基准测试中得到验证,与标准的预测编码相比,在深层窄网络架构中显示出显著的改进。
-
JAX3D 支持分层 NeRF,实现高级三维渲染和重建
研究人员开发了一种使用 JAX 和 jax3d 库创建分层神经辐射场(NeRF)的方法。该方法支持体积渲染、新视角合成和三维重建。本教程详细介绍了构建合成数据集、实现带有位置编码和分层采样的 NeRF,以及使用 JAX 的 JIT 编译和 Adam 优化进行训练的过程。结果通过 PSNR 等指标以及深度、不透明度和提取几何形状的视觉输出来评估。
-
新智能体HORIZON通过分层信念建模增强多智能体导航能力
研究人员开发了HORIZON,一种专为Lux AI第三赛季竞赛设计的层级智能体,该竞赛要求在部分可观察的多智能体导航场景中进行适应。该智能体采用多方面方法,包括空间感知、信念跟踪、图注意力以及探索策略,将短期控制与长期推理分开。HORIZON使用JAX模拟器中的Proximal Policy Optimization进行训练,与现有基线相比,在胜率和适应能力方面均有显著提升。
-
中国推出AI4S科学计算平台及台风追踪系统
太初元启发布了其AI4S计算平台,该平台专为气象学和量子力学等领域的科学研究而设计。平台采用专有的异构AI芯片,并支持多种主流框架和CPU,旨在用国产硬件赋能科学计算。此次发布还包括一项重要更新——TecoWeatherNext台风追踪系统,该系统提供了集成的数据采集、AI预测和台风识别功能。
-
新框架利用人工智能优化化学品输运过程
研究人员开发了一个新的可微分混合建模框架,旨在提高化学品输运过程的准确性和优化水平。该框架集成了JAX有限体积求解器和神经网络组件,可以直接从实验数据中学习本构定律和初始条件,克服了传统模型的局限性。该系统的可微分性还允许通过直接调整实验设置以获得期望的结果来实现过程优化,显示出在涉及质量、能量和动量输运的各种化学分离应用中的潜力。
-
GRADSOLVE库加速GPU上ODE梯度计算
一个名为GRADSOLVE的新型开源JAX库已被开发出来,用于加速NVIDIA GPU上常微分方程(ODE)集成精确梯度的计算。该库解决了现有GPU软件在快速ODE求解和高效梯度计算之间进行权衡的性能瓶颈。GRADSOLVE通过记录求解器步骤然后执行固定步重放来实现更快的梯度计算,与Diffrax等标准方法相比,微分速度显著提高。
-
新方法从3D人体模型中提取生物力学姿态
研究人员开发了一种新方法,可以从单个RGB图像中提取生物力学上准确的关节角度,解决了当前3D人体恢复技术的局限性。该方法通过添加一个生物力学预测头来扩展SAM 3D Body基础模型。为了在没有生物力学标签数据的情况下训练这个头,他们使用了自监督蒸馏,通过优化逆运动学拟合来匹配来自无标签图像的网格预测。该模型使用JAX实现,并结合Equinox与MuJoCo一起使用,在SAM-3D-Body数据集上进行了训练,并在MoVi、BioCV…
-
新的EMR-HyperNEAT方法通过张量化加速神经演化
研究人员开发了EMR-HyperNEAT,这是一种新颖的神经演化方法,可显著加速演化大规模神经网络基底的过程。该新方法通过预先评估所有分辨率下的所有位置,克服了先前基于四叉树技术的局限性,实现了并行处理并获得了显著的加速。实验表明,在XOR任务上GPU性能提高了12-34倍,并在各种基准测试中提高了求解率。
-
Google Cloud 在 TPU 上启用 vLLM 以支持 Qwen3 长上下文嵌入
Google Cloud 为优化嵌入推理的张量处理单元(TPU)引入了原生的 vLLM 支持,旨在用于生产检索系统。此次更新侧重于增强 Qwen3-Embedding-8B 和 Qwen3-VL-Embedding-8B 等模型的长上下文和多模态嵌入,解决了内存压力和池化正确性等挑战。Google 的工程努力包括优化张量对齐、编译预热以及混合 StepPool 设计,以确保数学正确性和高性能,其中一种配置实现了超过 83,000 to…