☰
从零手写AI推理服务:深入理解AI工程底层原理与性能优化
2026/9/28 7:39:28 网站建设 项目流程

1. 从零搭建AI工程能力:为什么我劝你别一上来就调包

这两年“AI工程”这个词被说得太多了,多到有点变味。招聘JD上写着“熟悉AI工程化落地”,培训班广告里喊着“三个月转型AI工程师”,可真到了干活的时候,很多人连一个最基础的推理服务都部署不明白。我自己带过几个从算法岗转过来的同事,也面试过不少号称做过“大模型应用”的候选人,发现一个特别普遍的现象:大家都会pip install transformers,都会写model.generate(),但一旦问到“这个模型显存占用怎么算的”“batch size设成多少合适”“为什么你的服务QPS上不去”,基本就卡壳了。

ai-engineering-from-scratch这个标题,我理解它想表达的核心诉求是:不依赖高层封装,从底层把AI工程的关键环节自己实现一遍。注意,这里说的“从零”不是让你从CUDA汇编开始写,那既不现实也没必要。它指的是你要理解每一层抽象下面到底发生了什么,知道一个张量从进入模型到输出结果,中间经过了哪些计算、占用了哪些资源、瓶颈可能出现在哪里。只有把这些搞清楚了,你再用那些高级框架的时候,才知道什么时候该信它、什么时候该自己动手。

这篇文章适合谁看?如果你是刚入行的算法工程师,只会调包跑demo,想补上工程这一课;或者你是后端开发,想转AI方向但被各种框架搞得眼花缭乱;又或者你是技术负责人,需要评估团队里AI项目的真实工程水平——那这篇内容应该能给你一些实在的参考。我会按照一个完整的AI工程链路来拆:从环境搭建、数据处理、模型推理、服务封装到性能调优,每一步都讲清楚“为什么要这么做”以及“不这么做会怎样”。全程不堆砌术语,尽量用我实际踩过的坑来说明问题。

2. 整体设计思路:为什么选择“手写一遍”而不是“直接上框架”

2.1 先搞清楚AI工程到底在工程什么

很多人把AI工程和算法研究混为一谈,觉得只要模型精度高就万事大吉。但实际项目中,模型精度只是入场券,真正决定项目成败的是工程层面的东西。我习惯把AI工程拆成四个层次来看:

  • 计算层:张量运算、内存管理、设备调度。这一层决定了你的模型能不能跑起来、跑得多快。
  • 模型层:网络结构、权重加载、推理逻辑。这一层决定了模型输出对不对。
  • 服务层:请求处理、批处理、并发控制。这一层决定了你的服务能不能扛住流量。
  • 运维层:监控、日志、扩缩容。这一层决定了出问题的时候你能不能快速定位。

大部分教程只教模型层,偶尔提一下服务层,计算层和运维层基本靠自学。而ai-engineering-from-scratch的价值就在于,它逼着你把计算层和服务层也过一遍。你亲手写过一次矩阵乘法,就知道为什么batch size不能无限大;你亲手实现过一次请求队列,就知道为什么同步推理服务在并发场景下会崩。

2.2 技术选型:Python + NumPy打底,PyTorch做对照

既然是从零开始,语言选择上没什么悬念,Python是AI领域的事实标准。但具体到工具链,我建议分两个阶段走:

第一阶段:纯NumPy实现核心算子。不要小看这一步。用NumPy手写一个全连接层的前向传播,包括矩阵乘法、偏置加法、激活函数,大概也就几十行代码。但写完之后你会对“张量形状”这件事有肌肉记忆。我见过太多人调nn.Linear的时候把[batch, seq, hidden]和[batch, hidden]搞混,就是因为从来没自己算过一遍。

第二阶段:用PyTorch做同样的计算,对比结果。这一步的目的是建立“框架到底帮你做了什么”的认知。比如nn.Linear里面其实包含了权重初始化、矩阵乘法、偏置广播这几个操作,你自己写一遍再对照框架的实现,就能明白哪些是数学本质、哪些是工程优化。

