Skip to content

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;

fddatafile_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;

kv 指標指向 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 的處理:

  1. 取出該 head 的 q 向量
  2. 與所有過去位置的 k 向量計算 dot product score
  3. Score 除以 sqrt(head_size) 做 scaling
  4. Softmax 得到注意力權重
  5. 用權重加總所有位置的 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]  // 對每個 token

encode (Lines 452-571)

編碼流程:

  1. 將輸入字串以 UTF-8 codepoint 為單位分解為 token(支援 byte_fallback)
  2. 反覆尋找並合併最佳相鄰 token pair(依據 BPE 分數)
  3. 直到無法找到可合併的 pair

decode (Lines 418-429)

解碼流程:

  1. 從 vocab 取出 token 對應的字串
  2. 特殊處理: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]

每輪對話:

  1. 使用者輸入 → 編碼為 tokens
  2. Transformer 生成回應
  3. 遇到 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 模式)