拓冰建站拓冰建站
首页 / 资讯中心 / 正文

Trax fastmath 详解:一套后端可切换的 GPU/TPU 加速数学 API

深度学习机器学习【免费下载链接】traxTrax — Deep Learning with Clear Code and Speed项目地址https://gitcode.com/gh_mirrors/tr/trax点击查看免费下载Trax 的trax.fastmath模块是整个框架的数学计算底座它以 NumPy 风格的接口封装了卷积、池化、自动微分、并行映射等加速运算并通过统一的后端抽象在 JAX、TensorFlowtf-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调用底层数学运算而不在业务代码里直接绑定某一个框架。模块的 docstringtrax/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.pyfrom 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输入元素取指数求和后再取 logL91-L93expit/sigmoid计算 sigmoidexpit函数两者等价erf计算误差函数卷积与池化conv广义卷积avg_pool/max_pool/sum_pool平均池化 / 最大池化 / 求和池化规约与选择top_kTop-k 选择sort_key_val沿维度对 key 排序并对 value 施加相同置换控制流scan扫描使循环函数在加速器上运行更快map将函数映射到前导数组轴上fori_loop从lower到upper的编译型整数循环L151-L179cond加速器上的条件计算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.pydef 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) - ainit_val是初始 carry 值。此外trax/fastmath/__init__.py从trax.fastmath.numpy额外导出了嵌套结构工具nested_map、nested_map_multiarg、nested_stack、nested_zip、tree_flatten、tree_leaves、tree_unflattentrax/fastmath/init.py并在通配导入 ops 后使它们可直接以fastmath.nested_map(...)使用。后端选择机制gin、set_backend 与 use_backendops.py用一张字典把三种后端映射到各自的实现字典trax/fastmath/ops.py_backend_dict { Backend.JAX: JAX_BACKEND, Backend.NUMPY: NUMPY_BACKEND, Backend.TFNP: TF_BACKEND, }Backend枚举定义了三个合法取值L40-L44Backend.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可清除。函数参数namebackend()自身带默认值namejax且标注了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_dilationNone池化统一走 _pooling_general 调用lax.reduce_windowmax_pool用lax.max、初值-infsum_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_gradjax.custom_transformsdefvjp_all与 _custom_vjpjax.custom_vjpdefvjp接入。TensorFlow 后端tensorflow-numpyTF_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: fL171即 TF 后端下激活重计算不生效TODO 表明支持方案仍在评估。设备计数用max(len(tf_np_extensions.accelerators()), 1)保证无加速器时也返回至少 1。纯 NumPy 后端调试/单测NUMPY_BACKEND 是最小实现np就是原生numpyjit为恒等logsumexp取自scipy.specialexpit是1/(1exp(-x))的 lambda。随机数函数如 random_uniform故意忽略传入的 rng直接调用np.random.*random_split返回一组NoneL75。get_prng 则把标量种子拆成两个uint32拼成 JAX 风格的 2 元素 key保持 PRNG 接口的形状兼容。它的abstract_eval是 np_abstract_eval把每个输入替换成同形状全零张量后真跑一遍函数来推断输出形状——这是从源码结构看的朴素形状推断策略意味着该后端的 dry-run 必须能在零值输入上无副作用地执行完。关键实现中的降级与回退策略ops.py的多个入口对后端能力不齐做了显式兜底理解这些回退路径对跨后端开发很重要fori_loop回退到scanL171-L179若后端字典里没有fori_loopJAX 与 TF 后端都没有独立实现则构造一个把(i, x)推进为(i1, body_fn(i, x))的 scanned 函数用scan(..., lengthupper - lower)等价执行。value_and_grad的合成回退L261-L278后端未提供时用grad与原始fn合成has_auxTrue路径返回((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-L358TF 后端字典里该键被注释掉了见 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, level0, ignore_nonesTrue)L81-L114对任意 dict/list/tuple 嵌套结构逐叶应用f保留原始类型包括 namedtuplelevel控制停在第几层nested_zip(objs)/nested_stack(objs, axis0, np_modulenp)L146-L193先把结构叶子两两 zip再在level1处用np_module.stack堆叠——np_module参数允许调用方传入jax.numpy使结果落在加速器上tree_flatten/tree_leaves/tree_unflatten(flat, tree, copy_from_treeNone)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_setfastmath.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_strwith 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 的参数化写法为自己的算子补充跨后端一致性测试。赞分享深度学习机器学习【免费下载链接】traxTrax — Deep Learning with Clear Code and Speed项目地址https://gitcode.com/gh_mirrors/tr/trax点击查看免费下载相关推荐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锁视图库为开发者上一篇Litestar 依赖注入实战分层声明、Provide 包装器与 yield 清理机制全解析下一篇如何轻松实现VLC视频点击控制Pause Click插件的完整解决方案创作声明:本文部分内容由AI辅助生成(AIGC),仅供参考
分享:

看完干货,该让你的企业上线了

免费需求沟通 · 48 小时内出具建站方案 · 河南本地可上门