Zig 机器学习推理:张量、GEMM 与 int8 量化

用 Zig 从零实现一个可用的神经网络推理引擎:张量表示与 NCHW 内存布局、分块加 SIMD 的 GEMM 内核、激活函数与 softmax 的数值稳定实现、im2col 卷积、int8 对称量化与整数 GEMM,以及 ONNX 模型加载思路与与 Python 训练侧的对接方式。

引言

Python 训练、C++/Rust 部署是深度学习落地的常见分工,原因是推理侧需要的是低延迟、低内存、无运行时依赖——这恰恰是系统语言的强项,也正是 Zig 的舒适区。用 Zig 写推理引擎有三个直接收益:comptime 可以把算子特化成针对固定形状的循环、@Vector 让你不必引入任何 BLAS 就能写出可用的 GEMM、而交叉编译让同一份代码同时产出 x86-64 与 aarch64 的二进制。

本文不依赖任何第三方框架,从零搭建一个可跑的推理引擎:定义张量与内存布局、实现分块 SIMD 的矩阵乘、补齐激活函数与 softmax、用 im2col 把卷积转成 GEMM、引入 int8 量化把吞吐再翻几倍,最后讨论 ONNX 模型如何加载以及如何与 Python 训练侧对接。

前置:SIMD 与向量化、内存管理与分配器模式。


目录


1. 推理引擎的组成

一个能跑通 MLP 与 CNN 的推理引擎,拆开只有五件事:

组件职责难点
张量数据容器:形状、步长、dtype、内存布局选择与视图语义
算子GEMM、卷积、激活、归一化、softmax内存访问模式与数值稳定
图算子之间的依赖与执行顺序拓扑排序与内存复用
权重加载从文件读参数并对齐布局字节序、对齐、量化反量化

本文的重点是前两项与量化——图调度与线程划分会在最后一节给出可落地的做法。工程上建议按「先跑通、再跑快」的顺序推进:先写一版纯标量实现验证数值正确(用 PyTorch 的输出做基准),再逐步替换热点为 SIMD 与分块版本。


2. 张量表示与内存布局

张量就是一个带元数据的连续缓冲区。用 shape 加 strides 表达视图,是支持转置、切片、广播的前提:

const std = @import("std");

pub const Tensor = struct {
    data: []f32,          // 行主序连续存储
    shape: [4]usize,      // 最多 4 维:N/C/H/W
    strides: [4]usize,    // 每个维度跨过的元素数
    ndim: u8,

    pub fn init(allocator: std.mem.Allocator, shape: []const usize) !Tensor {
        var dims = [_]usize{ 1, 1, 1, 1 };
        @memcpy(dims[0..shape.len], shape);
        var strides = [_]usize{ 1, 1, 1, 1 };
        var acc: usize = 1;
        var i: usize = shape.len;
        while (i > 0) {            // 行主序:最后一维步长为 1,向前累乘
            i -= 1;
            strides[i] = acc;
            acc *= dims[i];
        }
        return .{ .data = try allocator.alloc(f32, acc), .shape = dims, .strides = strides, .ndim = @intCast(shape.len) };
    }
};

布局的选择直接决定性能:

布局含义谁在用
NCHW批次、通道、高、宽PyTorch 默认,卷积核连续
NHWC批次、高、宽、通道TensorFlow / TFLite,逐像素访问友好

卷积在 NCHW 下,同一个输出像素要跨通道取值,跨距大;在 NHWC 下,一个输出像素的所有输入通道是连续的。推理引擎通常把 NCHW 权重在加载时重排成 NHWC 或分块格式,这一步(packing)的收益常常超过微调 GEMM 本身。

内存分配建议:权重用 page_allocator 或 mmap 只读映射(多进程共享),激活值用按层复用的 Arena;形状固定时可在启动时一次性算好每层缓冲区,运行期零分配。


3. 矩阵乘与 GEMM 内核

全连接层、注意力层、卷积的 im2col 之后,最终都落到 C[M,N] = A[M,K] * B[K,N]。朴素三重循环的瓶颈是每次只做一次乘加,且对 B 的访问跨距大。分块 + SIMD 的版本如下:

