You need to enable JavaScript to run this app.
优惠活动
大模型
产品
解决方案
定价
更多

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')

代码问题分析

  1. 模型加载时机不合理:模型在首次请求时才加载,Flask默认多线程模式下,若多个请求同时触发首次加载,可能临时创建多个模型实例,直接导致内存飙升
  2. 张量内存未彻底释放:虽然调用了torch.cuda.empty_cache()和gc.collect(),但output、predict等CPU张量未显式删除,仍占用内存;同时image.to(device)生成的GPU张量副本,原CPU张量未及时清理
  3. Flask默认模式额外开销:默认多线程模式下每个线程会占用额外内存,加上PyTorch的CUDA上下文初始化会占用500-800MB固定内存,这部分开销是模型无关的,直接缩小了ResNet-34和YOLOv5s的内存差距
  4. 权重加载与推理精度浪费: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

相关产品推荐
方舟 Agent Plan

超全模态模型 × Harness 升级,最新支持 Deepseek-V4.1-Flash、GLM-5.3 系列、Doubao-Seedream-5.0-pro、Kimi-K3 (部分), 限时 9.9 元起

最近更新时间:2026.07.07 23:27:26