尧图网站设计 尧图网站设计YAOTU DESIGN
ARTICLE DETAIL

资讯详情

深耕网站设计与一线实操的经验洞察。

JAX FAQ 实战精解:从 jit 语义陷阱到设备放置、梯度异常与性能调优

JAX FAQ 实战精解:从 jit 语义陷阱到设备放置、梯度异常与性能调优 机器学习深度学习【免费下载链接】jaxComposable transformations of PythonNumPy programs: differentiate, vectorize, JIT to GPU/TPU, and more项目地址https://gitcode.com/gh_mirrors/jax/jax点击查看免费下载导读本文以 JAX 官方 FAQdocs/faq.rst为主线系统梳理 JAX 使用者在实际开发中最常遇到的十余类问题为什么jit会改变函数行为与数值结果、类方法如何正确使用jit、如何控制数据与计算的设备放置、如何科学地做基准测试、什么是 tracer 与 buffer donation以及where产生 NaN 梯度、排序类函数梯度为零、CUDA 库加载失败等疑难杂症的根因与解法。读完本文你将掌握这些踩坑点背后的 JAX 设计原理结合仓库源码佐证并拿到可直接复制运行的修复代码从而写出行为可预期、性能可度量、迁移无坑的 JAX 程序。一、jit改变了我的函数行为全局状态与副作用1.1 现象加不加jit输出不一样如果你发现某个 Python 函数在加上jax.jit装饰器后行为发生了变化那么大概率是该函数使用了全局状态或带有副作用如print。官方 FAQ 给出了如下经典示例docs/faq.rsty 0 # jit # 加上 jit 后行为不同 def impure_func(x): print(Inside:, y) return x y for y in range(3): print(Result:, impure_func(y))不加jit时输出为Inside: 0 Result: 0 Inside: 1 Result: 2 Inside: 2 Result: 4加上jit后输出为Inside: 0 Result: 0 Result: 1 Result: 21.2 原理一次 Python 执行 编译缓存复用原因在于jax.jit的工作方式函数只用 Python 解释器执行一次trace 阶段此时print(Inside:, y)发生、全局变量y的第一个值 0被捕获进计算图之后函数被编译并缓存后续调用虽然传入了不同的x但使用的始终是第一次观察到的y值。从源码看jax.jit在 jax/_src/api.py 中要求传入的函数应当是pure function纯函数其文档字符串明确写道fun: Function to be jitted.funshould be a pure function.所谓纯函数即不读取外部可变状态、不产生副作用的函数。如果你的函数依赖全局变量或打印输出应把依赖值显式作为参数传入或用static_argnums标记为静态参数把print之类的副作用移到函数外部或改用 jax.debug.print 这类专门面向 trace 阶段的调试工具。更多关于纯净性与 trace 机制的讨论可参见仓库中的 Common Gotchas 教程。二、jit改变了输出的精确数值XLA 代数化简2.1 现象JIT 前后结果有细微差异有时你会惊讶地发现同一个函数在 JIT 前后输出的小数位不同 from jax import jit import jax.numpy as jnp def f(x): ... return jnp.log(jnp.sqrt(x)) x jnp.pi print(f(x)) 0.572365 print(jit(f)(x)) 0.57236492.2 原因XLA 编译器对运算的重排与省略这种细微差异来自XLA 编译器内部的优化编译期间 XLA 有时会重排或消去某些运算以提升整体效率。在本例中XLA 利用对数性质把log(sqrt(x))替换为数学上等价的0.5 * log(x)后者计算更高效。由于浮点运算只是实数运算的近似不同计算路径会带来细微差别。更极端的例子 def f(x): ... return jnp.log(jnp.exp(x)) x 100.0 print(f(x)) inf print(jit(f)(x)) 100.0非 JIT 的逐算子op-by-op模式下jnp.exp(x)先溢出返回inf因此结果也是inf而在 JIT 下XLA 识别出log是exp的逆运算直接把这两个算子从编译函数中消去返回输入本身——此时 JIT 反而得到了更接近真实数学结果的浮点近似。XLA 的全部代数化简规则并未完整文档化但若熟悉 C可直接查阅其源码实现algebraic_simplifier.ccTensorFlow/XLA 编译器内部的代数化简器了解优化类型。实践要点不要依赖 JIT 前后逐位一致的数值这在浮点语义下并不被保证。三、jit装饰的函数编译非常慢3.1 症状与根因如果jit函数首次调用耗时数十秒甚至更久、第二次调用却很快说明 JAX 花费了大量时间在trace 或编译上。这通常是因为你的函数在 JAX 内部表示jaxpr中生成了大量代码典型元凶是重度使用 Python 控制流如 Pythonfor循环。少量循环迭代用 Python 没问题但迭代次数很大时应改用 JAX 的结构化控制流原语如lax.scan、lax.cond、lax.fori_loop、lax.while_loop或者干脆不要让循环被jit包裹可以在循环内部继续调用jit装饰的函数。3.2 诊断工具jax.make_jaxpr如果你不确定是否属于这类问题可以对函数运行jax.make_jaxpr实现见 jax/_src/api.pyimport jax jax.make_jaxpr(my_function)(*sample_args)如果打印出的 jaxpr 有成百上千行那么慢编译基本可以确认来自 Python 循环展开。文档还提醒如果代码涉及大量形状不同的数组推荐用jax.numpy.where在固定形状的填充数组上完成计算从而避免逐形状分支带来的 trace 膨胀。若排除上述原因仍编译缓慢可以在 GitHub 上提交 issue。四、如何在类方法上使用jitjax.jit的大多数示例都是装饰独立函数装饰类方法则引入一个麻烦方法的第一个参数是self其类型是自定义类实例JAX 不知道如何处理该类型。以下简单类在直接jit时会报错 import jax.numpy as jnp from jax import jit class CustomClass: ... def __init__(self, x: jnp.ndarray, mul: bool): ... self.x x ... self.mul mul ... ... jit # ---- 怎样写才是对的 ... def calc(self, y): ... if self.mul: ... return self.x * y ... return y c CustomClass(2, True) c.calc(3) # doctest: SKIP --------------------------------------------------------------------------- TypeError Traceback (most recent call last) File stdin, line 1, in module TypeError: Argument CustomClass object at 0x7f7dd4125890 of type class CustomClass is not a valid JAX type.错误本质self被当作第一个参数传入 trace而CustomClass不是合法的 JAX 类型。FAQ 给出了三种可行策略。4.1 策略一JIT 编译的外部辅助函数把方法逻辑抽到类外、用常规方式 JIT 化 from functools import partial class CustomClass: ... def __init__(self, x: jnp.ndarray, mul: bool): ... self.x x ... self.mul mul ... ... def calc(self, y): ... return _calc(self.mul, self.x, y) partial(jit, static_argnums0) ... def _calc(mul, x, y): ... if mul: ... return x * y ... return y c CustomClass(2, True) print(c.calc(3)) 6注意这里用static_argnums0把mul一个 Python bool参与if分支判断标记为静态参数否则if mul:这种依赖具体值的分支在抽象 trace 下无法解析详见第七节。该策略简单、显式无需教会 JAX 认识CustomClass代价是方法逻辑与类分离。4.2 策略二把self标记为静态需谨慎用static_argnums0标记self class CustomClass: ... def __init__(self, x: jnp.ndarray, mul: bool): ... self.x x ... self.mul mul ... ... # WARNING: 这个示例是坏的见下文。不要复制粘贴 ... partial(jit, static_argnums0) ... def calc(self, y): ... if self.mul: ... return self.x * y ... return y首次调用c.calc(3)返回 6不再报错但随后若修改对象 c.mul False print(c.calc(3)) # 期望打印 3 6结果仍是 6原因静态对象会被用作 JIT内部编译缓存的字典键其hash(obj)、obj1 obj2与obj1 is obj2的一致性被假定成立。自定义对象默认的__hash__基于对象 ID因此 JAX 无法感知对象被修改也就不会触发重新编译。部分缓解办法是为类定义合适的__hash__与__eq__ class CustomClass: ... def __init__(self, x: jnp.ndarray, mul: bool): ... self.x x ... self.mul mul ... ... partial(jit, static_argnums0) ... def calc(self, y): ... if self.mul: ... return self.x * y ... return y ... ... def __hash__(self): ... return hash((self.x, self.mul)) ... ... def __eq__(self, other): ... return (isinstance(other, CustomClass) and ... (self.x, self.mul) (other.x, other.mul))这里对object.__hash__的约束可参见 Python 官方文档覆写__hash__时必须保证相等对象哈希一致、且对象不可变。只要你不修改对象这种方式可与 JIT 及其他变换正确配合一旦对象被原地修改如self.attr ...作为哈希键就会引发多种隐蔽问题——这也正是可变容器dict、list不定义__hash__、而不可变容器tuple定义__hash__的原因。如果类依赖原地修改请用策略三。4.3 策略三把CustomClass注册为 PyTree最灵活的方案是把类注册为自定义 PyTree精确声明哪些成员是静态的、哪些是动态的详见仓库的 PyTrees 指南 与 pytrees.md class CustomClass: ... def __init__(self, x: jnp.ndarray, mul: bool): ... self.x x ... self.mul mul ... ... jit ... def calc(self, y): ... if self.mul: ... return self.x * y ... return y ... ... def _tree_flatten(self): ... children (self.x,) # 数组 / 动态值 ... aux_data {mul: self.mul} # 静态值 ... return (children, aux_data) ... ... classmethod ... def _tree_unflatten(cls, aux_data, children): ... return cls(*children, **aux_data) from jax import tree_util tree_util.register_pytree_node(CustomClass, ... CustomClass._tree_flatten, ... CustomClass._tree_unflatten)这样处理之后所有此前的问题都得到解决 c CustomClass(2, True) print(c.calc(3)) 6 c.mul False # 修改会被检测到 print(c.calc(3)) 3 c CustomClass(jnp.array(2), True) # 不可哈希的 x 也支持 print(c.calc(3)) 6要点tree_flatten中返回的children是参与动态计算的数组/动态值aux_data是参与分支等控制流判断的静态值只要二者正确覆盖类的全部相关属性这种类型的对象就能直接作为 JIT 函数的参数无需任何特殊注解。这是官方推荐的、对带有数值成员和静态配置成员的类最健壮的写法。五、控制数据与计算在设备上的放置5.1 两条放置原则JAX 中计算跟随数据computation follows data。JAX 数组有两个放置属性数据所在的设备device数据是否提交committed给该设备未提交的数据有时称为sticky粘在设备上。默认情况下JAX 数组被未提交地放在默认设备jax.devices()[0]通常是第一个 GPU/TPU没有加速器时是 CPU上。默认设备可以通过jax.default_device上下文管理器临时覆盖通过环境变量JAX_PLATFORMS或 absl flag--jax_platforms设为cpu、gpu或tpu来为整个进程设置JAX_PLATFORMS还支持平台列表按优先级决定哪些平台可用。 from jax import numpy as jnp print(jnp.ones(3).devices()) # doctest: SKIP {CudaDevice(id0)}只涉及未提交数据的计算在默认设备上进行结果同样未提交到默认设备。5.2 用device_put显式放置并提交用带device参数的jax.device_put可把数据显式放置到指定设备此时数据变为已提交committed import jax from jax import device_put arr device_put(1, jax.devices()[2]) # doctest: SKIP print(arr.devices()) # doctest: SKIP {CudaDevice(id2)}规则如下计算只要包含已提交的输入就会发生在该已提交设备上结果也提交到同一设备对提交到多个不同设备的参数执行运算会抛出错误device_put不带 device 参数时数据若已在某设备上无论是否提交保持原状若不在任何设备上普通 Python/NumPy 值则未提交地放到默认设备jit函数与普通原语行为一致——跟随数据、并对提交到多个设备的输入报错。从源码看jax.device_putjax/_src/api.py支持传入Device、Sharding或TransferToMemoryKind并把整个 pytree 展平后逐叶绑定device_put_p原语其文档明确指出给定device时结果提交给该设备且该函数总是异步返回。 import jax from jax import device_put arr device_put(1, jax.devices()[2]) # doctest: SKIP print(arr.devices()) # doctest: SKIP {CudaDevice(id2)}5.3 代码级验证test_computation_follows_data仓库的 tests/multi_device_test.py 中test_computation_follows_data完整演示了这些规则默认构造的x jnp.ones(2)是未提交且位于devices[0]只含未提交数据的计算y jnp.sin(x x * 1) 2结果仍在devices[0]且未提交z jax.device_put(1, devices[1])后z已提交到devices[1]u z x已提交 未提交结果提交到devices[1]未提交数据可被移动到已提交设备对devices[2]与devices[3]上两个已提交数组做抛出ValueError错误消息为Received incompatible devices for jitted computationjit函数同样遵循以上全部行为。5.4 历史背景与实验性参数2021 年 3 月的 PR #6002 之前数组常量创建存在一些惰性优化如jax.device_put(jnp.zeros(...), jax.devices()[1])会在目标设备上直接创建数组而非先建后搬该优化后来为简化实现而被移除。自 2020 年 4 月起jax.jit带有一个device参数影响设备放置但该参数是实验性的很可能被移除或改变官方不推荐使用在 tests/multi_device_test.py 中也可看到对jit(..., devicedevices[4])行为的测试确认了其结果提交到指定设备的语义。六、科学地基准测试 JAX 代码6.1 与 NumPy 测速的四个关键差异把函数从 NumPy/SciPy 移植到 JAX 后测速时必须注意四点JAX 代码是 JIT 编译的大多数 JAX 代码都支持 JIT编译后运行会快得多。追求最大性能应在最外层函数调用上应用jax.jit。注意首次运行 JAX 代码都会更慢正在编译即使你自己的代码没用jit也一样因为 JAX 内建函数本身也是 JIT 编译的。JAX 采用异步派发需要调用.block_until_ready()才能确保计算真正完成参见 docs/async_dispatch.rst。JAX 默认只使用 32 位 dtype为公平对比应显式在 NumPy 中使用 32 位 dtype或在 JAX 中启用 64 位参见 Common Gotchas 的 Double precision 一节。CPU 与加速器之间的数据传输耗时如果只想测函数求值耗时应先把数据转移到目标设备上见第五节。6.2 一个可复制的微基准模板import numpy as np import jax.numpy as jnp import jax def f(x): # 被测函数NumPy 与 JAX 通用 return x.T (x - x.mean(axis0)) x_np np.ones((1000, 1000), dtypenp.float32) # 与 JAX 默认 dtype 一致 %timeit f(x_np) # 测量 NumPy 运行时间 %time x_jax jax.device_put(x_np) # 测量 JAX 设备传输时间 f_jit jax.jit(f) %time f_jit(x_jax).block_until_ready() # 测量 JAX 编译时间 %timeit f_jit(x_jax).block_until_ready() # 测量 JAX 运行时间在 Colab 的 GPU 上运行典型结果NumPy 每次求值 CPU 耗时16.2 msJAX 把 NumPy 数组拷贝到 GPU 耗时1.26 msJAX 编译函数耗时193 msJAX 在 GPU 上每次求值485 µs。即数据转移与编译完成后GPU 上重复求值的 JAX 比 NumPy 快约30 倍。但这样的对比公平吗——未必真正重要的性能是完整应用的运行性能它不可避免地包含传输与编译开销。而且我们刻意选了足够大的数组1000×1000与足够重的计算是矩阵乘法来摊销 JAX/加速器的额外开销若改用 10×10 的输入JAX/GPU 反而比 NumPy/CPU 慢约 10 倍100 µs vs 10 µs。结论微基准结论强依赖于问题规模与平台测速前务必考虑传输、编译与摊销成本。6.3 JAX 一定比 NumPy 快吗没有简单答案。宽泛地说NumPy立即执行eager、同步、只在 CPU 上运行JAX可立即执行也可在jit内编译后执行异步派发可运行在 CPU、GPU、TPU 上且各平台的性能特征差异巨大、持续演进。这些架构差异使有意义的直接对比变得困难也导致两者工程重心不同NumPy 把大量精力放在降低单个数组操作的每次调用派发开销上因为其计算模型无法避免该开销JAX 则通过 JIT 编译、异步派发、批量变换等方式绕开派发开销因此降低单次调用开销不是首要目标。归纳在 CPU 上对单个数组操作做微基准NumPy 通常因更低的每次操作派发开销而胜出在 GPU/TPU 上运行、或对 CPU 上更复杂的 JIT 编译操作序列做基准JAX 通常胜出。七、认识 JAX 中的各类值tracer、抽象值与具体值7.1 tracer 从哪来在对函数做变换的过程中JAX 会把某些函数参数替换为特殊的tracer 值。你可以用print亲眼看到def func(x): print(x) return jnp.cos(x) res jax.jit(func)(0.)上面的代码返回值1.是正确的但它还会打印出TracedShapedArray(float32[])。通常 JAX 在内部透明地处理这些 tracer例如实现jax.numpy各函数的数值原语就是如此所以jnp.cos能正常工作。7.2 tracer 的两类抽象 tracer 与具体 tracer更精确地说tracer 值为 JAX 变换函数的参数引入被static_argnums对应jax.jit或static_broadcasted_argnums对应jax.pmap等特殊参数指明的除外。通常涉及至少一个 tracer 的计算产生 tracer。与之相对的是常规 Python 值在 JAX 变换之外算出的值、由上述静态参数产生的值、或仅由其他常规值算出的值——无 JAX 变换时处处使用的就是它们。一个 tracer 携带一个抽象值abstract value例如带 shape 与 dtype 信息的ShapedArray这类 tracer 称为抽象 tracer。另一些 tracer如自动微分变换的参数引入的携带ConcreteArray抽象值其中包含常规数组的真实数据可用于解析条件分支这类称为具体 tracer由具体 tracer或它与常规值组合算出的 tracer 仍是具体 tracer。具体值既可以是常规值也可以是具体 tracer。在仓库源码 jax/_src/core.py 中可以确认这两种抽象值的实现ShapedArray只记录shape、dtype、weak_type、named_shape等形状/类型信息ConcreteArray继承ShapedArray并额外持有val字段保存真实数组数据。大多数从 tracer 算出的值仍是 tracer极少数情况下计算可以完全利用 tracer 携带的抽象值完成此时结果可能是常规值例如取 tracer 的 shapeShapedArray抽象值即可给出显式把具体 tracer 转换为常规类型如int(x)、x.astype(float)bool(x)在具体性允许时产生 Python bool——这个情形特别重要因为它频繁出现在控制流中。7.3 各变换引入的 tracer 类型一览变换 / 原语引入的 tracer备注jax.jit抽象 tracer除static_argnums指定的参数保持为常规值外jax.pmap抽象 tracer除static_broadcasted_argnums指定的参数外jax.vmap、jax.make_jaxpr、xla_computation抽象 tracer对所有位置参数jax.jvp、jax.grad具体 tracer例外处于外层变换中且实参本身是抽象 tracer 时autodiff 引入的也是抽象 tracerlax.cond、lax.while_loop、lax.fori_loop、lax.scan抽象 tracer处理其函数体时引入无论当前是否有 JAX 变换在进行7.4 对代码的启示依赖具体值的条件分支这些机制与只能操作常规 Python 值的代码密切相关例如基于数据做条件分支def divide(x, y): return x / y if y 1. else 0.要对该函数应用jax.jit必须指定static_argnums1让y保持为常规值——因为布尔表达式y 1.需要具体值常规值或具体 tracer显式写bool(y 1.)、int(y)、float(y)也一样。有趣的是jax.grad(divide)(3., 2.)可以直接工作jax.grad使用具体 tracer用y的具体值解析了条件分支。八、Buffer donation用缓冲区捐赠做内存高效的函数式更新8.1 基本思想JAX 执行计算时输入与输出都在设备上占用缓冲区。如果你知道某个输入在计算后不再需要且它与某个输出的形状与元素类型匹配就可以指定把该输入缓冲区捐赠给输出复用。这会把执行所需内存减少一个捐赠缓冲区的大小。典型模式params, state jax.pmap(update_fn, donate_argnums(0, 1))(params, state)可以把它理解为对不可变 JAX 数组做内存高效的函数式更新在计算边界内部 XLA 自己能做这种优化但在 jit/pmap 边界你需要向 XLA保证捐赠的输入缓冲区在调用后不再使用。8.2 使用方式与边界通过donate_argnums参数jax.jit、jax.pjit、jax.pmap均支持指定位置参数索引从 0 开始def add(x, y): return x y x jax.device_put(np.ones((2, 3))) y jax.device_put(np.ones((2, 3))) # 捐赠 y 的缓冲区。结果的形状与类型和 y 相同因此会复用其缓冲区。 z jax.jit(add, donate_argnums(1,))(x, y)从源码看donate_argnums的语义在 jax/_src/api.py 中有明确说明指定哪些位置参数的缓冲区可被计算覆盖并在调用方标记删除只要计算开始后不再需要这些缓冲区捐赠就是安全的XLA 可借此复用输入缓冲区存放结果以减少内存默认不捐赠任何缓冲区。注意事项关键字参数目前不支持捐赠下面这行代码实际上不会捐赠任何缓冲区params, state jax.pmap(update_fn, donate_argnums(0, 1))(paramsparams, statestate)捐赠的参数是 pytree 时其全部组件的缓冲区都会被捐赠def add_ones(xs: List[Array]): return [x 1 for x in xs] xs [jax.device_put(np.ones((2, 3))), jax.device_put(np.ones((3, 4)))] # 捐赠 xs 的全部缓冲区。输出与 xs 各元素形状类型相同因此复用这些缓冲区。 z jax.jit(add_ones, donate_argnums0)(xs)不允许捐赠后续还会使用的缓冲区否则 JAX 会报错因为y的缓冲区在捐赠后已失效# 捐赠 y 的缓冲区 z jax.jit(add, donate_argnums(1,))(x, y) w y 1 # 复用上面已被捐赠的 y 缓冲区 # RuntimeError: Invalid argument: CopyToHostAsync() called on invalid buffer捐赠的缓冲区未被使用时会收到警告例如捐赠数量多于输出可用的数量# 同时捐赠 x 和 y 的缓冲区结果只会用其中一个另一个用不上。 z jax.jit(add, donate_argnums(0, 1))(x, y) # UserWarning: Some donated buffers were not usable: f32[2,3]{1,0}没有输出形状与捐赠匹配时捐赠也会被闲置y jax.device_put(np.ones((1, 3))) # y 的形状与输出不同 z jax.jit(add, donate_argnums(1,))(x, y) # UserWarning: Some donated buffers were not usable: f32[1,3]{1,0}九、梯度相关的两个经典陷阱9.1 用where规避未定义值时梯度出现NaN你可能会用where来规避未定义值但若不谨慎反向微分时仍会得到NaNdef my_log(x): return jnp.where(x 0., jnp.log(x), 0.) my_log(0.) 0. # 正常 jax.grad(my_log)(0.) NaN简短解释在grad计算中jnp.log(x)在x 0处的伴随值adjoint是NaN它被累积进了jnp.where的伴随值。正确写法是保证部分定义函数内部也有一层jnp.where使伴随值始终有限def safe_for_grad_log(x): return jnp.log(jnp.where(x 0., x, 1.)) safe_for_grad_log(0.) 0. # 正常 jax.grad(safe_for_grad_log)(0.) 0. # 正常内层jnp.where有时需要在原有外层where之外额外再加一层def my_log_or_y(x, y): x 0 时返回 log(x)否则返回 y return jnp.where(x 0., jnp.log(jnp.where(x 0., x, 1.)), y)可进一步阅读仓库/社区中关于梯度穿过jnp.where且某分支为 NaN的相关 issue 讨论理解边界情形。9.2 基于排序/比较的函数梯度处处为零如果函数用依赖输入相对顺序的运算处理输入如max、greater、argsort你可能会惊讶地发现梯度处处为零。例如一个阶跃函数import jax import numpy as np import jax.numpy as jnp def f(x): return (x 0).astype(float) df jax.vmap(jax.grad(f)) x jnp.array([-1.0, -0.5, 0.0, 0.5, 1.0]) print(ff(x) {f(x)}) # f(x) [0. 0. 0. 1. 1.] print(fdf(x) {df(x)}) # df(x) [0. 0. 0. 0. 0.]表面看输出随输入变化梯度为何为零令人困惑但零其实是正确结果微分衡量的是输入无穷小变化引起的输出变化。对x 1.0无论把x微扰变大还是变小输出都保持1.0因此grad(f)(1.0)应为零对小于零的区间同理。真正棘手的点是x 0向上微扰会改变输出——无穷小输入变化产生有限输出变化这意味着梯度未定义。此时采用另一种测量方式向下微扰输出不变梯度为零。JAX 与其他 autodiff 系统倾向于这样处理不连续点若正、负两个方向的梯度不一致但其中一个有定义而另一个没有就采用有定义的那个。在此定义下该函数梯度处处为零。问题的根源是函数在x 0处不连续。f本质上是Heaviside 阶跃函数可以用Sigmoid 函数作为平滑替代sigmoid 在远离零点时近似等于阶跃函数同时把x 0处的不连续替换为平滑可微曲线def g(x): return jax.nn.sigmoid(x) dg jax.vmap(jax.grad(g)) x jnp.array([-10.0, -1.0, 0.0, 1.0, 10.0]) with np.printoptions(suppressTrue, precision2): print(fg(x) {g(x)}) # g(x) [0. 0.27 0.5 0.73 1. ] print(fdg(x) {dg(x)}) # dg(x) [0. 0.2 0.25 0.2 0. ]jax.nn子模块jax/nn/init.py还提供了其他常见基于排序函数的平滑版本jax.nn.softmax可替代jax.numpy.argmaxjax.nn.soft_sign可替代jax.numpy.signjax.nn.softplus或jax.nn.squareplus可替代jax.nn.relu等。十、如何把 JAX Tracer 转成 NumPy 数组在运行时检查被变换的 JAX 函数会发现数组值变成了jax.core.Tracer对象jax.jit def f(x): print(type(x)) return x f(jnp.arange(5))输出class jax.interpreters.partial_eval.DynamicJaxprTracerDynamicJaxprTracer等 tracer 类的实现在 jax/_src/interpreters/partial_eval.py。简短回答不可能把 Tracer 转成 NumPy 数组——因为 tracer 是具有给定 shape 与 dtype 的一切可能值的抽象表示而 NumPy 数组是该抽象类的一个具体成员。更详细的 tracer 机制讨论可参见 thinking_in_jax 教程的 JIT mechanics 一节。把 Tracer 转回数组的需求通常来自另一层目标——在运行时访问计算中间值。FAQ 给出的三条路径想在运行时打印被 trace 的值以便调试考虑用jax.debug.print实现见 jax/_src/debugging.py想在变换函数内调用非 JAX 代码考虑用jax.pure_callback想在运行时输入/输出数组缓冲区如从文件加载数据、把数组内容写到磁盘考虑用jax.experimental.io_callback。更全面的运行时回调用法参见 External callbacks 教程。十一、为什么一些 CUDA 库加载/初始化失败JAX 解析动态库时使用通常的动态链接器搜索模式JAX 设置RPATH指向 pip 安装的 NVIDIA CUDA 包在 JAX 内的相对位置若已安装则优先使用。如果ld.so在常规搜索路径上找不到你的 CUDA 运行时库就必须把相应路径显式加入LD_LIBRARY_PATH。最省事的做法是安装nvidia-*-cu12系列 pip 包——它们包含在标准的jax[cuda_12]安装选项中。即便运行时库可被发现偶尔仍有加载/初始化问题常见原因是运行时 CUDA 库初始化可用内存不足JAX 会为加速执行预分配当前可用显存的一大块有时会导致 CUDA 运行时库初始化没有足够显存。这在以下场景尤其常见同时运行多个 JAX 实例、JAX 与同样会预分配显存的 TensorFlow 一起运行、或 GPU 正被其他进程重度占用。拿不准时可以降低预分配重试把XLA_PYTHON_CLIENT_MEM_FRACTION从默认的.75调低或设置XLA_PYTHON_CLIENT_PREALLOCATEfalse。更完整的内存管理机制参见 JAX GPU 内存分配文档。总结FAQ 背后的统一思维模型纵观整个 FAQ多数意外其实都源于 JAX 的核心设计假设可归纳为三条心智模型纯函数 trace 模型jit只看一次 Python 执行全局状态、副作用、基于具体值的分支都必须显式处理静态参数、PyTree 注册、where重写否则行为或数值会偏离预期计算跟随数据 异步执行设备放置、committed/uncommitted、.block_until_ready()、buffer donation 都围绕数据在哪、何时真正执行展开基准测试与多设备编程必须据此设计抽象值替代真实值tracer 携带的是 shape/dtype 而非数据ConcreteArray是例外这让自动微分、向量化、编译得以实现也让转 NumPy 数组排序函数零梯度等问题有了确定性的答案。掌握这三条再对照 docs/faq.rst 中每一节的复现示例与仓库源码jit/device_put定义见 jax/_src/api.pyShapedArray/ConcreteArray见 jax/_src/core.py设备放置测试见 tests/multi_device_test.py就能在遇到任何JAX 行为不符合直觉的问题时快速定位根因。赞分享机器学习深度学习【免费下载链接】jaxComposable transformations of PythonNumPy programs: differentiate, vectorize, JIT to GPU/TPU, and more项目地址https://gitcode.com/gh_mirrors/jax/jax点击查看免费下载相关推荐JAX 常见问题FAQ实战指南从 jit 副作用到 NaN 梯度的完整排查手册JAX 常见问题FAQ实战指南从 jit 副作用到 NaN 梯度的完整排查手册 JAX 是一套对 Python NumPy 程序进行可组合变换微分、人工智能机器学习深度学习编译器高性能计算Eve框架配置实战从常见陷阱到性能优化Eve框架配置实战从常见陷阱到性能优化 你是否曾在构建REST API时遇到这样的困扰明明配置看起来正确但API行为却出乎意料或者在生产环境中发现性能瓶后端Web框架告别JavaScript数字陷阱bignumber.js从精度设置到异常处理的全面解决方案告别JavaScript数字陷阱bignumber.js从精度设置到异常处理的全面解决方案 你是否曾因JavaScript浮点数精度问题导致财务计算错误是否后端上一篇【亲测免费】 50个Android Kotlin项目100天教程下一篇CMS-Hunter 开源项目教程创作声明:本文部分内容由AI辅助生成(AIGC),仅供参考
返回列表