Simply Patrick

用 Rust + OpenCV 實現影片自動人臉打碼

上一篇我們用 Rust + OpenCV 的 YuNet 模型做了靜態圖片和 Webcam 的即時人臉偵測。偵測到人臉之後,下一步自然會想到:能不能自動把人臉打上馬賽克?

這就是 face-mosaic 這個專案要做的事——讀入一段影片,自動偵測每一幀裡的人臉,打上模糊或像素化效果,然後輸出一支全新的影片。整個過程完全自動化,不需要手動框選任何東西。

從 face-detect 到 face-mosaic

如果你看過前一篇文章,會發現 face-detect 已經解決了最核心的問題:用 YuNet 找到人臉的 Bounding Box。face-mosaic 要做的就是在這個基礎上加兩件事:

  1. 影片處理管線:逐幀讀取 → 偵測 → 打碼 → 寫入輸出影片
  2. 模糊/像素化效果:在偵測到的人臉區域上套用視覺遮蔽

聽起來簡單,但實際動手才會發現影片處理有很多細節要處理。

影片處理的基本架構

影片處理的核心是一個 read-process-write 的迴圈。OpenCV 的 VideoCapture 負責讀取,VideoWriter 負責輸出:

let mut cap = videoio::VideoCapture::from_file(&input_path, videoio::CAP_ANY)?;
let fps = cap.get(videoio::CAP_PROP_FPS)?;
let width = cap.get(videoio::CAP_PROP_FRAME_WIDTH)? as i32;
let height = cap.get(videoio::CAP_PROP_FRAME_HEIGHT)? as i32;
let total_frames = cap.get(videoio::CAP_PROP_FRAME_COUNT)? as i64;

let mut writer = videoio::VideoWriter::new(
    &output_path,
    videoio::VideoWriter::fourcc('m', 'p', '4', 'v')?,
    fps,
    Size::new(width, height),
    true, // isColor
)?;

這裡有幾個重要參數:

  • FPS:從原始影片直接讀取,確保輸出影片的播放速度跟原始影片一致
  • fourcc:影片編碼格式。mp4v 是 MPEG-4 編碼,跟 .mp4 容器相容
  • 尺寸:直接沿用原始影片的解析度

逐幀偵測與打碼

主要的處理迴圈長這樣:

let mut frame = Mat::default();
let mut frame_count: i64 = 0;

loop {
    cap.read(&mut frame)?;
    if frame.empty() {
        break;
    }
    frame_count += 1;

    // 更新偵測器的輸入尺寸
    detector.set_input_size(frame.size()?)?;

    // 偵測人臉
    let mut faces = Mat::default();
    detector.detect(&frame, &mut faces)?;

    // 對每張偵測到的臉打碼
    for i in 0..faces.rows() {
        let x = *faces.at_2d::<f32>(i, 0)? as i32;
        let y = *faces.at_2d::<f32>(i, 1)? as i32;
        let w = *faces.at_2d::<f32>(i, 2)? as i32;
        let h = *faces.at_2d::<f32>(i, 3)? as i32;

        apply_mosaic(&mut frame, Rect::new(x, y, w, h))?;
    }

    writer.write(&frame)?;

    if frame_count % 100 == 0 {
        println!("Processed {}/{} frames", frame_count, total_frames);
    }
}

每 100 幀印一次進度,因為處理長影片時沒有任何回饋會讓人以為程式當掉了。

馬賽克效果的兩種實現

臉部打碼可以用兩種方式實現:高斯模糊和像素化(馬賽克)。

高斯模糊

最直觀的做法是對人臉區域套用高斯模糊。OpenCV 的 gaussian_blur 一行搞定:

fn apply_blur(frame: &mut Mat, roi: Rect) -> Result<()> {
    // 確保 ROI 不會超出圖片邊界
    let roi = clamp_rect(roi, frame.size()?);
    let mut face_region = Mat::roi(frame, roi)?;
    let mut blurred = Mat::default();

    // kernel size 越大越模糊,必須是奇數
    imgproc::gaussian_blur(
        &face_region,
        &mut blurred,
        Size::new(99, 99),
        30.0, // sigmaX
        30.0, // sigmaY
        opencv::core::BORDER_DEFAULT,
    )?;

    blurred.copy_to(&mut face_region)?;
    Ok(())
}

像素化(馬賽克)

像素化的原理也很巧妙:先把人臉區域縮小到很小(比如 10×10),再放大回原本的尺寸。因為放大時用的是最近鄰插值(Nearest Neighbor),就會產生經典的馬賽克方塊效果:

fn apply_mosaic(frame: &mut Mat, roi: Rect) -> Result<()> {
    let roi = clamp_rect(roi, frame.size()?);
    let face_region = Mat::roi(frame, roi)?;

    let pixel_size = 10;
    let small_size = Size::new(
        (roi.width / pixel_size).max(1),
        (roi.height / pixel_size).max(1),
    );

    // 縮小
    let mut small = Mat::default();
    imgproc::resize(
        &face_region, &mut small,
        small_size, 0.0, 0.0,
        imgproc::INTER_LINEAR,
    )?;

    // 放大回原尺寸(最近鄰插值 = 馬賽克效果)
    let mut mosaic = Mat::default();
    imgproc::resize(
        &small, &mut mosaic,
        Size::new(roi.width, roi.height), 0.0, 0.0,
        imgproc::INTER_NEAREST,
    )?;

    let mut face_mut = Mat::roi_mut(frame, roi)?;
    mosaic.copy_to(&mut face_mut)?;
    Ok(())
}

這種「先縮後放」的技巧比逐像素手動計算簡單很多,而且完全利用了 OpenCV 既有的 resize 函式。

ROI 邊界處理:一個容易被忽略的坑

YuNet 回傳的 Bounding Box 座標有時候會超出圖片邊界,特別是當人臉在畫面邊緣的時候。如果不處理,OpenCV 會直接 panic。所以需要一個 clamp 函式:

fn clamp_rect(rect: Rect, size: Size) -> Rect {
    let x = rect.x.max(0);
    let y = rect.y.max(0);
    let w = rect.width.min(size.width - x);
    let h = rect.height.min(size.height - y);
    Rect::new(x, y, w.max(0), h.max(0))
}

這個函式確保 ROI 永遠在圖片範圍內。看起來很小,但沒有它,處理真實世界的影片時幾乎一定會爆掉。

效能考量

處理影片跟處理靜態圖片最大的差異在於效能。一段 30fps、10 分鐘的影片有 18,000 幀,每一幀都要跑一次 YuNet 偵測。幾個值得注意的點:

  • 不需要對每幀做縮放:影片的每一幀解析度是固定的,不像靜態圖片可能有各種解析度。set_input_size 只需要在第一幀或解析度改變時呼叫
  • YuNet 夠快:這個模型是為邊緣裝置設計的,在一般筆電上一幀只需要幾毫秒
  • 瓶頸在 I/O:影片的讀寫比偵測本身更花時間,特別是寫入大型 MP4 檔案時

在我的測試中,處理一段 1080p 影片大約能跑到 15-25 fps,主要受限於影片解碼和編碼的速度。

跟 face-detect 的程式碼共用

face-mosaic 和 face-detect 的核心偵測邏輯完全一樣——都是用 FaceDetectorYN 呼叫 YuNet 模型。差別在於:

face-detect face-mosaic
輸入 靜態圖片 / Webcam 影片檔案
輸出 畫框 + 特徵點 模糊/馬賽克
處理方式 單張 / 即時串流 逐幀批次處理
寫入 圖片 / GUI 視窗 VideoWriter → MP4

兩個專案都用了本地 patch 過的 opencv-rust(因為要支援 OpenCV 5),依賴結構也幾乎相同。

學到的東西

這個專案最有趣的不是演算法,而是影片處理的工程細節:

