GROK / FIELD GUIDE

输入关键词搜索全部课程标题、摘要与正文。按 ↑ ↓ 选择,Enter 打开。

来源档案
CODE WALKTHROUGH·52 min

源码导读:从 run.py 追到一个新 token

以固定 commit 为坐标,认识 Grok-1 仓库每个文件、关键类型、调用链与适合继续阅读的断点。

源码确认官方声明合理推断未公开

版本与阅读方式

先从 run.py 读取不可变规格,再沿 runners.py 看服务循环,最后进入 model.py 的张量计算。若一开始从 1,400 行 model.py 顶部顺读,很容易迷失在分片规则里。

INTERACTIVE FIGUREGrok-1 JAX 示例:真实调用链
100%
run.pymain()配置模型与 mesh,启动生成器
runners.pyInferenceRunnertokenize · prefill · sample step
model.pyLanguageModelembedding → Transformer → tied decode
model.pyDecoderLayer × 64MHA block + MoE block
model.py:694MultiHeadAttentionRoPE · GQA · KVMemory
model.py:208/272Router / MoELayersoftmax · top_k · gated experts
checkpoint.pyrestore()把 checkpoint tensor 映射到分片状态
路径与符号来自 xai-org/grok-1 main@7050ed2。箭头表示主要调用关系,不是完整控制流图。

文件地图

文件关键符号责任
run.pymain()固化模型配置、mesh 与示例 prompt
model.pyLanguageModel, Transformer, DecoderLayer完整前向数据流
model.pyMultiHeadAttention, KVMemoryGQA、RoPE、mask 与缓存
model.pyRouter, MoELayer, DenseBlockTop-2 专家路由与 FFN
runners.pyModelRunner, InferenceRunnerJAX 编译、prefill、decode、batch slot
runners.pytop_p_filter, sample_tokenlogits 到新 token
checkpoint.pyload_tensors, restore从分片 checkpoint 恢复参数
tokenizer.modelSentencePiece model文本 ↔ token IDs

第一站:run.py 固定配置

main() 构造 LanguageModelConfigTransformerConfig。这里不是“推荐超参数”,而是 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__() 做四件事:

  1. 根据 pad token 建 input mask;
  2. InOutEmbed(tokens) 查表并缩放;
  3. Transformer(..., memory)
  4. 最终 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

真实 DecoderLayermodel.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 不等于普通笔记本能运行。

继续阅读的五个问题

  1. LM_PARTITION_RULES 怎样把 embedding、attention、router 和 expert 权重映射到 data/model 轴?
  2. RotaryEmbedding 的 offset 怎样与 KVMemory.step 同步?
  3. GQA reshape 后 einsum 的每个维度字母代表什么?
  4. moe_slow_matmul1/2 为什么一个需要 psum
  5. hk_prefill_memory 为什么手动覆盖 step,再把 slice 插回联合 batch memory?
本教材只把公开源码用于解释 Grok-1;当前闭源模型的内部结构,除非 xAI 明确披露,否则均标为未知。