### train_base_bert-base-chinese.py ```python import re import mlflow from sklearn.metrics import f1_score import pandas as pd import torch from sklearn.metrics import classification_report from sklearn.model_selection import train_test_split from torch import nn from torch.utils.data import Dataset, DataLoader from transformers import ( BertTokenizer, BertForSequenceClassification, AdamW, get_linear_schedule_with_warmup, ) # 初始化MLFlow mlflow.set_experiment("Weibo Sentiment Analysis v2") # 设置设备 device = torch.device("cuda" if torch.cuda.is_available() else "cpu") print(f"Using device: {device}") class WeiboDataset(Dataset): """优化后的数据集类""" def __init__(self, texts, labels, tokenizer, max_len=128): self.texts = texts self.labels = labels self.tokenizer = tokenizer self.max_len = max_len def __len__(self): return len(self.texts) def __getitem__(self, idx): text = str(self.texts.iloc[idx]) label = self.labels.iloc[idx] encoding = self.tokenizer.encode_plus( text, add_special_tokens=True, max_length=self.max_len, padding="max_length", truncation=True, return_attention_mask=True, return_tensors="pt", ) return { "input_ids": encoding["input_ids"].flatten(), "attention_mask": encoding["attention_mask"].flatten(), "labels": torch.tensor(label, dtype=torch.long), } class DataProcessor: """数据预处理管道""" @staticmethod def clean_text(text): """与预测代码完全一致的清洗逻辑""" if not isinstance(text, str): return "" text = re.sub(r"//@\S+[::]\s*", "", text) text = re.sub(r"#\S+#", "", text) text = re.sub(r"@\S+\s*", "", text) text = re.sub(r"https?://\S+", "", text) text = re.sub(r"//+", "", text) return re.sub(r"\s+", " ", text).strip() def process(self, file_path): data_frame = pd.read_csv( file_path, usecols=["label", "text"], dtype={"label": int, "text": str} ).dropna() data_frame["cleaned_text"] = data_frame["text"].apply(self.clean_text) sample_check = data_frame.sample(5, random_state=42) print("\n清洗后样本示例:") for idx, row in sample_check.iterrows(): print(f"原始: {row['text']}\n清洗: {row['cleaned_text']}\n") return data_frame[data_frame["cleaned_text"].str.len() > 0] class FGM: """对抗训练模块(调整epsilon)""" def __init__(self, model): self.model = model self.backup = {} def attack(self, epsilon=0.12): for name, param in self.model.named_parameters(): if param.requires_grad and "embeddings" in name: self.backup[name] = param.data.clone() if param.grad is not None: norm = torch.norm(param.grad) if norm > 0: r_at = epsilon * param.grad / norm param.data.add_(r_at) def restore(self): for name, param in self.model.named_parameters(): if name in self.backup: param.data = self.backup[name] self.backup.clear() class BertSentimentClassifier: def __init__(self, num_labels=2): self.tokenizer = BertTokenizer.from_pretrained("bert-base-chinese") self.model = BertForSequenceClassification.from_pretrained( "bert-base-chinese", num_labels=num_labels ).to(device) self.criterion = nn.CrossEntropyLoss() def create_data_loader(self, date_frame, batch_size=32): dataset = WeiboDataset( texts=date_frame["cleaned_text"], labels=date_frame["label"], tokenizer=self.tokenizer ) return DataLoader( dataset, batch_size=batch_size, num_workers=4, pin_memory=True, shuffle=True ) def train(self, train_loader_g, val_loader_g, epochs=4): optimizer = AdamW(self.model.parameters(), lr=2e-5, weight_decay=0.01) scheduler = get_linear_schedule_with_warmup( optimizer, num_warmup_steps=100, num_training_steps=len(train_loader_g) * epochs ) fgm = FGM(self.model) best_f1 = 0 with mlflow.start_run(): for epoch in range(epochs): self.model.train() total_loss = 0 for batch in train_loader_g: optimizer.zero_grad() inputs = { "input_ids": batch["input_ids"].to(device), "attention_mask": batch["attention_mask"].to(device), "labels": batch["labels"].to(device), } outputs = self.model(**inputs) loss = self.criterion(outputs.logits, inputs["labels"]) # 对抗训练 loss.backward() fgm.attack() # epsilon=0.12 # 对抗样本前向 outputs_adv = self.model(**inputs) loss_adv = self.criterion(outputs_adv.logits, inputs["labels"]) loss_adv.backward() # 累计梯度 fgm.restore() # 梯度裁剪 torch.nn.utils.clip_grad_norm_(self.model.parameters(), 1.0) optimizer.step() scheduler.step() total_loss += loss.item() # 验证 val_f1, val_report = self.evaluate(val_loader_g) avg_loss = total_loss / len(train_loader_g) # 记录指标 mlflow.log_metrics({ "train_loss": avg_loss, "val_f1": val_f1 }, step=epoch) print(f"Epoch {epoch+1}/{epochs}") print(f"Train Loss: {avg_loss:.4f} | Val F1: {val_f1:.4f}") print("Classification Report:") print(val_report) # 保存最佳模型 if val_f1 > best_f1: best_f1 = val_f1 self.model.save_pretrained("best_model") self.tokenizer.save_pretrained("best_model") mlflow.log_artifacts("best_model", "model") mlflow.pytorch.log_model(self.model, "final_model") def evaluate(self, data_loader): self.model.eval() all_pred_result = [] all_labels = [] with torch.no_grad(): for batch in data_loader: inputs = { "input_ids": batch["input_ids"].to(device), "attention_mask": batch["attention_mask"].to(device), } labels = batch["labels"].to(device) outputs = self.model(**inputs) pred_result = torch.argmax(outputs.logits, dim=1) all_pred_result.extend(pred_result.cpu().numpy()) all_labels.extend(labels.cpu().numpy()) report = classification_report( all_labels, all_pred_result, target_names=["Negative", "Positive"], digits=4 ) f1 = f1_score(all_labels, all_pred_result, average="macro") return f1, report if __name__ == "__main__": # 数据准备 processor = DataProcessor() df = processor.process("weibo_senti_100k.csv") # 拆分数据集 train_df, val_df = train_test_split( df, test_size=0.2, stratify=df["label"], random_state=42 ) # 初始化模型 classifier = BertSentimentClassifier() # 创建数据加载器 train_loader = classifier.create_data_loader(train_df, batch_size=32) val_loader = classifier.create_data_loader(val_df, batch_size=64) # 开始训练 classifier.train(train_loader, val_loader, epochs=4) print("Training finished!") ``` ### test_base_bert-base-chinese.py ```python import re import numpy as np import pandas as pd import torch from transformers import BertTokenizer, BertForSequenceClassification class WeiboClassifier: def __init__(self, model_path): self.device = torch.device("cuda" if torch.cuda.is_available() else "cpu") self.tokenizer = BertTokenizer.from_pretrained(model_path) self.model = BertForSequenceClassification.from_pretrained(model_path).to(self.device) self.model.eval() @staticmethod def clean_text(original_text): """与训练完全一致的清洗逻辑""" if not isinstance(original_text, str): return "" original_text = re.sub(r"//@\S+[::]\s*", "", original_text) original_text = re.sub(r"#\S+#", "", original_text) original_text = re.sub(r"@\S+\s*", "", original_text) original_text = re.sub(r"https?://\S+", "", original_text) original_text = re.sub(r"//+", "", original_text) original_text = re.sub(r"\s+", " ", original_text).strip() print(f"清洗后文本: {original_text}") return original_text def predict(self, predict_text, return_prob=False): processed_text = self.clean_text(predict_text) if not processed_text: return 0 if not return_prob else (0, [1.0, 0.0]) try: inputs = self.tokenizer( processed_text, max_length=128, padding="max_length", truncation=True, return_tensors="pt" ).to(self.device) with torch.no_grad(): outputs = self.model(**inputs) prob_result = torch.softmax(outputs.logits, dim=1).cpu().numpy()[0] # 取第一个样本的概率 pred_result = np.argmax(prob_result) # 取最大概率对应的类别 if return_prob: return pred_result, prob_result.tolist() return pred_result except Exception as e: print(f"预测错误: {str(e)}") return 0 if not return_prob else (0, [1.0, 0.0]) if __name__ == "__main__": # 示例测试 classifier = WeiboClassifier("best_model") # 从CSV读取测试数据 test_df = pd.read_csv("test_comments.csv") test_cases = [(row["comment"], row["label"]) for _, row in test_df.iterrows()] total_cases = len(test_cases) total_right = 0 total_error = 0 for text, expected in test_cases: print("\n" + "=" * 50) print(f"输入文本: {text}") pred, probs = classifier.predict(text, return_prob=True) print(f"预测概率: [Negative: {probs[0]:.4f}, Positive: {probs[1]:.4f}]") print(f"预测结果: {pred} | 期望: {expected} | {'正确' if pred == expected else '错误'}") if pred == expected: total_right += 1 else: total_error += 1 print("\n" + "=" * 50) print(f"测试结果: 总用例数: {total_cases}, 正确率: {total_right / total_cases:.4f}, 错误率: {total_error / total_cases:.4f}") ``` ## 项目分析 以下是对项目的详细流程分析和整体功能总结: ### 训练代码流程图 ```mermaid %% 训练代码流程图 graph TD A[开始] --> B[初始化MLFlow实验] B --> C[设置计算设备] C --> D[数据预处理] D --> D1[读取CSV文件] D1 --> D2[清洗文本] D2 --> D3[过滤空文本] D3 --> E[拆分训练集/验证集] E --> F[初始化BERT分类器] F --> G[创建数据加载器] G --> H[开始训练循环] H --> H1[对抗训练FGM] H1 --> H2[前向传播] H2 --> H3[计算损失] H3 --> H4[反向传播] H4 --> H5[对抗攻击] H5 --> H6[对抗样本反向传播] H6 --> H7[梯度裁剪] H7 --> H8[优化器更新] H8 --> I[验证评估] I --> J[记录指标到MLFlow] J --> K{是否最佳模型?} K -->|是| L[保存最佳模型] K -->|否| H L --> M[循环结束] M --> N[保存最终模型] N --> O[训练结束] ``` ### 预测代码流程图 ```mermaid %% 预测代码流程图 graph TD P[开始] --> Q[加载训练好的模型] Q --> R[读取测试数据] R --> S[遍历测试用例] S --> T[文本清洗] T --> U[文本编码] U --> V[模型预测] V --> W[计算概率] W --> X[结果比对] X --> Y{是否正确?} Y -->|是| Z1[正确计数] Y -->|否| Z2[错误计数] Z1 --> AA[统计准确率] Z2 --> AA AA --> AB[输出测试报告] AB --> AC[结束] ``` ### 整体代码功能分析 1. **架构设计**: - 采用模块化设计,分离数据预处理(DataProcessor)、模型训练(BertSentimentClassifier)、对抗训练(FGM)和预测(WeiboClassifier) - 支持完整的MLOps流程:实验跟踪(MLflow)、模型版本管理、评估指标记录 - 使用面向对象设计,提高代码复用性和可维护性 2. **核心技术**: - 基于BERT的文本分类(bert-base-chinese) - 对抗训练(FGM)提升模型鲁棒性 - 动态学习率调度(Linear Schedule with Warmup) - 混合精度训练(通过PyTorch自动管理) 3. **数据处理**: - 多层清洗逻辑(去除@用户、#话题#、URL等) - 严格一致的清洗流程(训练/预测使用相同方法) - 数据平衡处理(stratified split) - 空文本过滤和异常处理 4. **性能优化**: - DataLoader配置pin_memory和num_workers加速数据加载 - 梯度裁剪(grad_clip=1.0)防止梯度爆炸 - 权重衰减(weight_decay=0.01)避免过拟合 - 对抗训练提升模型泛化能力 5. **可观测性**: - MLflow跟踪训练损失、验证F1值 - 详细分类报告输出(precision/recall/f1) - 预测时显示清洗后文本和概率分布 - 完整的测试用例验证流程 6. **扩展性设计**: - 支持灵活调整的模型超参数(max_len/batch_size/epsilon等) - 方便替换其他预训练模型 - 可扩展的评估指标体系 - 支持概率输出和阈值调整 7. **异常处理**: - 预测时的类型检查(非字符串处理) - 空文本返回默认分类 - 梯度计算前的参数检查 - CUDA内存异常预防 8. **实际应用价值**: - 可直接部署为微博评论情感分析API - 支持实时情感监控系统 - 可用于用户反馈自动分类 - 为舆情分析提供基础能力 ### 改进想法(待完善) 1. **数据层面**: - 增加数据增强(同义词替换、随机mask) - 实现动态padding提升效率 - 添加类别权重处理不平衡数据 2. **模型层面**: - 尝试其他预训练模型(RoBERTa、ERNIE) - 加入Layer-wise Learning Rate Decay - 实现早停机制(Early Stopping) 3. **工程层面**: - 添加ONNX格式导出 - 实现模型热更新 - 增加Prometheus监控集成 - 添加API服务化封装 4. **可解释性**: - 增加Attention可视化 - 实现LIME解释器 - 添加错误样本分析模块