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

能否用余弦距离在Sklearn中对BERT嵌入做KMeans聚类?求方案与代码

BERT嵌入+KMeans聚类(余弦距离适配方案)

能不能用余弦距离做KMeans?

标准KMeans算法原生基于欧氏距离优化,但可以通过L2归一化嵌入向量间接实现余弦距离的聚类效果。数学上,两个L2归一化后的向量,欧氏距离的平方等于2*(1-余弦相似度),这意味着最小化欧氏距离等价于最大化余弦相似度,完全适配余弦距离的聚类需求。

解决方案步骤

  • 用sentence-transformers的bert-base-nli-mean-tokens生成文档的BERT嵌入
  • 对生成的嵌入做L2归一化处理
  • 使用sklearn的KMeans进行聚类(此时欧氏距离等价于余弦距离)

完整代码示例

首先安装依赖:

pip install sentence-transformers scikit-learn numpy

核心代码:

from sentence_transformers import SentenceTransformer
from sklearn.cluster import KMeans
from sklearn.preprocessing import normalize
import numpy as np

# 加载预训练模型并生成BERT嵌入
model = SentenceTransformer('bert-base-nli-mean-tokens')
# 示例文档列表
documents = [
    "机器学习是人工智能的一个分支",
    "KMeans是一种无监督聚类算法",
    "BERT模型擅长处理自然语言理解任务",
    "余弦距离常用于衡量文本向量的相似度",
    "无监督学习不需要标注数据",
    "Transformer架构是很多NLP模型的基础"
]
embeddings = model.encode(documents)

# 对嵌入做L2归一化
normalized_embeddings = normalize(embeddings, norm='l2')

# 用KMeans聚类(此时欧氏距离等价于余弦距离)
num_clusters = 2
kmeans = KMeans(n_clusters=num_clusters, random_state=42)
cluster_labels = kmeans.fit_predict(normalized_embeddings)

# 输出聚类结果
for idx, (doc, label) in enumerate(zip(documents, cluster_labels)):
    print(f"文档{idx+1}: {doc} | 聚类标签: {label}")

代码说明

  • 归一化步骤是核心:确保后续KMeans的欧氏距离计算等价于余弦距离的相似度比较
  • random_state设置为固定值保证聚类结果可复现
  • 可根据实际文档数量调整num_clusters参数

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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.06.29 11:42:05