ARTICLE DETAIL

资讯详情

深耕网站视觉设计与运营推广的一线实战洞察。

三元量化模型接入vLLM:从逆向加载到Kernel优化的完整实战

三元量化模型接入vLLM:从逆向加载到Kernel优化的完整实战 先说结论把一个三元量化模型跑在 vLLM 上最大的障碍往往不在 kernel而在于让 vLLM 的加载链路肯认这个「非主流」格式。我这次拿到的是一个权重只有 {-1, 0, 1} 三种取值的 embedding 模型目标是把它塞进 vLLM 做服务化推理。动手之前我以为写个自定义 CUDA kernel 就完事后来才发现从逆向 vLLM 的量化抽象、对齐磁盘权重与 GPU 权重布局再到 kernel 从「跑出数」到「跑得快」的三轮演进每一步都有意料之外的坑。这篇就是当时完整的落地记录适合正在做自定义量化接入 vLLM、或者想改造推理框架底层加载逻辑的朋友。里面不会只给结论还会把我实际推过的代码路径、对账脚本的思路、kernel 优化的取舍都摊开讲包括那些改到凌晨才发现的字节序问题和切分下标问题。看完你应该能少走一半弯路。1. 为什么非要把三元量化模型塞进 vLLM1.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 的向量化如果走原生 PyTorchbatch 一大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: ternaryvLLM 会直接抛 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_weightshape 是[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 规则。我这次模型的规则是每个权重映射为 2bit00 - 001 - 110 - -111 - 保留位按 0 处理每 4 个权重按小端顺序凑成一个uint8第一个权重占最低 2 位每连续 128 个权重共享一个 FP16 scaleblock 边界按 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 是uint8scale 是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 扩到 32x 和 scale 都缓存到 shared memorydecode 改为查表积攒起明显的访存收益。实测吞吐开始反超 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 最小可行的接入路径我在实际操作中走通的接入路径大致是这样实现一个TernaryQuantConfigquant_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_methodternary没在 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 memoryinner tile 扩大到 16 以上多卡推理结果错乱张量并行切分时没处理 block 边界对齐先解包成整数矩阵再切分切完重新打包kernel 报越界有 block 维度不是 4 或 128 的整数倍把边界补齐到 block_size 对齐用 mask 丢弃无效位置前向输出 shape 对不上vLLM 传入 3D/带 seq 维度的 tensorkernel 只处理 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 做融合也值得尝试。不过这些都是后话了先把基础链路跑稳收益就足够大了。
返回列表