首页
/ BitNet GPU 推理框架:从零构建 W2A8 的 2-bit 大模型 CUDA GEMV 内核

BitNet GPU 推理框架:从零构建 W2A8 的 2-bit 大模型 CUDA GEMV 内核

2026-09-05 12:31:29作者:齐冠琰

本文以 BitNet 官方仓库中的 gpu/README.md 为核心,完整讲解其针对 BitNet-b1.58-2B 模型定制的 W2A8(2-bit 权重 × 8-bit 激活)GEMV 推理内核:从环境安装、内核编译、性能测试,到模型权重转换与端到端推理的全流程命令,并结合 pack_weight.pyconvert_checkpoint.pymodel.pybitnet_kernels.h 的源码,深入解析权重置换、快速解码与 dp4a 指令三大核心优化。读完本文,你既能完整复现该 GPU 推理管线,也能理解低比特量化在 CUDA 层面的实现原理。

一、模块定位:为 BitNet 定制的低比特 GEMV 内核

gpu/ 目录是 BitNet 官方推理框架中面向 GPU 的推理内核实现,其定位在 gpu/README.md 中明确给出:

  • 支持 W2A8(2-bit 权重 × 8-bit 激活)GEMV 计算;
  • 提供低延迟的自定义 CUDA 内核;
  • 针对内存访问、解码与计算吞吐做了专项优化;
  • 专为 BitNet-b1.58-2B 模型(4T tokens 版本)定制。

需要特别注意其运行前提:

  1. 该内核仅针对 decode(逐 token 解码)阶段的 GEMV 优化。从 gpu/bitnet_kernels/bitnet_kernels.cu 可以看到,所有被实例化的内核都要求 M == 1(单行输入向量),未匹配的 (M, N, K) 组合只会打印 "required ladder gemm kernel" 提示。仓库中针对 2B / 3B / 4B 等尺寸预置了 10 种形状,例如 M=1, N=3840, K=2560M=1, N=13824, K=2560M=1, N=55296, K=5120 等。
  2. 编译脚本 gpu/bitnet_kernels/compile.sh 使用 -gencode=arch=compute_80,code=compute_80,即 面向 Ampere 架构(sm_80,如 A100)编译,这是 dp4a 低精度点积指令可用性的直接依赖。
  3. 模型侧目前仅在 gpu/convert_safetensors.py 中内置了 2B 一种模型配置(n_layer=30, n_head=20, dim=2560, vocab_size=128256, n_local_heads=5, intermediate_size=6912)。

从整体架构看,该模块采用"prefill 走 BF16、decode 走 W2A8 内核"的混合策略,这一设计在 gpu/generate.py 中体现得非常清晰:

model_args_prefill = fast.ModelArgs(use_kernel=False)   # prefill 用 BF16 矩阵乘
model_args_decode  = fast.ModelArgs(use_kernel=True)    # decode 用 W2A8 内核
...
prefill_model.load_state_dict(fp16_checkpoint, strict=True)  # model_state_fp16.pt
decode_model.load_state_dict(int2_checkpoint, strict=True)   # model_state_int2.pt

generate.py 会加载两套权重:BF16 的 model_state_fp16.pt 用于 prefill,int2 的 model_state_int2.pt 用于 decode,两套模型分别被捕获为独立的 CUDA Graph。

二、环境准备与内核编译

2.1 安装依赖

按照 gpu/README.md 给出的流程:

# (Recommended) Create a new conda environment
conda create --name bitnet-gpu "python<3.13"
conda activate bitnet-gpu

# Install dependencies
pip install -r requirements.txt

gpu/requirements.txt 中的依赖清单及其用途如下:

