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

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和上面的特征维度完全一致

这个表达式的计算过程:

  1. 两个输入共享维度c,所以einsum会对c维度做对应元素相乘后求和(也就是向量点积)
  2. 箭头->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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.07.10 01:03:26