基于 KuiperLLama 二次开发的大模型推理框架。原项目是一个全手写 CUDA 算子的 Llama/Qwen 推理框架(支持 Llama2/3、Qwen2.5、Qwen3 及 Int8 量化),InferLLama 在其基础上主要做了 GEMM kernel 的性能优化 与 可复现的基准测量。
- 优化 fp32 批量 GEMM kernel(prefill):把"每 block 只算 1 个输出元素"的旧 kernel 重写为 2D tile GEMM(每 block 算 64×64 输出子块,共享内存 stage + 寄存器分块 + float4 合并读),微基准 111.6 ms → 4.36 ms(约 25.6×),真实推理长 prompt prefill 3.53 s → 0.50 s(约 7×);
- 优化 int8 W8A8 tensor core kernel:B 片段由逐字节 load + shift 拼装改为一次
uint32直读(M%4==0对齐、小端与 pack 布局逐位一致),微基准 3.42 ms → 1.50 ms(约 2.28×),输出逐位不变; - 新增
bench/gemm_bench:直接链接libllama.so、通过公开接口调用真实 kernel 的 GEMM 微基准 driver,支持 fp32(--fp32)与 int8 W8A8 双路径,并与 fp64 CPU 参考 / 逐位基线校验。
InferLLama
├── kuiper/
│ ├── include/ # 头文件(base / model / op / sampler / tensor)
│ └── source/
│ ├── base/ # 基础库:配置、日志、CUDA 流等
│ ├── tensor/ # 张量 / 缓存 / Buffer
│ ├── model/ # 模型层:llama3 / qwen2 / qwen3 的层组织与 forward 调度
│ ├── op/ # 算子层:matmul / mha / rmsnorm / swiglu / embedding / encode …
│ │ └── kernels/ # kernel 层:CPU + CUDA 双后端
│ │ ├── cuda/ # CUDA:matmul、mha、rmsnorm、rope、swiglu、emb、argmax、add …
│ │ └── cpu/
│ └── sampler/ # 采样(top-k / top-p 等)
├── demo/ # 推理入口:llama_infer / qwen_infer / qwen3_infer / ppl_qwen3
├── bench/ # GEMM 微基准 driver(gemm_bench)
└── tools/ # 模型导出脚本(export_llama3 / export_qwen2 / export_qwen3 …)
简要分层:
- 模型层(
model/)解析配置、加载权重、逐层调用算子 forward,支持 Llama2/3、Qwen2.5、Qwen3; - 算子层(
op/)封装张量级计算,每个算子同时有 CPU 与 CUDA 实现;其中 prefill(批量,输入[N][hidden])与 decode(单 token GEMV)分别走不同 kernel; - kernel 层(
op/kernels/cuda/)手写 CUDA kernel:- prefill fp32:2D tile GEMM —— 共享内存 stage(A 用 float4 广播读、W 用 stride-17 防 bank 冲突)+ 每线程 4×4 寄存器分块;
- prefill int8:W8A8 走 tensor core
mma.sync.aligned.m16n8k32.row.col.s32.s8.s8.s32; - decode:内存带宽受限的 SIMT GEMV(fp32 直接乘,int8 为 W8A16 现场反量化)。
- 量化方案:weight-only —— 权重存 int8、每 64 个连续权重共享一个 fp32 scale(对称、无 zero-point);prefill 时激活动态按 tile 量化为 int8(W8A8),decode 时激活保持 fp32。
同配置实测:Qwen3-0.6B,长 prompt(1734 tokens,seq_len=2048),Nvidia RTX 3090(sm_86),CUDA 12.9,同 driver 同 prompt。
| 指标 | fp32 | int8 W8A8 | 加速比 |
|---|---|---|---|
| prefill(1734 tokens) | 0.497 s → 3490 tok/s | 0.406 s → 4274 tok/s | 1.22× |
| decode(314 steps) | 2.69 s → 116.6 tok/s | 2.28 s → 137.9 tok/s | 1.18× |
| 端到端(2048 tokens) | 4.17 s → 491 steps/s | 3.71 s → 552 steps/s | 1.13× |
说明:
- prefill 的 int8 优势(22%)小于纯 GEMM 的差距(微基准约 2.9×),因为推理中 attention / RMSNorm / SwiGLU / RoPE 等非 matmul 算子在两条路径上都是 fp32、且现在占了相当比重;
- decode 是 GEMV、原理上内存带宽受限(int8 权重字节少 4×),但当前 SIMT GEMV kernel 未达带宽饱和,实测提升约 18%;
- int8 的额外收益:权重显存由 3.0 GB → 1.25 GB;
- 这里的
steps/s与tok/s同义(一个 step 即一个 token 位置);decode行是纯生成速度,端到端行是 prefill+decode 的全程平均(含少量编码/采样开销)。
GEMM 微基准(bench/gemm_bench,形状 N=1472, M=4096, K=4096):
| 路径 | 时间 | 算力 |
|---|---|---|
| fp32(2D tile GEMM) | 4.36 ms | 11.3 TFLOPS |
| int8 mma(W8A8 tensor core) | 1.50 ms | 32.9 TFLOPS |
# 仅构建 libllama.so
cmake --build build --target llama -j$(nproc)
# 全量构建(含 demo 推理程序,demo 依赖 LLAMA3/QWEN2/QWEN3_SUPPORT 编译开关)
cmake --build build -j$(nproc)export LD_LIBRARY_PATH=lib:/home/kitty/deps/prefix/lib
./build/demo/qwen3_infer /home/kitty/deps/qwen0.6.bin \
/home/kitty/huggingface/Qwen3-1.7B/tokenizer.json 0 /tmp/long_prompt.txt
# 参数依次为:checkpoint tokenizer 是否量化(0=fp32,1=int8) [prompt 文件]./build/demo/qwen3_infer /home/kitty/deps/qwen0.6_qint8.bin \
/home/kitty/huggingface/Qwen3-1.7B/tokenizer.json 1 /tmp/long_prompt.txt运行结束会打印 prefill: prompt_tokens … time … 与 decode: steps … steps/s … 两行,即上表的对比数据(prompt 需满足 token 数 ≤ checkpoint 中的 seq_len,本例为 2048)。
# 编译 driver(链接 libllama.so)
nvcc -O3 -std=c++17 -arch=sm_86 \
-I kuiper/include -I kuiper/source -I /usr/local/cuda/include \
bench/gemm_bench.cu \
-L lib -lllama -L /usr/local/cuda/lib64 -lcudart -L /home/kitty/deps/prefix/lib \
-lglog -larmadillo -lpthread -ldl -o build/gemm_bench
# int8 W8A8(mma tensor core)
./build/gemm_bench -n 1472 -m 4096 -k 4096 -i 30
# fp32(2D tile GEMM)
./build/gemm_bench --fp32 -n 1472 -m 4096 -k 4096 -i 30
# int8 首次运行可加 --save 保存基线,之后自动做优化前后逐位对比- google glog / google gtest
- sentencepiece(Qwen3 也可直接用 HF 的
tokenizer.json) - armadillo + openblas
- CUDA Toolkit(本机测试环境:RTX 3090 / sm_86 / CUDA 12.9)
# Llama3.x
python3 tools/export.py Llama-3.2-1B.bin --hf=meta-llama/Llama-3.2-1B
# Qwen2.5
python3 tools/export_qwen2.py Qwen2.5-0.5B.bin --hf=Qwen/Qwen2.5-0.5B
# Qwen3:先用 tools/export_qwen3/load.py 把 HF 模型导出为 pth,
# 再用同目录 write_bin.py 导出 .bin;量化权重(int8+scale)可用 tools/export_qwen2.py 的量化导出路径。./build/demo/llama_infer Llama-3.2-1B.bin meta-llama/Llama-3.2-1B/tokenizer.json
./build/demo/qwen_infer Qwen2.5-0.5B.bin Qwen/Qwen2.5-0.5B/tokenizer.json
./build/demo/qwen3_infer qwen0.6.bin <tokenizer.json> 0 # 0=fp32,1=int8