BERT无标注文本多分类任务中计算余弦相似度出现NaN和维度错误求助
问题根源
- 第一个
NaN/无穷值报错:你仅处理了输入矩阵X的异常值,未处理dic_y中存储的y向量的异常值,计算时y携带的无效值触发了报错。 - 第二个
维度不匹配报错:sklearn.metrics.pairwise.cosine_similarity要求两个输入均为2D数组,你仅调整了X的维度,循环中调用的y仍为1D数组,所以即使X是2D也会触发维度错误。
修复方案
第一步:统一处理所有输入的异常值与维度
import numpy as np from sklearn import metrics # 处理X的异常值与格式 X = np.nan_to_num(X.astype(np.float32)) # 若X本身为1D(单样本),统一转成2D if X.ndim == 1: X = X.reshape(1, -1) # 处理dic_y中所有y的异常值与格式 processed_y = [] for y in dic_y.values(): y = np.nan_to_num(y.astype(np.float32)) # 1D的y统一转成2D if y.ndim == 1: y = y.reshape(1, -1) processed_y.append(y)
第二步:重新计算余弦相似度
similarities = np.array( [metrics.pairwise.cosine_similarity(X, y)[:, 0] for y in processed_y] ).T
可选调试步骤
如果仍有报错,可添加以下代码确认所有输入的合法性:
# 校验X的状态 print(f"X 维度: {X.shape}, 含NaN: {np.isnan(X).any()}, 含inf: {np.isinf(X).any()}") # 校验所有y的状态 for i, y in enumerate(processed_y): print(f"第{i}个y 维度: {y.shape}, 含NaN: {np.isnan(y).any()}, 含inf: {np.isinf(y).any()}")
内容的提问来源于stack exchange,提问作者William
相关产品推荐
相关产品推荐

