训练/推理 优化


因为大模型只是比较杂比较乱

我又不太擅长零散的记忆

所以试图把训练和推理的一些知识以一种顺序让自己能串联起来

训练

首先训练模型的步骤我们走一个完整流程

文本->张量->loss->梯度->参数更新->checkpoint->可部署模型

一开始我们有文本/样本,此时他们:
如果要做预训练,为一段段文本doc
如果做SFT,为prompt + answer
这时候他们都还是字符串

文本在开始训练前,需要进行清洗和格式化:
去重、统一编码、过滤低质内容之类的
如果是对SFT,则还要加上提示/标签、分隔符等

之后经过tokenizer字符串会变成整数序列token ids。序列中每个整数之后都会查表替换为一个token向量(下下一步)

多条样本会被打包成batch,进行:
padding、attention_mask等操作

之后则是跑前向计算,token ids被embedding为token向量,再被计算为logits(下一个token的打分):
input_ids[B,T]->hidden[B,T,D]
->多层transformer block ->输出层->logits[B,T,V] (V=词表大小)

得到logits了就可以计算loss:logits + labels -> scalar loss
用交叉熵(cross entropy):
对每个位置t,用logits[:,t,:]预测label[:,t]
mask掉padding,mask prompt(SFT)
平均得到一个标量loss

有了loss之后开始反向传播,毕竟loss就是反应模型训练程度的参数
loss.backward()后:每个可训练参数都有梯度grad
输出:参数梯度

接着用参数梯度更新参数:optimizer step
常见AdamW优化器做:动量/二阶矩更新、weight decay、更新参数
训练技巧也往往在这里出现:如梯度裁剪、学习率调度、梯度累积等

这样就训练了一步了,之后循环执行取 batch → 前向 → loss → 反向 → 更新
并监控:loss、困惑度perplexity、稳定性等

保存checkpoint(可恢复训练的模型状态),完整的checkpoint会包含:
模型权重、optimizer states、学习率调度器状态、当前step/随机种子
可用于断点续训、回滚、对比试验

导出可用模型:训练态->部署态
训练结束后,你通常做这些事情让它“可上线”:
只保留模型权重(不带 optimizer states)
如果是LoRA/Adapter:选择
合并(merge)到权重,或保留为可切换模块
(可选)量化(INT8/INT4)
(可选)转换推理引擎格式(TensorRT-LLM / vLLM 等)
配置 tokenizer、特殊 token、对话模板
输出:一个可推理的模型包(权重 + tokenizer + 配置)

最后推理沿着与上线

好,以上就是一个示例的完整训练流程了

接下来让我们看看有哪些训练优化技术可以用在里面

首先最开始我们可以选择自己要用什么精度训练,可以用fp32也可以用fp16/fp32混合精度

前向时候的优化技巧:
gradient checkpointing:前向仍然要算所有层的激活,但只保存少数“检查点”激活(比如每 N 层的输入/输出),中间层激活丢掉
算子融合:把多个本来要分开执行的小算子,合并成一个(或更少的)GPU kernel。还会用于在Attn/LayerNorm/RmsNorm/MLP/FFN/一串小算子

并行技巧:
DP:扩Batch,数据并行。同步梯度
TP:把一个大矩阵乘拆到多卡(模型太大装不下的时候用)。同步激活/中间结果(更频繁的通信)
PP:把层切成好几段,像流水线一样跑
ZeRO:把optimizer state、梯度、参数切分到多卡,省显存、训大模型,需要时聚合

优化器与训练技巧:
AdamW、学习率调度、梯度裁剪、数据质量/配比

参数高效微调:
Lora
Adapters:层间插小模块,只训小模块。
Prefix/Prompt tuning
QLoRA:权重 4bit 量化存着,训练 LoRA

训练章节我们按步骤梳理一下,防止混乱

文本->混合精度/数据配比
->token->batch维度DP
->前向->gradient checkpoint/算子融合/TP/PP/ZeRO
->反向->算子融合/AdamW/学习率调度/梯度裁剪/ZeRO
->ckpt->
->模型 Lora/Adapters/QLora/Adapters/Prompt tuning

训练部分依然还有一些相对零散的方法

结构层面:
参数共享(如 Transformer 中的某些变体)
低秩分解(LoRA、Adapters)
剪枝(Pruning)
知识蒸馏(Teacher → Student)

大概结论

不收敛 → 看学习率和优化器
过拟合 → 看正则和数据
太慢 → 看工程和精度
泛化差 → 看数据分布和训练策略

推理

模型:量化 / 蒸馏
输入和预处理

prefill阶段:FlashAttention / prefix cache / 编译

decode阶段:KV cache 管理(paged/sliding/quantize)/ 连续批处理 / speculative decoding

输出与停止条件

服务:动态批处理、多卡并行、内存管理

经验结论

OOM(显存爆)→ 先看 KV cache(上下文长度、并发、层数、dtype),再看权重量化

延迟高 → 看 decode 是否被带宽卡住(attention/KV 读),以及是否能用 speculative / 编译 / fused kernels

吞吐低 → 看 batching(dynamic/continuous)、KV cache 管理(paged)、并发策略

长上下文慢 → 基本就是 KV cache + attention 的代价,考虑 sliding window / chunking / 检索增强减少上下文


文章作者: Austin
版权声明: 本博客所有文章除特別声明外,均采用 CC BY 4.0 许可协议。转载请注明来源 Austin !
评论
 上一篇
2026-01-27 Austin
下一篇 
plan-now plan-now
2026-01-27 Austin
  目录