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

Python中使用JAX计算矩阵指数报错:期望输入为方阵

问题分析与解决方案

核心问题

  1. Pauli矩阵生成函数逻辑错误:你的pauli_matrix函数没有正确生成多量子比特的Pauli矩阵集合,错误的张量积操作导致后续计算的矩阵结构不符合expm的隐含要求。
  2. 缺失numpy导入:代码中使用了np但未导入,会先触发NameError。

修正后的代码

import numpy as np
from jax.scipy.linalg import expm
import jax.numpy as jnp

num_qubits = 2
# 补全numpy导入,生成theta张量
theta = jnp.asarray(np.pi * np.random.random((15, 2, 2, 2, 2, 2, 2, 2, 2)))

def pauli_matrix(num_qubits):
    _pauli = jnp.array([
        [[1, 0], [0, 1]],    # 单位矩阵I
        [[0, 1], [1, 0]],    # X门
        [[0, -1j], [1j, 0]], # Y门
        [[1, 0], [0, -1]]     # Z门
    ])
    # 生成所有长度为num_qubits的Pauli矩阵索引组合
    indices = jnp.array(jnp.meshgrid(*[jnp.arange(4)]*num_qubits)).T.reshape(-1, num_qubits)
    # 过滤掉全单位元的组合(索引全0)
    non_identity_mask = jnp.any(indices != 0, axis=1)
    non_identity_indices = indices[non_identity_mask]
    # 逐个计算张量积,生成多量子比特Pauli矩阵
    pauli_list = []
    for idx in non_identity_indices:
        mat = _pauli[idx[0]]
        for i in idx[1:]:
            mat = jnp.kron(mat, _pauli[i])
        pauli_list.append(mat)
    return jnp.array(pauli_list)

def SpecialUnitary(num_qubits, theta):
    assert theta.shape[0] == 15, "theta的第0轴长度必须等于非单位元Pauli矩阵数量"
    pauli_mats = pauli_matrix(num_qubits)
    A = jnp.tensordot(theta, pauli_mats, axes=[[0], [0]])
    print(f'{A.shape= }{pauli_mats.shape=}{theta.shape=}')
    return expm(1j * A / 2)

# 执行计算
result = SpecialUnitary(num_qubits, theta)
print(f"结果形状: {result.shape}")

关键说明

  • Pauli矩阵生成修正:原函数错误地将整个Pauli数组重复做张量积,修正后的函数通过生成所有可能的单量子比特Pauli矩阵组合,逐个计算张量积,得到正确的(15,4,4)形状的多量子比特Pauli矩阵集合。
  • expm兼容性:修正后A的形状仍为(2,2,2,2,2,2,2,2,4,4),jax.scipy.linalg.expm会自动对最后两个轴的每个4x4方阵计算矩阵指数,完全符合文档要求。
  • 运行验证:修正后的代码可以正常执行,输出的结果形状与A一致(最后两个轴保持4x4),对应每个theta参数组合生成的SU(4)矩阵。

内容的提问来源于stack exchange,提问作者kekcikon

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.06.30 17:56:09