LLaMA: Open and Efficient Foundation Language Models

摘要:

我们提出了 LLaMA,一组从 7B 到 65B 参数的基础语言模型。我们在数万亿 token 上训练我们的模型,并表明可以仅使用公开可用的数据集来训练最先进的模型,而无需依赖专有且无法获取的数据集。特别地,LLaMA-13B 在大多数基准上优于 GPT-3(175B),而 LLaMA-65B 与最佳模型 Chinchilla-70B 和 PaLM-540B 相比具有竞争力。我们向研究社区发布了所有模型。

1. Introduction

核心卖点:

  • 范式转变:从“训练最优”到“推理最优”
    • 不同于 Hoffmann 等 (2022) 的 Chinchilla 缩放定律(给定训练预算,找最优模型/数据配比),LLaMA 关注的是给定推理预算。
    • 结论:虽然训练一个大模型达到某性能可能更便宜,但训练更久的小模型在推理时成本更低。例如,7B 模型在 1T token 后性能仍在提升。
  • 更小规模,更强性能
    • 规模范围:7B ~ 65B。
    • LLaMA-13B 在大多数基准上超越 GPT-3 (175B),规模小 10 倍。
    • LLaMA-65B 与 Chinchilla-70B、PaLM-540B 性能相当。
  • 推理友好
    • LLaMA-13B 可在单张 V100 GPU 上推理,大幅降低部署成本,有助于民主化 LLM 研究。
  • 完全开源
    • 仅使用公开数据集训练,模型权重全部开源,与现有闭源模型(GPT-3、Chinchilla、PaLM)形成对比。
    • 与 OPT、GPT-NeoX、BLOOM 等开源模型相比,性能具有明显竞争力。

2. Approach

训练方法整体类似于 GPT-3。

2.1 Pre-training Data

多个开源数据集的混合:


共约 1.4T tokens,除了 Wikipedia 与 Books 数据集 tokens 被训练约 2 个 epoch,其他数据集都只训练一个 epoch。

2.2 Architecture

基于 decoder-only Transformer 架构,部分与原始 Transformer 不一样地方和来源如下:

  • Pre-normalization [GPT3],RMSNorm。
  • SwiGLU activation function [PaLM],\(d_{ffn}=\frac{2}{3}4d\)。
  • Rotary Embeddings [GPTNeo]。

2.3 Optimizer

使用 AdamW 优化器,超参数:\(\beta_{1}=0.9, \beta_{2}=0.95\) 。余弦学习率衰减。权重衰减为 0.1,梯度裁剪为 1.0。

2.4 Efficient implementation

  • 高效因果多头注意力:使用 xformers 实现,反向传播来自 FlashAttention。不存储 attention weights,不计算被因果掩码屏蔽的 key/query scores,降低显存和运行时间。
  • Checkpointing:手动实现 Transformer 层反向传播,保存计算昂贵的激活(如 linear 层输出),减少反向重计算量。
  • 并行策略:模型并行 + 序列并行,降低显存占用。
  • 通信重叠:尽可能重叠激活计算与 GPU 间 all_reduce 通信。
  • 训练吞吐:65B 模型,2048×A100-80GB,约 380 tokens/sec/GPU。
  • 训练时间:1.4T tokens 约 21 天。

实验结果等内容略。

3. 代码实现

官方代码链接:link

1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
22
23
24
25
26
27
28
29
30
31
32
33
34
35
36
37
38
39
40
41
42
43
44
45
46
47
48
49
50
51
52
53
54
55
56
57
58
59
60
61
62
63
64
65
66
67
68
69
70
71
72
73
74
75
76
77
78
79
80
81
82
83
84
85
86
87
88
89
90
91
92
93
94
95
96
97
98
99
100
101
102
103
104
105
106
107
108
109
110
111
112
113
114
115
116
117
118
119
120
121
122
123
124
125
126
127
128
129
130
131
132
133
134
135
136
137
138
139
140
141
142
143
144
145
146
147
148
149
150
151
152
153
154
155
156
157
158
159
160
161
162
163
164
165
166
167
168
169
170
171
172
173
174
175
176
177
178
179
180
181
182
183
184
185
186
187
188
189
190
191
192
193
194
195
196
197
198
199
200
201
202
203
204
205
206
207
208
209
210
211
212
213
214
215
216
217
218
219
220
221
222
223
224
225
226
227
228
229
230
231
232
233
234
235
236
237
238
239
240
241
242
243
244
245
246
247
248
249
250
251
252
253
254
255
256
257
258
259
260
261
# ---------- 模型配置 ----------
@dataclass
class ModelArgs:
dim: int = 4096 # d_model
n_layers: int = 32 # Transformer block 数量
n_heads: int = 32 # query 头数
n_kv_heads: Optional[int] = None # KV 头数,用于 GQA;None 时等于 n_heads
vocab_size: int = -1 # 由 tokenizer 决定
multiple_of: int = 256 # SwiGLU hidden dim 向上取整的倍数
ffn_dim_multiplier: Optional[float] = None # FFN hidden dim 缩放因子
norm_eps: float = 1e-5 # RMSNorm eps