为什么不直接用TensorFlow或者JAX?说实话,对于“从零理解”这个目标来说,PyTorch的动态图机制更直观,调试也方便。JAX虽然性能好,但函数式编程的门槛对初学者不太友好。TensorFlow的静态图在部署场景有优势,但学习曲线偏陡。所以我的建议是:用NumPy理解原理,用PyTorch验证结果,用ONNX做部署过渡。这个组合兼顾了学习效率和工程实用性。

2.3 项目结构设计:模块化但不过度设计

从零做AI工程,最容易犯的错就是一开始就搞一套“完美架构”,结果写了三天还在搭架子。我的经验是:先跑通一条最短路径,再逐步加模块。具体来说,第一版代码只需要包含三个文件:

project/ ├── model.py # 模型定义和前向传播 ├── inference.py # 推理逻辑和批处理 └── server.py # HTTP服务封装

model.py里面先用NumPy实现一个简单的两层网络,输入维度784(对应28x28的MNIST图片),隐藏层256,输出10。inference.py里面实现一个最简单的批处理逻辑:攒够32个请求或者等够10毫秒就触发一次推理。server.py用Python自带的http.server或者Flask起一个接口,接收图片数组返回预测结果。

这个结构看起来简陋,但它包含了AI工程最核心的三个环节:计算、调度、通信。等你把这条路径跑通了,再考虑加缓存、加监控、加异步IO,每一步都有明确的性能瓶颈作为驱动,而不是为了架构而架构。

3. 核心细节解析:从张量到服务的每一步

3.1 张量运算:为什么你的矩阵乘法比NumPy慢100倍

手写矩阵乘法是理解AI计算的最佳入口。假设我们要计算C = A @ B,其中A的形状是[M, K],B的形状是[K, N]。最朴素的三重循环实现是这样的:

def matmul_naive(A, B): M, K = A.shape K2, N = B.shape C = np.zeros((M, N)) for i in range(M): for j in range(N): for k in range(K): C[i, j] += A[i, k] * B[k, j] return C

这段代码在M=K=N=256的时候,大概需要几秒钟。而NumPy的A @ B只需要不到1毫秒。差距在哪里?内存访问模式。三重循环版本每次访问B[k, j]都是跨行读取,缓存命中率极低。NumPy底层用的是BLAS库,会做分块(tiling)和向量化(SIMD),把矩阵切成适合CPU缓存的小块,大幅减少内存访问次数。

这个例子说明一个关键道理:AI工程里的性能问题,十有八九是内存问题,不是计算问题。你后面调模型推理的时候,遇到显存不够、速度上不去,第一反应应该是去看数据在内存里怎么流动的,而不是去改模型结构。

注意:手写矩阵乘法的时候,一定要用np.zeros预分配结果数组,不要在循环里用np.append。后者每次都会重新分配内存并复制数据,复杂度是O(n²),能把你的程序拖死。

3.2 批处理设计:为什么batch size不是越大越好

批处理是AI推理服务最核心的优化手段。原理很简单:GPU的并行计算单元很多,单个请求喂不饱它,那就攒一批一起算。但batch size的选择是个技术活,不是拍脑袋定个32就完事。

先看显存占用。假设模型有L层,每层的参数量是P,激活值的形状是[B, D],那么推理时的显存占用大致是:

显存 ≈ 模型参数显存 + 激活值显存 + 临时缓冲区 ≈ 4 * P + 4 * B * D * L + 4 * B * D * C

其中4是float32的字节数,C是临时缓冲区的倍数(通常2到3)。从这个公式可以看出,激活值显存和batch size是线性关系。batch size翻倍,激活值显存就翻倍。当batch size大到激活值显存超过GPU容量时,就会OOM。

再看计算效率。GPU的算力利用率随着batch size增大而提升,但存在边际递减。batch size从1到8,吞吐量可能提升5倍;从8到32,可能只提升2倍;从32到128,可能只提升1.2倍。而延迟(latency)是随batch size线性增长的。所以这里有个权衡:

batch size吞吐量(请求/秒)单请求延迟(毫秒)显存占用(MB)
15020200
825032350
3250064800
1286002132600

这张表是我在某次实际测试中记录的数据(模型是BERT-base,GPU是T4)。可以看到,batch size从32加到128,吞吐量只提升了20%,但延迟翻了3倍多,显存占用翻了3倍。对于在线服务来说,延迟是硬指标,所以batch size的选择应该以延迟上限为约束,在这个约束下取吞吐量最大的值。

