如何使用TensorFlow实现余弦相似度计算
TensorFlow替换余弦相似度计算实现方案
完全可以用TensorFlow替换原有scikit-learn的余弦相似度计算环节,原有数据读取、特征拼接、CountVectorizer向量化的逻辑都不需要改动,仅替换相似度计算部分即可,计算结果和sklearn原生实现基本一致,大样本场景下还可以调用GPU加速计算。
核心实现逻辑
余弦相似度的本质是两个向量L2归一化后的点积,基于这个逻辑用TensorFlow实现可以完全避开API默认参数的坑:注意TF内置的tf.keras.losses.cosine_similarity默认返回负的余弦值,是为了适配损失函数最小化的训练目标,直接调用会得到反向的结果。
适配稀疏矩阵的通用版本
CountVectorizer输出的是scipy稀疏矩阵,大词表场景下转稠密矩阵会占用大量内存,优先用稀疏张量兼容版本:
import pandas as pd import numpy as np import tensorflow as tf from sklearn.feature_extraction.text import CountVectorizer # --------------- 原有前置逻辑完全保留 --------------- df = pd.read_csv('shops.csv', sep='|') df.columns = ['name', 'cate_1', 'cate_2', 'cate_3', 'dong', 'lon', 'lat'] df['cate_mix'] = df['cate_1'] + df['cate_2'] + df['cate_3'] df['cate_mix'] = df['cate_mix'].str.replace("/", " ") count_vect_category = CountVectorizer(min_df=0, ngram_range=(1,2)) place_category = count_vect_category.fit_transform(df['cate_mix']) # --------------------------------------------------- # scipy稀疏矩阵转TF稀疏张量工具函数 def scipy_sparse_to_tf_sparse(sparse_mat): sparse_mat = sparse_mat.tocoo() indices = np.column_stack((sparse_mat.row, sparse_mat.col)) return tf.SparseTensor( indices=indices, values=sparse_mat.data.astype(np.float32), dense_shape=sparse_mat.shape ) place_category_tf = scipy_sparse_to_tf_sparse(place_category) # 计算每个向量的L2范数,做归一化,加极小值避免除零错误 l2_norms = tf.sparse.reduce_sum(place_category_tf ** 2, axis=1, keepdims=True) ** 0.5 normalized_vecs = tf.sparse.divide(place_category_tf, l2_norms + 1e-10) # 归一化向量点积得到余弦相似度矩阵 place_simi_cate = tf.sparse.sparse_dense_matmul( normalized_vecs, tf.sparse.to_dense(normalized_vecs), adjoint_b=True ).numpy() # 排序逻辑和原有代码完全一致 place_simi_cate_sorted_ind = place_simi_cate.argsort()[:, ::-1]
小规模数据集简化版本
如果数据集规模小(商铺数<5万、词表规模<1万),可以直接转稠密矩阵计算,代码更简洁:
# 向量化结果转稠密数组 place_category_dense = place_category.toarray().astype(np.float32) # 直接调用TF内置的L2归一化API normalized_vecs = tf.math.l2_normalize(place_category_dense, axis=1) # 矩阵乘法得到相似度矩阵 place_simi_cate = tf.matmul(normalized_vecs, normalized_vecs, transpose_b=True).numpy() # 排序逻辑不变 place_simi_cate_sorted_ind = place_simi_cate.argsort()[:, ::-1]
注意事项
- 上述实现的计算结果和sklearn
cosine_similarity的数值差异在1e-6量级,属于浮点计算正常误差,不影响排序结果。 - 大样本场景下可以将张量加载到GPU显存计算,相比sklearn的CPU计算可以获得数倍到数十倍的速度提升。
- 计算得到的
place_simi_cate、place_simi_cate_sorted_ind格式和原有sklearn版本完全一致,后续业务逻辑不需要做任何适配。
内容的提问来源于stack exchange,提问作者min
相关产品推荐
相关产品推荐

