nchdlhbctm/TraceDetect-AI
0
1import pandas as pd
2import torch
3from torch.utils.data import Dataset, DataLoader
4from transformers import AutoTokenizer, AutoModelForSequenceClassification
5from torch.optim import AdamW
6from tqdm import tqdm
7
8
9# 1. 定义专属的 PyTorch 文本数据集
10class AITextDataset(Dataset):
11 def __init__(self, csv_file, tokenizer, max_len=128):
12 self.data = pd.read_csv(csv_file)
13 self.tokenizer = tokenizer
14 self.max_len = max_len
15
16 def __len__(self):
17 return len(self.data)
18
19 def __getitem__(self, index):
20 text = str(self.data.iloc[index, 0])
21 label = int(self.data.iloc[index, 1])
22
23 # 将汉字切成 token 序列
24 encoding = self.tokenizer(
25 text,
26 add_special_tokens=True,
27 max_length=self.max_len,
28 padding='max_length',
29 truncation=True,
30 return_attention_mask=True,
31 return_tensors='pt',
32 )
33 return {
34 'input_ids': encoding['input_ids'].flatten(),
35 'attention_mask': encoding['attention_mask'].flatten(),
36 'labels': torch.tensor(label, dtype=torch.long)
37 }
38
39
40def train_text():
41 device = torch.device("cuda" if torch.cuda.is_available() else "cpu")
42 print(f"💻 当前计算设备: {device}")
43 if device.type == 'cpu':
44 print("⚠️ 警告:当前使用 CPU 炼丹。NLP模型参数量巨大,这可能需要一些时间,请耐心等待...")
45
46 print("正在加载预训练的中文 BERT 分词器与模型权重...")
47
48 # 【核心修复】:换成了官方真实存在、最经典的 bert-base-chinese
49 model_name = "bert-base-chinese"
50 tokenizer = AutoTokenizer.from_pretrained(model_name)
51 model = AutoModelForSequenceClassification.from_pretrained(model_name, num_labels=2)
52 model = model.to(device)
53
54 print("正在封装数据...")
55 dataset = AITextDataset('./data/text_dataset.csv', tokenizer, max_len=128)
56 # 批量大小设为 8,防止 CPU 内存吃紧
57 dataloader = DataLoader(dataset, batch_size=8, shuffle=True)
58
59 optimizer = AdamW(model.parameters(), lr=2e-5)
60
61 epochs = 1
62
63 print("\n🚀 --- 开始文本模型微调 ---")
64 model.train()
65
66 for epoch in range(epochs):
67 progress_bar = tqdm(dataloader, desc=f"第 {epoch + 1}/{epochs} 轮", leave=True, colour='blue')
68 running_loss = 0.0
69
70 for batch in progress_bar:
71 optimizer.zero_grad()
72
73 input_ids = batch['input_ids'].to(device)
74 attention_mask = batch['attention_mask'].to(device)
75 labels = batch['labels'].to(device)
76
77 outputs = model(input_ids, attention_mask=attention_mask, labels=labels)
78 loss = outputs.loss
79
80 loss.backward()
81 optimizer.step()
82
83 running_loss += loss.item()
84 progress_bar.set_postfix({'loss': f"{loss.item():.4f}"})
85
86 print(f"✅ 第 {epoch + 1} 轮完成 | 平均 Loss: {running_loss / len(dataloader):.4f}")
87
88 # 保存咱们微调后的专属大模型权重
89 save_dir = "./finetuned_text_model"
90 model.save_pretrained(save_dir)
91 tokenizer.save_pretrained(save_dir)
92 print(f"\n🎉 炼丹成功!专属的文本鉴别模型已保存在: {save_dir} 文件夹中。")
93
94
95if __name__ == "__main__":
96 train_text()