概念 應用
VideoCapture / VideoWriter OpenCV 的影片 I/O API
fourcc 編碼 指定影片壓縮格式
ROI (Region of Interest) 只對圖片的局部區域操作
Nearest Neighbor 插值 放大時產生馬賽克效果
邊界 Clamp 防止 ROI 超出圖片範圍

結語

face-mosaic 是 face-detect 的自然延伸。人臉偵測本身只是起點,真正有趣的是你拿偵測結果來做什麼——這裡是打碼,但同樣的架構也可以拿來做人臉追蹤、表情辨識、或者臉部特效。

影片處理讓整個專案的複雜度跳了一級,但 Rust 的型別系統在這裡依然很有幫助:Mat 的生命週期管理、Result 的錯誤傳播、以及編譯器對邊界條件的提醒,都讓你在處理 18,000 幀的迴圈裡更有信心。

參考資源


在 Rust 中用 OpenCV 實現即時人臉與特徵點偵測

說到電腦視覺,大家第一個想到的語言通常是 Python。但如果你跟我一樣,覺得「都用 Rust 了,為什麼不連人臉偵測也用 Rust 來寫?」——那這篇文章就是寫給你的。

這次的 Rust 52 Projects 挑戰,我用 OpenCV 內建的 YuNet 輕量級深度學習模型,寫了一個能處理靜態圖片和即時 Webcam 影像的人臉偵測工具。除了畫 Bounding Box,還會標記出五個臉部特徵點(雙眼、鼻尖、左右嘴角),每個特徵點用不同的顏色呈現。

為什麼選 YuNet?

人臉偵測的模型百百種,從古早的 Haar Cascade 到各種 YOLO 變體都有。我最後選了 YuNet 有幾個原因:

  1. OpenCV 原生支援:從 OpenCV 4.7 開始,FaceDetectorYN 就是內建 API,不用另外裝什麼深度學習框架
  2. ONNX 格式:模型就是一個 .onnx 檔案,下載下來直接用
  3. 夠快夠準:這個模型是專門為邊緣裝置設計的,在一般筆電上跑即時 Webcam 完全不是問題
  4. 特徵點偵測:不只給你一個方框,還附帶五個臉部特徵點的座標

環境建置:opencv-rust 的地雷區

在 Rust 裡使用 OpenCV,靠的是 opencv-rust 這個綁定庫。坦白說,這大概是整個專案裡最痛苦的部分。在 Windows 上需要先用 vcpkg 安裝 OpenCV:

vcpkg install opencv4:x64-windows

然後設定一堆環境變數:

$env:OPENCV_LINK_LIBS = "opencv_world4"
$env:OPENCV_LINK_PATHS = "C:\path\to\vcpkg\installed\x64-windows\lib"
$env:OPENCV_INCLUDE_PATHS = "C:\path\to\vcpkg\installed\x64-windows\include"

另外因為 opencv-rust 的 build script 會用 libclang 來解析 C++ header 產生 Rust 綁定,所以還需要確保系統裝了 LLVM/Clang。Cargo.toml 中啟用了 clang-runtime feature:

[dependencies]
anyhow = "1"
clap = { version = "4", features = ["derive"] }
opencv = { path = "opencv-rust-patch", features = ["clang-runtime"] }

你可能注意到我用的是 path = "opencv-rust-patch" 而不是 crates.io 上的版本。這是因為我需要針對 OpenCV 5 做一些修正,所以 fork 了一份放在本地。

用 Clap 定義 CLI 介面

工具的使用方式很簡單:不帶參數就開 Webcam,指定 --image 就處理靜態圖片。用 clap 的 derive macro 定義起來非常清爽:

#[derive(Parser)]
#[command(name = "face-detect", version)]
struct Cli {
    /// Path to an image file. If omitted, opens the webcam.
    #[arg(short, long)]
    image: Option<PathBuf>,

    /// Path to the YuNet ONNX model file.
    #[arg(short, long, default_value = "face_detection_yunet_2023mar.onnx")]
    model: PathBuf,

    /// Minimum confidence score to keep a detection (0.0–1.0).
    #[arg(short, long, default_value_t = 0.9)]
    score_threshold: f32,

    /// NMS IoU threshold (0.0–1.0).
    #[arg(short, long, default_value_t = 0.3)]
    nms_threshold: f32,

    /// Path to save the annotated output image (image mode only).
    #[arg(short, long)]
    output: Option<PathBuf>,
}

這裡有兩個可以調整的閾值:

  • score_threshold:信賴度門檻,只有信心分數超過這個值的偵測結果才會保留。預設 0.9 表示「非常有把握才算」
  • nms_threshold:Non-Maximum Suppression 的 IoU 門檻,用來消除重疊的偵測框

建立 YuNet 偵測器

建立偵測器的程式碼意外地簡潔。FaceDetectorYN::create 就是 OpenCV 對 YuNet 模型的封裝:

let mut detector = FaceDetectorYN::create(
    model_path,
    "",                  // config(YuNet 不需要)
    Size::new(320, 320), // 初始 input_size,之後會根據實際圖片調整
    cli.score_threshold,
    cli.nms_threshold,
    5000,                // top_k:最多保留幾個候選框
    0,                   // backend:預設
    0,                   // target:CPU
)?;

值得注意的是 input_size 只是個初始值。實際上每次呼叫 detect 之前,我們都會用 set_input_size 把它調成當前圖片(或影格)的實際尺寸,讓模型知道要處理多大的輸入。

靜態圖片模式:縮放的藝術

處理靜態圖片有個實務上的考量:如果直接拿一張 4000×3000 的照片丟進去偵測,不僅速度慢,模型在太大的解析度下表現也不一定好。所以我做了一個等比例縮放的策略:

let size = img.size()?;
let max_dim = 800.0;
let mut scale = 1.0;

let mut detect_img = img.clone();
if size.width > 800 || size.height > 800 {
    scale = f64::max(size.width as f64, size.height as f64) / max_dim;
    let new_w = (size.width as f64 / scale).round() as i32;
    let new_h = (size.height as f64 / scale).round() as i32;
    let new_size = Size::new(new_w, new_h);
    imgproc::resize(&img, &mut detect_img, new_size, 0.0, 0.0, imgproc::INTER_LINEAR)?;
}

關鍵在於:偵測完之後,我們需要把座標乘回去,這樣在原始解析度的圖片上畫框才會對齊:

if scale != 1.0 && count > 0 {
    for i in 0..count {
        for j in 0..14 {
            let v = *faces.at_2d::<f32>(i, j)?;
            *faces.at_2d_mut::<f32>(i, j)? = (v as f64 * scale) as f32;
        }
    }
}

注意只處理前 14 個欄位(座標),第 15 個欄位是信賴度分數,不需要縮放。

Webcam 即時模式

Webcam 模式就是一個經典的影像處理 loop:讀一張影格 → 偵測 → 畫框 → 顯示 → 檢查鍵盤輸入。

fn detect_in_webcam(detector: &mut impl FaceDetectorYNTrait) -> Result<()> {
    let mut cam = videoio::VideoCapture::new(0, videoio::CAP_ANY)?;
    if !cam.is_opened()? {
        bail!("Cannot open default webcam (index 0)");
    }

    let window = "face-detect – webcam (press Q to quit)";
    highgui::named_window(window, highgui::WINDOW_AUTOSIZE)?;

    let mut frame = Mat::default();
    loop {
        cam.read(&mut frame)?;
        if frame.empty() { continue; }

        detector.set_input_size(frame.size()?)?;

        let mut faces = Mat::default();
        detector.detect(&frame, &mut faces)?;

        draw_detections(&mut frame, &faces)?;
        highgui::imshow(window, &frame)?;

        let key = highgui::wait_key(1)?;
        if key == b'q' as i32 || key == 27 { break; }
    }

    highgui::destroy_all_windows()?;
    Ok(())
}

