Skip to content

本地最小训练脚本

工程版项目有很多必要但会干扰学习的东西:

  • 配置文件
  • 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()不需要不需要

建议练习方式

不要第一遍就从空白文件开始。

更好的顺序是:

  1. 先照着骨架补全函数。
  2. 能跑通后,删掉代码重写一遍。
  3. 第二遍只看函数名,不看实现。
  4. 第三遍只看大纲,从空白文件写。

真正的目标不是背代码,而是形成顺序感:

text
读数据 -> 建标签映射 -> tokenizer -> Dataset -> DataLoader -> model -> optimizer -> train -> evaluate -> save