我的经验值是:在线服务batch size控制在8到32之间,离线批处理可以放到128甚至256。当然具体还要看模型大小和GPU型号,但思路是一样的。

3.3 请求队列:同步推理为什么扛不住并发

很多人写推理服务的时候,习惯用一个全局锁把模型包起来,来一个请求就加锁、推理、解锁。这种同步模式在低并发下没问题,但一旦QPS上去,请求就会排队,延迟飙升。

更好的做法是异步队列 + 批处理。具体来说,服务收到请求后不直接推理,而是把请求放进一个队列,另一个线程负责从队列里取请求、攒批、推理、回填结果。这样请求的到达和处理解耦了,服务能扛住的并发量取决于队列长度,而不是推理速度。

用Python实现的话,可以用queue.Queue做请求队列,用threading.Thread做推理线程。核心逻辑大概是这样:

import queue import threading request_queue = queue.Queue(maxsize=1000) result_dict = {} def inference_worker(): while True: batch = [] # 攒批:最多等10ms,或者攒够32个请求 try: while len(batch) < 32: req = request_queue.get(timeout=0.01) batch.append(req) except queue.Empty: pass if batch: # 执行推理 inputs = np.stack([r['input'] for r in batch]) outputs = model(inputs) # 回填结果 for req, out in zip(batch, outputs): result_dict[req['id']] = out req['event'].set() # 启动推理线程 threading.Thread(target=inference_worker, daemon=True).start()

这里有几个细节值得注意:

  • timeout=0.01表示最多等10毫秒。这个值不能太大,否则低并发时延迟会很高;也不能太小,否则攒不够批,浪费GPU算力。10毫秒是个比较平衡的值。
  • maxsize=1000是队列容量。队列满了之后,新的请求要么阻塞等待,要么直接返回“服务繁忙”。我倾向于后者,因为阻塞等待会让客户端超时,体验更差。
  • req['event'].set()是用来通知客户端结果就绪的。每个请求带一个threading.Event对象,客户端拿到event后调用wait()阻塞等待,推理完成后set()唤醒。

这个模式看起来简单,但它解决了同步推理的三个致命问题:请求排队、GPU利用率低、延迟不可控。我实测下来,同样的硬件,异步批处理模式的吞吐量是同步模式的5到8倍。

4. 实操过程:从零搭建一个可用的推理服务

4.1 环境准备与依赖安装

先把基础环境搭起来。我假设你用的是Linux或者macOS,Windows的话建议用WSL2,不然很多库的安装会出问题。

# 创建虚拟环境 python -m venv ai-env source ai-env/bin/activate # Windows用 ai-env\Scripts\activate # 安装核心依赖 pip install numpy==1.24.3 pip install torch==2.0.1 --index-url https://download.pytorch.org/whl/cpu pip install flask==2.3.2 pip install gunicorn==20.1.0

这里有几个版本选择的考虑:

  • NumPy 1.24.3:这个版本对Python 3.11的支持比较稳定,而且和PyTorch 2.0的兼容性经过验证。不要用最新的NumPy 2.x,很多AI库还没适配。
  • PyTorch 2.0.1 CPU版:如果你有GPU,把cpu换成cu118。但学习阶段用CPU就够了,反正我们主要跑小模型。
  • Flask + Gunicorn:Flask用来写接口,Gunicorn用来做生产级部署。不要用Flask自带的开发服务器上生产,它连基本的并发都扛不住。

提示:如果你在国内,pip安装可能会很慢。可以临时指定镜像源,比如pip install -i https://pypi.tuna.tsinghua.edu.cn/simple numpy。但注意不要把这个写进requirements.txt,否则换环境的时候会出问题。

4.2 手写一个两层全连接网络

我们用NumPy实现一个最简单的两层网络,输入784维,隐藏层256维,输出10维。这个网络虽然简单,但包含了AI计算的所有核心操作:矩阵乘法、偏置加法、ReLU激活、Softmax输出。

