为什么用 C++,以及 ML 框架是如何工作的 本书里的每一次 、每一个 、每一个 调用,底层其实都在执行 C++ 与 CUDA 代码。本节将拉开这层幕布:解释 ML 框架为什么这样搭建,为 Python 工程师快速补齐 C++ 基础,告诉你什么时候该写自定义 C++ 内核,以及如何把 C++ 绑定进 Python——这是「你写的代码」与「真正运行它的硬件」之间的桥梁。 你已经用了 15 章的篇幅写 Python。你 import 过 JAX,调用过 ,跑过训练循环,搭过模型。这一切感觉都很 Python。但真相是:真正发生的计算几乎没有一项是在 Python 里完成的。 当你在 PyTorch 里写 ,或在 JAX 里写 时,Python 几乎什么都没干。
本书里的每一次 jnp.matmul、每一个 torch.nn.Linear、每一个 np.dot 调用,底层其实都在执行 C++ 与 CUDA 代码。本节将拉开这层幕布:解释 ML 框架为什么这样搭建,为 Python 工程师快速补齐 C++ 基础,告诉你什么时候该写自定义 C++ 内核,以及如何把 C++ 绑定进 Python——这是「你写的代码」与「真正运行它的硬件」之间的桥梁。
你已经用了 15 章的篇幅写 Python。你 import 过 JAX,调用过 jax.grad,跑过训练循环,搭过模型。这一切感觉都很 Python。但真相是:真正发生的计算几乎没有一项是在 Python 里完成的。
当你在 PyTorch 里写 output = model(input),或在 JAX 里写 output = jnp.matmul(W, x) 时,Python 几乎什么都没干。它只是构造出一份「计算描述」(一张由操作组成的图),然后把它交给一个 C++/CUDA 后端,由后者真正干活。Python 是方向盘,C++ 才是发动机。
| Python | C++ | |
|---|---|---|
| 开发速度 | 快(动态类型、REPL、无需编译) | 慢(静态类型、头文件、编译耗时) |
| 执行速度 | 比 C 慢约 100 倍(解释执行、GIL) | 接近硬件速度(编译执行、无额外开销) |
| 内存控制 | 自动(GC),无法控制布局 | 手动,对每个字节精确控制 |
| 硬件访问 | 无(不支持 SIMD、GPU、自定义内存) | 全面支持(内联函数、CUDA、内联汇编) |
| 生态 | ML 生态丰富(notebook、可视化、数据) | 系统生态丰富(操作系统、驱动、引擎) |
关键洞见:让每种语言干它擅长的事。Python 负责那些「人的生产力重要」的部分(实验设计、超参搜索、数据探索);C++ 负责那些「机器性能重要」的部分(矩阵乘法、卷积、注意力内核)。
一次矩阵乘法 jnp.matmul(A, B),当 A 是 4096 \times 4096 时,大约要做 1370 亿次浮点运算。用纯 Python(嵌套循环)来跑,大约要 30 分钟。而用 AVX-512 SIMD 加多线程的优化 C++,只要大约 10 毫秒。这是 18 万倍的差距。任何 Python 层面的聪明技巧都填不平这个鸿沟。
用户代码(Python) ↓ Python API 层(torch.nn、jax.numpy、numpy) ↓ 调度 / JIT 编译器(torch.compile、XLA、NumPy 调度) ↓ C++ 内核库(ATen/PyTorch、XLA、BLAS/LAPACK) ↓ 硬件相关后端(CUDA、cuDNN、MKL、oneDNN、Metal) ↓ 硬件(CPU 的 SIMD 单元、GPU 核心、TPU 的 MXU)
NumPy 的核心是用 C 写的。当你调用 np.dot(A, B) 时,Python 会调一个 C 函数,后者再调 BLAS(Basic Linear Algebra Subprograms,基本线性代数子程序),通常是 Intel MKL 或 OpenBLAS。BLAS 是手工优化的 C 与 Fortran 代码,使用 SIMD 指令、缓存友好的访存模式和多线程。把矩阵乘法做到这么快,背后是几十年的优化积累。
NumPy 只支持 CPU,不支持 GPU。但在 CPU 上它极快,因为它把活儿交给了能找到的最好的 BLAS 实现。
PyTorch 的计算引擎叫 ATen(A Tensor Library,张量库),用 C++ 写成。ATen 实现了约 2000 个张量操作(add、matmul、conv2d、softmax……),每个操作都有 CPU 和 CUDA 两套后端。
当你调用 torch.matmul(A, B) 时:
torch.compile(PyTorch 2.0+)更进一步:它会追踪你的 Python 代码,构建计算图,然后用 Triton(GPU)或 C++/OpenMP(CPU)编译。编译后的代码会融合操作、消除 Python 开销,可能比 eager 模式快 2-5 倍。
JAX 会把 Python 函数编译成 XLA(Accelerated Linear Algebra,加速线性代数),这是 Google 为 ML 工作负载写的编译器。当你对函数加上 jax.jit 时:
这就是为什么 jax.jit 如此重要:没有它,每个操作都是一次 Python→C++ 的往返;有了它,整个函数就变成一个被编译好的单一内核。
// C++ 需要显式声明类型(不像 Python) int count = 0; // 32 位整数 float loss = 0.5f; // 32 位浮点 double lr = 3e-4; // 64 位浮点 bool training = true; // 布尔值 // 数组(固定大小,分配在栈上) float weights[1024]; // 1024 个 float,在内存中连续 // 指针:一个保存内存地址的变量 float* ptr = weights; // ptr 指向 weights 的第一个元素 float val = ptr[42]; // 通过指针算术访问第 42 个元素 // ptr[42] 等价于 *(ptr + 42)
// 函数声明:返回类型 函数名(参数类型 参数名) float relu(float x) { return x > 0.0f ? x : 0.0f; } // 按引用传递(避免拷贝大对象) void scale_vector(std::vector<float>& vec, float factor) { for (size_t i = 0; i < vec.size(); i++) { vec[i] *= factor; } } // const 引用:只读,不拷贝 float sum(const std::vector<float>& vec) { float total = 0.0f; for (float x : vec) { // 范围 for 循环(类似 Python 的 for x in vec) total += x; } return total; }
// 栈分配:快,生命周期自动(函数返回时自动释放) float buffer[256]; // 栈上 256 个 float // 堆分配:手动,生命周期可超出函数 float* data = new float[n]; // 在堆上分配 n 个 float // ... 使用 data ... delete[] data; // 你必须自己释放(没有垃圾回收器) // 现代 C++:智能指针(自动清理,类似 Python 的引用) #include <memory> auto data = std::make_unique<float[]>(n); // 出作用域时自动释放
// 一个适用于任意数值类型的函数 template <typename T> T add(T a, T b) { return a + b; } add<float>(1.5f, 2.5f); // 返回 4.0f add<int>(3, 4); // 返回 7
#include <vector> // 动态数组(类似 Python 的 list) #include <string> // 字符串类型 #include <unordered_map> // 哈希表(类似 Python 的 dict) #include <algorithm> // sort、find、transform 等 #include <cmath> // 数学函数 std::vector<float> vec = {1.0f, 2.0f, 3.0f}; vec.push_back(4.0f); // 追加 float first = vec[0]; // 索引 size_t len = vec.size(); // 长度 std::unordered_map<std::string, int> counts; counts["hello"] = 5; // 插入 if (counts.count("hello")) { } // 检查是否存在
框架里没有你需要的操作:一个新的激活函数、一种自定义的注意力模式、一个用现有操作组合无法表达的损失函数。
为了性能而融合操作:你的模型做 relu(layernorm(matmul(x, W) + b))。每个操作都会启动一个独立的内核,读写内存并同步。一个融合内核能一次性做完,避免内存往返。这能快 2-5 倍。
减少内存占用:自定义内核可以在不保存所有中间激活值的情况下计算梯度(内核层面的梯度检查点)。
针对新硬件:一个新加速器(如 Cerebras、Groq)可能还没有框架支持,你得直接写内核。
// my_ops.cpp #include <pybind11/pybind11.h> #include <pybind11/numpy.h> namespace py = pybind11; // 一个简单的自定义操作 py::array_t<float> custom_relu(py::array_t<float> input) { auto buf = input.request(); float* ptr = static_cast<float*>(buf.ptr); size_t n = buf.size; auto result = py::array_t<float>(n); float* out = static_cast<float*>(result.request().ptr); for (size_t i = 0; i < n; i++) { out[i] = ptr[i] > 0 ? ptr[i] : 0; } return result; } PYBIND11_MODULE(my_ops, m) { m.def("custom_relu", &custom_relu, "Custom ReLU operation"); }
# 编译 pip install pybind11 c++ -O3 -shared -std=c++17 -fPIC $(python3 -m pybind11 --includes) my_ops.cpp -o my_ops$(python3-config --extension-suffix)
# 从 Python 使用 import my_ops import numpy as np x = np.array([-1.0, 2.0, -3.0, 4.0], dtype=np.float32) y = my_ops.custom_relu(x) print(y) # [0. 2. 0. 4.]
// custom_op.cpp #include <torch/extension.h> torch::Tensor custom_gelu(torch::Tensor x) { return x * 0.5 * (1.0 + torch::erf(x / std::sqrt(2.0))); } PYBIND11_MODULE(TORCH_EXTENSION_NAME, m) { m.def("custom_gelu", &custom_gelu, "Custom GELU activation"); }
# 边编译边加载 from torch.utils.cpp_extension import load custom_ops = load( name="custom_ops", sources=["custom_op.cpp"], extra_cflags=["-O3"], ) x = torch.randn(1000) y = custom_ops.custom_gelu(x)
torch.utils.cpp_extension.load 一次调用就完成编译 C++ 代码、生成共享库、加载为 Python 模块三件事。这是在 PyTorch 中试验自定义 C++ 操作最简单的方式。JAX 使用 XLA 的 custom call。流程更复杂(你要向 XLA 注册一个 C 函数),但概念相同:写 C/C++、绑定、从 Python 调用。
对大多数 JAX 用户来说,Pallas(见第 05 节)是更好的选择:它让你用类似 Python 的语法写 GPU 内核,由 XLA 编译,全程不离开 JAX 生态。
本节讲清了 Python 与硬件之间的那一层。本章其余文件会深入下去:
这条进路对应着抽象阶梯:C++ 内联函数(最低层,控制力最强)→ CUDA(GPU 专属)→ Triton/Pallas(Python 风格,编译生成)→ JAX/PyTorch(最高层,全自动)。每上一层都用一部分控制力换一份便利。理解了底层,你才能更好地使用高层。
// task1_basics.cpp // 编译:g++ -O3 -o task1 task1_basics.cpp // 运行:./task1 #include <iostream> #include <chrono> #include <vector> int main() { const int N = 10'000'000; // C++ 允许把 ' 当作数字分隔符 std::vector<float> data(N); // 填充数组 for (int i = 0; i < N; i++) { data[i] = static_cast<float>(i) * 0.001f; } // 求和 auto start = std::chrono::high_resolution_clock::now(); float sum = 0.0f; for (int i = 0; i < N; i++) { sum += data[i]; } auto end = std::chrono::high_resolution_clock::now(); double elapsed = std::chrono::duration<double, std::milli>(end - start).count(); std::cout << "Sum: " << sum << std::endl; std::cout << "Time: " << elapsed << " ms" << std::endl; std::cout << "Elements: " << N << std::endl; std::cout << "Throughput: " << (N * sizeof(float)) / elapsed / 1e6 << " GB/s" << std::endl; return 0; }
// task2_relu.cpp // 编译:c++ -O3 -shared -std=c++17 -fPIC $(python3 -m pybind11 --includes) \ // task2_relu.cpp -o my_relu$(python3-config --extension-suffix) #include <pybind11/pybind11.h> #include <pybind11/numpy.h> namespace py = pybind11; py::array_t<float> cpp_relu(py::array_t<float> input) { auto buf = input.request(); float* ptr = static_cast<float*>(buf.ptr); int n = buf.size; auto result = py::array_t<float>(n); float* out = static_cast<float*>(result.request().ptr); for (int i = 0; i < n; i++) { out[i] = ptr[i] > 0.0f ? ptr[i] : 0.0f; } return result; } PYBIND11_MODULE(my_relu, m) { m.def("relu", &cpp_relu, "C++ ReLU"); }
# test_relu.py —— 编译好上面的 C++ 模块后运行 import numpy as np import time import my_relu # 编译出来的 C++ 模块 x = np.random.randn(10_000_000).astype(np.float32) # C++ ReLU start = time.time() for _ in range(100): y_cpp = my_relu.relu(x) cpp_time = (time.time() - start) / 100 # NumPy ReLU start = time.time() for _ in range(100): y_np = np.maximum(x, 0) np_time = (time.time() - start) / 100 print(f"C++ ReLU: {cpp_time*1000:.2f} ms") print(f"NumPy ReLU: {np_time*1000:.2f} ms") print(f"Match: {np.allclose(y_cpp, y_np)}")
// task3_layout.cpp // 编译:g++ -O3 -o task3 task3_layout.cpp #include <iostream> #include <chrono> #include <vector> int main() { const int N = 4096; std::vector<float> matrix(N * N, 1.0f); // 行优先访问:地址连续(缓存友好) auto start = std::chrono::high_resolution_clock::now(); float sum_row = 0.0f; for (int i = 0; i < N; i++) { for (int j = 0; j < N; j++) { sum_row += matrix[i * N + j]; // 步长为 1 的访问 } } auto end = std::chrono::high_resolution_clock::now(); double row_ms = std::chrono::duration<double, std::milli>(end - start).count(); // 列优先访问:步长为 N(对缓存不友好) start = std::chrono::high_resolution_clock::now(); float sum_col = 0.0f; for (int j = 0; j < N; j++) { for (int i = 0; i < N; i++) { sum_col += matrix[i * N + j]; // 步长为 N 的访问(缓存未命中!) } } end = std::chrono::high_resolution_clock::now(); double col_ms = std::chrono::duration<double, std::milli>(end - start).count(); std::cout << "Row-major (cache-friendly): " << row_ms << " ms" << std::endl; std::cout << "Col-major (cache-hostile): " << col_ms << " ms" << std::endl; std::cout << "Slowdown: " << col_ms / row_ms << "x" << std::endl; std::cout << "(Both sums: " << sum_row << ", " << sum_col << ")" << std::endl; return 0; }