const Vec8f = @Vector(8, f32);

/// C = A * B,A 为 M x K,B 为 K x N,均为行主序
pub fn gemm(A: []const f32, B: []const f32, C: []f32, M: usize, K: usize, N: usize) void {
    var i: usize = 0;
    while (i < M) : (i += 1) {
        const a_row = A[i * K ..][0..K];
        const c_row = C[i * N ..][0..N];
        var j: usize = 0;
        while (j + 8 <= N) : (j += 8) {
            var acc: Vec8f = @splat(0);
            for (a_row, 0..) |a_ik, k| {
                const b_vec: Vec8f = B[k * N + j ..][0..8].*;
                acc = @mulAdd(Vec8f, @as(Vec8f, @splat(a_ik)), b_vec, acc);
            }
            c_row[j..][0..8].* = acc;
        }
        while (j < N) : (j += 1) {
            var s: f32 = 0;
            for (a_row, 0..) |a_ik, k| s += a_ik * B[k * N + j];
            c_row[j] = s;
        }
    }
}

三个层次的优化,收益依次递减但都很关键:

  1. SIMD 内层:@mulAdd 映射到 FMA,一次算 8 个输出,约 6~9 倍。
  2. 分块(blocking):把 M/K/N 切成能装进 L1/L2 的小块,缓存未命中从 O(MNK) 降到 O(MNK/B),大矩阵下再快 2~3 倍。
  3. 权重重排(packing):把 B 按 K x 8 的小块预先打包成连续内存,内层访问变成纯顺序,且能配合预取。

工程上的折中是:先转置 B。B 按列访问是跨距访问,转置成 Bᵀ 后内层读写都连续,这一步用几十行代码就能拿到显著收益。

对多线程,把 M 维按行划分给 n_jobs 个线程即可——不同行之间没有数据依赖,且各自写入 C 的不同区域,天然没有假共享。用 std.Thread.Pool 提交、WaitGroup 等待,即可在 4 核上拿到约 3.5 倍加速。

提示:先用 -Doptimize=ReleaseFast 测出标量基线,再逐层加 SIMD 与分块。每次只改一件事,才能知道收益来自哪里。


4. 激活函数与逐元素算子

逐元素算子是最容易向量化的部分——没有归约、没有依赖,只要循环连续就能自动向量化。ReLU 用 @max 一行搞定:

pub fn relu(x: []f32) void {
    const V = @Vector(8, f32);
    const zero: V = @splat(0);
    var i: usize = 0;
    while (i + 8 <= x.len) : (i += 8) {
        x[i..][0..8].* = @max(@as(V, x[i..][0..8].*), zero);
    }
    while (i < x.len) : (i += 1) x[i] = @max(x[i], 0);
}

GELU 与 SiLU 是 Transformer 的主力激活。精确 GELU 要算 erf,标准库没有,工程上普遍用 tanh 近似:

pub fn geluApprox(x: f32) f32 {
    // 0.5 * x * (1 + tanh(0.7978845608 * (x + 0.044715 * x^3)))
    const c: f32 = 0.7978845608;
    return 0.5 * x * (1.0 + std.math.tanh(c * (x + 0.044715 * x * x * x)));
}

pub fn silu(x: f32) f32 { return x / (1.0 + @exp(-x)); }

Softmax 必须做数值稳定处理:直接算 exp(x) / sum(exp(x)) 在 x 较大时溢出成 inf,结果是 nan。标准做法是减去最大值:

pub fn softmax(logits: []f32) void {
    var max_v = logits[0];
    for (logits[1..]) |v| max_v = @max(max_v, v);   // 1) 求最大值
    var sum: f32 = 0;
    for (logits) |*v| {                             // 2) 减最大值后取指数
        v.* = @exp(v.* - max_v);
        sum += v.*;
    }
    const inv = 1.0 / sum;
    for (logits) |*v| v.* *= inv;                   // 3) 归一化
}
算子数值要点向量化难度
ReLU无极易(@max)
GELU用 tanh 近似替代 erf易(逐元素)
Softmax必须减最大值,否则溢出中(含两次归约)
LayerNorm均值与方差两次归约,注意 eps中

