arXivDaily arXiv每日学术速递 周一至周五更新
arXiv周末暂无论文更新,休息一下吧,周末愉快~~
arXiv 2609.28248cs.LGcs.MS

hyperbolix:JAX 中的双曲深度学习

hyperbolix: Hyperbolic Deep Learning in JAX

Timo Klein, Thomas Lang, Yllka Velaj, Sebastian Tschiatschek

首次发表
浏览论文内容

中文总结 AI 辅助

hyperbolix 是 JAX 中首个全面的双曲深度学习库,提供多种流形、层族和优化器,并通过无取消公式在远距离保持 float32 精度。

中文摘要 AI 辅助

我们提出了 hyperbolix,一个基于 Flax NNX 构建的、用于 JAX 中双曲深度学习的开源库。据我们所知,它是 JAX 中第一个全面的、通用的双曲深度学习库。它包含六个具有共同接口的流形:欧几里得空间、庞加莱球、双曲面、κ-立体投影模型、混合曲率乘积空间以及固有速度空间。我们实现了涵盖线性层、卷积、注意力、归一化、位置编码、回归和向量量化的层族。这些构建模块涵盖了从 Ganea 的原始双曲神经网络到最近的全双曲架构(如 Hypformer 和洛伦兹 ResNet)的各种方法。此外,hyperbolix 还包含作为 optax 变换实现的黎曼优化器、包装分布以及双曲降维技术。其 API 使用惯用的 JAX:流形是无状态的,曲率在调用时传递,而流形操作作用于单个点,通过 jax.vmap 实现批量操作。每个检查过的操作的精度都针对 float32 和 float64 进行了测试,测试方式是与源论文中的闭式 NumPy/SciPy 转录或有限差分进行比较。在双曲面上,两点操作(如距离)的标准公式在远离原点时会失去精度,因为它们会减去两个大且几乎相等的项。hyperbolix 用无取消公式替换了这些减法,这些公式在 float32 下,在先前实现返回 NaN 的距离处仍能保持准确。hyperbolix 在 MIT 许可下可用,网址为 https://github.com/timoklein/hyperbolix 。

英文摘要

We present hyperbolix, an open-source library for hyperbolic deep learning in JAX, built on Flax NNX. To our knowledge, it is the first comprehensive, general-purpose hyperbolic deep learning library in JAX. It includes six manifolds with a common interface: Euclidean space, the Poincaré ball, the hyperboloid, the $κ$-stereographic model, mixed-curvature product spaces, and the proper velocity space. We implement layer families that cover linear layers, convolutions, attention, normalization, positional encoding, regression, and vector quantization. These building blocks span methods ranging from Ganea's original hyperbolic neural networks to recent fully hyperbolic architectures such as Hypformer and Lorentzian ResNet. Additionally, hyperbolix contains Riemannian optimizers implemented as optax transformations, wrapped distributions, and hyperbolic dimensionality-reduction techniques. Its API uses idiomatic JAX: Manifolds are stateless, with curvature being passed at call time, while manifold operations act on single points, with jax.vmap enabling batch operations. The precision of every checked operation is tested against a closed-form NumPy/SciPy transcription from the source paper or a finite difference, for both float32 and float64. On the hyperboloid, standard formulas for two-point operations, such as the distance, lose precision far from the origin, because they subtract two large, nearly equal terms. hyperbolix replaces these subtractions with cancellation-free formulas that stay accurate in float32 at distances where prior implementations return NaN. hyperbolix is available under the MIT license at https://github.com/timoklein/hyperbolix .

发表机构

  • University of Vienna(维也纳大学)

机构由 AI 辅助整理,请以论文原文为准。

↑