基于SentenceTransformer的语义相似度聚类返回空值排查求助
问题分析与修正
函数中的错误点
- 未定义变量
init_max_size:函数中直接使用该变量但未声明或传入,运行时会触发NameError。 - 缩进逻辑混乱:
if top_k_values[i][-1] >= threshold:块内的代码缩进错误,导致后续聚类逻辑无法在满足条件时执行。 - 无效代码:
break语句后的new_cluster.append(idx)永远不会执行,因为break会直接跳出循环。 - 参数与样本量不匹配:你的数据集仅9条样本,若
min_community_size设置过大(比如默认的20),top_k_values[i][-1]会取所有样本中最小的相似度值,大概率低于阈值,导致无社区被提取。
修正后的聚类函数
def detect_clusters(embeddings, threshold=0.75, min_community_size=2): # 计算余弦相似度矩阵 cos_scores = util.pytorch_cos_sim(embeddings, embeddings) extracted_communities = [] num_samples = len(embeddings) for i in range(num_samples): # 获取当前样本与所有样本的相似度并排序 sim_scores = cos_scores[i].squeeze().tolist() sorted_indices = sorted(range(num_samples), key=lambda x: sim_scores[x], reverse=True) # 收集所有相似度符合阈值的样本索引 cluster = [idx for idx in sorted_indices if sim_scores[idx] >= threshold] # 仅保留满足最小社区规模的聚类 if len(cluster) >= min_community_size: extracted_communities.append(cluster) # 去重,避免重复聚类 unique_communities = [] seen = set() for cluster in extracted_communities: cluster_tuple = tuple(sorted(cluster)) if cluster_tuple not in seen: seen.add(cluster_tuple) unique_communities.append(cluster) return unique_communities
关键优化说明
- 移除原函数中复杂且易出错的top_k筛选逻辑,改为直接筛选所有符合阈值的样本。
- 修复未定义变量问题,简化参数逻辑。
- 添加聚类去重逻辑,避免重复输出相同聚类。
- 适配小样本数据集,调整默认参数更合理。
调用示例
# 针对9条样本设置合理参数 clusters = detect_clusters(embeddings, min_community_size=2, threshold=0.75) # 查看聚类对应的原始短语 for idx, cluster in enumerate(clusters): print(f"聚类 {idx+1}:") for sample_idx in cluster: print(df['Key_Phrases'].iloc[sample_idx]) print("---")
运行后可得到语义相似的聚类(例如样本7和8会被归为一类)。
内容的提问来源于stack exchange,提问作者shan
相关产品推荐
相关产品推荐