LayerNorm 的 eps 通常取 1e-5 或 1e-6,写成 1.0 / @sqrt(var + 1e-5)——漏掉 eps 会在方差接近 0 时产生 inf,这是从零实现时最常见的数值 bug。

心法:推理引擎的数值正确性靠对比验证:用同一组输入跑 PyTorch 与 Zig 实现,逐层比较最大绝对误差。fp32 下每层误差应在 1e-5 量级;如果某层突然到 1e-1,一定是布局或归约维度搞错了。


5. 卷积实现

卷积的朴素实现是六层循环,性能极差。工业界的通用做法是 im2col + GEMM:把输入按滑窗展开成矩阵,卷积就变成一次矩阵乘。

/// 把 NCHW 输入展开成 (N*OH*OW) x (C*KH*KW) 的矩阵
pub fn im2col(input: []const f32, out: []f32, dims: Dims) void {
    const cols = dims.c * dims.kh * dims.kw;
    const oh = (dims.h + 2 * dims.ph - dims.kh) / dims.sh + 1;
    const ow = (dims.w + 2 * dims.pw - dims.kw) / dims.sw + 1;
    for (0..dims.n) |bi| for (0..oh) |oy| for (0..ow) |ox| {
        const row = ((bi * oh) + oy) * ow + ox;
        var col: usize = 0;
        for (0..dims.c) |ci| for (0..dims.kh) |ky| for (0..dims.kw) |kx| {
            const iy = @as(isize, @intCast(oy * dims.sh + ky)) - @as(isize, @intCast(dims.ph));
            const ix = @as(isize, @intCast(ox * dims.sw + kx)) - @as(isize, @intCast(dims.pw));
            const ok = iy >= 0 and iy < dims.h and ix >= 0 and ix < dims.w;   // 越界补零
            out[row * cols + col] = if (ok)
                input[((bi * dims.c + ci) * dims.h + @as(usize, @intCast(iy))) * dims.w + @as(usize, @intCast(ix))]
            else
                0;
            col += 1;
        };
    };
}

之后调用 gemm,权重矩阵形状是 (C*KH*KW) x OC,输出形状 (N*OH*OW) x OC,再 reshape 回 NCHW 即可。

代价与优化:

  • im2col 会让内存膨胀 KH*KW 倍。3x3 卷积膨胀 9 倍,5x5 是 25 倍。对激活值本来就大的层,这个开销不可忽略。
  • 优化方向一是 分块 im2col:一次只展开若干行,与 GEMM 交织执行,把中间矩阵控制在 L2 缓存内。
  • 优化方向二是 直接卷积(direct conv):不展开,直接在滑窗上做向量化,对小卷积核通常更快。1x1 卷积本质就是 GEMM,直接调 gemm 即可。
卷积类型推荐实现理由
1x1直接 GEMM无滑窗,等价于矩阵乘
3x3(大通道)分块 im2col + GEMM复用高,GEMM 效率高
3x3(小通道)直接卷积im2col 膨胀比收益大
深度可分离逐通道卷积 + 1x1 GEMMMobileNet 系结构

提示:推理引擎支持的第一种卷积最好是 1x1——它既是 GEMM 的直接复用,又能验证「权重布局 + 输出 reshape」这条链路是否正确。


6. int8 量化

fp32 推理的瓶颈常常是内存带宽而非算力。int8 量化把权重和激活都压到 1 字节,内存流量降到 1/4,且整数乘加在 SIMD 里吞吐是浮点的 2~4 倍。对称量化是最简单的方案:

q = round(x / scale)          // 量化:fp32 -> int8,范围 [-127, 127]
x = q * scale                 // 反量化:int8 -> fp32
scale = max(abs(x)) / 127     // 每张量或每通道一个 scale
pub const Quantized = struct { data: []i8, scale: f32 };  // 每通道量化时 scale 为 []f32

pub fn quantize(x: []const f32, out: []i8, scale: f32) void {
    for (x, out) |v, *q| q.* = @intFromFloat(std.math.clamp(@round(v / scale), -127.0, 127.0));
}

pub fn dequantize(q: []const i8, out: []f32, scale: f32) void {
    for (q, out) |v, *o| o.* = @as(f32, @floatFromInt(v)) * scale;
}