依赖 版本约束 用途
torch >=2.2.0 CUDA 张量与 CUDA Graph 捕获(generate.py 使用 torch.cuda.CUDAGraph
xformers >=0.0.22 模型中的 RMSNormfmharope_padded(见 model.py 的导入)
transformers 无约束 HuggingFace 生态支持
einops 无约束 权重转换中的 rearrange 布局变换(convert_safetensors.py
sentencepiece 无约束 SentencePiece 分词器(tokenizer.model
fire 无约束 generate.py 的命令行参数解析
tiktoken / blobfile / flask 无约束 其他工具链依赖

2.2 编译 CUDA 内核

# Build the kernel
cd bitnet_kernels
bash compile.sh
cd ..

compile.sh 的完整命令为:

nvcc -std=c++17 -Xcudafe --diag_suppress=177 --compiler-options -fPIC -lineinfo --shared bitnet_kernels.cu -lcuda -gencode=arch=compute_80,code=compute_80 -o libbitnet.so

要点解读:

  • 产物是动态库 bitnet_kernels/libbitnet.so,随后由 model.py 通过 ctypes.CDLL('bitnet_kernels/libbitnet.so') 加载;
  • --shared + -fPIC 保证库可被 Python 进程直接 dlopen;
  • arch=compute_80 锁定了 Ampere 架构,dp4a 指令路径依赖 sm_80 及以上能力。

内核的唯一导出符号是 bitlinear_int8xint2,其 C 接口签名为(bitnet_kernels.cu#L3):

extern "C" void bitlinear_int8xint2(
    int8_t* input0, int8_t* input1, __nv_bfloat16* output0,
    __nv_bfloat16* s, __nv_bfloat16* ws, int M, int N, int K, cudaStream_t stream);

其中 input1 是已经按 4:1 打包的 int2 权重(所以 Python 侧用 K = input1.shape[1] * 4 还原真实的 K 维度,见 model.py#L31),s 是激活量化 scale,ws 是权重 scale 数组,输出为 BF16。

三、内核级性能测试

# Run performance tests
python test.py

gpu/test.py 的测试逻辑值得拆解,它是验证内核正确性与性能的标准入口:

  1. 形状覆盖test_list 覆盖 8 种 (N, K) 形状(test.py#L32-L41),与 2B 模型中 wqkv / w13 / w2 / wo 及词表映射等 GEMV 的实际维度一一对应,例如 (13824, 2560) 对应 w13(ffn_dim×2 = 6912×2)、(2560, 6912) 对应 w2
  2. 正确性校验:对随机 int8 权重与 int8 激活,先用 numpy 的 int32 矩阵乘作为参考(test.py#L50-L61),再通过 ctypes 调用编译出的 bitlinear_int8xint2,逐位断言 torch.all(out==out_np)。注意参考实现用的是未打包的原始权重,而内核输入是经过 convert_weight_int8_to_int2 打包后的 weight_compressed,两者结果完全一致,即打包/置换/交错过程在数值上是无损的。
  3. 性能对比:用 torch.utils.benchmark.Timer 各计时 50 次,对比 W2A8 内核与 torch.matmul(BF16)(test.py#L71-L86),输出形如:
Shape(13824, 2560), W2A8: 18.75us, torch BF16: 59.51us

运行结果即下文中"性能"一节 README 表格数据的来源,且 test.py 中的形状列表与表格完全吻合,可作为表格数据的可复现依据。

四、模型权重转换:从 HuggingFace 到 int2 内核格式

端到端推理前必须完成权重转换,README 给出的命令序列如下:

# Download and convert the BitNet-b1.58-2B model
mkdir checkpoints
huggingface-cli download microsoft/bitnet-b1.58-2B-4T-bf16 --local-dir ./checkpoints/bitnet-b1.58-2B-4T-bf16
python ./convert_safetensors.py --safetensors_file ./checkpoints/bitnet-b1.58-2B-4T-bf16/model.safetensors --output checkpoints/model_state.pt --model_name 2B
python ./convert_checkpoint.py --input ./checkpoints/model_state.pt
rm ./checkpoints/model_state.pt

这一步实际上是两级流水线,下面逐级拆解。

4.1 第一级:convert_safetensors.py — 布局重排与模块合并

convert_safetensors.py 把 HuggingFace 格式的 model.safetensors 转成 BitNet 推理框架内部命名的 model_state.pt,核心工作有两类:

(1)Q/K 投影的交错重排(interleaving)。HuggingFace 的 Q/K 权重是 (head_dim, dim) 平铺布局,而 BitNet 的 RoPE 实现按相邻两维成对旋转,因此需要把每两个 head_dim 维交叉重排:

def invert_convert_q(w: torch.Tensor, config: ModelArgs) -> torch.Tensor:
    return rearrange(w, '(h l d) i -> (h d l) i', h=config.n_head, l=2)

convert_safetensors.py#L43-L47)。

(2)GQA 头合并与权重拼接。逐层执行(convert_safetensors.py#L61-L90):

  • q_proj + k_proj + v_projlayers.{l}.attention.wqkv.weight(先各自做 Q/K 重排再在 dim 0 拼接);
  • gate_proj + up_projlayers.{l}.feed_forward.w13.weight
  • mlp.down_projfeed_forward.w2.weightself_attn.o_projattention.wo.weight
  • input_layernorm / post_attention_layernorm 分别映射为 attention_norm / ffn_norm
  • 额外取出 BitNet b1.58 特有的 attn_sub_normffn_sub_norm
  • tok_embeddings.weight 同时用作 output.weight(权重绑定)。

--model_name 参数通过 ModelArgs.from_name 解析,目前 transformer_configs 仅注册了 2B,因此该脚本当前只适用于 2B 配置。

4.2 第二级:convert_checkpoint.py — int2 量化打包 + BF16 备份

convert_checkpoint.py 读取上一步产物,同时输出两份 checkpoint(convert_checkpoint.py#L86-L90):

  • model_state_int2.pt:供 decode 阶段的 W2A8 内核使用;
  • model_state_fp16.pt:供 prefill 阶段的 BF16 模型使用。

权重量化逻辑在 convert_checkpoint.py#L23-L27

def quant_weight_int8(weight):
    s = 1.0 / weight.abs().mean().clamp_(min=1e-5)
    new_weight = (weight * s).round().clamp(-1, 1).to(torch.int8)
    new_scale = (1.0 / s).to(torch.bfloat16)
    return new_weight, new_scale.reshape(1)

即按权重的平均绝对值计算 scale,把 BF16 权重量化到 {-1, 0, 1}(int8 容器,后续压成 int2)。对 wqkv / w13 / w2 / wo 四类投影矩阵分别量化,每个矩阵带一个独立 scale,并以 weight_scale 键(长度 4 的 BF16 向量,不足部分补零)随权重一起保存。

convert_int8_to_int2 调用的正是 pack_weight.py 中的 convert_weight_int8_to_int2——下一节的"三大优化"全部发生在这个函数里。

五、三大核心优化:权重置换、快速解码与 dp4a

这是 gpu/README.md 的 "Optimizations" 一节给出的内核设计要点,下面结合源码逐条展开。

5.1 权重置换(Weight Permutation)

官方描述:权重矩阵被划分为 16×32 的块以优化内存访问模式;块内数值在内存中连续存储,并按特定顺序置换,便于高效访问和处理。详见 convert_checkpoint.py

16×32 恰好是 CUDA 中 WMMA/tensor-core 风格的 tile 粒度(wmma_N=16, wmma_K=32)。置换的具体映射在 pack_weight.py#L5-L14

def B_global_16x32_to_shared_load_16x32_layout(i, j):
    thread_id = i * 2 + j // 16
    row = (thread_id // 16) * 8 + (thread_id % 8)
    col = (j % 16) + 16 * ((thread_id % 16) // 8)
    return row, col

permutate_weight_fastest 将该映射向量化:对每个 16×32 块构建 (row, col) 查找表,然后用高级索引一次性从原始权重中抽取置换后的布局,输出形状为 (N//16, K//32, 16, 32) 的 int8 张量。

置换的目的从内核侧可以印证:在 bitnet_kernels.h#L60-L67 中,每个线程对 B 矩阵的读取是一个完整的 32-bit 标量加载*(int*)(B + ...)),其地址由 blockIdx.xthreadIdx.x/y 的固定线性函数给出。只有当权重在显存中按"线程读取顺序"预先摆放好时,这种读取才是合并、无 bank 冲突的。换句话说,置换把"内核期望的访存顺序"提前固化到了权重布局里,运行期无需任何索引计算。

5.2 快速解码(Fast Decoding):交错打包

官方描述:每 16 个 2-bit 值按如下交错模式打包进一个 32-bit 整数:[0, 4, 8, 12, 1, 5, 9, 13, 2, 6, 10, 14, 3, 7, 11, 15]。该布局通过"一次提取 4 个值到 int8"的方式加速解码。

打包链路由 convert_weight_int8_to_int2 串联四步:

weight = weight + 2                                  # ① {-1,0,1} -> {1,2,3} 无符号化
permutated_weight = permutate_weight_fastest(weight)  # ② 16×32 块置换
compressed_weight = compress_int2_to_int8(permutated_weight)  # ③ 每字节塞 4 个 2-bit 值
interleaved_weight = interleave_weight_int8(compressed_weight, 2)  # ④ int32 级交错

其中:

  • compress_int2_to_int8:每 4 个 2-bit 值右移后按位或进同一字节,权重体积缩小到 1/8;
  • interleave_weight_int8:把整个 int8 数组 reinterpret 成 int32,对每个 32-bit 字内的 16 个 2-bit 域执行跨字节重排。注意源注释中给出的位移动作表 [0, 8, 16, 24, 2, 10, 18, 26, ...]:重排后,每个字节 j 内的 4 个 2-bit 值恰好来自原数据中下标 {4j, 4j+1, 4j+2, 4j+3} 位置——这正是 README 中 [0, 4, 8, 12, ...] 交错模式的位级实现。

运行期解码在 bitnet_kernels.h#L23-L44decode_i2s_to_i8s 中,用一条 lop3.b32(三输入逻辑运算)指令把 32-bit 打包字一次展开为 4 个 int8:

static constexpr uint immLut = (0xf0 & 0xcc) | 0xaa;   // 0b11101010
static constexpr uint BOTTOM_MASK = 0x03030303;
asm volatile("lop3.b32 %0, %1, %2, %3, %4;\n"
             : "=r"(i8s[i])
             : "r"(i2s >> (2 * i)), "n"(BOTTOM_MASK), "n"(0x00000000), "n"(immLut));
i8s[i] = __vsubss4(i8s[i], 0x02020202);  // 减去偏移 2,恢复 {-1,0,1}

由于第 ④ 步交错后每个字节都只含来自同一字节的 4 个连续值(只是与别的字节交叉存放),解码只需 lop3(掩码 + LUT)加一次 __vsubss4 SIMD 减法即可拿到 4 个 int8 激活域值——这就是"一次提取 4 个值"的含义:一个 32-bit 加载对应 4 次 lop3,覆盖整个 16 元素 K 块,无任何分支。

5.3 dp4a 指令与规约

官方描述:使用 dp4a 指令加速低精度点积。该指令对两个 4 元素向量(各存于一个 32-bit 字中,元素为 int8)做点积,并累积到 32-bit 整数中,显著提升量化权重/激活下的 GEMV 吞吐。

主计算循环见 bitnet_kernels.h#L46-L82

// 每线程处理 K_per_loop = 16 个元素
*(int4*)(A_local + 0) = *(int4*)(A + ...);   // 16 字节合并加载 int8 激活
B_reshape_local[0]   = *(int*)(B + ...);    // 4 字节加载打包权重
decode_i2s_to_i8s(B_reshape_local, B_decode_local, 16);
#pragma unroll
for (int k_2_0 = 0; k_2_0 < 4; ++k_2_0)
    in_thread_C_local[0] = __dp4a(
        *(int *)&A_local[k_2_0 * 4],
        *(int *)&B_decode_local[k_2_0 * 4],
        in_thread_C_local[0]);

即:1 次 16B 激活加载 + 1 次 4B 权重加载 + 4 次 lop3 解码 + 4 次 dp4a 完成一个 16 元素子块的点积累积。整块 K 维用 K_block_size 个线程(x 方向)并行切分,块内再通过 4 轮 warp shuffle 规约汇总(bitnet_kernels.h#L76-L78):

for (int offset = K_block_size/2; offset > 0; offset /= 2)
    red_buf0[0] += __shfl_down_sync(__activemask(), red_buf0[0], offset, K_block_size);

最后由 threadIdx.x == 0 的线程按 out = acc * s[0] * ws[ws_idx] 写出 BF16 结果(bitnet_kernels.h#L81-L82),其中 s 即前文激活/权重量化的 scale。

ws_num 模板参数(bitnet_kernels.cu#L5-L33 中各实例取 1~6 不等)与 ws[out_idx / (N / ws_num)] 的索引方式看,不同 N 段的输出对应不同的权重 scale——这与 convert_checkpoint.py 中把 wq/wk/wv 各自量化、w13 中 w1/w3 各自量化再拼接的做法相吻合。

5.4 形状特化实例化

bitnet_kernels.cu 用一串 if (M == 1 && N == ... && K == ...) 分支将 ladder_int8xint2_kernel<M, N, K, ws_num, K_block_size, N_block_size> 的模板参数(块划分与 grid 尺寸)硬编码为编译期常量。以 2B 模型的典型形状为例:

N × K 实例化参数 <ws_num, K_block, N_block> grid
2560 × 2560(wo) 1, 8, 16 (160, 1, 1)
13824 × 2560(w13) 2, 8, 16 (864, 1, 1)
2560 × 6912(w2) 1, 8, 16 (160, 1, 1)
20480 × 3200 2, 8, 16 (1280, 1, 1)

线程块统一为 dim3(8, 16, 1) = 128 线程(与内核的 __launch_bounds__(128) 一致),每个 (x 方向 8 线程) 负责 K 块的一段,y 方向 16 个线程对应 N_block 的 16 个输出通道。这也解释了为什么测试/模型中的所有 (N, K) 必须与实例表精确匹配。

六、端到端推理

权重转换完成后即可启动交互式推理:

# Inference
python3 ./generate.py ./checkpoints/ --interactive --chat_format

generate.py 基于 fire 解析参数(generate.py#L322-L359),支持:

  • ckpt_dir(位置参数):包含 model_state_int2.ptmodel_state_fp16.pt 的目录;
  • --interactive:循环读取用户输入对话;
  • --chat_format:把底层 SentencePiece 分词器包装为对话模板格式(tokenizer.pyChatFormat);
  • --sampling:开启采样(temperature=0.7 / top_p=0.95,硬编码于 generate.py#L253-L256),默认走 argmax 贪心。

运行时行为要点(结合 generate.py 源码):

  1. 双模型 + 双 CUDA Graphcompile_prefillcompile_generate 分别对 prefill(BF16 全精度)与 decode(W2A8 内核)各做一次 warmup 并捕获 CUDA Graph,之后每个 prompt 的 prefill 与每个 decode step 都是 graph replay,最大限度压低 kernel launch 开销。
  2. KV cache:由 model.pymake_cachegen_bsz * (prompt_length + gen_length) 预分配,注意力计算走 xformers 的 flash 路径(model.py#L154-L156)。
  3. 逐 token 量化:decode 路径中,BitLinearKernel.quant_input 对每个 token 的隐藏状态做逐行 per-channel 量化:s = 127 / max|x|,再取整钳位到 [-128, 127] 转 int8,与 5.3 节内核的 s[0] 标度输入对应。该函数带 @torch.compile 装饰以融合量化算子。
  4. 环境变量开关:设置 NO_CUDA_GRAPHS 可禁用 CUDA Graph(generate.py#L343),便于排查问题。

推理过程中每轮会打印 prefill / decode 两个 phase 的吞吐统计(stats.py)以及当前显存占用 Memory used: ... GB

七、性能数据

以下数据来自 gpu/README.md 的 Performance 一节,测试环境为 NVIDIA A100 40GB;kernel 表格与 test.py 的输出可直接复现。

7.1 内核基准(GEMV 延迟)

Shape (N×K) W2A8 Latency (us) BF16 Latency (us) Speedup
2560 × 2560 13.32 18.32 1.38
3840 × 2560 14.90 18.87 1.27
13824 × 2560 18.75 59.51 3.17
2560 × 6912 14.49 37.78 2.61
3200 × 3200 14.61 19.08 1.31
4800 × 3200 13.09 21.84 1.67
3200 × 10240 19.64 60.79 3.10
20480 × 3200 30.99 112.39 3.63

规律清晰:K 越大(累加维度越长)加速比越高,最大达 3.63×。这与 GEMV 的访存瓶颈一致——int2 权重使显存带宽需求降为 BF16 的约 1/8,K 维越长,被压缩的字节数越多,带宽收益越显著;而小矩阵(2560×2560)受 kernel 启动等固定开销占比影响,加速比接近 1.3×。

7.2 端到端生成延迟

与同规模 BF16 模型(Gemma-2-2B,vLLM 后端)在 A100 40GB 上的对比:

Input Length Output Length BF16 Latency (ms) W2A8 Latency (ms) Speedup
64 16 187.64 57.40 3.27
64 32 353.50 112.22 3.15
64 64 683.23 221.08 3.09
256 16 183.14 61.24 2.99
256 32 353.14 115.47 3.06
256 64 684.24 224.16 3.05
512 16 208.99 68.06 3.07
512 32 354.33 122.72 2.89
512 64 709.65 231.82 3.06

端到端加速稳定在 3× 左右,且对输入/输出长度不敏感——decode 阶段占比越高(输出越长)时 W2A8 优势越明显,这与 decode 走 int2 内核、prefill 走 BF16 的混合设计自洽。

八、适用边界与注意事项

结合仓库事实,使用该模块时应注意:

  1. 仅 decode 阶段受益:内核要求 M == 1bitnet_kernels.cu#L3-L37),prefill 仍走 BF16 权重(model_state_fp16.pt),因此显存中同时驻留两套权重。
  2. 形状封闭集合:内核按 (N, K) 硬编码实例化,新增模型尺寸需自行在 bitnet_kernels.cu 中补充模板实例与 grid 配置并重新编译。
  3. 架构限定compile.sh 固定 compute_80,面向 Ampere 平台(README 性能数据均出自 A100 40GB)。
  4. 模型限定:权重转换脚本目前只内置 2B 配置(convert_safetensors.py#L9-L11),面向 BitNet-b1.58-2B-4T 的 BF16 safetensors 检查点;仓库根目录的 gpu/convert_checkpoint.py 依赖该脚本产出的键名(如 wqkvw13w2wo)。
  5. 分词器资源generate.py 硬编码加载 ./tokenizer.modelgenerate.py#L58),需在 gpu/ 目录下运行以命中 gpu/tokenizer.model

小结

gpu/ 模块展示了低比特 LLM 推理内核的一条完整工程路径:用权重置换把 WMMA tile 的访存模式固化进显存布局,用 int32 交错打包 + lop3 位运算实现无分支的 4 值/次解码,用 dp4a 把 int8×int2 点积交给硬件指令,再配合形状特化实例化CUDA Graph 压低启动开销,最终在 A100 上取得 kernel 级 1.3×~3.6×、端到端约 3× 的加速。相关入口文件包括 gpu/README.mdgpu/test.pygpu/convert_safetensors.pygpu/convert_checkpoint.pygpu/pack_weight.pygpu/model.pygpu/generate.pygpu/bitnet_kernels/bitnet_kernels.h,可按本文流程逐一对照验证。

登录后查看全文
热门项目推荐
相关项目推荐

项目优选

收起
kernelkernel
deepin linux kernel
C
33
18
ops-transformerops-transformer
本项目是CANN提供的transformer类大模型算子库,实现网络在NPU上加速计算。
C++
1.12 K
2.72 K
kernelkernel
openEuler内核是openEuler操作系统的核心,既是系统性能与稳定性的基石,也是连接处理器、设备与服务的桥梁。
C
528
588
ops-nnops-nn
本项目是CANN提供的神经网络类计算算子库,实现网络在NPU上加速计算。
C++
906
1.83 K
pytorchpytorch
作为 Ascend for PyTorch 社区的核心组件,TorchNPU 是昇腾专为 PyTorch 打造的深度学习适配插件,使 PyTorch 框架能够直接调用昇腾 NPU,为开发者提供昇腾 AI 处理器的超强算力。
Python
854
1.34 K
docsdocs
暂无描述
Markdown
891
5.79 K
jiuwenswarmjiuwenswarm
JiuwenSwarm 是一款基于openJiuwen开发的智能AI Agent,它能够将大语言模型的强大能力,通过你日常使用的各类通讯应用,直接延伸至你的指尖。
Python
3.53 K
1.01 K
ops-mathops-math
本项目是CANN提供的数学类基础计算算子库,实现网络在NPU上加速计算。
C++
1.34 K
1.45 K
cann-learning-hubcann-learning-hub
CANN 学习中心仓,支持在线互动运行、边学边练,提供教程、示例与优化方案,一站式助力昇腾开发者快速上手。
Jupyter Notebook
988
506
AscendNPU-IRAscendNPU-IR
AscendNPU-IR是基于MLIR(Multi-Level Intermediate Representation)构建的,面向昇腾亲和算子编译时使用的中间表示,提供昇腾完备表达能力,通过编译优化提升昇腾AI处理器计算效率,支持通过生态框架使能昇腾AI处理器与深度调优
C++
540
384