本地最小训练脚本
工程版项目有很多必要但会干扰学习的东西:
- 配置文件
- checkpoint
- early stop
- TensorBoard
- AMP 混合精度
- Trainer 封装
学习阶段可以先写一个单文件最小脚本,把训练主线完整跑通。
最小脚本应该包含哪些函数
建议保留这些函数:
python
get_device()
read_tsv()
ProductTitleDataset
evaluate()
main()分工如下:
| 函数 | 负责什么 |
|---|---|
get_device() | 选择 mps / cuda / cpu |
read_tsv() | 读取原始 train.txt / valid.txt |
ProductTitleDataset | 一条样本转成 BERT 输入 |
evaluate() | 在验证集上统计 loss 和 accuracy |
main() | 串起完整训练流程 |
ProductTitleDataset
这个类最关键。
它把一条原始数据:
python
{"label": "服装", "text": "男士运动鞋夏季透气"}转换成:
python
{
"input_ids": ...,
"attention_mask": ...,
"labels": 1,
}核心代码:
python
class ProductTitleDataset(Dataset):
def __init__(self, rows, tokenizer, label2id):
self.rows = rows
self.tokenizer = tokenizer
self.label2id = label2id
def __len__(self):
return len(self.rows)
def __getitem__(self, index):
row = self.rows[index]
encoded = self.tokenizer(row["text"], truncation=True, max_length=64)
encoded["labels"] = self.label2id[row["label"]]
return encoded训练循环
你最终必须能手写这一段:
python
for epoch in range(EPOCHS):
model.train()
for batch in train_loader:
batch = {key: value.to(device) for key, value in batch.items()}
outputs = model(**batch)
loss = outputs.loss
loss.backward()
optimizer.step()
optimizer.zero_grad()这就是核心能力。
验证循环
验证不更新参数:
python
model.eval()
with torch.no_grad():
for batch in valid_loader:
batch = {key: value.to(device) for key, value in batch.items()}
outputs = model(**batch)
preds = outputs.logits.argmax(dim=-1)训练和验证的差异:
| 阶段 | 模式 | 计算图 | 参数更新 |
|---|---|---|---|
| 训练 | model.train() | 需要 | 需要 |
| 验证 | model.eval() | 不需要 | 不需要 |
建议练习方式
不要第一遍就从空白文件开始。
更好的顺序是:
- 先照着骨架补全函数。
- 能跑通后,删掉代码重写一遍。
- 第二遍只看函数名,不看实现。
- 第三遍只看大纲,从空白文件写。
真正的目标不是背代码,而是形成顺序感:
text
读数据 -> 建标签映射 -> tokenizer -> Dataset -> DataLoader -> model -> optimizer -> train -> evaluate -> save