整数 GEMM 的关键是累加器必须用 i32 而不是 i8——两个 int8 相乘最大 127*127 ≈ 16129,累加 K 次很容易溢出:

const I8x16 = @Vector(16, i8);
const I32x16 = @Vector(16, i32);

/// int8 点积:输入 int8,累加到 i32
pub fn dotI8(a: []const i8, b: []const i8) i32 {
    var acc: I32x16 = @splat(0);
    var i: usize = 0;
    while (i + 16 <= a.len) : (i += 16) {
        const wa: @Vector(16, i16) = @as(I8x16, a[i..][0..16].*);  // 先拓宽到 i16
        const wb: @Vector(16, i16) = @as(I8x16, b[i..][0..16].*);
        acc += @as(I32x16, wa) * @as(I32x16, wb);                  // 再累加到 i32
    }
    var sum: i32 = @reduce(.Add, acc);
    while (i < a.len) : (i += 1) sum += @as(i32, a[i]) * @as(i32, b[i]);
    return sum;
}

量化推理的完整链路:

  1. 权重离线量化成 int8 + 每通道 scale,随模型一起存盘。
  2. 激活在运行时量化:x_q = round(x / scale_x),scale_x 可以离线用校准集统计(静态量化),也可以每批动态计算(动态量化)。
  3. 整数 GEMM 得到 i32 累加结果。
  4. 反量化:out = acc * (scale_a * scale_b),再乘权重整体的缩放因子。
量化方案精度损失实现复杂度适用
仅权重量化(W8A32)很小低内存受限、算力充足
静态量化(W8A8)小中服务端 CPU 推理首选
动态量化小中LSTM、Transformer 的激活
每通道量化比每张量更小中权重的标准做法

量化最容易被忽略的两点:一是 bias 不要量化(保持 fp32 或 i32,在反量化阶段加),二是 softmax 与 LayerNorm 保持 fp32——它们的输入动态范围大,量化后精度崩塌最明显。

心法:量化的验证方式是「端到端指标」而非「逐层误差」。逐层误差会累积放大,但只要最终 top-1 准确率下降在 1% 以内,方案就是可用的。


7. ONNX 模型加载

ONNX 是推理侧事实上的交换格式,本质是一个 protobuf 文件:ModelProto → GraphProto → NodeProto/TensorProto。加载分三步:

步骤内容工具
1. 解析 protobuf读出节点、张量、初始化器zig-protobuf 生成代码,或手写最小解析器
3. 权重重排把 NCHW 权重打包成引擎内部格式一次性离线完成并缓存

最小可用的加载器只需要覆盖你实际用到的算子。一个实用的做法是:先用 Python 把 ONNX 模型简化并导出成自定义的紧凑格式——把算子类型、形状、权重按顺序写成一个扁平二进制文件(含 magic + 版本号 + 张量表 + 权重块),Zig 侧只做顺序读取:

const ModelHeader = extern struct {
    magic: u32 = 0x5A4D4C31,   // "ZML1"
    version: u32,
    tensor_count: u32,
    node_count: u32,
};

pub fn loadModel(allocator: std.mem.Allocator, path: []const u8) !Model {
    const file = try std.fs.cwd().openFile(path, .{});
    defer file.close();
    var header: ModelHeader = undefined;
    _ = try file.readAll(std.mem.asBytes(&header));   // 直接读进结构体
    if (header.magic != 0x5A4D4C31) return error.BadMagic;
    _ = .{ allocator, header };                       // 再按 tensor_count/node_count 顺序读表与权重
    return .{};
}

const Model = struct {};

为什么推荐自定义格式而不是直接读 ONNX:

  • ONNX 的 protobuf 解析需要额外依赖与完整 schema,而推理侧用到的字段不到 10%。
  • 训练侧导出时顺手做权重重排与量化,推理侧启动更快、代码更少。
  • 自定义格式可以带版本号与校验和,升级时能明确报错而不是静默出错。

如果必须直接读 ONNX,注意 protobuf 是变长编码:varint 编码整数、字段顺序不保证、未知字段必须跳过。手写解析器时最容易错的是 varint 的续位判断(最高位为 1 表示还有后续字节)。

