- 深度学习
- 机器学习
【免费下载链接】trax
Trax — Deep Learning with Clear Code and Speed
Trax 的trax.fastmath模块是整个框架的数学计算底座:它以 NumPy 风格的接口封装了卷积、池化、自动微分、并行映射等加速运算,并通过统一的后端抽象在 JAX、TensorFlow(tf-numpy)和纯 NumPy 之间自由切换。本文基于文档页docs/source/trax.fastmath.rst所指向的trax.fastmath.ops模块及其三个后端实现,完整介绍该模块的公开 API 面、后端选择机制(含 gin 配置与上下文管理器)、各后端的实现细节与回退策略,以及测试对跨后端行为一致性的验证方式。读完后,你可以在 Trax 中正确地选择、切换后端,理解每一类 fastmath 操作的底层实现,并编写跨后端可移植的模型代码。
fastmath 在 Trax 中的定位
Trax 的设计目标是"清晰代码 + 速度"(Deep Learning with Clear Code and Speed),其层(trax/layers/)、模型(trax/models/)、优化器(trax/optimizers/)等上层代码都通过fastmath调用底层数学运算,而不在业务代码里直接绑定某一个框架。模块的 docstring(trax/fastmath/ops.py)开宗明义:
Trax accelerated math operations for fast computing on GPUs and TPUs. Trax uses either TensorFlow 2 or JAX as backend for accelerating operations.
文档页 docs/source/trax.fastmath.rst 只有一行 Sphinx 指令.. automodule:: trax.fastmath.ops,其生成内容的主体正是trax.fastmath.ops的全部公开 API 及 docstring——也就是本文接下来逐一展开的内容。
快速上手:像 NumPy 一样使用加速运算
ops.py模块 docstring 给出的标准用法(trax/fastmath/ops.py):
from trax import fastmath from trax.fastmath import numpy as np x = np.array([1.0, 2.0]) # Use like numpy. y = np.exp(x) # Common numpy ops are available and accelerated. z = fastmath.logsumexp(y) # Special operations available from fastmath.要点有两处:
fastmath.numpy是一个"惰性代理"。它不是某个具体框架的 numpy 模块,而是 NumpyBackend 类的实例,其__getattr__在每次属性访问时才调用backend()['np']转发请求。源码中的注释解释了原因:必须惰性调用backend(),否则在 import 阶段就会解析后端,早于 gin 配置的解析时机,导致无法通过配置文件切换后端(trax/fastmath/ops.py)。fastmath.random同样是代理对象。RandomBackend 暴露get_prng、split、fold_in、uniform、randint、normal、bernoulli七个接口,同样转发到当前后端,保证随机数语义跨后端一致。
公开 API 面:automodule 文档涵盖的全部操作
按功能归类,trax/fastmath/ops.py的公开函数(含 docstring)如下表,这是文档页automodule实际生成的内容:
| 类别 | 函数 | 说明(源自 docstring) |
|---|---|---|
| 特殊函数 | logsumexp | 输入元素取指数求和后再取 log(L91-L93) |
expit/sigmoid | 计算 sigmoid(expit)函数,两者等价 | |
erf | 计算误差函数 | |
| 卷积与池化 | conv | 广义卷积 |
avg_pool/max_pool/sum_pool | 平均池化 / 最大池化 / 求和池化 | |
| 规约与选择 | top_k | Top-k 选择 |
sort_key_val | 沿维度对 key 排序,并对 value 施加相同置换 | |
| 控制流 | scan | 扫描,使循环函数在加速器上运行更快 |
map | 将函数映射到前导数组轴上 | |
fori_loop | 从lower到upper的编译型整数循环(L151-L179) | |
cond | 加速器上的条件计算 | |
remat | 反向传播时重算一切以省内存(激活重计算) | |
| 索引操作 | index_update/index_add/index_min/index_max | 不可变数组的索引更新/累加/取小/取大 |
dynamic_slice/dynamic_slice_in_dim/dynamic_update_slice/dynamic_update_slice_in_dim | 动态切片与切片更新 | |
lt | 供未重载<的后端使用的 less-than | |
| 梯度 | stop_gradient | 前向恒等、反向置零 |
jit | 即时编译函数供加速器使用 | |
disable_jit | 关闭 JIT 编译,便于调试 | |
vmap/grad/value_and_grad/vjp | 向量化 / 梯度 / 值与梯度 / 向量-雅可比积 | |
custom_grad/custom_vjp | 为函数设置自定义梯度 / 自定义 VJP | |
| 并行 | pmap/psum | 多加速器并行映射 / 并行求和归约 |
| 形状与设备 | abstract_eval | 仅按参数签名求值,返回签名(形状推断) |
dataset_as_numpy | 将tf.data.Dataset转为 numpy 数组流 | |
global_device_count/local_device_count | 返回全部主机 / 本机上的加速器数量 | |
| 后端选择 | Backend(枚举)/set_backend/backend/use_backend/backend_name/is_backend | 见下一节 |
其中fori_loop的 docstring 明确给出了语义(trax/fastmath/ops.py):
def fori_loop(lower, upper, body_fn, init_val): val = init_val for i in range(lower, upper): val = body_fn(i, val) return vallower为闭区间下界,upper为开区间上界,body_fn类型为(int, a) -> a,init_val是初始 carry 值。
此外,trax/fastmath/__init__.py从trax.fastmath.numpy额外导出了嵌套结构工具:nested_map、nested_map_multiarg、nested_stack、nested_zip、tree_flatten、tree_leaves、tree_unflatten(trax/fastmath/init.py),并在通配导入 ops 后使它们可直接以fastmath.nested_map(...)使用。
后端选择机制:gin、set_backend 与 use_backend
ops.py用一张字典把三种后端映射到各自的实现字典(trax/fastmath/ops.py):
_backend_dict = { Backend.JAX: JAX_BACKEND, Backend.NUMPY: NUMPY_BACKEND, Backend.TFNP: TF_BACKEND, }Backend枚举定义了三个合法取值(L40-L44):Backend.JAX = 'jax'、Backend.TFNP = 'tensorflow-numpy'、Backend.NUMPY = 'numpy'。
后端解析遵循一个明确的优先级链,backend()(L405-L418)按以下顺序决定:
override_backend:由上下文管理器use_backend(name)设置的临时覆盖(L421-L435)。它在finally中恢复原值,保证即使被包裹的代码抛异常也能正确还原——源码注释特别提到这一 try-finally 设计就是为测试场景考虑的。use_backend接受字符串(如'tensorflow-numpy')或Backend枚举,非法名称由_assert_valid_backend_name抛ValueError。default_backend:由set_backend(name)设置的进程级默认(L389-L394),传None可清除。- 函数参数
name:backend()自身带默认值name='jax',且标注了@gin.configurable——这意味着可以在 gin 配置中写backend.name = 'numpy'来全局切换后端,这是 Trax 配置驱动风格(配合trax/trainer_flags.py等入口)的一部分。
backend_name()与is_backend(Backend.X)则用于查询当前实际生效的后端。
一个重要的配套开关是disable_jit()(L245-L248):它把模块级_disable_jit置为真,此后fastmath.jit(f)直接返回f本身而不走后端的jit。docstring 说明其用途是调试——JIT 编译会掩盖逐语句执行时的错误,关掉它可让异常直接暴露。
三个后端逐一拆解
JAX 后端(默认)
JAX_BACKEND 是一个'name': 'jax'的实现字典,要点包括:
'np': jnp,即fastmath.numpy在 JAX 后端下就是jax.numpy;- 卷积由 jax_conv 包装
lax.conv_general_dilated实现,要求显式传入dimension_numbers(用'I'/'O'/'C'/'W'/'H'/'D'编码数据格式),且不允许输入扩张(lhs_dilation=None); - 池化统一走 _pooling_general 调用
lax.reduce_window:max_pool用lax.max、初值-inf;sum_pool用lax.add、初值0.;avg_pool在求和后由 _normalize_by_window_size 再用一次reduce_window数出每个窗口实际覆盖的样本数(以正确处理边界 padding),然后除回去——而不是简单除以pool_size; - 形状推断
abstract_eval由 jax_abstract_eval 实现:内部调用jax.eval_shape,再把结果用tnp.nested_map(signature, ...)逐叶转换为 Trax 的ShapeDtype(来自 trax/shapes.py); - 随机数全部来自
jax.random,其中random_get_prng被jax.jit包了一层(L205)以避免每次取 key 的编译开销;jax_randint 单独包装以把默认dtype固定为int32(与jax_random.randint的默认不同); - 索引操作统一映射为 JAX 的不可变
.at[]语法,如'index_add': lambda x, idx, y: jnp.asarray(x).at[idx].add(y)(L192-L195); - 自定义梯度经 _custom_grad(
jax.custom_transforms+defvjp_all)与 _custom_vjp(jax.custom_vjp+defvjp)接入。
TensorFlow 后端(tensorflow-numpy)
TF_BACKEND 的'np'指向trax.tf_numpy.numpy(即 Trax 自带的 TF2 NumPy 兼容层),运算大量来自 trax/tf_numpy/extensions.py。值得注意的实现细节:
jit被 _tf_jit 包装:会注入xla_forced_compile标志(可由set_tf_xla_forced_compile全局开关控制),并剥离 TF 不识别的donate_argnums参数;pmap同理(_tf_pmap)。_tf_grad支持argnums非 0 的情形:通过交换第 0 个与第argnums个参数、求导后再换回来实现(L110-L127)。random_fold_in没有直接对应物,_fold_in 用rng + sum(d)后 split 近似jax.random.fold_in——源码中的 TODO 提示该等价性尚未做严格的随机性质验证,属于使用时的已知限制。remat目前是空操作('remat': lambda f: f,L171),即 TF 后端下激活重计算不生效,TODO 表明支持方案仍在评估。- 设备计数用
max(len(tf_np_extensions.accelerators()), 1),保证无加速器时也返回至少 1。
纯 NumPy 后端(调试/单测)
NUMPY_BACKEND 是最小实现:'np'就是原生numpy,jit为恒等,logsumexp取自scipy.special,expit是1/(1+exp(-x))的 lambda。随机数函数(如 random_uniform)故意忽略传入的 rng,直接调用np.random.*;random_split返回一组None(L75)。get_prng 则把标量种子拆成两个uint32拼成 JAX 风格的 2 元素 key,保持 PRNG 接口的形状兼容。
它的abstract_eval是 np_abstract_eval:把每个输入替换成同形状全零张量后真跑一遍函数来推断输出形状——这是"从源码结构看"的朴素形状推断策略,意味着该后端的 dry-run 必须能在零值输入上无副作用地执行完。
关键实现中的降级与回退策略
ops.py的多个入口对"后端能力不齐"做了显式兜底,理解这些回退路径对跨后端开发很重要:
fori_loop回退到scan(L171-L179):若后端字典里没有'fori_loop'(JAX 与 TF 后端都没有独立实现),则构造一个把(i, x)推进为(i+1, body_fn(i, x))的 scanned 函数,用scan(..., length=upper - lower)等价执行。value_and_grad的合成回退(L261-L278):后端未提供时,用grad与原始fn合成;has_aux=True路径返回((res, aux), g)的元组形式。custom_vjp的nondiff_argnums兼容层(L291-L336):后端有custom_vjp时直接透传;否则校验nondiff_argnums必须是从 0 开始的连续前缀(只支持(0,)、(0, 1)这类形式,否则抛ValueError),然后退化到custom_grad实现,并用闭包处理非可微参数。源码中的 TODO 指出统一两种 API、最终移除nondiff_argnums是演进方向。dataset_as_numpy回退到 JAX 实现(L354-L358):TF 后端字典里该键被注释掉了(见 trax/fastmath/tf.py 的 TODO),因此实际总是走 trax/fastmath/jax.py 中基于tfds.as_numpy加dense_to_ragged_batch批量化的版本,TF 1.x 缺该 API 时再退化为逐样本迭代。jit的全局禁用开关:如前所述,disable_jit()后所有后端共享这一行为。
嵌套结构工具:让树状张量与后端解耦
trax/fastmath/numpy.py中的树工具与具体后端无关(仅依赖 dict/list/tuple/namedtuple),被__init__.py提升到包级:
nested_map(f, obj, level=0, ignore_nones=True)(L81-L114):对任意 dict/list/tuple 嵌套结构逐叶应用f,保留原始类型(包括 namedtuple),level控制停在第几层;nested_zip(objs)/nested_stack(objs, axis=0, np_module=np)(L146-L193):先把结构叶子两两 zip,再在level=1处用np_module.stack堆叠——np_module参数允许调用方传入jax.numpy使结果落在加速器上;tree_flatten/tree_leaves/tree_unflatten(flat, tree, copy_from_tree=None)(L196-L262):自定义的拍平/取叶/还原三件套。tree_unflatten的copy_from_tree参数支持从参考树拷贝"不关心的元素",docstring 举例:模型权重树中无权重层以()占位,用copy_from_tree=[()]即可从只含可训练权重的文件恢复完整模型——这是 Trax 序列化(如 trax/optimizers/trainer.py 保存/恢复权重)依赖的基础工具之一。
测试如何验证跨后端一致性
trax/fastmath/ops_test.py 中的BackendTest直接验证了上述机制的行为,可作为使用示例:
- gin 切换后端(test_backend_imports_correctly、test_numpy_backend_delegation):先断言默认后端下
backend['np']就是jnp;再gin.parse_config_files_and_bindings(None, "backend.name = 'numpy'")后断言它变成原生numpy,并且fastmath.numpy.isinf、fastmath.numpy.inf随之指向新后端——这正是NumpyBackend惰性代理存在的意义,每个测试setUp里都先gin.clear_config()防止串扰。 - 程序化设置(test_backend_can_be_set):
fastmath.set_backend('tensorflow-numpy')后backend_name()返回新值,set_backend(None)恢复'jax'。 - 跨后端语义一致性(test_fori_loop):用
parameterized.named_parameters在 JAX 与 TFNP 两个后端下分别执行fori_loop(2, 5, lambda i, x: x + i, 1),断言结果恒等于1 + 2 + 3 + 4——同一个 API 在两种后端下数值一致。 - 上下文管理器(test_use_backend_str):
with fastmath.use_backend('tensorflow-numpy'):内backend_name()为'tensorflow-numpy',退出后还原;既支持字符串也支持Backend枚举。 - 注册完整性(test_names_match):断言
_backend_dict中每个后端对象自带'name'字段与枚举值一致,且每个枚举成员都登记在字典中——防止新增后端时漏注册。
小结
trax.fastmath用一张"后端字典 + 惰性代理 + 优先级链"的组合,把 JAX、tf-numpy 与纯 NumPy 三种实现统一到一套 NumPy 风格 API 之下:默认后端是jax,可用 gin 配置(backend.name = 'numpy')、set_backend或use_backend上下文三种方式切换;特殊函数、卷积池化、控制流(scan/map/fori_loop/cond/remat)、索引、自动微分与多设备并行(pmap/psum)等公开操作都经过能力探测与回退处理,fori_loop→scan、value_and_grad合成、custom_vjp→custom_grad等降级路径使上层代码无需感知后端差异。编写模型层或研究新算子时(参考 trax/layers/core.py 等通过 fastmath 实现的层),应始终经由trax.fastmath而非直接 import 某个框架;调试时可用disable_jit()与纯 NumPy 后端定位问题,并参照 trax/fastmath/ops_test.py 的参数化写法为自己的算子补充跨后端一致性测试。
- 深度学习
- 机器学习
【免费下载链接】trax
Trax — Deep Learning with Clear Code and Speed
相关推荐
openJiuwen Agent Store 案例拆解:TripWise 如何用一套前端驾驭 5 种可切换 AI 后端
openJiuwen Agent Store 案例拆解:TripWise 如何用一套前端驾驭 5 种可切换 AI 后端 openJiuwen Agent Sto
示例工程ChatGLM-6B Mac部署指南:MPS后端GPU加速配置详解
ChatGLM 6B Mac部署指南:MPS后端GPU加速配置详解 ChatGLM 6B作为一款开源的双语对话语言模型,在Mac设备上通过MPS后端实现GPU加
大模型人工智能交互助手本地部署微调NLPPinLockView布局优化技巧:响应式设计与多设备适配终极指南
PinLockView布局优化技巧:响应式设计与多设备适配终极指南 PinLockView是一个简洁、极简且高度可定制的Android PIN锁视图库,为开发者
创作声明:本文部分内容由AI辅助生成(AIGC),仅供参考