You need to enable JavaScript to run this app.
优惠活动
大模型
产品
解决方案
定价
更多

如何用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

相关产品推荐
方舟 Agent Plan

超全模态模型 × Harness 升级,最新支持 Deepseek-V4.1-Flash、GLM-5.3 系列、Doubao-Seedream-5.0-pro、Kimi-K3 (部分), 限时 9.9 元起

最近更新时间:2026.07.18 09:07:31