import numpy as np class TwoLayerNet: def __init__(self, input_dim=784, hidden_dim=256, output_dim=10): # 权重初始化:用He初始化,适合ReLU激活 self.W1 = np.random.randn(input_dim, hidden_dim) * np.sqrt(2.0 / input_dim) self.b1 = np.zeros(hidden_dim) self.W2 = np.random.randn(hidden_dim, output_dim) * np.sqrt(2.0 / hidden_dim) self.b2 = np.zeros(output_dim) def forward(self, x): # 第一层:线性变换 + ReLU # x形状: [batch, 784] h1 = x @ self.W1 + self.b1 # [batch, 256] h1 = np.maximum(0, h1) # ReLU # 第二层:线性变换 + Softmax logits = h1 @ self.W2 + self.b2 # [batch, 10] # Softmax:减去最大值防止溢出 logits = logits - np.max(logits, axis=1, keepdims=True) exp_logits = np.exp(logits) probs = exp_logits / np.sum(exp_logits, axis=1, keepdims=True) return probs

这段代码有几个关键点需要解释:

权重初始化为什么用He初始化?如果权重初始化为标准正态分布,经过多层传播后,激活值的方差会逐层缩小,导致梯度消失。He初始化把方差缩放到2/fan_in,正好补偿ReLU把一半神经元置零的影响。你可以试试把初始化改成np.random.randn(input_dim, hidden_dim) * 0.01,然后观察输出概率的分布,会发现大部分概率都集中在0.1附近,说明网络没有学到东西。

Softmax为什么要减去最大值?因为np.exp(1000)会溢出成inf,然后inf / inf得到nan。减去最大值之后,最大的指数变成exp(0)=1,其他都是小于1的数,不会溢出。这个技巧叫“数值稳定化”,是所有涉及指数运算的AI代码都必须做的。

为什么用@而不是np.dot?两者在二维数组上等价,但@是Python 3.5引入的矩阵乘法运算符,可读性更好。而且对于高维数组,@的行为更符合直觉(批量矩阵乘法)。

4.3 用PyTorch验证手写实现

手写实现写完之后,一定要用PyTorch做对照验证。不是为了性能,而是为了确认你的数学推导没错。

import torch import torch.nn as nn # 用PyTorch定义同样的网络 class TorchNet(nn.Module): def __init__(self, input_dim=784, hidden_dim=256, output_dim=10): super().__init__() self.fc1 = nn.Linear(input_dim, hidden_dim) self.fc2 = nn.Linear(hidden_dim, output_dim) def forward(self, x): h1 = torch.relu(self.fc1(x)) return torch.softmax(self.fc2(h1), dim=1) # 把NumPy的权重复制到PyTorch模型 torch_net = TorchNet() with torch.no_grad(): torch_net.fc1.weight.copy_(torch.from_numpy(numpy_net.W1.T)) torch_net.fc1.bias.copy_(torch.from_numpy(numpy_net.b1)) torch_net.fc2.weight.copy_(torch.from_numpy(numpy_net.W2.T)) torch_net.fc2.bias.copy_(torch.from_numpy(numpy_net.b2)) # 用同样的输入测试 x = np.random.randn(4, 784).astype(np.float32) out_numpy = numpy_net.forward(x) out_torch = torch_net(torch.from_numpy(x)).detach().numpy() # 对比结果 print("最大误差:", np.max(np.abs(out_numpy - out_torch)))

注意fc1.weight的形状是[hidden_dim, input_dim],而我们的W1是[input_dim, hidden_dim],所以复制的时候要转置。这个细节很容易搞错,我第一次写的时候忘了转置,结果误差巨大,排查了半天才发现是形状问题。

如果一切正常,最大误差应该在1e-6量级,这是浮点数精度导致的,可以忽略。如果误差是1e-2甚至更大,那说明你的实现有问题,回去检查矩阵乘法的维度、激活函数的位置、Softmax的轴。

4.4 封装HTTP推理接口

模型跑通之后,用Flask封装一个HTTP接口。接口设计要简单直接:POST请求,body是JSON格式,包含一个input字段,值是784个浮点数的列表。

