Skip to content
当前页大纲

Mojo 案例:加速神经网络推理(端到端)

学习模块的第十一篇——综合实战案例。场景很真实:你有一个用 Python 训练好的小模型,需要在低延迟场景下做推理(机器人控制环、边缘设备、在线服务热路径)。传统答案是「用 C++ 重写 + pybind11 胶水」,本案例展示 Mojo 路线:Python 训练 → Mojo 推理内核 → 一致性验证 → 性能对比,把前十篇的技术全部串起来。

一、案例背景:为什么是「推理加速」

训练和推理的性能诉求完全不同:

训练推理(在线/边缘)
批量大小大(几百上千)小(经常是 1)
延迟要求不敏感微秒级敏感
典型瓶颈GPU 吞吐单条调用的固定开销

问题在这:单条样本推理时,Python 的解释器开销比计算本身还贵。numpy 单样本调用要跨 Python 边界、构造临时对象、动态分发——底层 BLAS 再快也被胶水吃光。

这正是 Mojo 的甜区:把推理内核编译成原生代码,零解释器开销 + 显式 SIMD。用鸢尾花(iris)分类做演示——麻雀虽小,dense 层 + ReLU + 输出层的结构和大模型推理是同构的。

二、模型与数据

经典配置:

text
输入 (4 维: 花萼长/宽、花瓣长/宽)
  ↓ Dense(4→16) + ReLU        ← 隐藏层,行主序权重
  ↓ Dense(16→3) + argmax      ← 三分类输出

两个刻意的设计决策(真实工程思维):

  1. 隐藏层取 16:是 SIMD 宽度 8 的整数倍——教学版免去尾部处理
  2. 推理只做 argmax 不做 softmax:softmax 不改变最大值的位置,概率只有训练/可视化才需要。省掉一整个 exp 循环,结果完全等价——这是推理内核的经典免费优化

三、步骤一:Python 训练并导出权重

纯 numpy 手写训练循环(约 40 行,无框架依赖):

python
# train.py
import numpy as np
from sklearn.datasets import load_iris

rng = np.random.default_rng(42)
X, y = load_iris(return_X_y=True)
X = (X - X.mean(0)) / X.std(0)          # 标准化

D, H, C = 4, 16, 3                       # 输入/隐藏/输出维度
W1 = rng.normal(0, 0.5, (D, H)); b1 = np.zeros(H)
W2 = rng.normal(0, 0.5, (H, C)); b2 = np.zeros(C)

for epoch in range(200):
    h = np.maximum(0, X @ W1 + b1)                       # 前向
    logits = h @ W2 + b2
    probs = np.exp(logits - logits.max(1, keepdims=True))
    probs /= probs.sum(1, keepdims=True)
    dlogits = probs                                      # 反向
    dlogits[range(len(y)), y] -= 1
    dlogits /= len(y)
    gW2 = h.T @ dlogits; gb2 = dlogits.sum(0)
    dh = dlogits @ W2.T * (h > 0)
    gW1 = X.T @ dh; gb1 = dh.sum(0)
    W1 -= 0.1 * gW1; b1 -= 0.1 * gb1
    W2 -= 0.1 * gW2; b2 -= 0.1 * gb2

acc = (np.argmax(np.maximum(0, X @ W1 + b1) @ W2 + b2, 1) == y).mean()
print(f"训练准确率: {acc:.2%}")            # 约 96%~98%(种子而异)

def fmt(a):                                # 导出为 Mojo 常量
    return "[" + ", ".join(f"{v:.6f}" for v in a.ravel()) + "]"
print("var W1: List[Float64] =", fmt(W1))
print("var b1: List[Float64] =", fmt(b1))
print("var W2: List[Float64] =", fmt(W2))
print("var b2: List[Float64] =", fmt(b2))

跑一下,把打印出的四行常量复制进 Mojo 代码——小模型内嵌权重是最朴素的「部署」,没有序列化格式、没有加载代码(生产环境换成 npz 文件 + 互操作读取即可,思路见第九篇)。

注意权重布局:fmtravel() 展平成行主序——W1[i * H + j] 表示输入 i 到神经元 j 的权重。这个布局选择直接决定下一步 SIMD 代码好不好写。

四、步骤二:Mojo 推理内核

dense 层:复用第十篇的 matmul 骨架

单样本 dense 就是「向量 × 矩阵」。关键洞察和矩阵乘法一致:固定输入维度 i 时,输出 out[:] += x[i] * W[i, :] 是一段连续内存,天生适合向量指令:

mojo
from memory import UnsafePointer

# y = x @ W + b,W 行主序: w[i * out_dim + j]
# 要求 out_dim 是 width 的整数倍(教学版约定)
def dense_simd(x: UnsafePointer[Float64],
               w: UnsafePointer[Float64],
               b: UnsafePointer[Float64],
               out: UnsafePointer[Float64],
               in_dim: Int, out_dim: Int):
    # 先把偏置铺进去:out = b
    for j in range(0, out_dim, 8):
        out.store[width=8](j, b.load[width=8](j))

    # 累加每个输入分量的贡献:out += x[i] * W[i, :]
    for i in range(in_dim):
        var xi = x[i]
        for j in range(0, out_dim, 8):
            var vw = w.load[width=8](i * out_dim + j)   # W 第 i 行,8 个连读
            var vo = out.load[width=8](j)
            out.store[width=8](j, vo + xi * vw)