提示:权重加载必须校验对齐。用 @alignCast 把字节缓冲转成 []f32 前,先确认缓冲区按 4 字节对齐;否则在 ARM 上会直接触发对齐异常,而不是像 x86 那样默默降速。


8. 性能基准与工程化

推理性能的度量有三个层次,缺一不可:

指标含义目标
端到端延迟一次推理的总耗时用户体验
吞吐每秒处理多少样本(batch 支持)服务成本

基准的正确写法:预热若干轮(让缓存与频率稳定)、跑足够多次、用 std.time.Timer 测单调时间、把结果喂给 std.mem.doNotOptimizeAway 防止被优化掉;报告 timer.read() / iters 作为单次延迟。

Roofline 视角:先算清楚你的算子属于「计算受限」还是「内存受限」。以 fp32 GEMM 为例,若 CPU 单核峰值约 50 GFLOP/s、内存带宽约 30 GB/s,算术强度(FLOP/Byte)阈值约为 1.7。M=K=N=1024 的 GEMM 算术强度是 2MNK / (4(MN+MK+NK)) ≈ 85,远超阈值,属于计算受限——优化方向是 SIMD 与分块。而逐元素激活(如 ReLU)算术强度只有 0.25,属于内存受限——加宽向量没用,要靠算子融合减少内存往返。

工程化的六条建议:

  1. 算子融合:GEMM + bias + ReLU 合成一个内核,省掉两次全量内存往返,逐元素部分常能快 2~3 倍。
  2. 固定形状:推理服务通常 batch 与序列长度固定,把形状写进 comptime 参数可以让编译器完全展开循环。
  3. 权重共享与线程复用:多进程用 mmap 只读映射同一份权重,请求复用固定的 std.Thread.Pool。
  4. 对齐分配:激活缓冲区按 64 字节对齐(缓存行),权重按 32 字节对齐。
  5. 数值回归测试:把 PyTorch 的输出存成二进制基准,每次改动后逐层比对最大误差。

心法:先分清计算受限还是内存受限,再决定优化方向。对着内存受限的算子加宽 SIMD 向量,是最常见的无效优化。


9. 速查表

需求手段
张量定义连续 []f32 + shape[4] + strides[4]
布局选择PyTorch 权重 NCHW,推理前重排为 NHWC 或分块
GEMM 内核@mulAdd + 8 宽向量 + 分块 + 转置 B
多线程 GEMM按 M 行划分,std.Thread.Pool + WaitGroup
ReLU@max(v, @splat(0))
GELU 近似0.5x(1+tanh(0.7978845608(x+0.044715x³)))
Softmax先减最大值再取指数,防溢出
LayerNorm1/sqrt(var + 1e-5),eps 不可省
卷积1x1 走 GEMM,3x3 走分块 im2col
对称量化q = round(x/scale),scale = max(abs)/127
int8 点积i16 中间量 + i32 累加器,最后 @reduce(.Add, ...)
量化例外bias 不量化,softmax/LayerNorm 保持 fp32
模型加载训练侧导出紧凑二进制(magic + 版本 + 张量表)
性能判定先算算术强度,区分计算受限与内存受限
基准预热 + std.time.Timer + doNotOptimizeAway
正确性验证与 PyTorch 输出逐层比对,fp32 误差应在 1e-5

10. 一句话记忆

Zig 推理引擎的骨架是「张量 + GEMM + 量化」:张量用 shape/strides 表达布局,GEMM 靠 @mulAdd 与分块吃满 FMA,卷积用 im2col 化归为 GEMM,int8 量化把内存流量砍到四分之一——先分清计算受限还是内存受限,再决定优化方向。


延伸阅读

继续阅读

探索更多技术文章

浏览归档,发现更多关于系统设计、工具链和工程实践的内容。

全部文章 返回首页

「系统编程」更多文章

  1. Zig 时间、日期与时区处理:std.time 与 epoch 换算
  2. Zig 插件系统与动态加载:C ABI 契约、热重载与错误隔离
  3. Zig HTTP 客户端与 REST 集成:std.http.Client 实战