Skip to content

run.c 中文註釋版 — Llama 2 純 C 推理引擎

原始檔案:run.c(973 行)

本檔案為原始 run.c 的中文註釋版本,保留所有原始程式碼,
並在關鍵段落加入中文說明。
c
/* Inference for Llama-2 Transformer model in pure C */
/* 以純 C 語言實作 Llama-2 Transformer 模型推理 */

#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"       // Windows 平台相容
#else
    #include <unistd.h>
    #include <sys/mman.h>  // POSIX 記憶體映射
#endif

1. 資料結構 — Transformer 模型定義

c
// Transformer 模型超參數(儲存在 checkpoint 檔案頭部)
typedef struct {
    int dim;            // Transformer 維度
    int hidden_dim;     // FFN 隱藏層維度
    int n_layers;       // 層數
    int n_heads;        // Query 頭數
    int n_kv_heads;     // Key/Value 頭數(GQA 用,可少於 n_heads)
    int vocab_size;     // 詞彙表大小(通常 32000)
    int seq_len;        // 最大序列長度
} Config;

// 所有模型權重指標(指向 mmap 映射的記憶體位置)
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 投影權重
    float* wk;                       // Key 投影權重(維度較小,GQA)
    float* wv;                       // Value 投影權重
    float* wo;                       // Attention 輸出投影
    float* w1;                       // FFN w1
    float* w2;                       // FFN w2
    float* w3;                       // FFN w3(SwiGLU 門控分支)
    float* rms_final_weight;         // 最終 RMSNorm
    float* wcls;                     // 分類器權重(可與嵌入表共享)
} TransformerWeights;

// 執行期間的激活值緩衝區 —— 每次 forward 會重新使用
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 向量 (kv_dim,)
    float *v;           // Value 向量 (kv_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;

// Transformer 頂層結構:config + weights + state + mmap 資訊
typedef struct {
    Config config;
    TransformerWeights weights;
    RunState state;
    int fd;              // mmap 檔案描述符
    float* data;         // mmap 資料指標
    ssize_t file_size;   // 檢查點檔案大小
} Transformer;

2. 記憶體管理

c
void malloc_run_state(RunState* s, Config* p) {
    // kv_dim = GQA 調整後的 Key/Value 維度
    int kv_dim = (p->dim * p->n_kv_heads) / p->n_heads;
    s->x = calloc(p->dim, sizeof(float));
    s->xb = calloc(p->dim, sizeof(float));
    s->xb2 = calloc(p->dim, sizeof(float));
    s->hb = calloc(p->hidden_dim, sizeof(float));
    s->hb2 = calloc(p->hidden_dim, sizeof(float));
    s->q = calloc(p->dim, sizeof(float));
    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));
    s->att = calloc(p->n_heads * p->seq_len, sizeof(float));
    s->logits = calloc(p->vocab_size, sizeof(float));
    if (!s->x || !s->xb || !s->xb2 || !s->hb || !s->hb2 || !s->q
     || !s->key_cache || !s->value_cache || !s->att || !s->logits) {
        fprintf(stderr, "malloc failed!\n");
        exit(EXIT_FAILURE);
    }
}
// 使用 calloc(而非 malloc)讓 valgrind 快樂

void memory_map_weights(TransformerWeights *w, Config* p, float* ptr, int shared_weights) {
    // 指標推進:依 checkpoint 檔案格式,將權重依序指向正確位置
    int head_size = p->dim / p->n_heads;
    unsigned long long n_layers = p->n_layers;  // 用 64-bit 避免大模型溢位
    w->token_embedding_table = ptr;
    ptr += p->vocab_size * p->dim;
    w->rms_att_weight = ptr;
    ptr += n_layers * p->dim;
    w->wq = ptr;
    ptr += n_layers * p->dim * (p->n_heads * head_size);
    w->wk = ptr;
    ptr += n_layers * p->dim * (p->n_kv_heads * head_size);
    w->wv = ptr;
    ptr += n_layers * p->dim * (p->n_kv_heads * head_size);
    w->wo = ptr;
    ptr += n_layers * (p->n_heads * head_size) * p->dim;
    w->rms_ffn_weight = ptr;
    ptr += n_layers * p->dim;
    w->w1 = ptr;
    ptr += n_layers * p->dim * p->hidden_dim;
    w->w2 = ptr;
    ptr += n_layers * p->hidden_dim * p->dim;
    w->w3 = ptr;
    ptr += n_layers * p->dim * p->hidden_dim;
    w->rms_final_weight = ptr;
    ptr += p->dim;
    // 跳過預計算的 RoPE 表格(C 版即時計算,不使用)
    ptr += p->seq_len * head_size / 2; // skip freq_cis_real
    ptr += p->seq_len * head_size / 2; // skip freq_cis_imag
    w->wcls = shared_weights ? w->token_embedding_table : ptr;
    // 權重綁定(weight tying):分類器 = 嵌入表,或獨立指標
}