激活与输出:小维度不值得向量化

隐藏层 16 维走 SIMD;输出层只有 3 类——宽度 8 的向量指令装不满 3 个数,向量化反而亏,老老实实标量循环:

mojo
def relu(v: UnsafePointer[Float64], size: Int):
    for i in range(size):
        if v[i] < 0.0:
            v.store(i, 0.0)

def dense_scalar(x: UnsafePointer[Float64],
                 w: UnsafePointer[Float64],
                 b: UnsafePointer[Float64],
                 out: UnsafePointer[Float64],
                 in_dim: Int, out_dim: Int):
    for j in range(out_dim):
        var acc = b[j]
        for i in range(in_dim):
            acc += x[i] * w[i * out_dim + j]
        out.store(j, acc)

def argmax(v: UnsafePointer[Float64], size: Int) -> Int:
    var best = 0
    var best_val = v[0]
    for i in range(1, size):
        if v[i] > best_val:
            best_val = v[i]
            best = i
    return best

组装:单样本推理

mojo
def predict(x: UnsafePointer[Float64],          # 4 维输入
            w1: UnsafePointer[Float64], b1: UnsafePointer[Float64],
            w2: UnsafePointer[Float64], b2: UnsafePointer[Float64],
            hidden: UnsafePointer[Float64],      # 16 维工作缓冲
            logits: UnsafePointer[Float64]) -> Int:
    dense_simd(x, w1, b1, hidden, 4, 16)         # Dense(4→16),SIMD
    relu(hidden, 16)                              # ReLU
    dense_scalar(hidden, w2, b2, logits, 16, 3)   # Dense(16→3),标量
    return argmax(logits, 3)                      # 跳过 softmax,直接 argmax

注意 hidden/logits 是调用方传入的工作缓冲——推理内核不分配内存,热路径上零 malloc,这是低延迟代码的基本素养。

五、步骤三:验证正确性

把训练脚本打印的四行权重常量粘进来(此处值仅示意,以你训练输出为准):

mojo
from memory import UnsafePointer

def main():
    var W1: List[Float64] = [0.123456, ...]   # 4×16 = 64 个值
    var b1: List[Float64] = [0.0, ...]        # 16 个值
    var W2: List[Float64] = [0.123456, ...]   # 16×3 = 48 个值
    var b2: List[Float64] = [0.0, ...]        # 3 个值

    # 一条标准化后的 iris 样本(setosa 的典型值)
    var x: List[Float64] = [-0.9, 1.02, -1.34, -1.31]

    # 分配工作缓冲
    var hidden = List[Float64](16)
    var logits = List[Float64](3)

    var pred = predict(x.ptr(), W1.ptr(), b1.ptr(), W2.ptr(), b2.ptr(),
                       hidden.ptr(), logits.ptr())
    print("预测类别:", pred)      # 0 = setosa

验证标准:同一批样本,Mojo 推理的准确率必须和 Python 侧完全一致(都是约 96%~98%)。任何一条对不上,先查权重布局(行主序 ravel)和标准化参数——这两个是经典翻车点。

.ptr() 拿到底层指针后,从这一行往下就是纯原生代码——没有解释器、没有对象包装、没有动态分发。

六、步骤四:性能对比

单样本推理正是 Python 的痛点场景。用第八篇的 benchmark 思路对比两种实现:

  • Python 侧np.maximum(0, X @ W1 + b1) @ W2 + b2 逐条调用
  • Mojo 侧predict() 逐条调用

预期量级(具体数值因机器而异):

场景Python + numpyMojo 内核
单样本推理数十微秒(解释器开销主导)数百纳秒(纯机器码)
批量 GEMM(如 512×512)BLAS 出战,通常更快单内核不占优

第二条要诚实:大批量矩阵乘是 BLAS 的主场,别硬碰。Mojo 推理内核的价值在批量=1 的在线场景——机器人 1kHz 控制环里的策略网络、边缘设备的逐帧感知、推荐服务的单请求打分。这些场景的延迟预算是微秒级,Python 的固定开销是不可承受之重,而它们恰恰是你写一份 Mojo 内核就能解决的地方。

七、工程要点回顾

这个案例浓缩了几条可迁移的实战经验:

  1. 权重布局先行:行主序/列主序的决定权在你手里,选那个让 SIMD 连续读的
  2. 按维度选实现:16 维用 SIMD、3 维用标量——向量化不是信仰,是算账
  3. 推理砍尾巴:argmax 替代 softmax,能省的 exp 一个不留
  4. 热路径零分配:工作缓冲复用,malloc 不进循环
  5. 一致性验证是底线:换实现不换结果,先对齐再谈优化

八、这个案例之后

到这里你已经走完「Mojo 语言学习 → 数值内核 → 端到端推理」的完整链路。往前的方向:

  • 接上你的模型:把 miniGPT(见 从零搭建小型语言模型)的 embedding + 前向层换成 Mojo 内核,就是一个小型推理引擎的雏形
  • GPU 内核:语言层的功夫是通用的,GPU 编程随 MAX 平台提供
  • 服务化:配合 Python 互操作把内核包成服务,团队无感升级

系列导航环境搭建变量类型函数structtrait所有权编译期参数SIMDPython互操作矩阵乘法本篇(推理案例)

延伸阅读Mojo 1.0:AI 时代的新编程语言? —— 语言背景与发展判断

MIT License.