训练循环
会员专享拆开 LightningCLI、PretrainModule、优化器和调度器
上一节已经把文本样本整理成了 input_ids、labels 和 attention_mask。从这一节开始,这个 batch 会进入真正的训练循环:模型前向、计算 loss、反向传播、优化器更新、学习率调度、日志记录和 checkpoint 保存。
cookllm-bento 的预训练循环可以先看成下面这条链路:
Pretrain Training Loop
How configs become a running Lightning training job.
1
Shell script
compose trainer, model and data configs
fit command
2
LightningCLI
instantiate PretrainModule and PretrainDataModule
objects
3
DataLoader batch
input_ids, labels and attention_mask
batch
4
training_step
forward BentoLM and return language modeling loss
loss
5
Optimizer step
AdamW update after gradient accumulation
weights
6
Callbacks
log metrics, validate, sample text and save checkpoints
logs
训练入口
预训练入口文件很薄:
def main():
LightningCLI(PretrainModule, PretrainDataModule, save_config_callback=None)它主要做三件事:
- 把项目根目录加入
sys.path,让src包可以被正常导入。 - 设置
torch.set_float32_matmul_precision("medium"),让 Ampere 及之后的 GPU 可以使用 TF32 加速部分矩阵计算。 - 用 LightningCLI 把
PretrainModule、PretrainDataModule和 LightningTrainer组装成一次训练任务。
这里没有手写复杂的 argparse。训练参数主要来自 YAML 配置和命令行覆盖,这也是后面做不同实验时最重要的组织方式。
登录以继续阅读
这是一篇付费内容,请登录您的账户以访问完整内容。
CookLLM文档