Pallas Jax复数向量乘法报错,求多输出内核实现方法
求解JAX Pallas复数向量乘法内核的多输出实现问题
问题场景
编写复数向量乘法的Pallas内核时遇到类型转换错误,代码如下:
from functools import partial import jax from jax.experimental import pallas as pl import jax.numpy as jnp import numpy as np def cdot_vectors_kernel(x_real_ref, x_imag_ref, y_real_ref, y_imag_ref, o_ref): x_real = pl.load(x_real_ref, (slice(None),)) x_imag = pl.load(x_imag_ref, (slice(None),)) y_real = pl.load(y_real_ref, (slice(None),)) y_imag = pl.load(y_imag_ref, (slice(None),)) o_real = x_real * y_real - x_imag * y_imag o_imag = x_real * y_imag + x_imag * y_real o = jnp.array(o_real + 1j * o_imag) pl.store(o_ref, (slice(None), ), o) @jax.jit def cdot_vectors(x_real: jax.Array, x_imag: jax.Array, y_real: jax.Array, y_imag: jax.Array) -> jax.Array: return pl.pallas_call( cdot_vectors_kernel, out_shape=jax.ShapeDtypeStruct(x_real.shape, jnp.complex64) )(x_real, x_imag, y_real, y_imag) array1 = jnp.array([2+3j, 1-1j, 5+1j]) array2 = jnp.array([1+2j, 1-2j, 2+1j]) cdot_vectors(array1.real, array1.imag, array2.real, array2.imag)
运行时触发错误:
NotImplementedError: cannot cast Value(%33 = "arith.addf"(%31, %32) <{fastmath = #arith.fastmath<none>}> : (tensor<3xf32>, tensor<3xf32>) -> tensor<3xf32>) to tensor<3xcomplex<f32>>
用户推测拆分实虚部输出可能解决问题,但不清楚Pallas如何实现多输出。
解决思路
1. 错误根源
当前内核中试图将实数运算结果合并为复数数组后存储,Pallas底层暂不支持这种直接的实数张量到复数张量的类型转换,因此报错。
2. 方案一:实现多输出内核,外部合并复数
Pallas支持多输出,只需在out_shape中传入多个ShapeDtypeStruct组成的列表,内核函数的末尾参数对应各输出的引用,分别调用pl.store即可:
修改后的代码示例:
from functools import partial import jax from jax.experimental import pallas as pl import jax.numpy as jnp import numpy as np def cdot_vectors_kernel(x_real_ref, x_imag_ref, y_real_ref, y_imag_ref, o_real_ref, o_imag_ref): x_real = pl.load(x_real_ref, (slice(None),)) x_imag = pl.load(x_imag_ref, (slice(None),)) y_real = pl.load(y_real_ref, (slice(None),)) y_imag = pl.load(y_imag_ref, (slice(None),)) o_real = x_real * y_real - x_imag * y_imag o_imag = x_real * y_imag + x_imag * y_real # 分别存储实部和虚部结果 pl.store(o_real_ref, (slice(None), ), o_real) pl.store(o_imag_ref, (slice(None), ), o_imag) @jax.jit def cdot_vectors(x_real: jax.Array, x_imag: jax.Array, y_real: jax.Array, y_imag: jax.Array) -> jax.Array: # 定义两个输出的形状和类型 out_real, out_imag = pl.pallas_call( cdot_vectors_kernel, out_shape=[ jax.ShapeDtypeStruct(x_real.shape, jnp.float32), jax.ShapeDtypeStruct(x_real.shape, jnp.float32) ] )(x_real, x_imag, y_real, y_imag) # 外部合并为复数数组 return out_real + 1j * out_imag array1 = jnp.array([2+3j, 1-1j, 5+1j]) array2 = jnp.array([1+2j, 1-2j, 2+1j]) # 调用测试 result = cdot_vectors(array1.real, array1.imag, array2.real, array2.imag) print(result)
3. 方案二:直接传入复数数组简化代码
Pallas原生支持复数数组的加载与存储,无需手动拆分实虚部,代码更简洁:
from functools import partial import jax from jax.experimental import pallas as pl import jax.numpy as jnp import numpy as np def cdot_vectors_kernel(x_ref, y_ref, o_ref): # 直接加载复数数组 x = pl.load(x_ref, (slice(None),)) y = pl.load(y_ref, (slice(None),)) # 直接执行复数乘法 pl.store(o_ref, (slice(None), ), x * y) @jax.jit def cdot_vectors(x: jax.Array, y: jax.Array) -> jax.Array: return pl.pallas_call( cdot_vectors_kernel, out_shape=jax.ShapeDtypeStruct(x.shape, jnp.complex64) )(x, y) array1 = jnp.array([2+3j, 1-1j, 5+1j]) array2 = jnp.array([1+2j, 1-2j, 2+1j]) # 直接传入复数数组调用 result = cdot_vectors(array1, array2) print(result)
内容的提问来源于stack exchange,提问作者bsaoptima
相关产品推荐
相关产品推荐

