crackrammer/ShieldBERT-Base-Chinese-Sensitive
09
1"""2评估脚本:在测试集上评估模型性能3"""4 5import os6import json7import argparse8 9import numpy as np10from transformers import BertTokenizer, BertForSequenceClassification, Trainer, TrainingArguments11from datasets import load_dataset12from sklearn.metrics import (13 accuracy_score,14 precision_score,15 recall_score,16 f1_score,17 confusion_matrix,18 classification_report,19)20 21 22def compute_metrics(eval_pred):23 logits, labels = eval_pred24 preds = np.argmax(logits, axis=-1)25 return {26 "accuracy": accuracy_score(labels, preds),27 "precision": precision_score(labels, preds, average="binary"),28 "recall": recall_score(labels, preds, average="binary"),29 "f1": f1_score(labels, preds, average="binary"),30 }31 32 33def main():34 parser = argparse.ArgumentParser(description="评估敏感词过滤模型")35 parser.add_argument("--model_path", type=str, default="output/best_model")36 parser.add_argument("--test_file", type=str, default="data/test.csv")37 parser.add_argument("--batch_size", type=int, default=64)38 args = parser.parse_args()39 40 # 加载配置41 config_path = os.path.join(args.model_path, "filter_config.json")42 with open(config_path, "r", encoding="utf-8") as f:43 config = json.load(f)44 45 max_length = config.get("max_length", 128)46 label_names = [config["label_map"]["0"], config["label_map"]["1"]]47 48 # 加载模型和 tokenizer49 print(f"加载模型: {args.model_path}")50 tokenizer = BertTokenizer.from_pretrained(args.model_path)51 model = BertForSequenceClassification.from_pretrained(args.model_path)52 53 # 加载测试集54 dataset = load_dataset("csv", data_files={"test": args.test_file})55 56 def tokenize_fn(examples):57 return tokenizer(examples["text"], padding="max_length", truncation=True, max_length=max_length)58 59 dataset = dataset.map(tokenize_fn, batched=True, remove_columns=["text"])60 dataset = dataset.rename_column("label", "labels")61 dataset.set_format("torch")62 print(f"测试集: {len(dataset['test'])} 条")63 64 # 评估65 training_args = TrainingArguments(66 output_dir="/tmp/eval_output",67 per_device_eval_batch_size=args.batch_size,68 report_to="none",69 )70 trainer = Trainer(model=model, args=training_args, compute_metrics=compute_metrics)71 72 # 预测73 predictions = trainer.predict(dataset["test"])74 preds = np.argmax(predictions.predictions, axis=-1)75 labels = predictions.label_ids76 77 acc = accuracy_score(labels, preds)78 precision = precision_score(labels, preds, average="binary")79 recall = recall_score(labels, preds, average="binary")80 f1 = f1_score(labels, preds, average="binary")81 cm = confusion_matrix(labels, preds)82 83 print(f"\n{'='*60}")84 print("模型评估结果")85 print(f"{'='*60}")86 print(f"准确率 (Accuracy): {acc:.4f}")87 print(f"精确率 (Precision): {precision:.4f}")88 print(f"召回率 (Recall): {recall:.4f}")89 print(f"F1 值 (F1-Score): {f1:.4f}")90 91 print(f"\n--- 混淆矩阵 ---")92 print(f"{'':>12} 预测正常 预测敏感")93 print(f"{'实际正常':>10} {cm[0][0]:>6} {cm[0][1]:>6}")94 print(f"{'实际敏感':>10} {cm[1][0]:>6} {cm[1][1]:>6}")95 96 print(f"\n--- 分类报告 ---")97 report = classification_report(labels, preds, target_names=label_names, digits=4)98 print(report)99 100 # 保存结果101 results = {102 "accuracy": acc,103 "precision": precision,104 "recall": recall,105 "f1": f1,106 "confusion_matrix": cm.tolist(),107 }108 output_path = os.path.join(os.path.dirname(args.model_path), "eval_results.json")109 with open(output_path, "w", encoding="utf-8") as f:110 json.dump(results, f, ensure_ascii=False, indent=2)111 print(f"\n评估结果已保存至: {output_path}")112 113 114if __name__ == "__main__":115 main()116 