torch.einsum如何从3D张量与2D张量生成4D张量?
理解Mamba-SSM中delta与A的einsum运算逻辑
先明确涉及的张量维度定义:
delta:[b, d, l],其中b=batch_size,d=d_inner,l=sequence_lengthA:[d, n],其中n=d_state
原代码的einsum运算逻辑
torch.einsum('bdl,dn->bdln', delta, A)的核心是对每个batch、每个d_inner维度,执行序列向量与状态向量的外积,具体拆解:
- 固定某一个样本(b维度)和某一个inner特征维度(d维度):
delta[b, d, :]是长度为l的1D张量(对应该样本该特征维度下的整个序列)A[d, :]是长度为n的1D张量(对应该特征维度下的状态参数)- 两者做外积,得到一个
[l, n]的2D张量(序列中每个位置与每个状态维度一一相乘)
- 把所有b和d对应的
[l, n]张量堆叠起来,最终得到[b, d, l, n]的4D张量,再经过torch.exp()得到deltaA。
你可以用广播乘法的等价代码来直观理解:
# 与原einsum完全等价的实现 delta_expanded = delta.unsqueeze(-1) # shape: [b, d, l, 1] A_expanded = A.unsqueeze(0).unsqueeze(2) # shape: [1, d, 1, n] deltaA = torch.exp(delta_expanded * A_expanded)
对比你能理解的矩阵乘法版einsum
torch.einsum('bdl,dn->bln', delta, A)是收缩d维度的矩阵乘法:
- 固定某一个样本(b维度)时,
delta[b, :, :]是[d, l]的矩阵,转置后为[l, d],与A[d, n]做矩阵乘法,得到[l, n]的矩阵 - 最终堆叠所有样本的结果,得到
[b, l, n]的3D张量
两者的核心区别在于:原代码没有收缩d维度,而是保留d维度的同时,对每个d维度下的序列和状态参数做外积,因此输出多了一个d维度,从3D变成4D。
内容的提问来源于stack exchange,提问作者Matsuri247
相关产品推荐
相关产品推荐

