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

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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.06.28 15:57:04