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

TensorFlow中如何高效实现余弦相似度生成指定维度张量?

问题与解决方案

我需要通过矩阵M与X的运算生成维度为32 x 576 x 2的输出张量,两者形状如下:

M.shape: (576, 2, 2048)
X.shape: (32, 2048)

运算为元素级余弦相似度,即特征向量𝑥与向量M_j,k的余弦相似度(公式为:余弦相似度 = (x·y)/(||x||₂ × ||y||₂))。

当前代码实现存在错误(其中BATCH_SIZE=32,C=576,V=2):

@tf.function
def call(self, X):
    M = self.kernel

    norm_M = tf.norm(M, ord=2, axis=2)
    norm_X = tf.norm(X, ord=2, axis=1)

    l_r = (some scalar value, separate to this question)

    # Compute cosine similarity between X and M
    # as a matrix with dimensionality:
    # BATCH_SIZE x C x V
    feature_batch_size = tf.shape(X)[0]
    c = tf.shape(M)[0]
    v = tf.shape(M)[1]
    output_matrix = tf.zeros([feature_batch_size, c, v])
    output_matrix = tf.Variable(output_matrix, trainable=False)
    for row in tf.range(feature_batch_size):
      for column in tf.range(c):
        for channel in tf.range(v):
          a = tf.tensordot(M[column][channel], X[row], 1)
          b = norm_M[column][channel] * norm_X[row]
          output_matrix[row][column][channel] = a / b

    return [output_matrix, l_r]

代码在output_matrix[row][column][channel] = a / b行执行失败,原因是无法对tf.Variable的单个元素直接赋值。

请问:

  1. 有没有更优的实现方式,无需三层嵌套循环且兼容tf.Function图功能,来完成该运算并生成目标输出矩阵?
  2. 如果没有更优方式,如何正确对tf.Variable的单个元素进行赋值?

额外信息:

norm_M.shape: (576, 2)
norm_X.shape: (32,)

解决方案

1. 无循环的高效实现(推荐)

TensorFlow的广播机制和矩阵运算可以完全替代嵌套循环,同时完美兼容tf.function,效率远高于循环实现。具体实现如下:

@tf.function
def call(self, X):
    M = self.kernel

    norm_M = tf.norm(M, ord=2, axis=2)
    norm_X = tf.norm(X, ord=2, axis=1)
    l_r = (some scalar value, separate to this question)

    # 扩展维度实现广播,计算点积
    X_expanded = tf.expand_dims(tf.expand_dims(X, axis=1), axis=1)  # shape: (32, 1, 1, 2048)
    M_expanded = tf.expand_dims(M, axis=0)  # shape: (1, 576, 2, 2048)
    dot_product = tf.reduce_sum(X_expanded * M_expanded, axis=-1)  # shape: (32, 576, 2)

    # 扩展范数维度,计算分母乘积
    norm_X_expanded = tf.expand_dims(tf.expand_dims(norm_X, axis=1), axis=1)  # shape: (32, 1, 1)
    norm_M_expanded = tf.expand_dims(norm_M, axis=0)  # shape: (1, 576, 2)
    denominator = norm_X_expanded * norm_M_expanded  # shape: (32, 576, 2)

    # 计算余弦相似度,加epsilon避免除以0
    cos_sim = dot_product / (denominator + 1e-8)

    return [cos_sim, l_r]

逻辑说明:

  • 通过tf.expand_dims给张量扩展维度,让X和M、范数张量可以通过广播完成批量运算
  • 点积通过元素相乘后在特征维度求和得到
  • 最后直接用广播后的点积除以范数乘积,得到目标形状的余弦相似度矩阵

2. 正确的tf.Variable元素赋值方式(不推荐,仅作参考)

如果一定要保留循环逻辑,不能直接通过索引赋值,需使用tf.tensor_scatter_nd_update来更新单个元素。修改后的循环部分代码如下:

@tf.function
def call(self, X):
    M = self.kernel

    norm_M = tf.norm(M, ord=2, axis=2)
    norm_X = tf.norm(X, ord=2, axis=1)
    l_r = (some scalar value, separate to this question)

    feature_batch_size = tf.shape(X)[0]
    c = tf.shape(M)[0]
    v = tf.shape(M)[1]
    output_matrix = tf.Variable(tf.zeros([feature_batch_size, c, v]), trainable=False)

    for row in tf.range(feature_batch_size):
      for column in tf.range(c):
        for channel in tf.range(v):
          a = tf.tensordot(M[column][channel], X[row], 1)
          b = norm_M[column][channel] * norm_X[row]
          value = a / (b + 1e-8)
          # 构造索引并更新变量
          indices = tf.constant([[row, column, channel]])
          output_matrix.assign(tf.tensor_scatter_nd_update(output_matrix, indices, [value]))

    return [output_matrix, l_r]

注意:这种三层循环在tf.function中会被编译为图,但运算效率远低于广播实现,尤其是当维度较大时,不建议使用。

内容的提问来源于stack exchange,提问作者magmacollaris

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.08.08 10:15:12