ARTICLE DETAIL

资讯详情

深耕编程入门与网站建设的一线实战洞察。

JAX `jax.extend.linear_util` 模块解析:WrappedFun、变换堆栈与记忆化缓存机制

JAX `jax.extend.linear_util` 模块解析:WrappedFun、变换堆栈与记忆化缓存机制 JAXjax.extend.linear_util模块解析WrappedFun、变换堆栈与记忆化缓存机制【免费下载链接】jaxComposable transformations of PythonNumPy programs: differentiate, vectorize, JIT to GPU/TPU, and more项目地址: https://gitcode.com/GitHub_Trending/ja/jax导读jax.extend.linear_util是 JAX 内部实现中一套用于定义可与变换transformation组合的函数的基础设施它将一个普通 Python 函数包装成WrappedFun通过生成器generator式的变换堆栈在调用时统一改写参数与返回值并利用哈希相等语义支撑函数级记忆化memoization。本文以仓库中 docs/jax.extend.linear_util.rst 列出的公开 API 为骨架结合 jax/_src/linear_util.py 的源码实现系统讲解wrap_init、WrappedFun、transformation/transformation2/transformation_with_aux/transformation_with_aux2、cache、merge_linear_aux与StoreException的设计原理、调用方式和实际应用场景。读完本文你将理解 JAX 的jit、vmap、grad等变换底层是如何叠加与记忆化的并能在自己的扩展代码中正确使用这套机制。一、模块定位jax.extend与linear_util的角色1.1 从公共 API 到内部机制jax.extend是 JAX 向扩展开发者暴露内部机制的统一入口。在 jax/extend/init.py 中linear_util与backend、core、lowering、mlir、random、sharding等模块一起被导出。需要注意它的 API 政策与公共 API 不同该模块不提供跨版本的兼容性保证破坏性变更会通过 JAX 的 changelog 公布因此只适合了解原理或在扩展场景中使用不适合写入对外稳定的公共接口。1.2jax.extend.linear_util与jax._src.linear_util的关系仓库中有两份linear_utiljax/_src/linear_util.py —— 真正的实现约 500 行包含全部类与函数jax/extend/linear_util.py —— 薄薄的转发层用from jax._src.linear_util import ...原样再导出符号。从 jax/extend/linear_util.py 可以看到StoreException、WrappedFun、cache、merge_linear_aux、transformation、transformation_with_aux、transformation2、transformation_with_aux2都是直接转发的只有wrap_init是例外——它被包了一层免 DebugInfo的兼容版本见下文 4.3 节。这一等价关系也由测试 tests/extend_test.py 中的test_symbols用assertIs逐项验证。二、核心概念WrappedFun 与变换堆栈2.1 一个函数 一串变换 WrappedFun模块的模块级 docstringjax/_src/linear_util.py给出了最核心的定义一个WrappedFun对象表示函数f同时携带一系列嵌套的变换transformation。这些变换在调用时作用于位置参数与关键字参数在返回时作用于函数的返回值。典型用法是from jax._src import linear_util as lu from jax._src import api_util # 把 f 包装成可施加变换的 WrappedFun wf lu.wrap_init(f, debug_infoapi_util.debug_info(test, f, (), {}))WrappedFun的构造参数见 jax/_src/linear_util.py包括字段含义f被变换的原始函数f_transformed叠加了全部变换后的可调用对象call_wrapped实际调用的就是它transforms(gen, gen_static_args)元组的元组表示要施加的变换栈stores各变换辅助输出auxiliary output的存放槽位类型为Store/EqualStore/Noneparams以(name, param)形式记录的额外关键字参数调用f时与变换后的 kwargs 一起传入in_type可选的输入类型core.InputType用于缓存键debug_info关于被包装函数及其参数、结果的调试信息2.2 调用时的数据流后进先出结果反向WrappedFun.call_wrapped(*args, **kwargs)jax/_src/linear_util.py只是转发给self.f_transformed。真正的数据流遵循变换栈语义参数方向动态参数和关键字参数先被最后施加的变换处理再逐层向前最后到达原始函数f结果方向f的返回值先被最先施加的变换处理再逐层向后。文档中的一句话概括了这种对称性如果有多个变换它们构成一个栈。参数先被最后应用的变换处理结果先被最先应用的变换处理。这正是jit(vmap(f))这类组合变换能够正确叠加的原因。2.3 变换的生成器协议一个变换 yield 两次一个变换被实现为生成器函数它接受零个或多个静态位置参数实例化变换时给定外加要变换的位置/关键字参数并且恰好yield两次lu.transformation_with_aux def trans1(static_arg, *dynamic_args, **kwargs): # ... 根据 static_arg 与参数计算新的参数 ... # 第一次 yield给出变换后的 (args, kwargs)并从 send() 拿回函数调用的结果 results yield (new_dynamic_args, new_kwargs) # ... 对 results 做后处理 ... # 第二次 yield给出 (变换后的结果, 辅助输出) yield new_results, auxiliary_output配合transformation的非 aux 版本第二次 yield 只需给出变换后的结果即可。静态参数在包装时确定WrappedFun会通过functools.partial把它们固化为f_transformed的一部分见 jax/_src/linear_util.py 的wrap方法。三、如何施加变换四个 transformation 入口docs/jax.extend.linear_util.rst的 autosummary 列出了四个变换入口它们的分工如下3.1transformation2无辅助输出的推荐路径transformation2(gen, fun, *gen_static_args)被curry装饰jax/_src/linear_util.py即可以先只给gen得到接受fun的函数。它调用fun.wrap(gen, gen_static_args, None)把新变换压入栈顶不产生辅助输出。3.2transformation向后兼容的无辅助版本transformation在源码中明确标注 Backwards compat onlyjax/_src/linear_util.py。它把用户生成器gen包进一个适配生成器gen2next(gen_inst)取第一次 yield 的(args_, kwargs_)调用底层f后通过gen_inst.send(...)送回结果再把生成器最终产出作为返回值。新代码应优先使用transformation2。3.3transformation_with_aux2带辅助输出的推荐路径transformation_with_aux2(gen, fun, *gen_static_args, use_eq_storeFalse, unk_namesFalse)返回(新的 WrappedFun, out_thunk)jax/_src/linear_util.py内部创建一个Store或use_eq_storeTrue时的EqualStore用于存放辅助输出out_thunk是一个零参闭包调用它即可取出辅助输出unk_namesTrue时会把DebugInfo中的参数名/结果路径置为未知with_unknown_names适用于名字无法追踪的场景。3.4transformation_with_aux向后兼容的带辅助输出版本同样标注 Backwards compat onlyjax/_src/linear_util.py。区别在于适配生成器gen2额外接收一个store参数把生成器第二次 yield 出的aux存入 store 后再返回结果。它最终也是委托给transformation_with_aux2完成压栈。四、辅助输出的存储与合并Store、StoreException 与 merge_linear_aux4.1Store单次写入的槽位Storejax/_src/linear_util.py是一个极简的单槽容器store(val)向槽位写入值若槽位已占用则抛出StoreException(Store occupied)val属性读取值若槽位为空则抛出StoreException(Store empty)reset()仅在调试等异常场景下清空槽位。EqualStorejax/_src/linear_util.py放宽了规则允许重复写入相等的值否则抛出StoreException(Store occupied with not-equal value)。这为缓存命中时重复回填辅助输出提供了便利。4.2StoreExceptionStoreException只是Exception的空子类jax/_src/linear_util.py用于区分存储槽位被占用/为空等与控制流无关的异常。4.3merge_linear_aux合并两个辅助输出 thunk在变换组合中经常出现两个分支各有一个辅助输出 thunk但最终只能保留一个的情况。merge_linear_aux(aux1, aux2)jax/_src/linear_util.py正是为此设计恰好一个 store 被占用返回(True, out1)或(False, out2)布尔值标识来自哪个分支两个 store 都空抛StoreException(neither store occupied)两个 store 都占用抛StoreException(both stores occupied)。这保证了最多取一个分支的辅助输出是条件分支如lax.cond类场景中传递 aux 的标准做法。五、记忆化cache装饰器与缓存键5.1 为什么 WrappedFun 可以当字典键WrappedFun实现了值语义的__hash__与__eq__jax/_src/linear_util.py只有f、transforms、params、in_type、debug_info全部相等时才认为两个WrappedFun等价。注意其 docstring 的告诫生成器的静态/动态参数以及辅助输出数据都必须是不可变的因为它们会被存进函数记忆化表。5.2cache的行为cache(call, *, explainNone)是装饰器jax/_src/linear_util.py对第一个参数是 WrappedFun 的可调用对象做记忆化缓存按fun.f分组weakref.WeakKeyDictionary避免长期持有函数引用缓存键为(fun.transforms, fun.params, fun.in_type, args, config.trace_context())其中trace_context()用于区分不同的追踪tracing上下文防止在不同变换环境下错误复用缓存命中时除了返回结果还会调用fun.populate_stores(stores)把缓存中保存的 stores 回填进当前 WrappedFun保证辅助输出仍可读取未命中时执行call(fun, *args)把(结果, fun.stores)存入缓存explain回调配合config.explain_cache_misses配置可用于记录缓存未命中的原因含耗时附带memoized_fun.evict_function(f)与memoized_fun.cache_clear()两个管理接口并注册进全局缓存统计register_cache。jax._src.api.py中清理所有lu.cache实例的代码jax/_src/api.py正说明jit、pmap等入口都依赖这一记忆化基础设施来避免重复编译。六、实际调用链与验证6.1 在 JAX 公共 API 中的位置jax.jit的底层入口正是用lu.wrap_init(fun, debug_infodbg)将用户函数包装成WrappedFunjax/_src/api.py随后由lu.cache记忆化其追踪与编译结果。lu.transformation_with_aux等则在ad.py反向模式自动微分、batching.pyvmap、partial_eval.py部分求值等解释器中被大量使用——它们是grad、vmap、jit能无限组合的底层原因。6.2 如何自行验证仓库中已有现成的验证测试 tests/extend_test.pytest_symbols断言jex.linear_util.StoreException、WrappedFun、cache、merge_linear_aux、transformation、transformation_with_aux与内部实现是同一个对象assertIs。运行方式python -m pytest tests/extend_test.py -k test_symbols你也可以在 Python 中直接体验变换栈import jax.extend as jex from jax._src import api_util def f(x): return x 1 wf jex.linear_util.wrap_init( f, debug_infoapi_util.debug_info(test, f, (), {})) print(wf) # 会打印变换栈与核心函数名 print(wf.call_wrapped(1)) # 26.3wrap_init的两个版本与废弃警告内部实现lu.wrap_init(f, paramsNone, *, debug_info)jax/_src/linear_util.py要求显式传入DebugInfo而jax.extend.linear_util.wrap_init提供了不要求DebugInfo的兼容版本jax/extend/linear_util.py缺省时它会调用_missing_debug_info并发出DeprecationWarning提示改用api_util.debug_info()构造规范的DebugInfo对象再传入。因此新代码应始终通过api_util.debug_info()构造并显式传递 debug_info。七、使用注意事项小结不可变性作为变换静态参数或辅助输出的数据必须不可变因为它们会进入记忆化缓存键生成器协议自定义变换必须恰好yield两次第一次交参数、收回结果第二次交结果与 aux优先新接口transformation2/transformation_with_aux2是当前推荐入口transformation/transformation_with_aux仅为向后兼容保留缓存感知WrappedFun的值语义哈希让同一变换栈的函数能复用缓存而trace_context()进入缓存键保证了不同追踪上下文的隔离兼容性风险jax.extend整体不承诺跨版本兼容生产代码请锁定 JAX 版本并关注 CHANGELOG.md。参考文件索引API 文档骨架docs/jax.extend.linear_util.rst对外转发层jax/extend/linear_util.py核心实现jax/_src/linear_util.py模块导出与 API 政策jax/extend/init.py等价性测试tests/extend_test.py公共 API 中的使用示例jax/_src/api.py【免费下载链接】jaxComposable transformations of PythonNumPy programs: differentiate, vectorize, JIT to GPU/TPU, and more项目地址: https://gitcode.com/GitHub_Trending/ja/jax创作声明:本文部分内容由AI辅助生成(AIGC),仅供参考
返回列表