如何针对任意维度张量的最后两维计算Kronecker积?
实现通用维度张量最后两维的Kronecker积
方法一:利用广播与维度重塑(最通用)
Kronecker积可拆解为元素级广播相乘后合并对应维度,该方法适配任意d≥2的张量:
- 对A的最后两维分别插入单例维度,形状变为
(*batch_dims, A_{d-1}, 1, A_d, 1) - 对B的最后两维分别插入单例维度,形状变为
(*batch_dims, 1, B_{d-1}, 1, B_d) - 执行元素乘法后,形状变为
(*batch_dims, A_{d-1}, B_{d-1}, A_d, B_d) - 合并对应维度,得到最终形状
(*batch_dims, A_{d-1}*B_{d-1}, A_d*B_d)
JAX代码实现:
import jax.numpy as jnp def my_kron(A, B): # 扩展维度实现广播匹配 A_exp = A[..., :, None, :, None] B_exp = B[..., None, :, None, :] # 元素级相乘 kron_prod = A_exp * B_exp # 计算目标形状并重塑 batch_shape = kron_prod.shape[:-4] new_last_dims = (kron_prod.shape[-4] * kron_prod.shape[-3], kron_prod.shape[-2] * kron_prod.shape[-1]) return kron_prod.reshape(*batch_shape, *new_last_dims)
方法二:结合vmap与维度扁平化
通过扁平化前d-2维为单batch维度,再用批量映射处理2D矩阵的Kronecker积:
import jax import jax.numpy as jnp def my_kron_vmap(A, B): # 提取前d-2维的形状 batch_shape = A.shape[:-2] # 扁平化前d-2维为单个batch维度 A_flat = A.reshape(-1, *A.shape[-2:]) B_flat = B.reshape(-1, *B.shape[-2:]) # 批量执行Kronecker积 kron_flat = jax.vmap(jnp.kron)(A_flat, B_flat) # 恢复原维度结构 new_last_dims = (A.shape[-2] * B.shape[-2], A.shape[-1] * B.shape[-1]) return kron_flat.reshape(*batch_shape, *new_last_dims)
验证示例
以d=3的场景测试:
A = jnp.ones((2, 3, 3)) B = jnp.ones((2, 4, 4)) C1 = my_kron(A, B) C2 = my_kron_vmap(A, B) print(C1.shape) # 输出 (2, 12, 12) print(jnp.allclose(C1, C2)) # 输出 True
两种方法均支持任意d≥2的张量计算,方法一无需维度扁平化,逻辑更直接;方法二贴合批量处理思维,易于理解。
内容的提问来源于stack exchange,提问作者bayes2021
相关产品推荐
相关产品推荐