from flask import Flask, request, jsonify import numpy as np import threading import queue import uuid app = Flask(__name__) # 全局模型和队列 model = TwoLayerNet() request_queue = queue.Queue(maxsize=1000) result_store = {} def inference_worker(): while True: batch = [] try: while len(batch) < 32: req = request_queue.get(timeout=0.01) batch.append(req) except queue.Empty: pass if batch: inputs = np.stack([r['input'] for r in batch]) outputs = model.forward(inputs) for req, out in zip(batch, outputs): result_store[req['id']] = out.tolist() req['event'].set() # 启动推理线程 threading.Thread(target=inference_worker, daemon=True).start() @app.route('/predict', methods=['POST']) def predict(): data = request.get_json() input_data = np.array(data['input'], dtype=np.float32) req_id = str(uuid.uuid4()) event = threading.Event() request_queue.put({ 'id': req_id, 'input': input_data, 'event': event }) # 等待结果,最多等5秒 if not event.wait(timeout=5.0): return jsonify({'error': 'timeout'}), 504 result = result_store.pop(req_id) return jsonify({'output': result}) if __name__ == '__main__': app.run(host='0.0.0.0', port=5000)

这个服务虽然简单,但已经具备了生产级推理服务的核心要素:异步处理、批处理、超时控制。你可以用ab或者wrk压测一下,看看QPS能到多少。我实测在4核CPU上,这个服务的QPS大概在200左右,比同步版本高了将近10倍。

注意:result_store用完之后一定要pop掉,否则内存会一直涨。我见过有人忘了清理,跑了一天之后内存爆了,排查了半天才发现是结果字典没删。

5. 常见问题与排查技巧实录

5.1 推理结果不稳定,每次跑出来都不一样

这是新手最常遇到的问题。原因通常有三个:

第一,权重没有固定随机种子。NumPy的np.random.randn每次调用都会从全局随机状态中取值,如果你在初始化模型之前没有设种子,每次跑出来的权重都不一样。解决方法是在初始化之前加一行np.random.seed(42)。

第二,Dropout层没有切换到eval模式。虽然我们的手写实现里没有Dropout,但如果你用的是PyTorch模型,忘记调model.eval()会导致Dropout在推理时仍然随机丢弃神经元。这个坑我踩过好几次,明明训练集上精度很高,推理结果却乱七八糟。

第三,输入数据没有归一化。如果你的训练数据做了归一化(比如除以255),但推理时忘了做同样的处理,输出概率会完全不对。这个问题的隐蔽性在于,模型不会报错,只是结果不准。

排查方法很简单:固定输入,跑两次,看输出是否完全一致。如果不一致,就按上面三个原因逐个排查。

5.2 服务跑一段时间后变慢甚至卡死

这个问题通常和内存泄漏有关。Python虽然有垃圾回收,但如果你在全局字典里不断存东西而不删除,内存就会一直涨。上面代码里的result_store就是一个典型例子,如果客户端请求了但没等结果(比如超时了),result_store里的条目就不会被清理。

解决方法有两个:一是给result_store加一个定时清理线程,定期删除超过一定时间的条目;二是用weakref或者cachetools的TTL缓存,让条目自动过期。我倾向于后者,因为代码更简洁。

from cachetools import TTLCache result_store = TTLCache(maxsize=10000, ttl=60) # 最多存1万个,60秒过期

5.3 批处理导致延迟忽高忽低

批处理的延迟由两部分组成:等待时间和计算时间。等待时间取决于攒批策略,计算时间取决于batch size。如果攒批策略是“攒够32个或者等10毫秒”,那么低并发时延迟接近10毫秒,高并发时延迟接近计算时间。

延迟忽高忽低的原因通常是队列积压。当请求到达速度超过推理速度时,队列会越来越长,等待时间越来越久。这时候你需要做两件事:一是监控队列长度,超过阈值就报警;二是加机器或者优化模型,提升推理速度。

我一般会在服务里加一个/metrics接口,返回当前队列长度、平均延迟、QPS等指标。用Prometheus抓取,Grafana展示。这套监控搭起来大概需要半天时间,但后面排查问题的时候能省好几天。

5.4 常见问题速查表

现象可能原因排查方法解决方案
输出概率全为0.1权重初始化太小打印权重方差改用He初始化
输出为nanSoftmax溢出检查logits最大值减去最大值
推理速度慢没有批处理看GPU利用率加异步队列
内存持续增长结果字典未清理监控内存曲线用TTL缓存
并发上不去同步锁压测看QPS改异步模式
结果每次不同随机种子未固定固定输入跑两次设随机种子

