crackrammer/ShieldBERT-Base-Chinese-Sensitive
09
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 