CoolFace
Modelpublic

crackrammer/ShieldBERT-Base-Chinese-Sensitive

sourceHugging Faceapache-2.0updated 7mo agoView on Hugging Face
0likes9downloads
train.py149 linesDownload Raw Back to root
1"""2训练脚本:使用 HuggingFace Trainer 微调 BERT 进行敏感词二分类3 4使用 BertForSequenceClassification + Trainer API,5支持自动混合精度、梯度累积、学习率调度等。6"""7 8import os9import json10import argparse11import numpy as np12import torch13 14from transformers import (15    BertTokenizer,16    BertForSequenceClassification,17    TrainingArguments,18    Trainer,19)20from datasets import load_dataset21from sklearn.metrics import accuracy_score, precision_score, recall_score, f1_score22 23 24def compute_metrics(eval_pred):25    """计算评估指标"""26    logits, labels = eval_pred27    preds = np.argmax(logits, axis=-1)28    return {29        "accuracy": accuracy_score(labels, preds),30        "precision": precision_score(labels, preds, average="binary"),31        "recall": recall_score(labels, preds, average="binary"),32        "f1": f1_score(labels, preds, average="binary"),33    }34 35 36def main():37    parser = argparse.ArgumentParser(description="训练敏感词过滤模型")38    parser.add_argument("--model_name", type=str, default="bert-base-chinese")39    parser.add_argument("--train_file", type=str, default="data/train.csv")40    parser.add_argument("--val_file", type=str, default="data/val.csv")41    parser.add_argument("--output_dir", type=str, default="output")42    parser.add_argument("--epochs", type=int, default=3)43    parser.add_argument("--batch_size", type=int, default=32)44    parser.add_argument("--lr", type=float, default=2e-5)45    parser.add_argument("--max_length", type=int, default=128)46    parser.add_argument("--warmup_ratio", type=float, default=0.1)47    parser.add_argument("--weight_decay", type=float, default=0.01)48    args = parser.parse_args()49 50    # 加载 tokenizer51    print(f"加载 tokenizer: {args.model_name}")52    tokenizer = BertTokenizer.from_pretrained(args.model_name)53 54    # 加载数据集55    print("加载数据集...")56    dataset = load_dataset(57        "csv",58        data_files={"train": args.train_file, "validation": args.val_file},59    )60    print(f"训练集: {len(dataset['train'])} 条")61    print(f"验证集: {len(dataset['validation'])} 条")62 63    # Tokenize64    def tokenize_fn(examples):65        return tokenizer(66            examples["text"],67            padding="max_length",68            truncation=True,69            max_length=args.max_length,70        )71 72    dataset = dataset.map(tokenize_fn, batched=True, remove_columns=["text"])73    dataset = dataset.rename_column("label", "labels")74    dataset.set_format("torch")75 76    # 加载模型77    print(f"加载模型: {args.model_name}")78    model = BertForSequenceClassification.from_pretrained(79        args.model_name,80        num_labels=2,81    )82 83    # 训练参数84    best_model_dir = os.path.join(args.output_dir, "best_model")85    training_args = TrainingArguments(86        output_dir=args.output_dir,87        num_train_epochs=args.epochs,88        per_device_train_batch_size=args.batch_size,89        per_device_eval_batch_size=args.batch_size * 2,90        learning_rate=args.lr,91        weight_decay=args.weight_decay,92        warmup_ratio=args.warmup_ratio,93        eval_strategy="epoch",94        save_strategy="epoch",95        load_best_model_at_end=True,96        metric_for_best_model="f1",97        greater_is_better=True,98        logging_steps=50,99        fp16=False,100        bf16=False,101        save_total_limit=2,102        report_to="none",103    )104 105    # Trainer106    trainer = Trainer(107        model=model,108        args=training_args,109        train_dataset=dataset["train"],110        eval_dataset=dataset["validation"],111        compute_metrics=compute_metrics,112    )113 114    # 训练115    print(f"\n{'='*60}")116    print(f"开始训练")117    print(f"Epochs: {args.epochs}, Batch Size: {args.batch_size}, LR: {args.lr}")118    print(f"{'='*60}\n")119 120    trainer.train()121 122    # 保存最佳模型123    print(f"\n保存最佳模型至: {best_model_dir}")124    trainer.save_model(best_model_dir)125    tokenizer.save_pretrained(best_model_dir)126 127    # 保存配置信息128    config = {129        "model_name": args.model_name,130        "max_length": args.max_length,131        "num_labels": 2,132        "label_map": {"0": "正常", "1": "敏感"},133    }134    with open(os.path.join(best_model_dir, "filter_config.json"), "w", encoding="utf-8") as f:135        json.dump(config, f, ensure_ascii=False, indent=2)136 137    # 最终评估138    final_metrics = trainer.evaluate()139    print(f"\n{'='*60}")140    print(f"训练完成!")141    print(f"F1: {final_metrics.get('eval_f1', 'N/A'):.4f}")142    print(f"Accuracy: {final_metrics.get('eval_accuracy', 'N/A'):.4f}")143    print(f"模型已保存至: {best_model_dir}")144    print(f"{'='*60}")145 146 147if __name__ == "__main__":148    main()149