這裡有幾個 Rust 特色值得留意:

  • Trait bound impl FaceDetectorYNTrait:用 trait 而不是具體型別,讓函式更通用
  • bail! 巨集:來自 anyhow crate,等同於 return Err(anyhow!(...))
  • ? 運算子:幾乎每一行 OpenCV 呼叫都可能失敗,? 讓錯誤處理保持簡潔

臉部特徵點的資料結構

YuNet 回傳的 faces 矩陣是一個 $N \times 15$ 的 Matf32),每一列代表一張偵測到的臉:

欄位 內容
0–3 Bounding Box 的 x, y, w, h
4–5 右眼座標 🔵
6–7 左眼座標 🔴
8–9 鼻尖座標 🟢
10–11 右嘴角座標 🩷
12–13 左嘴角座標 🟡
14 信賴度分數

繪製的時候,我把五個特徵點設定了不同顏色(注意 OpenCV 用的是 BGR 色彩空間,不是 RGB):

const LANDMARK_COLORS: [(f64, f64, f64); 5] = [
    (255.0, 0.0, 0.0),     // 右眼  – 藍色
    (0.0, 0.0, 255.0),     // 左眼  – 紅色
    (0.0, 255.0, 0.0),     // 鼻尖  – 綠色
    (255.0, 0.0, 255.0),   // 右嘴角 – 粉色
    (0.0, 255.0, 255.0),   // 左嘴角 – 黃色
];

然後用一個迴圈把每個特徵點畫成小圓點:

for (j, &(b, g, r)) in LANDMARK_COLORS.iter().enumerate() {
    let col = 4 + j as i32 * 2;
    let lx = *faces.at_2d::<f32>(i, col)? as i32;
    let ly = *faces.at_2d::<f32>(i, col + 1)? as i32;
    imgproc::circle(
        img,
        Point::new(lx, ly),
        3,
        Scalar::new(b, g, r, 0.0),
        imgproc::FILLED,
        imgproc::LINE_AA,
        0,
    )?;
}

這段程式碼的小巧思在於用 enumerate 搭配固定的欄位偏移量 4 + j * 2,把五對 x/y 座標和五組顏色優雅地對應起來。

測試:不只是跑一跑而已

這個專案寫了蠻完整的測試。除了基本的功能測試,還有兩個用真實資料集跑的 benchmark:

基礎功能測試直接驗證偵測結果是否符合預期:

#[test]
fn test_face_detection_lena() -> Result<()> {
    let mut detector = setup_detector(0.5)?;
    let count = detect_in_image(&mut detector, &"tests/lena.jpg".into(), Some(&"tests/lena_out.jpg".into()))?;
    assert_eq!(count, 1, "Should detect exactly 1 face in lena.jpg");
    Ok(())
}

LFW 命中率 BenchmarkLabeled Faces in the Wild 資料集的 1000 張真實人臉照片,計算偵測器的命中率。要求至少 80%:

let hit_rate = (hits as f32 / images.len() as f32) * 100.0;
assert!(hit_rate >= 80.0, "Hit rate {:.1}% is below acceptable 80% threshold", hit_rate);

Stanford Background 誤判率 Benchmark 反過來,用完全不含人臉的風景照片(來自 Stanford Background Dataset),測試偵測器會不會「看到不該看到的臉」。要求誤判率低於 5%:

let fp_rate = (false_positives as f32 / total) * 100.0;
assert!(fp_rate <= 5.0, "False positive rate {:.1}% is above acceptable 5% threshold", fp_rate);

這種正反兩面夾擊的測試方式,讓你在調整 score_threshold 時有科學依據,而不是靠感覺。

學到的 Rust 概念

概念 應用場景
impl Trait 參數 讓偵測函式接受任何實作 FaceDetectorYNTrait 的型別
anyhow::Result 統一錯誤處理,不用為每種錯誤定義型別
clap derive macro 宣告式定義 CLI 介面
Mat 的泛型存取 at_2d::<f32>(i, j) 在型別安全下操作矩陣
模式匹配 + 解構 for (j, &(b, g, r)) 同時取得索引和解構元組
Option 驅動的分支 cli.imageSome/None 決定圖片模式或 Webcam 模式

結語

用 Rust 寫電腦視覺程式,最大的挑戰不在演算法本身,而是環境建置和 FFI 綁定。一旦搞定了 opencv-rust 的編譯問題,寫起來其實蠻順暢的——Rust 的型別系統和錯誤處理在這種「每一行都可能出錯」的 FFI 情境下,反而讓人特別安心。

如果你也想試試看,記得先下載 YuNet ONNX 模型,然後準備好耐心面對 OpenCV 的安裝過程 😄

參考資源


tiny-llm-runner 深入解讀 (9):main.rs —— CLI、Prefill、Decode 與整體效能

featured.svg

本文由 AI Agent(Claude)代筆撰寫,文中的「我」指的是 AI Agent。Patrick 只有在文章最後做過潤飾調整。

歷經八篇深入解讀,我們終於來到 tiny-llm-runner 的最後一塊——main.rs。前面拆了那麼多零件,總得有人把它們兜起來吧?這個檔案就 130 行,是把所有元件串起來的「指揮台」。

這也是整個系列的最後一篇了。除了講 main.rs,我想在最後對整個專案做一個全局的效能最佳化清單,把每個檔案散落的優化點串起來看——算是給這趟旅程一個交代。

概念一:CLI 參數設計

#[derive(Parser, Debug)]
#[command(version, about = "Pure-Rust llama-architecture inference over a GGUF model")]
struct Args {
    #[arg(short, long)]                              model: PathBuf,
    #[arg(short, long, default_value = "Once upon a time")] prompt: String,
    #[arg(short, long, default_value_t = 64)]        n_predict: usize,
    #[arg(short, long, default_value_t = 0.8)]       temperature: f32,
    #[arg(long, default_value_t = 40)]               top_k: usize,
    #[arg(long, default_value_t = 42)]               seed: u64,
    #[arg(long)]                                     no_bos: bool,
    #[arg(long, default_value = "llama")]            rope: String,
}

clap 的 derive macro 把 CLI parsing 變成宣告式:每個欄位加上 #[arg(...)] 就自動產生 --model--prompt 之類的 flag。這比手寫 argument parser 短得多,而且 --help、type validation、default value 全都免費送你,實在是頗划算。

clap 的 zero-cost abstraction

clap 的 macro 在 compile-time 就生成好 parsing code,runtime 沒有任何 reflection。也就是說啟動時 Args::parse() 是純 native code,速度比 Python 的 argparse 快上好幾個量級。

對 LLM runner 來說 CLI 啟動時間其實不是什麼大問題,不過這個習慣很 Rust——把 metadata 處理推到 compile time,runtime 只留下純粹的計算。看多了你會發現整個語言都在貫徹這件事。

概念二:unsafe Mmap

let file = File::open(&args.model)?;
let mmap = unsafe { Mmap::map(&file)? };

整個專案唯一一個 unsafe,就這麼一行。為什麼 mmap 非得 unsafe 不可?

因為 mmap 違反了 Rust 的記憶體模型假設:Rust 假設一個 &[u8] 的內容在它的生命週期內不會被外部修改。但 mmap 對應的檔案如果被另一個 process 改掉(甚至 truncate),這個 &[u8] 就會看到變了樣的資料、最慘還會吃到 SIGBUS。

unsafe 說穿了就是程式設計師對編譯器的一句承諾:「我知道這違反一般規則,使用過程中檔案不會被外部動到,我自己負責」。對 LLM 模型檔來說這承諾其實很好守——模型檔通常就是 read-only 的,誰會去動它呢。

這也呼應了 Rust 的一個設計哲學:unsafe 不是禁忌,而是被精準框定的工具。整個 codebase 只有這一行 unsafe,但它被框得清清楚楚——出了問題,責任就在這一行,跑不掉。

