如何用Oryx对带额外参数的JAX仿射变换函数求逆?
使用Oryx反转带参数的JAX仿射变换
问题背景
需要反转一个基于JAX实现的带可调参数params的仿射变换函数affine,该函数实现x → y的映射。目标是得到一个以params为自变量的逆函数,输入params和y后返回原输入x。
原仿射变换代码:
import jax.numpy as jnp def affine(params, x): return x * params['scale'] + params['shift'] params = dict(scale=1.5, shift=-1.) x_in = jnp.array(3.) y_out = affine(params, x_in)
直接调用oryx.core.inverse会报错:
import oryx oryx.core.inverse(affine)(params, y_out)
错误信息:AssertionError: length mismatch: [1, 3],原因是Oryx默认会尝试反转所有输入参数,但我们仅需对x(函数的第二个参数)求逆,params是固定参数而非求逆对象。
解决方案
方法1:指定求逆参数位置(推荐)
利用oryx.core.inverse的argnums参数,明确告知Oryx仅对函数的第二个参数(x,索引为1)进行求逆:
import oryx # 定义逆函数,指定仅对第2个参数求逆 inverse_affine = oryx.core.inverse(affine, argnums=1) # 调用逆函数,传入params和y_out recovered_x = inverse_affine(params, y_out) print(recovered_x) # 输出: 3.0
这种方式直接定义了通用的逆函数,可接受任意合法的params(需满足scale≠0,保证仿射变换可逆),调用时只需传入params和目标y值即可。
方法2:通过部分应用函数封装
使用functools.partial将params绑定到原函数,生成仅接受x的新函数后再求逆,最后封装为接受params和y的逆函数:
import oryx from functools import partial def inverse_affine(params, y): # 绑定params到affine,生成仅接受x的函数 affine_with_params = partial(affine, params) # 对该函数求逆并传入y return oryx.core.inverse(affine_with_params)(y) # 测试调用 recovered_x = inverse_affine(params, y_out) print(recovered_x) # 输出: 3.0
说明
两种方法都能正确恢复原输入x,且逆函数支持任意合法的params取值。其中方法1更简洁直接,利用Oryx的原生参数指定功能实现需求。
内容的提问来源于stack exchange,提问作者Hylke
相关产品推荐
相关产品推荐