max_batch_size: int = 32
max_seq_len: int = 2048


# ---------- RMSNorm ----------
class RMSNorm(nn.Module):
"""
相比 LayerNorm,去掉 mean-centering 和 bias,只保留缩放。
RMSNorm(x) = x / sqrt(mean(x^2) + eps) * weight
"""
def __init__(self, dim: int, eps: float = 1e-6):
super().__init__()
self.eps = eps
self.weight = nn.Parameter(torch.ones(dim))

def _norm(self, x):
# 在 fp32 下计算,数值稳定
return x * torch.rsqrt(x.pow(2).mean(-1, keepdim=True) + self.eps)

def forward(self, x):
return self._norm(x.float()).type_as(x) * self.weight


# ---------- Rotary Position Embedding (RoPE) ----------
def precompute_freqs_cis(dim: int, end: int, theta: float = 10000.0):
"""
预计算 RoPE 的复数频率 e^(i * t * theta_k)。
dim: head_dim
end: 最大位置数
返回: [end, dim//2] 的 complex64 张量
"""
freqs = 1.0 / (theta ** (torch.arange(0, dim, 2)[: (dim // 2)].float() / dim))
t = torch.arange(end, device=freqs.device)
freqs = torch.outer(t, freqs).float() # [end, dim//2]
freqs_cis = torch.polar(torch.ones_like(freqs), freqs) # e^(i*angle)
return freqs_cis


def reshape_for_broadcast(freqs_cis: torch.Tensor, x: torch.Tensor):
"""将 freqs_cis 形状调整为可与 x 广播。"""
ndim = x.ndim
assert 0 <= 1 < ndim
assert freqs_cis.shape == (x.shape[1], x.shape[-1])
shape = [d if i == 1 or i == ndim - 1 else 1 for i, d in enumerate(x.shape)]
return freqs_cis.view(*shape)


def apply_rotary_emb(xq, xk, freqs_cis):
"""
对 query 和 key 应用旋转位置编码。
xq, xk: [B, T, n_heads, head_dim]
freqs_cis: [T, head_dim//2] complex
做法: 把 head_dim 视为复数对 (2i, 2i+1),乘以 e^(i*angle) 实现旋转。
"""
# head_dim -> complex: [.., head_dim//2]
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 = reshape_for_broadcast(freqs_cis, xq_)
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 head 复制 ----------
def repeat_kv(x: torch.Tensor, n_rep: int) -> torch.Tensor:
"""
当 n_kv_heads < n_heads 时,把 KV heads 复制 n_rep 次以匹配 Q heads。
x: [B, T, n_kv_heads, head_dim] -> [B, T, n_kv_heads * n_rep, head_dim]
"""
bs, slen, n_kv_heads, head_dim = x.shape
if n_rep == 1:
return x
return (
x[:, :, :, None, :]
.expand(bs, slen, n_kv_heads, n_rep, head_dim)
.reshape(bs, slen, n_kv_heads * n_rep, head_dim)
)


# ---------- 多头注意力 (含 GQA + KV Cache + RoPE) ----------
class Attention(nn.Module):
def __init__(self, args: ModelArgs):
super().__init__()
# GQA: KV 头数可小于 Q 头数
self.n_kv_heads = args.n_heads if args.n_kv_heads is None else args.n_kv_heads
self.n_heads = args.n_heads
self.n_rep = self.n_heads // self.n_kv_heads # KV 复制次数
self.head_dim = args.dim // args.n_heads

# 无 bias,与 LLaMA 论文一致
self.wq = nn.Linear(args.dim, args.n_heads * self.head_dim, bias=False)
self.wk = nn.Linear(args.dim, self.n_kv_heads * self.head_dim, bias=False)
self.wv = nn.Linear(args.dim, self.n_kv_heads * self.head_dim, bias=False)
self.wo = nn.Linear(args.n_heads * self.head_dim, args.dim, bias=False)

# KV Cache 预分配,避免推理时反复申请显存
self.cache_k = torch.zeros(
(args.max_batch_size, args.max_seq_len, self.n_kv_heads, self.head_dim)
)
self.cache_v = torch.zeros(
(args.max_batch_size, args.max_seq_len, self.n_kv_heads, self.head_dim)
)

def forward(self, x, start_pos, freqs_cis, mask):
bsz, seqlen, _ = x.shape
xq, xk, xv = self.wq(x), self.wk(x), self.wv(x)

# [B, T, n_heads, head_dim]
xq = xq.view(bsz, seqlen, self.n_heads, self.head_dim)
xk = xk.view(bsz, seqlen, self.n_kv_heads, self.head_dim)
xv = xv.view(bsz, seqlen, self.n_kv_heads, self.head_dim)

# RoPE 只作用于 Q 和 K
xq, xk = apply_rotary_emb(xq, xk, freqs_cis=freqs_cis)

# 写入 KV Cache
self.cache_k = self.cache_k.to(xq)
self.cache_v = self.cache_v.to(xq)
self.cache_k[:bsz, start_pos:start_pos + seqlen] = xk
self.cache_v[:bsz, start_pos:start_pos + seqlen] = xv

# 读取历史 + 当前
keys = self.cache_k[:bsz, :start_pos + seqlen]
values = self.cache_v[:bsz, :start_pos + seqlen]

# GQA: KV heads 复制到与 Q heads 数量一致
keys = repeat_kv(keys, self.n_rep)
values = repeat_kv(values, self.n_rep)

# 转成 [B, n_heads, T, head_dim]
xq = xq.transpose(1, 2)
keys = keys.transpose(1, 2)
values = values.transpose(1, 2)

# 缩放点积注意力
scores = torch.matmul(xq, keys.transpose(2, 3)) / math.sqrt(self.head_dim)
if mask is not None:
scores = scores + mask
scores = F.softmax(scores.float(), dim=-1).type_as(xq)

output = torch.matmul(scores, values) # [B, n_heads, T, head_dim]
output = output.transpose(1, 2).contiguous().view(bsz, seqlen, -1)
return self.wo(output)


# ---------- SwiGLU FFN ----------
class FeedForward(nn.Module):
"""
SwiGLU: FFN(x) = W2( SiLU(W1 x) * W3 x )
hidden_dim 通常取 4*dim 的 2/3,且向上取整到 multiple_of 的倍数。
"""
def __init__(self, dim, hidden_dim, multiple_of, ffn_dim_multiplier=None):
super().__init__()
hidden_dim = int(2 * hidden_dim / 3)
if ffn_dim_multiplier is not None:
hidden_dim = int(ffn_dim_multiplier * hidden_dim)
# 向上取整到 multiple_of 的倍数
hidden_dim = multiple_of * ((hidden_dim + multiple_of - 1) // multiple_of)

self.w1 = nn.Linear(dim, hidden_dim, bias=False)
self.w2 = nn.Linear(hidden_dim, dim, bias=False)
self.w3 = nn.Linear(dim, hidden_dim, bias=False)

def forward(self, x):
# 门控:SiLU(w1(x)) * w3(x),再投影回 dim
return self.w2(F.silu(self.w1(x)) * self.w3(x))


# ---------- Transformer Block ----------
class TransformerBlock(nn.Module):
"""
Pre-Norm 结构:
h = x + Attention(RMSNorm(x))
out = h + FFN(RMSNorm(h))
"""
def __init__(self, layer_id, args):
super().__init__()
self.n_heads = args.n_heads
self.dim = args.dim
self.head_dim = args.dim // args.n_heads
self.attention = Attention(args)
self.feed_forward = FeedForward(
dim=args.dim,
hidden_dim=4 * args.dim,
multiple_of=args.multiple_of,
ffn_dim_multiplier=args.ffn_dim_multiplier,
)
self.layer_id = layer_id
self.attention_norm = RMSNorm(args.dim, eps=args.norm_eps)
self.ffn_norm = RMSNorm(args.dim, eps=args.norm_eps)

def forward(self, x, start_pos, freqs_cis, mask):
# 残差 + Pre-Norm Attention
h = x + self.attention(self.attention_norm(x), start_pos, freqs_cis, mask)
# 残差 + Pre-Norm FFN
out = h + self.feed_forward(self.ffn_norm(h))
return out


# ---------- 完整 Transformer ----------
class Transformer(nn.Module):
def __init__(self, params: ModelArgs):
super().__init__()
self.params = params
self.vocab_size = params.vocab_size
self.n_layers = params.n_layers

self.tok_embeddings = nn.Embedding(params.vocab_size, params.dim)

self.layers = nn.ModuleList()
for layer_id in range(params.n_layers):
self.layers.append(TransformerBlock(layer_id, params))

self.norm = RMSNorm(params.dim, eps=params.norm_eps)
self.output = nn.Linear(params.dim, params.vocab_size, bias=False)

# 预计算 RoPE 频率,*2 是为了支持更长序列的微调
self.freqs_cis = precompute_freqs_cis(
self.params.dim // self.params.n_heads,
self.params.max_seq_len * 2,
)

@torch.inference_mode()
def forward(self, tokens, start_pos):
"""
tokens: [B, T]
start_pos: 当前 token 在序列中的起始位置(用于 KV Cache 和 RoPE)
"""
_bsz, seqlen = tokens.shape
h = self.tok_embeddings(tokens)
self.freqs_cis = self.freqs_cis.to(h.device)
freqs_cis = self.freqs_cis[start_pos:start_pos + seqlen]

# 因果掩码:仅当 seqlen > 1 时需要
mask = None
if seqlen > 1:
mask = torch.full((seqlen, seqlen), float("-inf"), device=tokens.device)
mask = torch.triu(mask, diagonal=1)
# 拼接 cache 部分:左边 start_pos 列全 0(可看见历史)
mask = torch.hstack([
torch.zeros((seqlen, start_pos), device=tokens.device),
mask
]).type_as(h)

for layer in self.layers:
h = layer(h, start_pos, freqs_cis, mask)
h = self.norm(h)
output = self.output(h).float()
return output

补充:MHA/GQA/MQA

特性 MHA (Multi-Head Attention) GQA (Grouped Query Attention) MQA (Multi-Query Attention)
全称 Multi-Head Attention Grouped Query Attention Multi-Query Attention
核心思想 Q、K、V 头数相同 Q 头多,K/V 头少,多个 Q 头共享一组 K/V Q 头多,K/V 头只有 1 组
Q 头数 \(n_{heads}\) \(n_{heads}\) \(n_{heads}\)
KV 头数 \(n_{heads}\) \(n_{kv\_heads}\)(\(1 < n_{kv\_heads} < n_{heads}\)) \(1\)
KV 共享比例 \(1:1\) \(n_{rep} = n_{heads} / n_{kv\_heads}\) \(n_{rep} = n_{heads}\)
KV Cache 大小 最大 中等,减少 \(n_{rep}\) 倍 最小,减少 \(n_{heads}\) 倍
推理速度 慢 较快 最快
显存/带宽开销 高 中 低
质量 最好 接近 MHA 可能下降
训练稳定性 好 好 可能需调参
代表模型 GPT-2、GPT-3、LLaMA 7B/13B LLaMA 2 34B/70B、LLaMA 3、Mistral PaLM、Falcon
适用场景 训练优先、小规模推理 大模型推理、显存/带宽敏感 极致推理速度、显存受限
备注 KV Cache 瓶颈明显 与 PagedAttention 兼容,vLLM 优化 KV Cache 极小,但质量需验证

说明:\(n_{heads}\) 为 Query 头数,\(n_{kv\_heads}\) 为 Key/Value 头数。GQA 当 \(n_{kv\_heads}=n_{heads}\) 时退化为 MHA,当 \(n_{kv\_heads}=1\) 时退化为 MQA。


LLaMA: Open and Efficient Foundation Language Models
https://arcsin2.cloud/posts/2026/10/1425595528/
作者
arcsin2
发布于
2026年10月5日
许可协议