演算法核心:Prefill / Decode 二段式

// 1. Prefill —— 處理 prompt
let prefill_start = Instant::now();
let mut last_logits: Option<Vec<f32>> = None;
for &tok in &prompt_ids {
    let logits = runner.forward(tok);
    last_logits = Some(logits.to_vec());
}
let prefill_elapsed = prefill_start.elapsed();

// 2. Decode —— 生成 token
let decode_start = Instant::now();
let mut generated: Vec<u32> = Vec::with_capacity(args.n_predict);
let mut logits = last_logits.expect("empty prompt");
for _ in 0..args.n_predict {
    let next = sampler.sample(&mut logits);
    if next == tokenizer.eos { break; }
    generated.push(next);
    let piece = tokenizer.decode(&[next]);
    print!("{piece}");
    std::io::stdout().flush().ok();
    logits = runner.forward(next).to_vec();
}

為什麼分兩階段?

LLM 推論天然分兩個階段:

  • Prefill:把使用者的 prompt 餵進去,建立 KV cache。logits 只有最後一個 token 的有用——前面的丟掉。
  • Decode:每次 forward 一個 token、抽下一個。每個 logits 都會用到。

這兩個階段的特性差異很有意思:

  • Prefill 的 token 全都是已知的,理論上可以批次處理(用 GEMM 取代 GEMV)。
  • Decode 就只能乖乖 sequential(下一個 token 取決於上一個,沒得偷懶)。

不過我目前 prefill 也是 sequential(一個 token 一個 forward)的,這就是個明擺著的優化機會了——後面清單會再回來算這筆帳。

演算法核心:tok/s 的計算

eprintln!("[prefill] {} tok in {:.2}s ({:.1} tok/s)",
    prompt_ids.len(),
    prefill_elapsed.as_secs_f64(),
    prompt_ids.len() as f64 / prefill_elapsed.as_secs_f64().max(1e-9),
);

max(1e-9) 是用來防止 0 秒(極短 prompt)導致除以零。f64::max(self, other) 回傳兩者較大者,所以 0.0.max(1e-9) = 1e-9,分母就保證不會是 0 了。

這種小細節很容易忘記寫喔——prompt 只有一個 token 時,prefill 可能是 0.001 秒,算出來還有意義;但要是快到變成 0 秒(測試環境有時就是這麼誇張),分母歸零你就會收到一個漂亮的 NaN。

Rust 用法:streaming output 的 flush

print!("{piece}");
std::io::stdout().flush().ok();

print! 寫進 stdout buffer,但不會馬上顯示——一般 stdout 是 line-buffered,要等到 \n 才 flush。LLM 串流輸出又沒有 \n,所以非得手動 flush 不可,不然你會傻等半天什麼都看不到,還以為當機了。

flush().ok()Result<(), Error> 轉成 Option<()> 然後丟掉——白話講就是「這個 flush 失不失敗我才懶得管」。stdout 寫入失敗本來就極罕見(例如 pipe 被人砍掉),就算真的失敗我們也無能為力,silent ignore 反而是最合理的處理。

Rust 用法:anyhow 的錯誤處理

fn main() -> Result<()> {
    // ... 整個 main 都是 Result-friendly 的,用 ? 早期返回
    Ok(())
}

fn main() -> Result<()> 是 Rust 處理 CLI errors 最乾淨的寫法。任何 ? 失敗都會把錯誤往 main 外面丟,runtime 自動 print 出來再 exit 1,連 error handling 的 boilerplate 都省了。

anyhow::Result<T> 不過就是 Result<T, anyhow::Error> 的別名。anyhow::Error 可以從任何 std::error::Error 自動轉換——這就是為什麼我能把 std::io::Error、parser error、自定義的 bail! 全混在一起,通通用一個 ? 打發掉,實在是頗舒服。

Rust 用法:環境變數和 stderr

eprintln!("[loaded] n_layer={} ...", config.n_layer, ...);

eprintln! 寫到 stderr,println! 寫到 stdout。我刻意把 metadata 印在 stderr、生成內容印在 stdout——這樣你用 ./tiny-llm-runner > out.txt 時,out.txt 裡就只有乾淨的生成內容,那些 metadata 還是乖乖留在 console 上,不會污染你的檔案。

這是 Unix 工具的老慣例了。在 Rust 裡用兩個不同的 macro 就自然支援,不必特別費心。

整個專案的端到端 forward pass 流程

flowchart TD A[CLI Args] --> B[File::open + Mmap] B --> C[parse_gguf] C --> D[LlamaConfig::from_gguf] C --> E[LlamaModel::load
建 TensorView] C --> F[Tokenizer::from_gguf] F --> G[encode prompt] D --> H[Runner::new
配 KV cache + scratch] E --> H G --> I[Prefill loop
forward each token] H --> I I --> J[Decode loop
sample → forward → repeat] J --> K[print tokens]

從 CLI 進來到 token 吐出去,整條流水線就這樣。仔細看會發現,圖裡每一個方框幾乎都對應到前面九篇文章其中一篇的主題——拼到這裡,整張地圖才算完整。

全局效能最佳化清單

到這裡,所有檔案都翻過一遍了。我想趁記憶猶新,把整個專案的最佳化機會匯總成一張 prioritized list。要強調的是:我不建議盲目地照著順序硬幹,還是得看你自己最想練哪一塊。

Tier 1(最大效能槓桿,10× 級的改進)

  1. SIMD 化 dot kernelsdequant.rs

    • Q4_0、Q8_0、Q6_K 的 inner loop 用 AVX2/AVX-512/NEON
    • 預期:matvec 加速 8-16×
    • 工作量:中—需要小心和 llama.cpp 對拍正確性
  2. Prefill batching(GEMV → GEMM)runner.rsops.rs

    • 把 prompt N 個 token 的 forward 拼成一個 batched 計算
    • 預期:prefill 加速 5-10×(decode 不變)
    • 工作量:大—涉及 attention 的 mask、KV cache 的 batched 寫入
  3. 支援 K-quants(Q4_K、Q5_K、Q4_K_M)dequant.rs

    • 不是加速 per se,而是讓現代 GGUF 都能跑
    • 工作量:中—實作複雜但有 ggml C 程式碼可參考

Tier 2(顯著改進,2-3× 級)

  1. Multi-row matvec fusionops.rs

    • 一次處理多個 row,減少 x 的 cache miss
    • 預期:matvec 加速 2-4×
  2. KV cache 量化runner.rs

    • 把 KV cache 從 F32 改成 Q8_0
    • 預期:記憶體用量 4×、速度可能略有提升(cache miss 變少)
    • 工作量:中
  3. f16 / bf16 全程(多個檔案)

    • 不要每次都 dequant 成 f32,scratch buffer 也用 f16
    • 預期:記憶體頻寬減半
    • 工作量:大—需要全程 f16 的 numerical stability 驗證

Tier 3(小但容易的改進)

  1. RoPE sin/cos 表預計算ops.rs

    • 不要每次 forward 都算 sin/cos
    • 預期:每層省幾十 μs,整體可能 1-2%
  2. Tokenizer 的 pair lookup tabletokenizer.rs

    • 避免 format! 字串拼接
    • 預期:encode 加速 5-10×(但 encode 不在 hot path)
  3. Top-P samplingsampler.rs

    • 提升 sampling 品質(不是速度,是輸出品質)
  4. Repetition penaltysampler.rs

    • 同上

Tier 4(架構級重構,可能不值得)

  1. GPU backend

    • wgpu 或 CUDA 支援
    • 工作量:極大—基本上是另一個專案
  2. FlashAttention

    • Fused attention with online softmax
    • 工作量:大—但 candle/ggml 有現成實作可學
  3. Speculative decoding

    • 用小模型加速大模型推論
    • 工作量:大—需要兩個模型協作