void read_checkpoint(char* checkpoint, Config* config, TransformerWeights* weights,
                     int* fd, float** data, ssize_t* file_size) {
    FILE *file = fopen(checkpoint, "rb");
    fread(config, sizeof(Config), 1, file);           // 讀取 Config 頭
    int shared_weights = config->vocab_size > 0 ? 1 : 0;
    config->vocab_size = abs(config->vocab_size);     // 負數表示不共享權重
    fseek(file, 0, SEEK_END);
    *file_size = ftell(file);
    fclose(file);
    *fd = open(checkpoint, O_RDONLY);
    *data = mmap(NULL, *file_size, PROT_READ, MAP_PRIVATE, *fd, 0);
    // mmap 將整個 checkpoint 映射到記憶體
    float* weights_ptr = *data + sizeof(Config)/sizeof(float);
    memory_map_weights(weights, config, weights_ptr, shared_weights);
}

3. 神經網路計算區塊

c
// RMSNorm:x / sqrt(mean(x²) + eps) * weight
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]);
}

// Softmax:exp(x_i - max) / sum(exp(x_j - max)),含數值穩定處理
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;
}

// 矩陣 × 向量乘法:W(d,n) @ x(n) → xout(d)
// 全程最耗時的函數(by far the most amount of time is spent inside this little function)
void matmul(float* xout, float* x, float* w, int n, int d) {
    int i;
    #pragma omp parallel for private(i)   // OpenMP 平行化
    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;
    }
}

4. Transformer 前向傳播

c
float* forward(Transformer* transformer, int token, int pos) {
    // token: 當前輸入的 token ID
    // pos: 當前在序列中的位置(從 0 開始)

    Config* p = &transformer->config;
    TransformerWeights* w = &transformer->weights;
    RunState* s = &transformer->state;
    float *x = s->x;
    int dim = p->dim;
    int kv_dim = (p->dim * p->n_kv_heads) / p->n_heads;  // GQA KV 維度
    int kv_mul = p->n_heads / p->n_kv_heads;               // 每個 KV head 對應的 Q head 數
    int hidden_dim = p->hidden_dim;
    int head_size = dim / p->n_heads;

    // Step 1: Token Embedding — 查表取得 token 對應的向量
    float* content_row = w->token_embedding_table + token * dim;
    memcpy(x, content_row, dim*sizeof(*x));

    // Step 2: 逐層處理
    for(unsigned long long l = 0; l < p->n_layers; l++) {

        // 2a. Attention 前的 RMSNorm
        rmsnorm(s->xb, x, w->rms_att_weight + l*dim, dim);

        // 2b. KV Cache 定位:k, v 指標指向當前位置
        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;

        // 2c. QKV 線性投影
        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);

        // 2d. RoPE 相對位置編碼:對 q 和 k 進行複數旋轉
        // 每個連續的 (i, i+1) 維度對視為複數的實部和虛部
        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; // 仍在 KV 維度內則同時旋轉 q 和 k
            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;
            }
        }

        // 2e. 多頭注意力
        int h;
        #pragma omp parallel for private(h)
        for (h = 0; h < p->n_heads; h++) {
            float* q = s->q + h * head_size;          // 此 head 的 q
            float* att = s->att + h * p->seq_len;      // 此 head 的注意力分數

            // 計算與所有過去位置的注意力分數
            for (int t = 0; t <= pos; t++) {
                // GQA: (h / kv_mul) 決定使用哪個 KV head
                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);  // Softmax 歸一化

            // 加權求和 Value
            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];
            }
        }

        // 2f. Attention 輸出投影 + 殘差連接
        matmul(s->xb2, s->xb, w->wo + l*dim*dim, dim, dim);
        for (int i = 0; i < dim; i++) x[i] += s->xb2[i];

        // 2g. FFN 前的 RMSNorm
        rmsnorm(s->xb, x, w->rms_ffn_weight + l*dim, dim);

        // 2h. SwiGLU FFN: w2(SiLU(w1(x)) * w3(x))
        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 = x * sigmoid(x)
            val *= s->hb2[i];                       // 門控相乘
            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];
    }

    // Step 3: 最終輸出
    rmsnorm(x, x, w->rms_final_weight, dim);         // 最終 RMSNorm
    matmul(s->logits, x, w->wcls, p->dim, p->vocab_size);  // 分類器
    return s->logits;
}

