Appearance
run.c 逐行解讀
1. 前置處理與標頭檔 (Lines 1-15)
c
/* Inference for Llama-2 Transformer model in pure C */
#include <stdio.h>
#include <stdlib.h>
#include <ctype.h>
#include <time.h>
#include <math.h>
#include <string.h>
#include <fcntl.h>
#if defined _WIN32
#include "win.h"
#else
#include <unistd.h>
#include <sys/mman.h>
#endif這是整個檔案的依賴引入。注意 zero dependencies 哲學——只用了 C 標準函式庫。mmap(記憶體映射)用於高效讀取權重檔,Windows 下以 win.h 模擬 POSIX 介面。
2. 資料結構定義 (Lines 17-75)
Config — 模型超參數
c
typedef struct {
int dim; // Transformer 維度(如 288)
int hidden_dim; // FFN 隱藏層維度
int n_layers; // 層數
int n_heads; // Query 頭數
int n_kv_heads; // Key/Value 頭數(可少於 query 頭數,GQA)
int vocab_size; // 詞彙表大小
int seq_len; // 最大序列長度
} Config;這個 struct 定義了模型的「藍圖」。所有的維度資訊儲存在 checkpoint 檔案開頭,讀取後就決定了整個模型的形狀。
TransformerWeights — 權重指標
c
typedef struct {
float* token_embedding_table; // Token 嵌入表 (vocab_size, dim)
float* rms_att_weight; // Attention RMSNorm 權重 (layer, dim)
float* rms_ffn_weight; // FFN RMSNorm 權重 (layer, dim)
float* wq; // Query 投影權重 (layer, dim, dim)
float* wk; // Key 投影權重 (layer, dim, kv_dim)
float* wv; // Value 投影權重 (layer, dim, kv_dim)
float* wo; // Output 投影權重 (layer, dim, dim)
float* w1; // FFN w1 (layer, hidden_dim, dim)
float* w2; // FFN w2 (layer, dim, hidden_dim)
float* w3; // FFN w3 (layer, hidden_dim, dim)
float* rms_final_weight; // 最終 RMSNorm 權重 (dim,)
float* wcls; // 分類器權重(可與嵌入表共享)
} TransformerWeights;每個指標指向 mmap 映射的權重資料中的特定位置。命名約定:
wq/wk/wv/wo:對應 Attention 的 Q/K/V/O 線性投影w1/w2/w3:對應 FFN 中的三個線性層(SwiGLU 需要 w1 和 w3)rms_*:RMSNorm 的縮放參數
RunState — 執行期緩衝區
c
typedef struct {
float *x; // 當前時間步的激活值 (dim,)
float *xb; // 殘差分支內的暫存 (dim,)
float *xb2; // 額外暫存 (dim,)
float *hb; // FFN 隱藏層暫存 (hidden_dim,)
float *hb2; // FFN 隱藏層暫存 (hidden_dim,)
float *q; // Query 向量 (dim,)
float *k; // Key 向量 (dim,)
float *v; // Value 向量 (dim,)
float *att; // 注意力分數緩衝區 (n_heads, seq_len)
float *logits; // 輸出 logits (vocab_size,)
float* key_cache; // KV Cache (layer, seq_len, kv_dim)
float* value_cache; // KV Cache (layer, seq_len, kv_dim)
} RunState;這些 buffer 在 malloc_run_state() 中用 calloc 分配。calloc 會將記憶體歸零,有助於 valgrind 偵測未初始化的使用。
Transformer — 頂層整合
c
typedef struct {
Config config;
TransformerWeights weights;
RunState state;
int fd;
float* data; // mmap 資料指標
ssize_t file_size;
} Transformer;fd、data、file_size 是為了在 free_transformer() 中正確清理 mmap 資源。
3. 記憶體管理 (Lines 77-177)
malloc_run_state
c
void malloc_run_state(RunState* s, Config* p) {
int kv_dim = (p->dim * p->n_kv_heads) / p->n_heads;
s->x = calloc(p->dim, sizeof(float));
// ... 分配所有 buffer
s->key_cache = calloc(p->n_layers * p->seq_len * kv_dim, sizeof(float));
s->value_cache = calloc(p->n_layers * p->seq_len * kv_dim, sizeof(float));
// ...
}關鍵計算:kv_dim = (dim * n_kv_heads) / n_heads
當 n_kv_heads < n_heads(GQA),KV 的維度比 Q 小。例如 dim=288, n_heads=6, n_kv_heads=4 → kv_dim = 288*4/6 = 192。
memory_map_weights
c
void memory_map_weights(TransformerWeights *w, Config* p, float* ptr, int shared_weights) {
int head_size = p->dim / p->n_heads;
unsigned long long n_layers = p->n_layers;
w->token_embedding_table = ptr;
ptr += p->vocab_size * p->dim;
w->rms_att_weight = ptr;
ptr += n_layers * p->dim;
// ... 依序推進 ptr
}這段程式碼將 .bin 檔中的資料依序分配到各個權重指標。注意 n_layers 使用 unsigned long long 是為了避免 13B+ 大模型下的整數溢位。
4. 神經網路區塊 (Lines 179-362)
RMSNorm (Lines 182-195)
c
void rmsnorm(float* o, float* x, float* weight, int size) {
float ss = 0.0f;
for (int j = 0; j < size; j++) {
ss += x[j] * x[j];
}
ss /= size;
ss += 1e-5f;
ss = 1.0f / sqrtf(ss);
for (int j = 0; j < size; j++) {
o[j] = weight[j] * (ss * x[j]);
}
}RMSNorm 公式:
RMSNorm(x) = weight * x / sqrt(mean(x^2) + eps)與 LayerNorm 的差異:
- LayerNorm:減均值 → 除標準差 → 縮放
- RMSNorm:除 RMS → 縮放(省略了減均值的步驟)
這使得 RMSNorm 計算更簡單,且在 LLM 中效果相當。
Softmax (Lines 197-215)
c
void softmax(float* x, int size) {
float max_val = x[0];
for (int i = 1; i < size; i++) {
if (x[i] > max_val) max_val = x[i];
}
float sum = 0.0f;
for (int i = 0; i < size; i++) {
x[i] = expf(x[i] - max_val);
sum += x[i];
}
for (int i = 0; i < size; i++) {
x[i] /= sum;
}
}標準的 softmax 實作,先減去最大值確保數值穩定性。三個 pass:找 max → exp 並加總 → 歸一化。
Matmul (Lines 217-229)
c
void matmul(float* xout, float* x, float* w, int n, int d) {
int i;
#pragma omp parallel for private(i)
for (i = 0; i < d; i++) {
float val = 0.0f;
for (int j = 0; j < n; j++) {
val += w[i * n + j] * x[j];
}
xout[i] = val;
}
}作者標註:「by far the most amount of time is spent inside this little function」。這是一個 W(d,n) @ x(n,) → xout(d,) 的矩陣-向量乘法。OpenMP 可用於將 outer loop 分配到多個執行緒。
forward() — 主要前向傳播 (Lines 231-362)
這是最關鍵的函數,實現了整個 Transformer 的單步前向傳播。
Step 1: Token Embedding Lookup (Lines 244-246)
c
float* content_row = w->token_embedding_table + token * dim;
memcpy(x, content_row, dim*sizeof(*x));根據 token ID 從嵌入表中取出對應的向量,複製到 x。
Step 2: Layer Loop (Lines 249-354)
對每一層執行:
2a. Attention RMSNorm (Line 252)
c
rmsnorm(s->xb, x, w->rms_att_weight + l*dim, dim);對 x 做 RMSNorm,結果存入 xb。
2b. KV Cache 指標設定 (Lines 254-257)
c
int loff = l * p->seq_len * kv_dim;
s->k = s->key_cache + loff + pos * kv_dim;
s->v = s->value_cache + loff + pos * kv_dim;k 和 v 指標指向 KV Cache 中當前位置(目前 token 的位置 pos)。之後的 QKV matmul 會直接將結果寫入 Cache。
2c. QKV Matmul (Lines 260-262)
c
matmul(s->q, s->xb, w->wq + l*dim*dim, dim, dim);
matmul(s->k, s->xb, w->wk + l*dim*kv_dim, dim, kv_dim);
matmul(s->v, s->xb, w->wv + l*dim*kv_dim, dim, kv_dim);注意 Q 的維度是 dim,K 和 V 的維度是 kv_dim(GQA 優化)。
2d. RoPE (Lines 265-279)
c
for (int i = 0; i < dim; i+=2) {
int head_dim = i % head_size;
float freq = 1.0f / powf(10000.0f, head_dim / (float)head_size);
float val = pos * freq;
float fcr = cosf(val), fci = sinf(val);
int rotn = i < kv_dim ? 2 : 1;
for (int v = 0; v < rotn; v++) {
float* vec = v == 0 ? s->q : s->k;
float v0 = vec[i], v1 = vec[i+1];
vec[i] = v0 * fcr - v1 * fci;
vec[i+1] = v0 * fci + v1 * fcr;
}
}RoPE 將 (q, k) 向量視為複數,按位置 pos 旋轉。每個 head 內部連續的兩個維度組成一個複數對。
rotn 變數處理 GQA:當 i < kv_dim 時(表示仍在 KV 的維度範圍內),同時旋轉 q 和 k;否則只旋轉 q。
2e. 多頭注意力 (Lines 282-319)
c
for (h = 0; h < p->n_heads; h++) {
float* q = s->q + h * head_size;
float* att = s->att + h * p->seq_len;
for (int t = 0; t <= pos; t++) {
float* k = s->key_cache + loff + t * kv_dim + (h / kv_mul) * head_size;
float score = 0.0f;
for (int i = 0; i < head_size; i++) score += q[i] * k[i];
score /= sqrtf(head_size);
att[t] = score;
}
softmax(att, pos + 1);
float* xb = s->xb + h * head_size;
memset(xb, 0, head_size * sizeof(float));
for (int t = 0; t <= pos; t++) {
float* v = s->value_cache + loff + t * kv_dim + (h / kv_mul) * head_size;
float a = att[t];
for (int i = 0; i < head_size; i++) xb[i] += a * v[i];
}
}每個 head 的處理:
- 取出該 head 的 q 向量
- 與所有過去位置的 k 向量計算 dot product score
- Score 除以
sqrt(head_size)做 scaling - Softmax 得到注意力權重
- 用權重加總所有位置的 v 向量
注意 (h / kv_mul) * head_size:當使用 GQA 時,多個 query head 共享同一個 KV head。
2f. Attention Output (Lines 322-327)
c
matmul(s->xb2, s->xb, w->wo + l*dim*dim, dim, dim);
for (int i = 0; i < dim; i++) x[i] += s->xb2[i];Attention 輸出經過 wo 投影後,透過殘差連接加到 x。
2g. FFN RMSNorm (Line 330)
c
rmsnorm(s->xb, x, w->rms_ffn_weight + l*dim, dim);2h. SwiGLU FFN (Lines 334-353)
c
matmul(s->hb, s->xb, w->w1 + l*dim*hidden_dim, dim, hidden_dim);
matmul(s->hb2, s->xb, w->w3 + l*dim*hidden_dim, dim, hidden_dim);
for (int i = 0; i < hidden_dim; i++) {
float val = s->hb[i];
val *= (1.0f / (1.0f + expf(-val))); // SiLU
val *= s->hb2[i]; // 逐元素乘 w3(x)
s->hb[i] = val;
}
matmul(s->xb, s->hb, w->w2 + l*dim*hidden_dim, hidden_dim, dim);
for (int i = 0; i < dim; i++) x[i] += s->xb[i];SwiGLU 的 PyTorch 等價為:w2(SiLU(w1(x)) * w3(x))
SiLU 激活函數:silu(x) = x * sigmoid(x) = x / (1 + exp(-x))
再次殘差連接。
Step 3: 最終輸出 (Lines 356-361)
c
rmsnorm(x, x, w->rms_final_weight, dim);
matmul(s->logits, x, w->wcls, p->dim, p->vocab_size);
return s->logits;最終的 RMSNorm 後經由分類器權重(與嵌入表共享)產生 logits。
5. BPE Tokenizer (Lines 364-571)
資料結構
c
typedef struct {
char** vocab; // token 字串陣列
float* vocab_scores; // token 分數(用於 BPE 合併優先權)
TokenIndex *sorted_vocab; // 排序後的 vocabulary(用於二分搜尋)
int vocab_size;
unsigned int max_token_length;
unsigned char byte_pieces[512]; // 所有單位元組字串的查找表
} Tokenizer;build_tokenizer (Lines 385-409)
讀取 .bin tokenizer 檔,格式為:
[max_token_length: int]
[score: float][len: int][bytes: len] // 對每個 tokenencode (Lines 452-571)
編碼流程:
- 將輸入字串以 UTF-8 codepoint 為單位分解為 token(支援 byte_fallback)
- 反覆尋找並合併最佳相鄰 token pair(依據 BPE 分數)
- 直到無法找到可合併的 pair
decode (Lines 418-429)
解碼流程:
- 從 vocab 取出 token 對應的字串
- 特殊處理:BOS token 後的開頭空白、
<0xNN>原始位元組
6. Sampler (Lines 573-714)
隨機數產生:xorshift* (Lines 680-689)
c
unsigned int random_u32(unsigned long long *state) {
*state ^= *state >> 12;
*state ^= *state << 25;
*state ^= *state >> 27;
return (*state * 0x2545F4914F6CDD1Dull) >> 32;
}xorshift* 是一種快速、高品質的 PRNG,狀態僅需一個 64-bit 整數。
採樣策略 (Lines 691-714)
c
int sample(Sampler* sampler, float* logits) {
if (sampler->temperature == 0.0f) {
next = sample_argmax(logits, sampler->vocab_size);
} else {
for (int q=0; q<sampler->vocab_size; q++) logits[q] /= sampler->temperature;
softmax(logits, sampler->vocab_size);
float coin = random_f32(&sampler->rng_state);
if (sampler->topp <= 0 || sampler->topp >= 1) {
next = sample_mult(logits, sampler->vocab_size, coin);
} else {
next = sample_topp(logits, sampler->vocab_size, sampler->topp, sampler->probindex, coin);
}
}
return next;
}三種模式:
- 溫度 = 0:貪婪解碼,取最高機率 token
- 溫度 > 0, top-p 關閉:依 softmax 分布隨機採樣
- 溫度 > 0, top-p 啟用:只從累積機率達 p 的最小集合採樣
7. Generation Loop (Lines 728-783)
c
void generate(Transformer *transformer, Tokenizer *tokenizer, Sampler *sampler, char *prompt, int steps) {
// 1. 編碼 prompt 為 tokens
// 2. 主迴圈:
// a. forward() → logits
// b. 若仍在 prompt 範圍,強制使用 prompt token
// c. 否則 sample() 下一個 token
// d. 遇到 BOS(=1) 終止
// e. decode + print
// 3. 回報 tok/s
}8. Chat 模式 (Lines 798-884)
c
void chat(Transformer *transformer, Tokenizer *tokenizer, Sampler *sampler,
char *cli_user_prompt, char *cli_system_prompt, int steps) {實作 Llama 2 的對話格式:
[INST] <<SYS>>
系統提示
<</SYS>>
使用者輸入 [/INST]每輪對話:
- 使用者輸入 → 編碼為 tokens
- Transformer 生成回應
- 遇到 EOS(=2) 結束 Assistant 回覆 → 回到使用者輸入
9. CLI 主程式 (Lines 889-973)
c
int main(int argc, char *argv[]) {
// 1. 解析命令列參數(手動 argparse)
// 2. 建立 Transformer、Tokenizer、Sampler
// 3. 依 mode 執行 generate() 或 chat()
// 4. 清理資源
}支援的參數:
-t:溫度 (default 1.0)-p:top-p (default 0.9)-s:隨機種子-n:生成步數 (default 256)-i:輸入 prompt-z:自訂 tokenizer 路徑-m:模式(generate / chat)-y:system prompt(chat 模式)