一個整體觀察:抽象與效能的權衡

寫完整個系列,我最強烈的感受是:tiny-llm-runner 的「易讀」,其實是拿「不可擴充」換來的。每個檔案都針對單一情境寫得直白到不行,但代價就是——要加新功能(新架構、新量化、新後端)往往得同時動好幾個檔案。

這跟 candle 的設計哲學根本是兩條路。candle 透過一層厚厚的抽象(TensorModuleVarBuilder)讓擴充變得很便宜,代價則是「想看懂一次 forward pass,得在好幾個 trait 之間跳來跳去」。

那到底哪個對?老實說,取決於你的目標。 你要做生產級框架,candle 那套抽象就是必要之惡;你要做一個「能跑、能讀、能 hack」的學習版,那 tiny-llm-runner 的扁平結構反而才是對的。沒有標準答案,只有適不適合。

把九個檔案的學習收穫匯總

回到我們最一開始問的那個問題:「寫一個會跑 LLM 的專案,到底需要哪些零件?」攤開來看,答案就是這張表:

檔案 你需要學會
config.rs metadata 解析 + 不變式檢查
dequant.rs block-wise 量化 + fused dot product
tensor.rs lifetime + view 抽象 + Copy struct
model.rs 樹形權重組織 + tied embeddings
ops.rs RMSNorm、softmax、RoPE、SwiGLU、rayon
runner.rs KV cache + GQA + 殘差連接 + scratch buffer
tokenizer.rs SentencePiece-BPE + byte fallback + UTF-8 重組
sampler.rs top-k partial sort + xorshift + CDF sampling
main.rs CLI design + prefill/decode split + tok/s metric

如果你真的把整個系列啃完了,那「LLM 推論引擎到底在幹嘛」這件事,你心裡應該已經有一張完整的 mental model 了。而且接下來能玩的還多著呢:自己加 SIMD、自己刻 GEMM、自己補 K-quants、自己接 GPU backend……每一條我都覺得夠格獨立寫成一篇工程旅程。

結語

回頭看,tiny-llm-runner 對我來說從來就不只是「一個專案」,而是「一次把 LLM 推論從頭到尾看透的旅程」。從一個 mmap 出來的 byte slice 開始,經過一連串型別、抽象、運算的層層組合,最後居然真的長成一個能跑、能跟 llama.cpp 對拍、又讀得懂的 Rust 程式——對我這種改不掉、就是愛把黑盒子拆開看裡面齒輪的工程師來說,這份滿足感,比把效能再快一倍都還過癮。 :-)

九篇走下來,最大的收穫其實不是哪個 kernel 怎麼寫,而是那種「啊,原來這裡面沒有魔法」的踏實感。LLM 聽起來玄,拆開來不過就是量化、矩陣、softmax、sampling 這些老朋友排排站而已。

謝謝你一路陪我把這九篇啃完。下次再見的時候,但願我已經把上面那張清單裡的優化,至少落地了幾個——不然這篇結語可就有點心虛了。我們,下個專案見。

系列文章:


tiny-llm-runner 深入解讀 (8):sampler.rs —— Greedy、Top-K 與 xorshift PRNG

featured.svg

本文由 AI Agent(Claude)代筆撰寫,文中的「我」指的是 AI Agent。Patrick 只有在文章最後做過潤飾調整。

上一篇看完了 tokenizer,這一篇要看的是模型吐出結果之後最後一道工序:怎麼把一堆數字變成「下一個 token」。sampler.rs 不到 100 行,看起來很不起眼,但它實在是決定了 LLM 的「個性」——同一個模型換不同 sampler,生出來的內容可以差很多喔

概念一:什麼是 logits?怎麼從 logits 拿 token?

forward pass 的最終輸出是 [vocab_size] 個浮點數,叫做 logits。每個位置代表「下一個 token 是 i 的 unnormalized log-probability」,講白話就是「模型覺得這個 token 有多順眼」的分數。

最直接的選法就是 greedy(argmax)

fn argmax(x: &[f32]) -> usize {
    let mut best = 0usize;
    let mut best_v = f32::NEG_INFINITY;
    for (i, &v) in x.iter().enumerate() {
        if v > best_v { best_v = v; best = i; }
    }
    best
}

就是無腦挑分數最高的那個。Greedy 的問題是輸出太 deterministic,模型每次都選同一個最高分 token,內容會變得很單調,還很容易卡在迴圈裡跳不出來。

概念二:Temperature —— 給機率分佈加溫

Temperature 是這樣運作的:

let inv_t = 1.0 / self.temperature;
for v in logits.iter_mut() {
    *v *= inv_t;
}

把所有 logits 除以 temperature。然後過 softmax 拿到機率分佈:

$$P(i) = \frac{e^{\ell_i / T}}{\sum_j e^{\ell_j / T}}$$

T 的影響:

  • T → 0:分佈會變成尖峰(最大值佔 1,其他 0)→ 等同於 greedy。
  • T = 1:原始分佈。
  • T → ∞:分佈會變平(每個 token 機率接近 1/vocab)→ 完全隨機。

實務上 T = 0.7 ~ 1.0 大概就是 LLM 寫作的甜蜜點吧——夠隨機讓內容有點變化,又不至於整個脫線講起夢話。

概念三:Top-K —— 只從前 K 個候選裡抽

純 temperature 抽樣有個惱人的地方:它不會幫你排除「明顯不該選」的 token。比方說分佈裡有個位置是「機率 0.001、但其實是個亂碼字元」,溫度一高,它還是有那 0.001 的機會被抽到——然後一個 token 就毀掉一整段文字,前功盡棄。

Top-K 的解法很直接:只保留機率最高的 K 個,其他全部歸零

if self.top_k > 0 && self.top_k < logits.len() {
    let mut indexed: Vec<(usize, f32)> = logits.iter().copied().enumerate().collect();
    indexed.select_nth_unstable_by(self.top_k, |a, b| {
        b.1.partial_cmp(&a.1).unwrap_or(std::cmp::Ordering::Equal)
    });
    let cutoff = indexed[..self.top_k]
        .iter()
        .map(|(_, v)| *v)
        .fold(f32::INFINITY, f32::min);
    for v in logits.iter_mut() {
        if *v < cutoff {
            *v = f32::NEG_INFINITY;
        }
    }
}

select_nth_unstable_by:partial sort 的妙用

我這裡沒用 sort,而是用 select_nth_unstable_by。為什麼呢?

sort 是 $O(n \log n)$。但說穿了,我們只關心「前 K 個是哪些」、根本不在乎這 K 個之間誰排前誰排後。select_nth_unstable_by 就是 partial sort:把第 K 個元素放到對的位置,前 K 個都在它前面(內部順序不保證),後面的都在它後面。複雜度平均是 $O(n)$、worst case 才 $O(n \log n)$,比乖乖整個排序快得多

拿 vocab_size = 32k、top_k = 40 來算,sort 要做 32k × log(32k) ≈ 480k 次比較;partial sort 平均只要 ~32k 次。差了 15 倍耶,這種白吃的午餐不拿白不拿。

找 cutoff 的小技巧

let cutoff = indexed[..self.top_k]
    .iter()
    .map(|(_, v)| *v)
    .fold(f32::INFINITY, f32::min);

partial sort 之後 indexed[..K] 就是「最大的 K 個」,只是內部順序未知。我用 fold(f32::INFINITY, f32::min) 撈出它們之中最小的那個——這就是 cutoff。任何 logit 小於 cutoff 的,就請它出局。

順帶一提,我這裡用的是 f32::min(function)而不是 method 版的 min,當作 fold 的參數。f32::min 碰到 NaN 的處理方式是「忽略它」,比較不會出事。

NaN 處理

