研究样本量对文本分类器性能影响的函数运行异常,求代码错误排查
代码错误排查与修复
核心错误点
- create_model未使用传入参数,硬编码读取全局数据集:函数定义时接收了
train_docs/train_y/test_docs/test_y四个入参,但内部所有读取训练、测试数据的逻辑全部调用全局的train、test变量,导致sample_size_impact传入的不同大小训练子集完全不生效,每次运行都用全量训练数据。 - TfidfVectorizer重复fit:代码中先后两次调用
tfidf_vect.fit_transform()处理训练文本,属于冗余操作,会造成不必要的性能损耗。 - 返回值处理错误:
create_model返回的是(auc_score, prc_score)二元组,sample_size_impact直接将元组存入列表绘图,会同时绘制两条不符合预期的曲线。 - 冗余/笔误问题:
sample_size_impact中定义的t_size变量未被使用,x轴标签拼写错误(Smple Size→Sample Size);SVM分支冗余初始化了SVC对象后又覆盖为LinearSVC;朴素贝叶斯分支的绘图标题错误标注为SVM模型。
修复后代码
1. 修复create_model函数
def create_model(train_docs, train_y, test_docs, test_y, \ model_type='svm', stop_words=None, min_df=1, print_result = True, algorithm_para=1.0): tfidf_vect = TfidfVectorizer(stop_words=stop_words, min_df=min_df) # 替换为传入的训练集拟合tfidf,删除重复fit逻辑 X_train = tfidf_vect.fit_transform(train_docs) X_test = tfidf_vect.transform(test_docs) # 替换为传入的标签 y_train = train_y y_test = test_y if 'svm' in model_type: # 删除冗余的SVC初始化 clf = svm.LinearSVC(C=algorithm_para).fit(X_train, y_train) predicted = clf.predict(X_test) labels = sorted(np.unique(y_train)) precision, recall, fscore, support = precision_recall_fscore_support(y_test, predicted, labels=labels) if print_result==True: print("labels: ", labels) print("precision: ", precision) print("recall: ", recall) print("f-score: ", fscore) print("support: ", support) predict_p = clf._predict_proba_lr(X_test) y_pred = predict_p[:,1] fpr, tpr, thresholds = roc_curve(y_test,y_pred, pos_label=1) precision, recall, thresholds = precision_recall_curve(y_test, y_pred, pos_label=1) auc_score = auc(fpr, tpr) prc_score = auc(recall, precision) if print_result==True: print("AUC: {:.2%}".format(auc_score), "PRC: {:.2%}".format(prc_score)) plt.figure() plt.plot(fpr, tpr, color='darkorange', lw=2) plt.plot([0, 1], [0, 1], color='navy', lw=2, linestyle='--') plt.xlim([0.0, 1.0]) plt.ylim([0.0, 1.05]) plt.xlabel('False Positive Rate') plt.ylabel('True Positive Rate') plt.title('AUC of SVM Model') plt.show() plt.figure() plt.plot(recall, precision, color='darkorange', lw=2) plt.xlim([0.0, 1.0]) plt.ylim([0.0, 1.05]) plt.xlabel('Recall') plt.ylabel('Precision') plt.title('Precision_Recall_Curve of SVM Model') plt.show() else: clf = MultinomialNB(alpha=algorithm_para).fit(X_train, y_train) predicted = clf.predict(X_test) labels = sorted(np.unique(y_train)) precision, recall, fscore, support = precision_recall_fscore_support(y_test, predicted, labels=labels) if print_result==True: print("labels: ", labels) print("precision: ", precision) print("recall: ", recall) print("f-score: ", fscore) print("support: ", support) predict_p = clf.predict_proba(X_test) y_pred = predict_p[:,1] fpr, tpr, thresholds = roc_curve(y_test,y_pred, pos_label=1) precision, recall, thresholds = precision_recall_curve(y_test, y_pred, pos_label=1) auc_score = auc(fpr, tpr) prc_score = auc(recall, precision) if print_result==True: print("AUC: {:.2%}".format(auc_score), "PRC: {:.2%}".format(prc_score)) plt.figure() plt.plot(fpr, tpr, color='darkorange', lw=2) plt.plot([0, 1], [0, 1], color='navy', lw=2, linestyle='--') plt.xlim([0.0, 1.0]) plt.ylim([0.0, 1.05]) plt.xlabel('False Positive Rate') plt.ylabel('True Positive Rate') plt.title('AUC of Naive Bayes Model') # 修正标题错误 plt.show() plt.figure() plt.plot(recall, precision, color='darkorange', lw=2) plt.xlim([0.0, 1.0]) plt.ylim([0.0, 1.05]) plt.xlabel('Recall') plt.ylabel('Precision') plt.title('Precision_Recall_Curve of Naive Bayes Model') # 修正标题错误 plt.show() return auc_score, prc_score
2. 修复sample_size_impact函数
def sample_size_impact(train_docs, train_y, test_docs, test_y): auc_list_svm = [] sample_sizes = [] # 按500步长遍历样本量 max_sample = len(train_docs) for i in range(int(max_sample/500)): current_size = (i+1)*500 # 截取对应大小的训练子集 sub_train_docs = train_docs[:current_size] sub_train_y = train_y[:current_size] # 只取返回的auc_score auc_score_svm, _ = create_model(sub_train_docs, sub_train_y, test_docs, test_y, \ model_type='svm', stop_words = 'english', min_df = 1, print_result=False, algorithm_para=1.0) auc_list_svm.append(auc_score_svm) sample_sizes.append(current_size) plt.figure() plt.plot(sample_sizes, auc_list_svm, color='darkorange') plt.xlabel('Sample Size') # 修正拼写错误 plt.ylabel('AUC') plt.title('Sample Size Impact on SVM Performance') plt.show()
3. 调用示例
# 提前提取数据传入函数,避免硬编码 train_docs = train['text'].values train_y = train['label'].values test_docs = test['text'].values test_y = test['label'].values sample_size_impact(train_docs, train_y, test_docs, test_y)
内容的提问来源于stack exchange,提问作者Vahid
相关产品推荐
相关产品推荐