5. BPE Tokenizer

c
typedef struct {
    char** vocab;            // 所有 token 字串
    float* vocab_scores;     // 每個 token 的分數(BPE 合併用)
    TokenIndex *sorted_vocab; // 排序後的 vocab(二分搜尋用)
    int vocab_size;
    unsigned int max_token_length;
    unsigned char byte_pieces[512]; // 單位元組字串表
} Tokenizer;

// 解碼:token ID → 可列印字串
char* decode(Tokenizer* t, int prev_token, int token) {
    char *piece = t->vocab[token];
    // BOS token (1) 後去掉前導空白
    if (prev_token == 1 && piece[0] == ' ') { piece++; }
    // 處理原始位元組編碼 <0xNN>
    unsigned char byte_val;
    if (sscanf(piece, "<0x%02hhX>", &byte_val) == 1) {
        piece = (char*)t->byte_pieces + byte_val * 2;
    }
    return piece;
}

// 編碼:字串 → token ID 序列(BPE 合併演算法)
void encode(Tokenizer* t, char *text, int8_t bos, int8_t eos,
            int *tokens, int *n_tokens) {
    // 1. 插入 BOS token (1)(如需要)
    // 2. 將 UTF-8 字串分解為初始 token 序列
    // 3. 反覆合併最高分數的相鄰 token pair
    // 4. 插入 EOS token (2)(如需要)
}

6. Sampler(採樣器)

c
// 三種採樣策略:
// 1. Argmax:取最高機率 token(溫度 = 0)
int sample_argmax(float* probabilities, int n) { ... }

// 2. Multinomial:依機率分布隨機採樣
int sample_mult(float* probabilities, int n, float coin) { ... }

// 3. Top-p (nucleus):從累積機率達 p 的最小集合採樣
int sample_topp(float* probabilities, int n, float topp,
                ProbIndex* probindex, float coin) { ... }

// xorshift* PRNG — 快速偽隨機數產生器
unsigned int random_u32(unsigned long long *state) {
    *state ^= *state >> 12;
    *state ^= *state << 25;
    *state ^= *state >> 27;
    return (*state * 0x2545F4914F6CDD1Dull) >> 32;
}

7. 生成主迴圈

c
void generate(Transformer *transformer, Tokenizer *tokenizer,
              Sampler *sampler, char *prompt, int steps) {
    // 1. 編碼 prompt → tokens
    // 2. 主迴圈 (pos = 0..steps):
    //    a. forward(token, pos) → logits
    //    b. 如在 prompt 範圍內 → 強制下一個 prompt token
    //      否則 → sample(logits) 產生下一個 token
    //    c. 遇到 BOS(=1) 終止生成
    //    d. decode + print token
    // 3. 輸出 tok/s 效能報告
}

8. Chat 對話模式

c
void chat(Transformer *transformer, Tokenizer *tokenizer, Sampler *sampler,
          char *cli_user_prompt, char *cli_system_prompt, int steps) {
    // 使用 Llama 2 Chat 格式:
    // [INST] <<SYS>>\n{system_prompt}\n<</SYS>>\n\n{user_prompt} [/INST]
    // 輪流:使用者輸入 → Assistant 生成 → 使用者輸入 → ...
    // Assistant 遇到 EOS(=2) 結束回答
}

9. 命令列介面

c
int main(int argc, char *argv[]) {
    // 預設參數
    char *checkpoint_path = NULL;   // 模型路徑(必要)
    char *tokenizer_path = "tokenizer.bin";
    float temperature = 1.0f;
    float topp = 0.9f;
    int steps = 256;
    char *prompt = NULL;
    char *mode = "generate";  // generate | chat

    // 手動命令列參數解析
    // -t 溫度 | -p top-p | -s 隨機種子
    // -n 步數 | -i 提示 | -z tokenizer 路徑
    // -m 模式 | -y system prompt

    // 初始化 Transformer + Tokenizer + Sampler
    // 執行 generate() 或 chat()
    // 清理資源
}