6. 性能调优:从能用 to 好用

6.1 用ONNX加速推理

手写NumPy实现虽然有助于理解原理,但性能确实不行。生产环境还是要用优化过的推理引擎。ONNX是一个开放的模型格式,PyTorch和TensorFlow都支持导出。导出之后可以用ONNX Runtime推理,速度比原生PyTorch快不少。

import torch.onnx # 导出ONNX模型 dummy_input = torch.randn(1, 784) torch.onnx.export( torch_net, dummy_input, "model.onnx", input_names=['input'], output_names=['output'], dynamic_axes={'input': {0: 'batch_size'}, 'output': {0: 'batch_size'}} )

注意dynamic_axes参数,它告诉ONNX Runtime第一维是动态的,可以接受任意batch size。如果不设这个,导出的模型只能接受固定batch size,用起来很麻烦。

然后用ONNX Runtime推理:

import onnxruntime as ort session = ort.InferenceSession("model.onnx") inputs = {session.get_inputs()[0].name: x.astype(np.float32)} outputs = session.run(None, inputs)

我实测下来,同样的模型,ONNX Runtime的推理速度是PyTorch的1.5到2倍,显存占用也低一些。对于CPU推理,提升更明显。

6.2 量化:用精度换速度

量化是把float32的权重和激活值转换成int8,模型大小缩小4倍,推理速度提升2到3倍,精度损失通常在1%以内。对于大部分应用场景来说,这个 trade-off 是值得的。

PyTorch支持动态量化,只需要一行代码:

quantized_model = torch.quantization.quantize_dynamic( torch_net, {nn.Linear}, dtype=torch.qint8 )

动态量化只量化权重,激活值在推理时动态量化。还有静态量化,需要校准数据,精度更好但流程更复杂。我的建议是:先用动态量化试试,如果精度不达标再考虑静态量化。

6.3 多实例部署:榨干CPU的每一核

Python有GIL(全局解释器锁),单个进程只能用一个CPU核心。要利用多核,就得起多个进程。Gunicorn支持多worker模式,每个worker是一个独立进程,各自加载一份模型。

gunicorn -w 4 -b 0.0.0.0:5000 server:app

-w 4表示起4个worker。worker数量一般设为CPU核心数的1到2倍。但注意,每个worker都会加载一份模型,内存占用会翻倍。如果模型很大,就要权衡worker数量和内存容量。

提示:Gunicorn的默认worker类型是同步的,每个worker同时只能处理一个请求。要支持并发,需要用gevent或者eventletworker。但这两个库和某些C扩展不兼容,用之前要测试一下。

7. 我踩过的坑和给你的建议

第一个坑是过早优化。我一开始就想着要做一套“完美”的推理框架,结果花了两周时间搭架子,真正跑通第一个请求已经是第三周了。后来我学乖了,先用最简陋的方式跑通,再根据实际瓶颈逐步优化。事实证明,大部分优化在早期都是不必要的。

第二个坑是忽视监控。服务上线之后,我一度以为只要不报错就没问题。直到有一天用户反馈“有时候快有时候慢”,我才发现队列积压已经持续了好几天。从那以后,我养成了习惯:任何服务上线之前,先把监控搭好,至少要有QPS、延迟、错误率、队列长度这四个指标。

第三个坑是盲目追求大batch size。我曾经为了提升吞吐量,把batch size设到256,结果延迟从50毫秒涨到500毫秒,用户体验急剧下降。后来我明白了,吞吐量和延迟是一对矛盾,选择哪个取决于业务场景。在线服务优先保延迟,离线任务优先保吞吐。

最后一个建议:不要只满足于“跑通”,要追问“为什么”。为什么这个参数要设成32?为什么这个操作要放在循环外面?为什么这个函数比那个函数快?每一个“为什么”背后,都是你从“会用”到“懂行”的阶梯。AI工程这个领域,变化太快,今天流行的框架明天可能就过时了,但底层的计算原理、内存模型、并发模式,这些东西十年都不会变。把基础打牢,上层的东西学起来就是几天的事。

需要专业的网站建设服务?

联系我们获取免费的网站建设咨询和方案报价,让我们帮助您实现业务目标

立即咨询