Skip to content
当前页大纲

Mojo 实战:手写矩阵乘法并对比性能

学习模块收官篇。把前面九篇的东西串起来——用 Mojo 实现矩阵乘法,从朴素版本一路优化,并用 benchmark 量化每步收益。矩阵乘法是深度学习的「Hello World」,也是检验一门语言数值能力的标准考场。

一、任务定义

C = A × B,其中 A 是 M×K,B 是 K×N,C 是 M×N。用一维数组按行存储(和 numpy 的默认布局一致):

text
A[i, k] = a[i * K + k]
B[k, j] = b[k * N + j]
C[i, j] = c[i * N + j]

二、版本一:朴素三重循环

和 Python 翻译过来的一样,只是加了类型:

mojo
def matmul_naive(a: List[Float64], b: List[Float64],
                 m: Int, n: Int, k: Int) -> List[Float64]:
    var c = List[Float64]()
    for _ in range(m * n):
        c.append(0.0)

    for i in range(m):
        for j in range(n):
            var acc = 0.0
            for t in range(k):
                acc += a[i * k + t] * b[t * n + j]
            c[i * n + j] = acc
    return c

注意内层循环用了累加器而不是 c[i*n+j] += ...——每次都写回数组会多一次访存。这个版本已经是纯编译型代码,比 Python 快百倍量级,但离硬件极限还很远。

三、版本二:SIMD 向量化内层循环

内层循环对 b 的访问是连续的——换个角度:把 j 维度向量化。

观察:固定 i 和 t 时,C[i, :] += A[i,t] * B[t, :]——C 和 B 的这一行都是连续内存,天生适合 SIMD:

mojo
from memory import UnsafePointer

def matmul_simd(a: UnsafePointer[Float64],
                b: UnsafePointer[Float64],
                c: UnsafePointer[Float64],
                m: Int, n: Int, k: Int):
    var width = 8   # 向量宽度,可换成 simd_width_of 查询的值

    # 先清零
    for i in range(m * n):
        c.store(i, 0.0)

    for i in range(m):
        for t in range(k):
            var a_it = a[i * k + t]
            for j in range(0, n, width):
                # 一次取 B 的8个、C 的8个
                var vb = b.load[width=8](t * n + j)
                var vc = c.load[width=8](i * n + j)
                c.store[width=8](i * n + j, vc + a_it * vb)

两个关键变化:

  1. 循环重排k 循环提到 j 外面——每次内层迭代做「C 的一行 += 标量 × B 的一行」,两次连续访存
  2. load/store 向量化:一次搬 8 个 float64,乘加都在向量寄存器里

这一步在多数 CPU 上能再拿 4~8 倍。

四、版本三:进一步的方向(思路)

继续压榨还有三张牌,思路给你,实现留给练习:

  1. 寄存器分块:把 C 的一小块(如 8×8)整体驻留在寄存器/缓存中,A、B 分块流入——减少 C 的重复读写,这是经典 GEMM 优化的核心
  2. 多线程并行:外层 i 循环天然无依赖,可用 Mojo 的并行能力分摊到多核
  3. 编译期参数化宽度:把 width 变成编译期参数 [width: Int],用 vectorize 处理尾部,为不同 CPU 生成专属代码

走到版本三,对标 numpy(底层是数十年打磨的 BLAS)虽然还有差距,但已经进入了「一门语言、一份代码」就能触碰的区间——这在过去需要 CUDA/C++ 和专门的内核工程能力。

五、用 benchmark 量化

把三个版本包进基准测试(第八篇的 benchmark.run):

mojo
from benchmark import run

def main():
    var m = 256
    var n = 256
    var k = 256
    # ... 准备数据、初始化指针 ...

    var report_naive = run(lambda: matmul_naive(a_list, b_list, m, n, k))
    print("朴素版 (ms):", report_naive.mean / 1e6)

    var report_simd = run(lambda: matmul_simd(a_ptr, b_ptr, c_ptr, m, n, k))
    print("SIMD 版 (ms):", report_simd.mean / 1e6)

预期量级(256³ 规模,具体数值因机器而异):

版本相对性能
Python 三重循环1×(基准,分钟级)
Mojo 朴素版~100×
Mojo SIMD 版~400-800×
numpy (BLAS)~1000×+

重点不是打败 BLAS,而是看清:从 Python 到 Mojo 朴素版的一百倍,你只是换了个语言零成本获得;从朴素到 SIMD 的几倍,是理解硬件换来的——Mojo 把这条路上的每一级台阶都修好了。

六、和 Python 生态合流

实战的最后一步,把成果接回 Python(第九篇):

mojo
from std.python import Python
from std.python.numpy import copy_to_numpy_array

def main() raises:
    # ... 跑完 matmul_simd,结果在 c_ptr 指向的缓冲 ...

    var np = Python.import_module("numpy")
    # 结果转 numpy 后就能进 torch / sklearn / matplotlib 的世界

至此闭环:Python 的生态 + Mojo 的内核,这正是 Mojo 给 AI 开发者的完整故事。

七、学习路线总结

十篇走完,你已经具备:

能力对应篇目
环境与工具链1
类型系统与语法2、3
struct / trait 组织代码4、5
内存安全(所有权)6
编译期魔法(参数/comptime)7
性能编程(SIMD/benchmark)8
生态融合(Python 互操作)9
综合实战10

后续进阶方向:GPU 编程(配合 MAX 平台)、并行计算、用 Mojo 写推理服务内核。语言层的东西,你已经全部在手里了。


配套阅读Mojo 1.0:AI 时代的新编程语言? —— 语言背景、开源生态与发展判断。

MIT License.