本文不堆砌概念,而是用可运行的代码和数学推导,拆解大语言模型训练全链路中的关键算法节点,涵盖预训练、指令微调、RLHF、模型压缩与推理优化,适合有一定深度学习基础的工程师快速上手。
大语言模型早已不是“调包调参”的玩具。一个合格的LLM算法工程师,需要深刻理解:
本文将以 LLaMA-like 架构为基础,从零实现核心模块,并给出训练/推理的可执行代码片段。
当前主流LLM(LLaMA 3、Qwen)均采用Grouped Query Attention (GQA) 和Rotary Position Embedding (RoPE)。我们直接实现这两部分。
RoPE通过旋转矩阵将位置信息内积到Q和K中,公式为:
f(q,m)=q⋅eimθ,f(k,n)=k⋅einθf(q,m)=q⋅eimθ,f(k,n)=k⋅einθ
实际采用复数形式实现:
import torch
import torch.nn as nn
import math
def precompute_freqs_cis(dim: int, seq_len: int, theta: float = 10000.0):
# 计算频率
freqs = 1.0 / (theta ** (torch.arange(0, dim, 2)[: (dim // 2)].float() / dim))
t = torch.arange(seq_len, dtype=torch.float)
freqs = torch.outer(t, freqs) # (seq_len, dim/2)
freqs_cis = torch.polar(torch.ones_like(freqs), freqs) # 复数表示
return freqs_cis
def apply_rotary_emb(xq, xk, freqs_cis):
# xq, xk: (batch, seq_len, n_heads, head_dim)
# 将最后两维转为复数
xq_ = torch.view_as_complex(xq.float().reshape(*xq.shape[:-1], -1, 2))
xk_ = torch.view_as_complex(xk.float().reshape(*xk.shape[:-1], -1, 2))
freqs_cis = freqs_cis.unsqueeze(0).unsqueeze(2) # (1, seq_len, 1, dim/2)
xq_out = torch.view_as_real(xq_ * freqs_cis).flatten(3)
xk_out = torch.view_as_real(xk_ * freqs_cis).flatten(3)
return xq_out.type_as(xq), xk_out.type_as(xk)GQA在KV头上分组,减少KV cache内存。假设 n_kv_heads = n_heads // 4:
class GroupedQueryAttention(nn.Module):
def __init__(self, dim, n_heads, n_kv_heads, head_dim=128):
super().__init__()
self.n_heads = n_heads
self.n_kv_heads = n_kv_heads
self.head_dim = head_dim
self.wq = nn.Linear(dim, n_heads * head_dim, bias=False)
self.wk = nn.Linear(dim, n_kv_heads * head_dim, bias=False)
self.wv = nn.Linear(dim, n_kv_heads * head_dim, bias=False)
self.wo = nn.Linear(n_heads * head_dim, dim, bias=False)
def forward(self, x, freqs_cis, mask=None):
B, T, _ = x.shape
q = self.wq(x).view(B, T, self.n_heads, self.head_dim)
k = self.wk(x).view(B, T, self.n_kv_heads, self.head_dim)
v = self.wv(x).view(B, T, self.n_kv_heads, self.head_dim)
# 应用RoPE(仅对q和k)
q, k = apply_rotary_emb(q, k, freqs_cis[:T])
# 将kv重复到与q相同的头数(分组扩展)
k = k.repeat_interleave(self.n_heads // self.n_kv_heads, dim=2)
v = v.repeat_interleave(self.n_heads // self.n_kv_heads, dim=2)
# 标准scaled dot-product attention
scores = torch.matmul(q, k.transpose(2, 3)) / math.sqrt(self.head_dim)
if mask is not None:
scores = scores + mask # causal mask
attn = torch.softmax(scores, dim=-1)
out = torch.matmul(attn, v)
out = out.transpose(1, 2).contiguous().view(B, T, -1)
return self.wo(out)技术要点:GQA在推理时可显著降低显存占用(KV cache减少75%),而性能损失小于1%。
预训练的核心是因果语言建模,Loss为交叉熵。以下展示使用 torch.distributed 进行数据并行时的Loss聚合(避免各卡Loss不均衡):
import torch.distributed as dist
def pretrain_loss(logits, labels, ignore_index=-100):
# logits: (B, T, vocab_size), labels: (B, T)
shift_logits = logits[..., :-1, :].contiguous()
shift_labels = labels[..., 1:].contiguous()
loss_fct = nn.CrossEntropyLoss(ignore_index=ignore_index, reduction='none')
per_token_loss = loss_fct(shift_logits.view(-1, shift_logits.size(-1)),
shift_labels.view(-1))
# 全局平均(考虑不同卡上的有效token数)
num_valid = (shift_labels.view(-1) != ignore_index).sum()
total_loss = per_token_loss.sum()
# 分布式同步
if dist.is_initialized():
dist.all_reduce(total_loss, op=dist.ReduceOp.SUM)
dist.all_reduce(num_valid, op=dist.ReduceOp.SUM)
return total_loss / num_valid数据预处理:使用 transformers 的 tokenizer,注意 add_special_tokens 和 max_length 的动态填充(支持packing)。对于超长上下文,可采用 随机窗口截断 或 Flash Attention 以降低显存。
LoRA(Low-Rank Adaptation)是当前最流行的PEFT方法,其核心是冻结原权重,在旁路添加低秩矩阵:
W′=W+ΔW=W+BA,B∈Rd×r,A∈Rr×k,r≪d,kW′=W+ΔW=W+BA,B∈Rd×r,A∈Rr×k,r≪d,k
前向时:h = W x + B A x。仅更新A和B。
使用 peft 库的底层实现:
class LoRALinear(nn.Module):
def __init__(self, in_features, out_features, r=8, alpha=16, dropout=0.1):
super().__init__()
self.linear = nn.Linear(in_features, out_features, bias=False)
# 冻结原权重
self.linear.weight.requires_grad = False
self.lora_A = nn.Parameter(torch.randn(r, in_features) * 0.01)
self.lora_B = nn.Parameter(torch.zeros(out_features, r))
self.scaling = alpha / r
self.dropout = nn.Dropout(dropout)
def forward(self, x):
base_out = self.linear(x) # (..., out_features)
lora_out = (x @ self.lora_A.T) @ self.lora_B.T # (..., r) -> (..., out_features)
return base_out + self.dropout(lora_out) * self.scaling实际微调技巧:
RLHF通常包括三个阶段:SFT、奖励模型训练、PPO微调。此处重点展示PPO的损失函数和优势估计(GAE)。
给定奖励序列 rtrt 和价值预测 V(st)V(st),计算优势:
A^t=∑l=0∞(γλ)lδt+l,δt=rt+γV(st+1)−V(st)A^t=l=0∑∞(γλ)lδt+l,δt=rt+γV(st+1)−V(st)
def compute_gae(rewards, values, gamma=0.99, lam=0.95):
# rewards, values: (T,)
advantages = torch.zeros_like(rewards)
last_gae = 0
for t in reversed(range(len(rewards))):
delta = rewards[t] + gamma * (values[t+1] if t+1 < len(values) else 0) - values[t]
last_gae = delta + gamma * lam * last_gae
advantages[t] = last_gae
returns = advantages + values[:-1] # 实际return
return advantages, returnsPPO的核心是截断概率比,防止更新过大:
LPPO=E[min(rt(θ)A^t,clip(rt(θ),1−ϵ,1+ϵ)A^t)]LPPO=E[min(rt(θ)A^t,clip(rt(θ),1−ϵ,1+ϵ)A^t)]
加上KL惩罚(与参考模型):
def ppo_loss(log_probs, old_log_probs, advantages, ref_log_probs, kl_coef=0.1, clip_eps=0.2):
# log_probs: 当前策略下动作的概率log, shape (T,)
ratio = torch.exp(log_probs - old_log_probs)
surr1 = ratio * advantages
surr2 = torch.clamp(ratio, 1 - clip_eps, 1 + clip_eps) * advantages
policy_loss = -torch.min(surr1, surr2).mean()
# KL penalty (通常使用近似)
kl = (log_probs - ref_log_probs).mean()
return policy_loss + kl_coef * kl工程陷阱:
Ray 或 DeepSpeed 的混合引擎。推理时,KV Cache是关键。我们实现一个带缓存的解码函数:
class KVCache:
def __init__(self, max_batch_size, max_seq_len, n_kv_heads, head_dim, dtype=torch.float16):
self.k = torch.zeros(max_batch_size, max_seq_len, n_kv_heads, head_dim, dtype=dtype)
self.v = torch.zeros(max_batch_size, max_seq_len, n_kv_heads, head_dim, dtype=dtype)
self.seen = 0
def update(self, k, v, input_pos):
# k,v: (B, T, n_kv_heads, head_dim)
self.k[:, input_pos, :, :] = k
self.v[:, input_pos, :, :] = v
self.seen += k.size(1)
return self.k[:, :self.seen], self.v[:, :self.seen]连续批处理(Continuous Batching) 允许在序列生成过程中动态加入新请求,提升吞吐。可使用 vLLM 或 HuggingFace TGI,其核心是分页注意力(PagedAttention)。
对于投机解码(Speculative Decoding),使用一个小模型(draft model)快速生成若干token,再由大模型验证,加速比可达2x。代码示例:
def speculative_decode(draft_model, target_model, prefix, num_tokens):
draft_tokens = draft_model.generate(prefix, num_tokens) # 快速生成
# 用target模型并行验证
logits = target_model(prefix + draft_tokens[:-1])
# 接受满足概率比的token,回退重新生成
...下面给出一个使用 transformers + peft + trl 的完整微调脚本片段(基于Qwen2-7B):
from transformers import AutoModelForCausalLM, AutoTokenizer
from peft import LoraConfig, get_peft_model
from trl import SFTTrainer, DataCollatorForCompletionOnlyLM
model = AutoModelForCausalLM.from_pretrained("Qwen/Qwen2-7B", torch_dtype=torch.bfloat16)
tokenizer = AutoTokenizer.from_pretrained("Qwen/Qwen2-7B")
tokenizer.pad_token = tokenizer.eos_token
lora_config = LoraConfig(
r=16, lora_alpha=32, target_modules=["q_proj", "k_proj", "v_proj", "o_proj"],
lora_dropout=0.05, bias="none", task_type="CAUSAL_LM"
)
model = get_peft_model(model, lora_config)
trainer = SFTTrainer(
model=model,
tokenizer=tokenizer,
train_dataset=dataset,
args=TrainingArguments(
per_device_train_batch_size=4,
gradient_accumulation_steps=4,
learning_rate=2e-4,
max_steps=1000,
bf16=True,
logging_steps=10,
),
data_collator=DataCollatorForCompletionOnlyLM(
response_template="<|im_start|>assistant", tokenizer=tokenizer
),
)
trainer.train()注意:
trl的SFTTrainer内部已支持 packing 和序列截断,可大幅简化代码。
算法工程师不能只关注loss,还需评估事实性(Factuality)和一致性。常用指标:
TruthfulQA 或自建对抗数据集。代码示例:使用 evaluate 库计算BERTScore:
import evaluate
bertscore = evaluate.load("bertscore")
results = bertscore.compute(predictions=generated, references=target, lang="en")
print(results["f1"])本文从算法工程师的视角,覆盖了LLM全生命周期的关键代码模块。但真实生产环境还需要考虑:
最后一句忠告:大模型算法不仅是“调参”,更是对算力、数据和系统设计的综合权衡。建议读者在本地用7B级别模型跑通上述所有代码,再逐步迁移到百亿/千亿规模。
原创声明:本文系作者授权腾讯云开发者社区发表,未经许可,不得转载。
如有侵权,请联系 cloudcommunity@tencent.com 删除。