先说结论:把一个三元量化模型跑在 vLLM 上,最大的障碍往往不在 kernel,而在于让 vLLM 的加载链路肯认这个「非主流」格式。我这次拿到的是一个权重只有 {-1, 0, 1} 三种取值的 embedding 模型,目标是把它塞进 vLLM 做服务化推理。动手之前我以为写个自定义 CUDA kernel 就完事,后来才发现,从逆向 vLLM 的量化抽象、对齐磁盘权重与 GPU 权重布局,再到 kernel 从「跑出数」到「跑得快」的三轮演进,每一步都有意料之外的坑。
这篇就是当时完整的落地记录,适合正在做自定义量化接入 vLLM、或者想改造推理框架底层加载逻辑的朋友。里面不会只给结论,还会把我实际推过的代码路径、对账脚本的思路、kernel 优化的取舍都摊开讲,包括那些改到凌晨才发现的字节序问题和切分下标问题。看完你应该能少走一半弯路。
1. 为什么非要把三元量化模型塞进 vLLM
1.1 三元量化到底省了什么
三元量化(Ternary Quantization)属于极端量化的一种:权重不再用 FP16/BF16,也不像 INT8 那样落在 256 个取值上,而是直接圈到三个值:-1、0、1。每个权重只需要 2 bit 就能表示,理论上相比 FP16 权重可以省下 8 倍存储,相比 INT8 省 4 倍。
省的不只是内存。推理时的矩阵乘法,权重与激活的乘法本质上是符号选择:值为 1 就原样加,值为 -1 就减,值为 0 直接跳过。这比任何浮点乘法都快。代价也很明显,精度损失大,尤其是那些本来权重分布很密集的层,强行量化到三值会掉点严重。所以工业界常用的做法是加分组缩放因子,比如每 128 个权重共享一个 FP16 的 scale,用 scale 去弥补动态范围的丢失,这就是「三元量化 + 分组 scale」这种形态的由来。
但这类模型在社区和实际业务里并不少,尤其是内存敏感的边缘部署、向量检索前的 embedding 模型、以及一些对召回精度没那么苛刻但吞吐要求极高的场景。我这次手上的模型就是一个 0.6B 左右的 embedding 模型,权重全部打包成 2bit 格式,还带 per-block scale,推理时先反量化再喂给普通算子是可行,但那样等于放弃了压缩带来的内存收益。
1.2 最省事的路走不通,才决定硬啃 vLLM
有人可能会说,既然 vLLM 不支持,那就先用 PyTorch 写个脚本临时顶上嘛。我试过,问题很现实:embedding 模型最值钱的就是高并发批处理能力,单条 query 的向量化如果走原生 PyTorch,batch 一大,CPU 侧的 GIL 和 GPU 的 kernel 调度开销立刻变成瓶颈。vLLM 的 continuous batching、PagedAttention、以及现成的 OpenAI 兼容 API 都是我需要的。
但这意味着我要正面解决三件事:
- vLLM 根本不认识「ternary」这种 quant_method,加载权重时直接报错;
- 磁盘上打包好的 2bit 权重,必须经过对齐和切分后才能按 vLLM 的预期落到各张卡的显存里;
- 就算加载进去了,还得有一个能真正发挥三元量化算力优势的 kernel,否则性能可能比 FP16 还慢。
这三个问题分别对应逆向、对账、kernel 优化,也是这篇记录的三个主线。
2. 逆向:从 config 一路摸到 vLLM 的量化加载链路
2.1 先找到 vLLM 识别量化模型的入口
做逆向最忌讳上来就翻源码大海捞针。我的做法是拿一个主流的量化格式当参照物。vLLM 内部对 AWQ、GPTQ、FP8 这些量化方案的支持已经挺成熟,它们的实现骨架大差不差:一个QuantConfig子类负责描述量化参数;一个LinearMethod子类负责权重创建、加载、前向计算;模型加载器根据config.json里的quantization_config/quant_method字段决定把这个模型分发到哪个量化实现上。
我第一步是搞清楚自己手上的模型 config.json 写了什么。打开一看,果然quantization_config里挂着一个没人认识的quant_method: "ternary",vLLM 会直接抛 KeyError。这就好办了,逻辑上我需要做的就是在 vLLM 的量化注册表里加一个TernaryQuantConfig,并实现配套的LinearMethod。
顺着这个思路,我逆向的起点就圈定在三个位置:
- 量化方法的注册表:
SUPPORTED_QUANT_METHODS之类的地方; - 抽象基类的接口定义:
QuantizeMethodBase/LinearMethodBase; - 以及一个已经实现好的量化方案,比如 GPTQ,把它的
create_weights、process_weights_after_loading、apply、weight_loader四个方法逐一看透。
2.2 顺着加载链路追到 kernel 调用点
vLLM 的模型加载顺序大致是:LLMEngine初始化 ->ModelRunner->ModelLoader,由ModelLoader遍历模型 state_dict 的每个 tensor,根据层的类型分发给对应的weight_loader。对于普通 Linear 层,vLLM 按ColumnParallelLinear/RowParallelLinear的切分规则,把全量权重拷到参数里。对于量化层,weight_loader会被替换成量化方法自己实现的版本。
这里有个特别容易踩坑的点:默认的weight_loader假设权重 shape 跟原始线性层的[out_features, in_features]完全一致。但三元量化模型在 safetensors 里存的权重字段往往叫packed_ternary_weight,shape 是[out_features, in_features / 4],因为四个 2bit 权重被打包进一个 byte。shape 对不上,默认 loader 根本不知道怎么切分,更不要说逐块切 scale。所以必须完全接管weight_loader。
我还注意到,vLLM 对 QKV 这类融合权重有专门的shard_loader,因为同一个 tensor 在加载时要按多段 offset 切到同一个qkv_proj参数里。三元量化模型如果也做了 QKV 融合,packed 权重和 scale 的切分逻辑要额外小心,否则就会出现「权重切对了,scale 还是全量」这种隐蔽 bug。
2.3 逆向时真正好用的几个工具
静态读源码当然要读,但我强烈建议边跑边看。我当时就是写了一个最小启动脚本,只加载一个同结构的 FP16 模型,打印每个参数的 name、shape、以及 load 时的 dispatch 路径,然后把三元量化模型的 safetensors 字段列表拿来对齐。具体工具就三个:pdb打断点、print打关键变量、以及git grep在源码里全局搜符号。vLLM 抽象层级重,类名绕来绕去,纯靠肉眼看很容易晕,打断点看实际走到的分支反而最直接。
另外一个心得:一定要把逆向结论沉淀成文档。比如「哪个字段是 packed 权重、哪些层是普通 int8、哪些层是 FP16 阈值」这种信息,如果不写成一份quant_format.md,第二天你准忘。
3. 对账:磁盘布局、GPU 布局与张量并行切分
3.1 先把 packing 规则钉死
对账的前提是有一份精确到 bit 的 packing 规则。我这次模型的规则是:
- 每个权重映射为 2bit:
00 -> 0,01 -> 1,10 -> -1,11 -> 保留位按 0 处理; - 每 4 个权重按小端顺序凑成一个
uint8,第一个权重占最低 2 位; - 每连续 128 个权重共享一个 FP16 scale,block 边界按 K 维连续切。
这些规则看起来很简单,但如果文档没写清楚,或者实现时有人改了某个映射,后面所有对账都是白做。所以第一步不是写脚本,而是把这份规则用文字写到代码注释里,再写一个解包函数,让它只做一件事:从 packed bytes 还原出[-1, 0, 1]的整数矩阵。
import torch def unpack_ternary(packed: torch.Tensor) -> torch.Tensor: # packed shape: [N, K // 4], uint8 p = packed.contiguous().view(torch.uint8) # 取每个 byte 的 4 个 2bit 窗口 codes = torch.stack( [(p >> shift) & 0x03 for shift in (0, 2, 4, 6)], dim=-1 ) # 00 -> 0, 01 -> 1, 10 -> -1, 11 -> 0 values = torch.where(codes == 1, 1, torch.where(codes == 2, -1, 0)) return values.reshape(-1)这只是示意,实际工程里还要处理 shape 还原、内存连续、以及在torch.compile下的 trace 兼容。但核心思路是:先把解包逻辑和参考实现做成一对可交叉验证的函数。
3.2 三个回合的对账
我做了三层对账,每一层都是独立的校验,任何一层挂了都必须查清再往前走。
第一回合是 CPU 字节级对账。把 safetensors 里 load 出来的 packed tensor 解包,和一份从全量权重离线重建的{-1,0,1}参考矩阵做整数精确比对。注意这里不要用allclose,因为量化对齐必须是精确的;用torch.equal或者算 SHA256。如果对不上,就逐字节打印前 64 字节,看看是哪个 bit 窗口不对。常见的错法包括:大端小端反了、pack 顺序是先列后行、11被映射成别的值。
第二回合是 GPU 张量对账。确认 CPU 解包没问题后,把 packed 权重真正 load 到 GPU,用torch.equal对比 GPU 上的packedtensor 和 CPU 侧原始字节。这一步主要是防加载路径中的to(device)或view(dtype)操作破坏了数据布局。很多自定义 loader 在copy_时容易忽略 dtype,导致 byte 被按 float 解释,结果就是 GPU 上拿到一堆垃圾数。
第三回合是张量并行切分对账。vLLM 在多卡推理时,每个线性层会按张量并行规则切成多个 shard,分别住在不同卡上。ColumnParallelLinear按输出维度切,RowParallelLinear按输入维度切。问题来了:packed 权重是按行方向连续打包的,一行是 K 个权重,如果 K 不是 128 的整数倍,block scale 的边界和 shard 切分边界就会错位。我这次的模型权重 K 基本都是 1024、2048、4096 这种 2 的幂,所以 block 边界天然对齐,真是个好消息;但如果遇到反例,就必须在切分前先把 packed 权重解包成整数矩阵,按原始 shape 切好,再重新打包,而且 scale 也要同步切。
3.3 对账脚本的工程建议
不要用临时脚本,建议直接把三层校验写成一个validate_quant_weights.py,每次改完 loader 或者 kernel 都跑一遍。脚本输出三行状态:CPU_UNPACK: OK、GPU_LOAD: OK、TP_SHARD: OK。一旦日志变红,马上知道是哪一层出了问题。
另外,对账脚本一定要能覆盖不同 dtype 的字段对比。packed 是uint8,scale 是float16,零点是float32甚至int32,比对时要分别处理,不能一把梭转 float。
4. Kernel 优化:把 -1/0/1 变成真正的算力优势
4.1 先想清楚三条路线
拿到一个自定义量化模型后,最自然的思路是「先反量化成 FP16,然后调用 cuBLAS」。这条路实现起来最快,在 vLLM 里也最容易集成,因为它复用了所有现成的 FP16 kernel。但性能上很亏:反量化器要么在加载时把权重转成 FP16 放显存,那内存压缩收益全没了;要么在前向时逐 block 反量化,反而比原生 FP16 还多一次遍历。
第二条路线是写一个自定义推理 kernel,在 kernel 内部完成 decode + 乘加。权重保持 packed 状态,scale 单独传,kernel 读取原始 2bit 数据,在寄存器里解出 -1/0/1 参与计算。这条路能保住内存收益,也能避免载入 FP16 权重的大带宽开销。
第三条路线更激进:利用三元权重的符号特征,设计无乘法的累加路径。比如把 -1/0/1 拆成两个 bitmask,用整数加法实现累积,只有在 block 边界才乘 scale。这条路线收益最大,但代码复杂度也最高,尤其要处理 scale 的边界累加。
我最终选的是第三条路线的务实版:block 内用 int 累加符号,block 边界乘 scale。原因很直接,这样的 kernel 既保留了符号运算的速度,又不需要做复杂的 mask 分解,代码可维护性更好。
4.2 核心 kernel 的设计要点
给一个只展示思路的 CUDA 片段。假设我们处理的是形如y = x @ W^T的矩阵乘,其中x是[M, K]的 FP16 激活,W是[N, K]的三元权重矩阵,但实际存储是packed [N, K/4]的 uint8,外加scale [N, K/block_size]。
// block内累加符号后,再乘scale。为避免过度具体,这里只画关键逻辑。 __global__ void ternary_linear_kernel( const half* __restrict__ x, // [M, K] const unsigned char* __restrict__ w, // [N, K/4] const half* __restrict__ scale, // [N, K/128] half* __restrict__ out, // [M, N] int M, int N, int K) { __shared__ half x_tile[TILE_K]; __shared__ half scale_tile[TILE_N]; int row = blockIdx.y * blockDim.y + threadIdx.y; int col = blockIdx.x * blockDim.x + threadIdx.x; int acc = 0; // int累加符号 float acc_scaled = 0; // 跨block时的浮点累加 for (int k = 0; k < K; k += TILE_K) { // 1. 协作加载 x 的tile到shared memory // 2. 协作加载当前block的scale到shared memory // 3. 每个线程解出4个权重(一个byte), // 将 ±1 累加到对应的 int 累加器 // 4. 到达block边界时: // acc_scaled += acc * scale_tile[k / TILE_K]; // acc = 0; } out[row * N + col] = __float2half(acc_scaled); }几个关键决策:
- 符号累加用
int而不是float。三元值乘 1 或 -1 根本不需要浮点乘法,用整型加减法快得多,FP32 的累加精度也够用。 - scale 一定要提前放进 shared memory 或寄存器。第一版我直接读 global,结果每个 block 都会重复拉取,性能直接崩掉。
- 按 block_size 切分 inner loop,在 block 边界才做
acc * scale。这样把 scale 的乘法和显存访问频率降到了最低。 - 解码 2bit 到符号这个操作,我试过查表
char dict[4] = {0,1,-1,0},也试过三元运算符。在这个 kernel 里两者差异不大,编译器基本都会优化成 predication,关键是别在这个地方引入分支发散。
4.3 三轮优化实测记录
第一版只求正确。跑通后发现延迟比「先反量化再走 FP16 cuBLAS」还要慢 15% 左右,瓶颈一眼便知:每个 thread 都在做 decode + 多次 float 乘法,scale 从 global memory 反复读,shared memory 里几乎没有做数据复用。这版的意义是完成了正确性闭环。
第二版就是做标准 tiling。把 inner dimension 的 tile 从 8 扩到 32,x 和 scale 都缓存到 shared memory,decode 改为查表,积攒起明显的访存收益。实测吞吐开始反超 FP16 cuBLAS 路线。
第三版才是真正吃到了三元量化的红利。把内循环的累积全部切成 int 方式,只在碰到 block 边界时才做浮点 scale 乘法;同时用向量化指令一次加载多个uint8,减少访存指令数;再配合__launch_bounds__调高 occupancy。这一版的吞吐相比第二版又提升了大几十个百分点。
我在这颗 A100 上拿自己的 embedding 模型测,显存占用大约是 FP16 版本的 1/8,单 batch 延迟从 12ms 级别降到 9ms 级别,16 并发下的吞吐大约是从 500 到 800 tok/s 的档位。数字具体多少不重要,重要的是趋势:内核瓶颈从「浮点乘加」变成了「访存带宽」,而三元量化恰好把带宽需求压到了最低。
4.4 与 vLLM 的融合
kernel 写好后,剩下的就是把自定义算子挂到 vLLM 的前向路径里。我是用torch.library注册了一个ternary_linearop,然后在自定义LinearMethod.apply里调用。需要注意,vLLM 不同版本对原有算子的替换机制差异极大,有的版本会在_custom_ops里统一 switch,有的版本已经改走vllm._C编译路径,所以集成前一定要先查你所在版本的torch.ops.vllm注册方式。
另一个常见的坑是维度假设。很多 kernel 示例只处理二维[M, K] @ [K, N],但 transformer 里输入 tensor 可能是[num_tokens, hidden]甚至带序列维度,必须显式reshape成二维并contiguous(),否则 kernel 拿到的 stride 根本不是连续内存,计算结果直接错乱。
5. 接入 vLLM 的完整流程与常见坑
5.1 最小可行的接入路径
我在实际操作中走通的接入路径大致是这样:
- 实现一个
TernaryQuantConfig,quant_method = "ternary",字段包含block_size、zero_point是否启用等参数; - 把它注册到 vLLM 的量化方法目录里;
- 实现
TernaryLinearMethod,关键方法是create_weights、weight_loader、apply; - 把自定义 kernel 编译成 CUDA extension,并通过
torch.library暴露成 op; - 修改模型
config.json里的quantization_config,指向ternary; - 启动 vLLM,先跑单条推理验证数值,再上性能测试。
这里的核心是第 3 步。create_weights里不能再创建原始 shape 的 FP16 权重变量,而要创建packed参数和scale参数,否则后面加载器还是会按普通 FP16 来对待。weight_loader里则要处理张量并行切分:column 切分按输出维度切 packed 权重和 scale 的行,row 切分按输入维度切时要注意 block 边界对齐。
5.2 实测数据速览
我随手整理了一下我当时跑出来的对比数据,仅供参考,不同模型、不同卡差异会很大:
| 方案 | 显存占用 | 单batch延迟 | 16并发吞吐 |
|---|---|---|---|
| FP16 原模型 | 约 1400 MB | 约 12 ms | 约 500 tok/s |
| 加载时反量化 + FP16 cuBLAS | 约 1400 MB | 约 14 ms | 约 420 tok/s |
| 自定义三元 kernel | 约 350 MB | 约 9 ms | 约 780 tok/s |
看出关键点没有?「加载时反量化」路线不仅没省显存,还因为多了一次反量化以及不匹配的数据布局,把性能拉低了。自定义 kernel 的收益主要来自两块:内存占用大幅下降,访存压力骤降;符号累加避免了大量浮点乘加。这也就是为什么我坚持要在 kernel 层面吃掉 2bit 格式,而不是偷懒走反量化。
5.3 常见问题速查表
我把自己踩过的、以及群里朋友踩过的坑整理成表格,遇到类似错误可以直接对号入座:
| 现象 | 可能原因 | 解决方案 |
|---|---|---|
加载时报KeyError: quant_method=ternary | 没在 vLLM 的量化注册表里注册 | 补齐TernaryQuantConfig并注册后重新加载 |
| 加载时 shape 对不上 | 默认weight_loader不认识 packed 格式 | 完整实现自定义weight_loader,接管切分 |
| 输出全是 NaN 或无限大 | 解包映射选错,比如 2bit 顺序反了 | 回对账脚本,逐字节确认00/01/10/11含义 |
| 输出接近正确但不精确 | scale 的 block 顺序与 packed 权重的 block 顺序不一致 | 按 block 粒度重新对账 scale 和权重 |
| 显存没有下降 | create_weights里仍然创建了原始 FP16 参数 | 只保留 packed 权重参数和 scale 参数 |
| 性能比 FP16 还差 | scale 反复读 global;线程间解码分支发散 | 把 scale 缓存到 shared memory;inner tile 扩大到 16 以上 |
| 多卡推理结果错乱 | 张量并行切分时没处理 block 边界对齐 | 先解包成整数矩阵再切分,切完重新打包 |
| kernel 报越界 | 有 block 维度不是 4 或 128 的整数倍 | 把边界补齐到 block_size 对齐,用 mask 丢弃无效位置 |
| 前向输出 shape 对不上 | vLLM 传入 3D/带 seq 维度的 tensor,kernel 只处理 2D | 在调用前reshape并contiguous() |
这些里面最隐蔽的其实是第二个和第五个。shape 不对还好排查,显存没下降这个问题最阴,因为表面上看代码是跑通了,但你根本不知道加载链路里已经被哪个默认逻辑偷偷反量化成了 FP16。我当时是打印了module.weight.dtype和显存占用才发现端倪。
5.4 一个小技巧:如何在不改 vLLM 源码的情况下注入自定义量化
很多人对改 vLLM 源码有心理负担,怕升级后补丁全丢。其实 vLLM 的量化注册是模块级的,很多版本支持通过quant_config里额外字段动态加载你预先注册的类。我是把自定义量化代码做成一个独立安装包,在启动脚本里先import my_ternary_quant,它会向 vLLM 注册表注册ternary,然后再创建LLM。这样 vLLM 主仓库的改动为零,后续升级兼容也容易迁移。
6. 写在最后:几点真实体会
这次折腾完,我对「把自定义模型塞进现成推理框架」这类工作有了新的理解。所谓逆向,不是破解什么高深算法,而是把框架作者的设计意图读明白。vLLM 的量化抽象我并不认为它哪里写得很烂,相反,它把常见量化的共性抽取得相当好,真正的问题在于三元量化这种极端形态不在它预设的「形状不变」假设里。所有后续的对账和 kernel 工作,本质都是在和这个假设较劲。
如果让我重新做一遍,我会把对账脚本写得更早、更严格。第一版 kernel 出数值之后,我一度以为所有权重对齐都对了,结果切到多卡才爆出 scale 的 shard 错位。那些脚本现在我还留着,每轮改完 kernel 都会全量跑一次,确认没有破坏任何一层。
再分享一个小技巧:kernel 优化的时候,不要一上来就追求「无乘法」「纯位运算」这种极致方案,先把朴素 decode 版本的数值验证通过,再逐步替换热点路径。我当时第一版和第二版之间的区别只在于访存策略,第三版才引入了 int 符号累加。每一步都有可回退的基线,晚上睡觉都踏实一些。
这个方案后续其实还有不少扩展空间。比如三元权重和稀疏结构通常天然共存,可以在 kernel 里加一层稀疏 mask 跳过零块;scale 本身也是 FP16,还能进一步压成 INT8 动态量化;另外把 KV cache 的量化和这个 kernel 做融合也值得尝试。不过这些都是后话了,先把基础链路跑稳,收益就足够大了。