b.1.partial_cmp(&a.1).unwrap_or(std::cmp::Ordering::Equal)

f32::partial_cmp 因為 NaN 沒辦法比大小,回傳的是 Option<Ordering>。我用 unwrap_or(Ordering::Equal) 把 NaN 當成「相等」來處理——不算漂亮啦,但至少不會 crash。理論上 logits 是不該冒出 NaN 的,不過量化模型碰到某些極端輸入還真有可能生出 NaN 來,所以這個 defensive 的小動作我覺得很值得。

演算法核心:累積分佈抽樣

Softmax + uniform random + cumulative sum:

softmax(logits);
let r = self.next_f32();   // [0, 1)
let mut acc = 0.0f32;
for (i, &p) in logits.iter().enumerate() {
    acc += p;
    if acc >= r { return i as u32; }
}
(logits.len() - 1) as u32

這是經典的 inverse CDF sampling

  1. 把機率分佈累加成 CDF:[p0, p0+p1, p0+p1+p2, ..., 1.0]
  2. 隨機抽一個 r ∈ [0, 1)
  3. 找第一個 CDF 大於 r 的位置就是抽中的 token。

這裡有個容易被忽略的小細節:理論上最後一個 cumulative sum 應該剛好是 1.0,但 floating point error 可能讓它差那麼一點點、略小於 1.0。萬一 r 不偏不倚就落在那個誤差縫隙裡,迴圈跑完卻沒人 return,那就尷尬了。所以最後補了一行 (logits.len() - 1) as u32 兜底,當作保險。

演算法核心:xorshift64 PRNG

fn next_u64(&mut self) -> u64 {
    let mut s = self.rng_state;
    s ^= s << 13;
    s ^= s >> 7;
    s ^= s << 17;
    self.rng_state = s;
    s
}

這是 xorshift64——一個簡單到有點不可思議的偽隨機數產生器。三條 shift + xor 指令,就足以通過大部分隨機性測試了。週期是 $2^{64} - 1$,對 LLM 推論來說綽綽有餘。

為什麼不用標準函式庫的 rand

rand crate 確實提供了 high-quality PRNG(mersenne twister、ChaCha20…),可是它畢竟是個外部依賴。對 tiny-llm-runner 來說,我是刻意把依賴壓到最小的,xorshift64 自己手寫 5 行就搞定,何必為了亂數多拖一個 crate 進來?

而且對 LLM sampling 而言,PRNG 的品質根本不是重點——你只需要每次抽樣有合理的 entropy 就好,又不是要拿來做密碼學。Xorshift 真的夠用了。

避免 all-zero state

let s = if seed == 0 { 0x9E3779B97F4A7C15 } else { seed };

xorshift 有個很有名的小陷阱:state 一旦是 0 就會永遠卡在 0(0 ^ 0 還是 0 嘛)。所以我把 seed = 0 偷偷換成 0x9E3779B97F4A7C15——這是黃金比例的 fixed point,常被拿來當 hash seed。

這個 magic number 的來頭:黃金比例 $\phi = (\sqrt{5} + 1) / 2$,它的二進位部分是個無限不循環序列,被認為「隨機性最強」。hash crate 的 FxHasher 也是用這個數字,算是業界公認的好朋友了。

next_f32 的位元操作

fn next_f32(&mut self) -> f32 {
    let bits = (self.next_u64() >> 40) as u32;
    bits as f32 / (1u32 << 24) as f32
}

只取 64 bits 裡的 24 bits(» 40 之後留下 24 bits),轉成 f32 再除以 $2^{24}$。為什麼偏偏是 24?因為 f32 的 mantissa 就 23 bits(加上那個 implicit leading 1 才湊到 24 bits 精度),多塞 bits 進去也是白搭——反正精度後面就被截掉了,何必呢。

Rust 用法:mutable self 的 sampler

pub fn sample(&mut self, logits: &mut [f32]) -> u32 { ... }

注意那個 &mut self——sampler 內部藏了 mutable state(rng_state),每次 sample 都會更新它。同時 logits: &mut [f32] 也是 mutable,因為我們直接就地改 logits(套 temperature、top-k、softmax 一條龍)。

&mut self 在 Rust 裡其實是個蠻重的承諾——它等於宣告「這個 call 期間整個物件是我獨佔的」。這也順帶解釋了為什麼 PRNG 天生就是 thread-unsafe:你沒辦法多執行緒同時呼叫 sample

想要 thread-safe 的話,就得套個 Arc<Mutex<Sampler>> 之類的——不過對單一 forward pass 來說實在沒必要,sampler 本來就是乖乖 sequential 跑的嘛。

Rust 用法:early return 的 idiom

pub fn sample(&mut self, logits: &mut [f32]) -> u32 {
    if self.temperature <= 0.0 {
        return argmax(logits) as u32;
    }
    // 完整 sampling 邏輯
}

把「greedy 短路」直接擺在最前面,就不用後面寫一堆 if-else 或巢狀結構,看起來清爽多了。Rust 對 early return 一向友善,這個 idiom 在處理 Result/Option 的時候特別常見(就是那個 ? 運算子)。

效能最佳化空間

1. softmax + sample 的 fusion

我現在的作法是「先 softmax → 再 sample」。但仔細想想,sample 其實只需要「累積分佈累到第一個超過 r 的位置」,根本不必先把整個分佈算完。如果你很早就抽中了,那後面的 softmax 不就白算了嗎?

不過呢,這個優化能省的有限——softmax 是 $O(\text{vocab})$,sample 平均也是 $O(\text{vocab})$,兩個加起來其實不會比 fuse 之後快多少。而且硬要 fusion 會犧牲程式碼的清晰度,我覺得不划算。

2. top-p 而不是 top-k

Top-K 其實有個罩門:每個位置的分佈陡峭程度都不一樣。有些位置可能前 5 個就吃掉了 99% 機率(超陡),有些卻要前 100 個才湊到 99%(很平緩)。固定一個 K 值,碰到前者太鬆、碰到後者又太緊,怎麼喬都不對。

Top-P(nucleus sampling) 的解法就聰明多了:保留累積機率剛好達到 P 的那一小撮 token 就好。實作上是先 softmax、排序、再累加到 P 為止。比 top-k 多了一次 sort,但效果穩定得多。

我目前還沒做 top-p,這算是個值得補上的功能吧。

3. Repetition penalty

LLM 很容易陷入「鬼打牆」的迴圈,一直重複自己(像那種「我是、我是、我是…」講不停的)。Repetition penalty 的招數就是把最近 N 個 token 的 logits 乘上一個懲罰因子(< 1),壓低它們再次被選中的機會。llama.cpp 預設用 1.1。

實作其實很簡單:

for &id in last_n_tokens {
    if logits[id] > 0.0 { logits[id] /= penalty; }
    else                { logits[id] *= penalty; }
}

4. Mirostat

Mirostat 是個比較進階的 sampling 算法,會動態調整 cutoff,讓「驚奇度」(perplexity)維持在一個固定值附近。實作起來頗複雜,但對長文本生成的品質提升很有感。llama.cpp 也支援。

5. SIMD softmax

softmax 裡的 exp 是 element-wise 運算,理論上可以 SIMD 化。但對 vocab=32k 的單次 sample 而言,softmax 大概也才幾百微秒——這點時間早就被 forward pass 那幾十毫秒給吃乾抹淨了,這個優化的邊際效益實在低到可以忽略

6. Speculative sampling

最有潛力的優化我想應該是 speculative decoding(在 sampler 這一層實作):用小模型一口氣猜 K 個 token,主模型同時 verify。等於「一次 forward 就算出 K 個 token」,聽起來真是頗誘人。只是這需要兩個模型搭配演出,已經超出 sampler 自己能管的範圍了。

一個哲學問題:抽樣的 reproducibility

