ResNet-34部署Flask内存占用过高问题排查与优化求助
ResNet-34 Flask部署内存占用过高问题排查与优化建议
问题描述
- 用Flask部署ResNet-34做图像分类,启动服务器加载模型时内存约1.5G,首次预测后内存跃升至3G并维持该水平
- ResNet-34参数量远少于YOLOv5s,但两者部署后内存占用相近(YOLOv5s约3.2G),无法理解该现象
- 希望排查代码问题并找到降低内存占用的方法
部署代码
# initialize the flask app app = flask.Flask(__name__) app.model = None device = torch.device("cuda:0" if torch.cuda.is_available() else "cpu") data_transform = transforms.Compose( [transforms.Resize(256), transforms.CenterCrop(224), transforms.ToTensor(), transforms.Normalize([0.485, 0.456, 0.406], [0.229, 0.224, 0.225])]) # read class_indict json_path = './class_indices.json' assert os.path.exists(json_path), "file: '{}' dose not exist.".format(json_path) with open(json_path, "r") as f: class_indict = json.load(f) def load_model(): """Load the pre-trained model, you can use your model just as easily. """ model = resnet34(num_classes=7).to(device) weights_path = './resnet34.pth' assert os.path.exists(weights_path), "file: '{}' dose not exist.".format(weights_path) model.load_state_dict(torch.load(weights_path, map_location=device)) model.eval() return model def load_image(path): assert os.path.exists(path), "file: '{}' dose not exist.".format(path) img = Image.open(path).convert("RGB") # [N, C, H, W] img = data_transform(img) # expand batch dimension img = torch.unsqueeze(img, dim=0) return img @app.route("/resnet", methods=['get', 'post']) def predict(): if app.model is None: app.model = load_model() path = unquote(request.args.get("path", "")) if path is None: return jsonify({'error': 'Path not provided'}), 400 image = load_image(path) with torch.no_grad(): # predict class output = torch.squeeze(app.model(image.to(device))).cpu() predict = torch.softmax(output, dim=0) predict_cla = torch.argmax(predict).numpy() torch.cuda.empty_cache() del image gc.collect() return jsonify({ 'label': str(predict_cla), 'predicted_class': class_indict[str(predict_cla)], 'probability': float(predict[predict_cla].numpy()) }) if __name__ == '__main__': print("Loading PyTorch model and Flask starting server ...") print("Please wait until server has fully started") # start the classification service and wait for request app.run(port='5012')
代码问题分析
- 模型加载时机不合理:模型在首次请求时才加载,Flask默认多线程模式下,若多个请求同时触发首次加载,可能临时创建多个模型实例,直接导致内存飙升
- 张量内存未彻底释放:虽然调用了
torch.cuda.empty_cache()和gc.collect(),但output、predict等CPU张量未显式删除,仍占用内存;同时image.to(device)生成的GPU张量副本,原CPU张量未及时清理 - Flask默认模式额外开销:默认多线程模式下每个线程会占用额外内存,加上PyTorch的CUDA上下文初始化会占用500-800MB固定内存,这部分开销是模型无关的,直接缩小了ResNet-34和YOLOv5s的内存差距
- 权重加载与推理精度浪费:
torch.load默认加载全量数据(可能包含训练时的优化器状态),且未启用半精度推理,额外占用内存
优化方案
1. 提前加载模型,避免重复实例化
将模型加载移到服务启动前,确保全局只加载一次:
if __name__ == '__main__': print("Loading PyTorch model and Flask starting server ...") app.model = load_model() # 启动时提前加载模型 print("Model loaded, server starting...") app.run(port='5012')
同时删除predict()函数中的模型加载逻辑:
@app.route("/resnet", methods=['get', 'post']) def predict(): path = unquote(request.args.get("path", "")) if not path: # 修正空路径判断,原代码path不会为None,只会是空字符串 return jsonify({'error': 'Path not provided'}), 400 # 剩余代码不变
2. 彻底清理临时张量
预测完成后显式删除所有临时张量,避免内存泄漏:
with torch.no_grad(): output = torch.squeeze(app.model(image.to(device))).cpu() predict = torch.softmax(output, dim=0) predict_cla = torch.argmax(predict).numpy() # 先提取需要的结果值,释放张量引用 prob = float(predict[predict_cla].numpy()) label_str = str(predict_cla) # 显式删除所有临时张量 del output, predict, image gc.collect() torch.cuda.empty_cache() return jsonify({ 'label': label_str, 'predicted_class': class_indict[label_str], 'probability': prob })
3. 启用半精度推理
ResNet-34无需全精度(FP32),切换为FP16半精度可直接减少50%的模型内存占用:
修改load_model()函数:
def load_model(): model = resnet34(num_classes=7).to(device) weights_path = './resnet34.pth' model.load_state_dict(torch.load(weights_path, map_location=device, weights_only=True)) model.eval() # GPU环境下启用半精度 if device.type == 'cuda': model = model.half() return model
同步修改图像预处理,匹配模型精度:
def load_image(path): assert os.path.exists(path), "file: '{}' dose not exist.".format(path) img = Image.open(path).convert("RGB") img = data_transform(img) img = torch.unsqueeze(img, dim=0) # 匹配模型半精度 if device.type == 'cuda': img = img.half() return img
4. 优化Flask运行模式
- 测试环境可使用单线程模式减少内存开销:
app.run(port='5012', threaded=False) - 生产环境建议用Gunicorn+UWSGI替代Flask默认开发服务器,更好地控制进程/线程数量,避免不必要的内存浪费
5. 权重加载优化
使用weights_only=True只加载模型权重,避免加载训练时的优化器等冗余数据:
model.load_state_dict(torch.load(weights_path, map_location=device, weights_only=True))
6. 理解CUDA固定内存开销
CUDA初始化本身会占用500-800MB固定内存,这部分开销和模型无关。ResNet-34的模型内存(约1.5G)加上固定开销后,总内存就会和YOLOv5s的内存(模型内存+固定开销)接近,这是两者内存占用相近的核心原因。
内容的提问来源于stack exchange,提问作者Coding Rookie
相关产品推荐
相关产品推荐

