在Equinox中为Module定义自定义JVP/VJP规则
概述
在JAX生态系统中,自动微分是其核心功能之一。Equinox作为构建在JAX之上的神经网络库,同样支持自动微分操作。然而,在某些特殊场景下,我们可能需要自定义前向模式微分(JVP)或反向模式微分(VJP)的行为。本文将详细介绍在Equinox模块中如何实现这一需求。
问题背景
考虑一个常见的场景:我们有一个神经网络模块,它首先对输入进行缩放处理,然后再传递给子模块。标准的自动微分会计算整个计算图的梯度,但有时我们可能希望梯度计算只针对缩放后的输入,而不是原始输入。
标准实现的问题
使用Equinox的标准实现方式如下:
class ScaledModel(eqx.Module):
sub_model: eqx.Module
scale: float = eqx.field(static=True)
def __call__(self, x):
scaled_x = x / self.scale
return self.sub_model(scaled_x)
这种实现会计算从原始输入到最终输出的完整梯度,包括缩放操作的梯度部分。但有时我们可能希望梯度计算跳过缩放操作,直接针对缩放后的输入。
自定义JVP的实现方案
由于JAX的custom_jvp装饰器不是描述符(descriptor),无法正确处理类方法的self参数,因此我们需要采用间接的方式实现:
- 首先定义一个独立的函数,并用
eqx.filter_custom_jvp装饰 - 然后在模块的
__call__方法中调用这个函数
具体实现如下:
class ScaledModel(eqx.Module):
sub_model: eqx.Module
scale: float
def __call__(self, x):
return scaled_model_forward(self, x)
@eqx.filter_custom_jvp
def scaled_model_forward(model, x):
scaled_x = x / model.scale
return model.sub_model(scaled_x)
@scaled_model_forward.def_jvp
def scaled_model_jvp(primals, tangents):
model, x = primals
primal_out = model.sub_model(x / model.scale)
_, tangent_out = jax.jvp(
model.sub_model,
(x / model.scale,),
(tangents[1] / model.scale,)
)
return primal_out, tangent_out
技术细节解析
-
装饰器选择:使用
eqx.filter_custom_jvp而非普通的jax.custom_jvp,因为前者能正确处理Equinox模块的过滤机制。 -
参数处理:在JVP函数中,
primals包含模型实例和输入数据,需要分别处理。 -
梯度计算:我们显式地计算子模块在缩放后输入处的梯度,并跳过对原始输入的梯度计算。
应用场景
这种技术特别适用于以下场景:
- 输入预处理需要从梯度计算中排除
- 需要实现特殊的梯度流动规则
- 构建具有自定义微分行为的复合模块
注意事项
-
确保自定义微分规则与数学定义一致,避免引入数值不稳定。
-
对于复杂的模块组合,需要仔细测试梯度计算的正确性。
-
考虑使用
jax.value_and_grad等工具验证自定义梯度的正确性。
总结
在Equinox中实现自定义微分规则需要绕过JAX对类方法的限制,通过外部函数的方式实现。这种方法既保持了Equinox模块的清晰结构,又提供了对微分行为的精确控制。掌握这一技术可以极大地扩展Equinox在复杂微分场景下的应用能力。
GLM-5智谱 AI 正式发布 GLM-5,旨在应对复杂系统工程和长时域智能体任务。Jinja00
GLM-5-w4a8GLM-5-w4a8基于混合专家架构,专为复杂系统工程与长周期智能体任务设计。支持单/多节点部署,适配Atlas 800T A3,采用w4a8量化技术,结合vLLM推理优化,高效平衡性能与精度,助力智能应用开发Jinja00
jiuwenclawJiuwenClaw 是一款基于openJiuwen开发的智能AI Agent,它能够将大语言模型的强大能力,通过你日常使用的各类通讯应用,直接延伸至你的指尖。Python0195- QQwen3.5-397B-A17BQwen3.5 实现了重大飞跃,整合了多模态学习、架构效率、强化学习规模以及全球可访问性等方面的突破性进展,旨在为开发者和企业赋予前所未有的能力与效率。Jinja00
AtomGit城市坐标计划AtomGit 城市坐标计划开启!让开源有坐标,让城市有星火。致力于与城市合伙人共同构建并长期运营一个健康、活跃的本地开发者生态。01
awesome-zig一个关于 Zig 优秀库及资源的协作列表。Makefile00