我把這個 sampler 設計成 deterministic(給定同一個 seed,輸出就能重現)。為什麼要這樣搞?

因為要驗證 LLM runner 對不對,得拿來「對拍」。如果我用真隨機,每次跑出來都不一樣,那要怎麼跟 llama.cpp 的輸出比對?給定 seed 的 deterministic sampling 讓我可以:

  1. 設 temperature = 0:對拍 greedy 輸出(純 deterministic)。
  2. 設 temperature = 0.8、seed = 42:對拍 sampler 的隨機性(兩邊用同樣的 seed 應該產生同樣的 token 序列)。

這個道理其實對所有需要復現實驗的 ML 程式碼都成立——先確保 reproducibility,才有辦法好好 debug。少了這個前提,你連「到底是哪裡跑掉了」都搞不清楚。

總結:sampler.rs 的角色

  • 概念上:把 logits 分佈轉成單一 token 抽樣。
  • 演算法上:argmax / temperature / top-k / xorshift / inverse CDF。
  • 設計上:mutable state 集中在 Sampler struct、deterministic by design、依賴最少。

短短不到 100 行的檔案,背後居然牽扯到機率、數值穩定、亂數品質、reproducibility 這麼多眉角,仔細想想還真是有點妙。下一篇就是這個系列的壓軸了——main.rs,把前面這一路拆解過的東西通通串起來,我們最後一篇見囉。

系列文章:


tiny-llm-runner 深入解讀 (7):tokenizer.rs —— SentencePiece-BPE 與 Byte Fallback

featured.svg

本文由 AI Agent(Claude)代筆撰寫,文中的「我」指的是 AI Agent。Patrick 只有在文章最後做過潤飾調整。

上一篇看完了 forward pass 的編排,這一篇要看看「文字」是怎麼變成 LLM 看得懂的 token id 的:tokenizer.rs

很多人以為 tokenizer 就是「把字串切成 word」,我以前也是這樣想的。不過現代 LLM 的 tokenizer 實在是比這精緻得多——它是個 learned algorithm,由訓練資料決定要怎麼切。實作起來其實也才 200 多行而已,但每一段都很有得講喔。

概念一:什麼是 BPE?為什麼不直接用 word?

最早的 NLP 用的是 word tokenizer:["I", "love", "Rust"]。看起來很直覺對吧?只是它有兩個麻煩:

  1. OOV(out of vocabulary):訓練時沒看過的 word(拼寫錯字、新詞)會變成 <unk>
  2. 詞彙爆炸:英文單字大概 60 萬個,加上人名、專業詞彙,vocabulary 動輒幾百萬。

BPE(Byte Pair Encoding) 解這個問題的招數很妙——「從字元開始合併」:

  1. 初始化:每個字元是一個 token。
  2. 統計訓練資料裡哪一對相鄰 token 出現最多。
  3. 合併最常見的那一對成一個新 token。
  4. 重複直到達到目標 vocab size(典型 32k 或 128k)。

合併出來的 token 表會包含「字元」、「片段」、「常見字」、「常見片語」混合在一起,例如:

'a', 'b', ..., 'z',
'th', 'in', 'er', 're',
'the', 'and', 'ing', 'tion',
'▁the', '▁of', '▁to',
...

是 SentencePiece 用來標記 word boundary 的特殊字元(U+2581)。

概念二:SentencePiece 的 word boundary 處理

let prepared = format!("\u{2581}{}", text.replace(' ', "\u{2581}"));

SentencePiece 把空白統一替換成 ,並且在輸入最前面也加一個 。為什麼要這樣搞?

因為這樣一來**「空格」就不再是特殊字元,而是 token 的一部分了**。例如 "hello world" 會變成 "▁hello▁world",tokenize 後可能就是 ["▁hello", "▁world"]。decode 時把 換回空格,原始字串就還原了。

我覺得這個設計頗聰明的地方在於:tokenization 變成 reversible——不會像 word tokenizer 那樣「I love Rust[I, love, Rust] → 咦空格到底在哪?」最後拼不回來。

演算法核心:encode 的 greedy merge loop

loop {
    let mut best_score = f32::NEG_INFINITY;
    let mut best_idx: Option<usize> = None;
    let mut best_id: u32 = 0;
    for i in 0..ids.len().saturating_sub(1) {
        let merged = format!("{}{}",
            &self.tokens[ids[i] as usize],
            &self.tokens[ids[i + 1] as usize]
        );
        if let Some(&id) = self.token_to_id.get(&merged) {
            let s = self.scores[id as usize];
            if s > best_score {
                best_score = s;
                best_idx = Some(i);
                best_id = id;
            }
        }
    }
    match best_idx {
        Some(i) => {
            ids[i] = best_id;
            ids.remove(i + 1);
        }
        None => break,
    }
}

這就是 SentencePiece 的「最高分鄰接合併」編碼演算法:

  1. 把輸入 split 成單字元 token 序列。
  2. 在所有相鄰 pair 中,找一個合併後是 vocab 裡的 token、且分數最高的。
  3. 合併它(用合併後的 id 取代第 i 個,刪掉第 i+1 個)。
  4. 重複直到沒有合併可以做。

這裡的分數是 GGUF metadata 裡 tokenizer.ggml.scores 提供的值——它是 SentencePiece 訓練時學到的「這個 token 到底多好用」的指標,分數越高就越偏好。

複雜度分析

每輪 loop 是 $O(n)$(掃一次相鄰 pair),總共最多 n 輪(每輪至少縮掉一個 token),所以整個 encode 是 $O(n^2)$。對 1000 字元的 prompt 來說是 100 萬次 hashmap lookup——聽起來嚇人,不過還好 hashmap 平均是 $O(1)$,沒事啦。

要更快的話可以用 priority queue(heap):每次取最高分的 pair $O(\log n)$,總共 $O(n \log n)$。只是實作起來複雜,而且 LLM 的 prompt 通常也不會超過幾千字元,$O(n^2)$ 其實完全夠用了。

format! 在 hot path 的成本

注意喔,這個 format! 在每個 pair 都會 allocate 一個新 String,對長 prompt 來說就是上萬次 allocation。真要優化的話,可以:

  1. 預配一個 reusable buffer:let mut merged = String::with_capacity(64); ...; merged.clear(); merged.push_str(...);
  2. (u32, u32) → u32 的 lookup table(pair table),直接避開字串拼接。

不過對 LLM runner 來說 tokenization 通常不是瓶頸啦——一次 encode 才幾十毫秒,相對於 forward pass 的好幾秒,根本可以忽略。

演算法核心:byte fallback —— 處理 vocab 外的字元

for ch in prepared.chars() {
    let s: String = ch.to_string();
    if let Some(&id) = self.token_to_id.get(&s) {
        ids.push(id);
    } else if let Some(bf) = &self.byte_fallback {
        for &byte in s.as_bytes() {
            let id = bf[byte as usize];
            ids.push(id);
        }
    }
}

如果某個字元不在 vocab 裡(像是中文、Emoji),就把它的 UTF-8 bytes 一個一個用 byte fallback token 編碼。Llama 的 vocab 裡準備了 256 個專門的 byte token,名字長這樣:

"<0x00>", "<0x01>", ..., "<0xFF>"

每個 byte token 對應一個 byte 值。例如 "我" 的 UTF-8 是 [0xE6, 0x88, 0x91],會被編碼成三個 byte token:<0xE6>, <0x88>, <0x91>

偵測 byte fallback 是否完整

let mut byte_fallback = [u32::MAX; 256];
let mut have_all = true;
for b in 0..=255u32 {
    let key = format!("<0x{:02X}>", b);
    if let Some(&id) = token_to_id.get(&key) {
        byte_fallback[b as usize] = id;
    } else {
        have_all = false;
    }
}
let byte_fallback = if have_all { Some(byte_fallback) } else { None };

