CoolFace
Apppublic

nchdlhbctm/TraceDetect-AI

sourceHugging Faceapache-2.0updated 4mo agoView on Hugging Face
0likes
train_text_model.py96 linesDownload Raw Back to root
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()