torch.einsum API工作原理及特定 einsum 调用的相似度类型咨询
torch.einsum 工作机制及相似度计算解析
一、torch.einsum 核心工作机制
- einsum基于爱因斯坦求和约定,通过字符串就能定义张量间的维度映射、求和逻辑,不用手动写循环或复杂维度变换。
- 核心规则:
- 逗号分隔多个输入张量的维度标识(比如
ac,bc) - 箭头
->后是输出张量的维度 - 多个输入里重复出现的维度会被自动求和缩并,箭头后保留的维度就是输出的维度顺序
- 逗号分隔多个输入张量的维度标识(比如
二、"ac,bc->ab" 计算逻辑拆解
先明确输入张量的维度含义:
norm_max_func_embedding:维度(a, c),a是函数样本数量(这里代码里是1个max函数),c是embedding的特征维度norm_nl_embedding:维度(b, c),b是自然语言描述的样本数量,c和上面的特征维度完全一致
这个表达式的计算过程:
- 两个输入共享维度
c,所以einsum会对c维度做对应元素相乘后求和(也就是向量点积) - 箭头
->ab指定输出是a×b的矩阵:输出矩阵的第i行第j列元素,就是第一个张量第i行向量和第二个张量第j行向量的点积
用公式直白表示:
$$output[i][j] = \sum_{k=1}^c norm_max_func_embedding[i][k] \times norm_nl_embedding[j][k]$$
三、这是哪种相似度计算?
代码里先对所有embedding做了torch.nn.functional.normalize(..., p=2, dim=1)——这是L2归一化,把每个样本的特征向量缩放到单位长度(向量的L2范数为1)。
在L2归一化后,两个向量的点积就等价于余弦相似度:
余弦相似度的公式是 $\text{cosine}(u, v) = \frac{u \cdot v}{||u||_2 ||v||_2}$,当两个向量都是单位向量时,分母为1,点积直接等于余弦相似度,取值范围[-1,1],值越大表示向量越相似。
所以这段代码里的einsum,实际是在计算函数embedding与自然语言embedding的余弦相似度矩阵:每行对应一个函数,每列对应一个自然语言描述,元素值就是两者的相似程度。
四、结合代码的实际用途
代码里分别生成了max、min两个函数的embedding,再和自然语言embedding计算相似度矩阵,目的就是找出和每个函数最匹配的自然语言描述——比如max函数的embedding,应该和描述“取两个数里的最大值”的自然语言embedding得到更高的相似度得分。
内容的提问来源于stack exchange,提问作者Exploring
相关产品推荐
相关产品推荐