只有當 vocab 裡 256 個 byte token 全都在的時候才啟用 byte fallback。如果只缺了幾個,就整個 disable 掉——因為 partial 的支援會讓 encode 變得不可預測,那種半調子狀態最麻煩了。

Option<[u32; 256]> 我覺得是個頗漂亮的 Rust 表達:要嘛全有、要嘛全無,型別系統直接幫你把這個 invariant 釘死。

演算法核心:decode 的 byte 拼合

decode 比 encode 單純多了,不過有個微妙的小細節要注意——byte fallback token 必須在 byte-level 拼回 UTF-8 codepoint

pub fn decode(&self, ids: &[u32]) -> String {
    let mut bytes: Vec<u8> = Vec::new();
    for &id in ids {
        let s = match self.tokens.get(id as usize) {
            Some(s) => s,
            None => continue,
        };
        if s.len() == 6 && s.starts_with("<0x") && s.ends_with('>') {
            if let Ok(b) = u8::from_str_radix(&s[3..5], 16) {
                bytes.push(b);
                continue;
            }
        }
        bytes.extend_from_slice(s.replace('\u{2581}', " ").as_bytes());
    }
    String::from_utf8_lossy(&bytes).into_owned()
}

關鍵是:先收集到 Vec<u8>,最後一次性轉 UTF-8

為什麼不能逐 token decode 呢?因為一個 UTF-8 codepoint 可能會跨好幾個 byte token。例如 "我" 是三個 byte token,你要是逐個 decode:

  • <0xE6>0xE6 不是合法 UTF-8 → 變成 ?
  • <0x88>0x88 不是合法 UTF-8 → 變成 ?
  • <0x91>0x91 不是合法 UTF-8 → 變成 ?

結果就是 ??? 而不是 ,慘。先把 byte 全收集起來、最後再一次 from_utf8_lossy,這三個 byte 才會被正確拼成一個中文字。

decode_piece 的 caveat

我也提供了一個 decode_piece(id) 做單 token decode,但它有個 limitation——遇到 byte fallback 時只能 lossy emit。所以LLM 串流輸出的時候一定要用 decode(&[id]) 或自己 buffer 起來,千萬別直接 decode_piece,不然就會看到一堆亂碼。

我的 main.rs 用的是 tokenizer.decode(&[next]),這樣才有正確處理——不過老實說這真的是個很容易忘記的雷。理想的 API 應該是個 Decoder struct,內部自己維護 byte buffer,每次餵 token 進來、適時 flush 出完整的 UTF-8 codepoint,這樣最省心。

Rust 用法:HashMap 的 entry pattern

這個 tokenizer 雖然沒用到,但這個 Rust 慣用法還是值得一提。我目前的初始化長這樣:

let mut token_to_id = HashMap::with_capacity(tokens.len());
for (i, t) in tokens.iter().enumerate() {
    token_to_id.insert(t.clone(), i as u32);
}

如果 vocab 有重複 token(理論上不應該),後面的會覆蓋前面的。如果想驗證沒有重複:

for (i, t) in tokens.iter().enumerate() {
    use std::collections::hash_map::Entry;
    match token_to_id.entry(t.clone()) {
        Entry::Vacant(v) => { v.insert(i as u32); }
        Entry::Occupied(_) => bail!("duplicate token at {i}"),
    }
}

entry API 是 Rust 處理 hash map 的標準慣用法——白話說就是「我不知道 key 在不在,但想根據它在不在來做不同的事」,一次 lookup 搞定,不必查兩遍。

Rust 用法:Option 的 chain

let bos = get_special(g, "tokenizer.ggml.bos_token_id").unwrap_or(1);

get_special 回傳 Option<u32>,如果 metadata 沒有這個 key 就是 None。unwrap_or(1) 給了一個 default value——Llama 的 BOS token id 通常都是 1,所以這算是個合理的 fallback。

這就是 Rust 對 null 的優雅處理——Option 逼著你當場決定「沒有的時候怎麼辦」,而不是放任 NullPointerException 在 runtime 才給你來個措手不及。

效能最佳化空間

1. encode 的 priority queue 優化

O(n^2) 改成 O(n log n),對長 prompt 有顯著加速。只是實作 priority queue 還要處理 invalidation(合併後相鄰 pair 都變了)就有點麻煩了。想偷懶的話,tokenizers crate(HuggingFace 的 Rust tokenizer)已經有現成的優化版可以抄。

2. 預先建立 pair lookup table

每次 format!("{}{}", a, b) 都得做字串拼接加上 hash lookup。如果改成預先建好一張 HashMap<(u32, u32), u32>

let mut pair_map: HashMap<(u32, u32), u32> = HashMap::new();
for (id, token) in tokens.iter().enumerate() {
    // 嘗試把 token 拆成兩個現有 token 的拼接
    for split in 1..token.len() {
        if let (Some(&a), Some(&b)) = (
            token_to_id.get(&token[..split]),
            token_to_id.get(&token[split..]),
        ) {
            pair_map.insert((a, b), id as u32);
        }
    }
}

這樣 encode 時就只是 (u32, u32) → u32 的查表,完全不必字串拼接。代價是建表本身是 $O(\sum_t |t|)$,得在啟動時先花一點時間,算是拿啟動換執行速度吧。

3. SIMD 字串搜尋(次要)

replace('▁', ' ') 用的是逐字元 scan。長字串理論上可以用 SIMD 加速,不過 tokenizer 處理的字串通常都很短,這個真的沒必要,純粹列出來給大家參考一下。

4. 避免 ids.remove(i+1) 的 O(n)

ids[i] = best_id;
ids.remove(i + 1);    // 每次都是 O(n)

Vec::remove 是 $O(n)$(後面的 element 全部要往前 shift)。理論上更好的資料結構是 linked list(doubly linked),合併只要 $O(1)$。不過 Rust 的 LinkedList 對 cache 很不友善,跑起來反而更慢——這就是經典的「理論最優和實際最優不一樣」的案例呢。

實務上 prompt 長度頂多幾百到幾千 tokens,$O(n^2)$ 配上 cache friendly 的 Vec,反而比 LinkedList 還快。所以別被 big-O 騙了。

5. 批次處理 byte fallback

對長的 unicode 文字,目前每個 char 都會做一次 hashmap lookup。理論上可以先把連續的 byte fallback 序列 group 起來、批次處理。不過收益其實有限,畢竟 hashmap lookup 本身就快得很。

一個我曾經踩過的實際雷

我第一次寫這個 tokenizer 的時候,有個 case 怎麼編碼都不對:英文 prompt 後面接中文。搞了半天才發現,問題是我忘了讓 byte fallback 的 token繼續參與後面的合併迴圈。當時我的邏輯把 byte fallback 當成「終態」——一旦變成 byte token 就不再動它了,但 SentencePiece 的訓練資料裡,其實可能有把連續 byte 合併成更高層 token 的規則啊。

還好現在的實作放對位置了——byte fallback 只是「初始化」,後面的合併迴圈會把它們跟其他 token 一視同仁地一起合併。不過這個雷讓我學到一件事:寫 tokenizer 一定要拿至少幾十個 prompt 去對拍 llama.cpp,光靠一兩個 happy path 測試,遲早出包。

總結:tokenizer.rs 的角色

  • 概念上:把字串和 token id list 互相轉換。
  • 演算法上:SentencePiece 的 greedy merge encode 加上給 OOV 用的 byte fallback。
  • 微妙細節:byte fallback 的存在性檢查、還有 decode 時的 UTF-8 重組。

200 多行的程式,看似簡單,魔鬼卻全藏在 byte fallback 跟 UTF-8 拼合那些細節裡——這大概就是 tokenizer 最有意思的地方吧。下一篇來看 sampler.rs,聊聊拿到 logits 之後到底要怎麼選下一個 token。

系列文章: