在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在复杂微分场景下的应用能力。
kernelopenEuler内核是openEuler操作系统的核心,既是系统性能与稳定性的基石,也是连接处理器、设备与服务的桥梁。C046
MiniMax-M2.1从多语言软件开发自动化到复杂多步骤办公流程执行,MiniMax-M2.1 助力开发者构建下一代自主应用——全程保持完全透明、可控且易于获取。Python00
kylin-wayland-compositorkylin-wayland-compositor或kylin-wlcom(以下简称kywc)是一个基于wlroots编写的wayland合成器。 目前积极开发中,并作为默认显示服务器随openKylin系统发布。 该项目使用开源协议GPL-1.0-or-later,项目中来源于其他开源项目的文件或代码片段遵守原开源协议要求。C01
PaddleOCR-VLPaddleOCR-VL 是一款顶尖且资源高效的文档解析专用模型。其核心组件为 PaddleOCR-VL-0.9B,这是一款精简却功能强大的视觉语言模型(VLM)。该模型融合了 NaViT 风格的动态分辨率视觉编码器与 ERNIE-4.5-0.3B 语言模型,可实现精准的元素识别。Python00
GLM-4.7GLM-4.7上线并开源。新版本面向Coding场景强化了编码能力、长程任务规划与工具协同,并在多项主流公开基准测试中取得开源模型中的领先表现。 目前,GLM-4.7已通过BigModel.cn提供API,并在z.ai全栈开发模式中上线Skills模块,支持多模态任务的统一规划与协作。Jinja00
agent-studioopenJiuwen agent-studio提供零码、低码可视化开发和工作流编排,模型、知识库、插件等各资源管理能力TSX0123
Spark-Formalizer-X1-7BSpark-Formalizer 是由科大讯飞团队开发的专用大型语言模型,专注于数学自动形式化任务。该模型擅长将自然语言数学问题转化为精确的 Lean4 形式化语句,在形式化语句生成方面达到了业界领先水平。Python00