Python中使用JAX计算矩阵指数报错:期望输入为方阵
问题分析与解决方案
核心问题
- Pauli矩阵生成函数逻辑错误:你的
pauli_matrix函数没有正确生成多量子比特的Pauli矩阵集合,错误的张量积操作导致后续计算的矩阵结构不符合expm的隐含要求。 - 缺失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
相关产品推荐
相关产品推荐

