能否在jax.numpy.ndarray.at调用中向jax.numpy.ufunc传递参数?
能否在jax.numpy.ndarray.at调用中向jax.numpy.ufunc传递参数?
当然可以,但不能直接通过x.at[:].apply(umyadd)的默认方式传递,需要通过包装函数或参数绑定的方式实现。以下是具体解决方法:
问题分析
你示例中的x.at[:].apply(umyadd)会使用myadd的默认参数b=1,因为apply方法只会将数组的每个元素作为第一个参数传入umyadd,无法直接传递额外参数。
解决方法
1. 使用Lambda包装函数
通过lambda将额外参数绑定到ufunc上,将其转换为单参数函数:
import jax.numpy as jnp def myadd(a,b=1): return a+b umyadd = jnp.frompyfunc(myadd,2,1,identity=0) x = jnp.arange(4) # 传递b=2 result = x.at[:].apply(lambda a: umyadd(a, 2)) print(result) # 输出: [2 3 4 5]
2. 使用functools.partial绑定参数
借助functools.partial提前绑定额外参数,生成新的单参数函数:
import jax.numpy as jnp from functools import partial def myadd(a,b=1): return a+b umyadd = jnp.frompyfunc(myadd,2,1,identity=0) umyadd_with_b3 = partial(umyadd, b=3) x = jnp.arange(4) result = x.at[:].apply(umyadd_with_b3) print(result) # 输出: [3 4 5 6]
关于ufunc.at的补充
你示例中使用的umyadd.at(x, jnp.arange(x.size), x, inplace=False)其实是ufunc原生的原地更新语法,这里的第三个参数x就是传递给myadd的第二个参数b。如果你的ufunc需要更多参数,这种方式不够灵活,而apply结合包装函数的方式更适合多参数场景。
内容的提问来源于stack exchange,提问作者DavidJ
相关产品推荐
相关产品推荐

