ARTICLE DETAIL

资讯详情

深耕网站视觉设计与运营推广的一线实战洞察。

从零手写AI推理服务:深入理解AI工程底层原理与性能优化

从零手写AI推理服务:深入理解AI工程底层原理与性能优化 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这段代码在MKN256的时候大概需要几秒钟。而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吞吐量请求/秒单请求延迟毫秒显存占用MB1502020082503235032500648001286002132600这张表是我在某次实际测试中记录的数据模型是BERT-baseGPU是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(maxsize1000) result_dict {} def inference_worker(): while True: batch [] # 攒批最多等10ms或者攒够32个请求 try: while len(batch) 32: req request_queue.get(timeout0.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(targetinference_worker, daemonTrue).start()这里有几个细节值得注意timeout0.01表示最多等10毫秒。这个值不能太大否则低并发时延迟会很高也不能太小否则攒不够批浪费GPU算力。10毫秒是个比较平衡的值。maxsize1000是队列容量。队列满了之后新的请求要么阻塞等待要么直接返回“服务繁忙”。我倾向于后者因为阻塞等待会让客户端超时体验更差。req[event].set()是用来通知客户端结果就绪的。每个请求带一个threading.Event对象客户端拿到event后调用wait()阻塞等待推理完成后set()唤醒。这个模式看起来简单但它解决了同步推理的三个致命问题请求排队、GPU利用率低、延迟不可控。我实测下来同样的硬件异步批处理模式的吞吐量是同步模式的5到8倍。4. 实操过程从零搭建一个可用的推理服务4.1 环境准备与依赖安装先把基础环境搭起来。我假设你用的是Linux或者macOSWindows的话建议用WSL2不然很多库的安装会出问题。# 创建虚拟环境 python -m venv ai-env source ai-env/bin/activate # Windows用 ai-env\Scripts\activate # 安装核心依赖 pip install numpy1.24.3 pip install torch2.0.1 --index-url https://download.pytorch.org/whl/cpu pip install flask2.3.2 pip install gunicorn20.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 GunicornFlask用来写接口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_dim784, hidden_dim256, output_dim10): # 权重初始化用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, axis1, keepdimsTrue) exp_logits np.exp(logits) probs exp_logits / np.sum(exp_logits, axis1, keepdimsTrue) 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_dim784, hidden_dim256, output_dim10): 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), dim1) # 把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(maxsize1000) result_store {} def inference_worker(): while True: batch [] try: while len(batch) 32: req request_queue.get(timeout0.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(targetinference_worker, daemonTrue).start() app.route(/predict, methods[POST]) def predict(): data request.get_json() input_data np.array(data[input], dtypenp.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(timeout5.0): return jsonify({error: timeout}), 504 result result_store.pop(req_id) return jsonify({output: result}) if __name__ __main__: app.run(host0.0.0.0, port5000)这个服务虽然简单但已经具备了生产级推理服务的核心要素异步处理、批处理、超时控制。你可以用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(maxsize10000, ttl60) # 最多存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}, dtypetorch.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工程这个领域变化太快今天流行的框架明天可能就过时了但底层的计算原理、内存模型、并发模式这些东西十年都不会变。把基础打牢上层的东西学起来就是几天的事。
返回列表