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

如何针对任意维度张量的最后两维计算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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.08.19 15:25:17