版本与阅读方式
- 仓库:
xai-org/grok-1 - 分支:
main - Commit:
7050ed204b8206bb8645c7b7bbef7252f79561b0 - Commit 日期:2024-03-19
- 本站核验:2026-07-26
先从 run.py 读取不可变规格,再沿 runners.py 看服务循环,最后进入 model.py 的张量计算。若一开始从 1,400 行 model.py 顶部顺读,很容易迷失在分片规则里。
文件地图
| 文件 | 关键符号 | 责任 |
|---|---|---|
run.py | main() | 固化模型配置、mesh 与示例 prompt |
model.py | LanguageModel, Transformer, DecoderLayer | 完整前向数据流 |
model.py | MultiHeadAttention, KVMemory | GQA、RoPE、mask 与缓存 |
model.py | Router, MoELayer, DenseBlock | Top-2 专家路由与 FFN |
runners.py | ModelRunner, InferenceRunner | JAX 编译、prefill、decode、batch slot |
runners.py | top_p_filter, sample_token | logits 到新 token |
checkpoint.py | load_tensors, restore | 从分片 checkpoint 恢复参数 |
tokenizer.model | SentencePiece model | 文本 ↔ token IDs |
第一站:run.py 固定配置
main() 构造 LanguageModelConfig 和 TransformerConfig。这里不是“推荐超参数”,而是 checkpoint 形状契约;随意修改层数、hidden size 或专家数会让权重无法对应。
# 真实源码节选(Grok-1, commit 7050ed2)
TransformerConfig(
emb_size=48 * 128,
widening_factor=8,
key_size=128,
num_q_heads=48,
num_kv_heads=8,
num_layers=64,
num_experts=8,
num_selected_experts=2,
)
随后 InferenceRunner(... local_mesh_config=(1, 8)) 表示示例预期把本地主设备组织为 data × model mesh,并调用 initialize() 与 run()。
第二站:runners.py 管理生命周期
ModelRunner.initialize() 建 mesh、创建 Haiku transform 和参数 sharding;load_or_init() 调 checkpoint.restore()。InferenceRunner.initialize() 加载 tokenizer,并把 forward、new memory、prefill 和 sample step 包装为 pjit。
InferenceRunner.run() 维护批处理槽位。新请求先 tokenize 和 prefill;活跃请求每轮执行 sample step;结束请求被 decode 并清空。这说明公开示例不只是 model(prompt),而是一个最小生成调度器。
第三站:LanguageModel 包住 Transformer
LanguageModel.__call__() 做四件事:
- 根据 pad token 建 input mask;
InOutEmbed(tokens)查表并缩放;- 调
Transformer(..., memory); - 最终 norm 后用同一 embedding 转置得到 logits。
Transformer.__call__() 建下三角 causal mask,循环 num_layers,把每层新 KVMemory 收集回来。这里能确认这是 decoder-only 自回归栈。
第四站:DecoderLayer 的两个残差块
下面是教学简化,保持主数据流但省略源码中的额外 RMSNorm、分片和 dtype:
# 教学简化(不可直接替代官方实现)
def decoder_layer(h, mask, memory):
attn, next_memory = attention(rms_norm(h), mask, memory)
h = h + attn
expert_mix = moe(rms_norm(h), padding_mask)
return h + expert_mix, next_memory
真实 DecoderLayer 从 model.py:1011 附近开始。建议分别给 MultiHeadAttention.__call__() 和 MoELayer._inference_call() 设置阅读断点,再回到上层理解张量形状。
第五站:checkpoint 不是一个普通文件
权重目录 checkpoints/ckpt-0 包含分片 pickle tensor。checkpoint.restore() 根据当前 state shapes、mesh 和 sharding 选择/组合 tensor。代码里还有 shared memory 加速本机多进程读取的辅助函数。
314B 参数即使 8-bit 也约 314 GB 的纯权重数量级,未计 scale、运行时和 cache;README 因此提醒需要足够 GPU 内存。下载 checkpoint 不等于普通笔记本能运行。
继续阅读的五个问题
LM_PARTITION_RULES怎样把 embedding、attention、router 和 expert 权重映射到 data/model 轴?RotaryEmbedding的 offset 怎样与KVMemory.step同步?- GQA reshape 后
einsum的每个维度字母代表什么? moe_slow_matmul1/2为什么一个需要psum?hk_prefill_memory为什么手动覆盖 step,再把 slice 插回联合 batch memory?