CoolFace
Modelpublic

crackrammer/ShieldBERT-Base-Chinese-Sensitive

sourceHugging Faceapache-2.0updated 7mo agoView on Hugging Face
0likes9downloads
evaluate.py116 linesDownload Raw Back to root
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