生产环境BERT模型推理速度慢的排查与优化方案咨询
优化方案
1. 核心优化:批量处理(最大化GPU利用率)
当前代码单条文本逐个推理,完全没发挥Tesla T4的并行计算能力,这是速度瓶颈的核心原因。GPU擅长批量处理,一次性送入多文本能大幅提升推理效率。
修改predict_fn支持批量文本
import torch from torch.nn import functional as F def predict_fn(mdl, device, tokenizer_fn, texts): # 批量编码所有文本 encoded_text = tokenizer_fn( texts, add_special_tokens=True, return_token_type_ids=True, padding=True, return_attention_mask=True, truncation=True, return_tensors='pt', ) with torch.no_grad(): input_ids = encoded_text['input_ids'].to(device) token_type_ids = encoded_text['token_type_ids'].to(device) attention_mask = encoded_text['attention_mask'].to(device) output = mdl(input_ids, token_type_ids, attention_mask) probs = F.softmax(output[0], dim=1) # 返回所有文本的正类概率数组 return probs[:, 1].cpu().detach().numpy()
修改get_prediction和get_results实现全批量处理
def get_prediction(texts): # 三个模型同时批量推理所有文本 t_vals = predict_fn(model_t, device, tokenizer, texts) s_vals = predict_fn(model_s, device, tokenizer, texts) i_vals = predict_fn(model_i, device, tokenizer, texts) return t_vals, s_vals, i_vals @app.route("/predict", methods=["POST"]) def get_results(): data = request.get_json(force=True) texts = list(data.values()) # 一次处理所有文本,替代循环单条处理 t_vals, s_vals, i_vals = get_prediction(texts) return jsonify({ "i_scores": i_vals.tolist(), "t_scores": t_vals.tolist(), "s_scores": s_vals.tolist(), })
2. 模型推理效率优化
开启模型评估模式
模型加载后需设置为eval()模式,关闭训练时的dropout等冗余操作:
在model.py的init_model函数中,加载权重后添加:
model_t.eval() model_s.eval() model_i.eval()
启用FP16半精度推理
Tesla T4支持半精度计算,可减少显存占用并提升推理速度,修改predict_fn的推理块:
with torch.no_grad(), torch.cuda.amp.autocast(): input_ids = encoded_text['input_ids'].to(device) token_type_ids = encoded_text['token_type_ids'].to(device) attention_mask = encoded_text['attention_mask'].to(device) output = mdl(input_ids, token_type_ids, attention_mask)
可选:导出为TorchScript/ONNX格式
生产环境可将模型导出为TorchScript或ONNX,进一步压缩推理延迟:
以TorchScript为例,在init_model中添加(需提前准备符合输入格式的样本张量):
# 准备样本输入(需和实际输入维度一致) sample_input = tokenizer("sample text", return_tensors='pt').to(device) input_ids_sample = sample_input['input_ids'] token_type_ids_sample = sample_input['token_type_ids'] attention_mask_sample = sample_input['attention_mask'] # 导出并冻结模型 model_t = torch.jit.trace(model_t, (input_ids_sample, token_type_ids_sample, attention_mask_sample)) model_t = torch.jit.freeze(model_t)
3. Flask服务优化
关闭Debug模式
Debug模式会增加额外开销,生产环境必须关闭:
if __name__ == "__main__": app.run(debug=False, host='0.0.0.0', port=8000)
使用生产级WSGI服务器
Flask自带服务器为单线程,无法处理并发请求,建议用Gunicorn+Uvicorn:
安装依赖:
pip install gunicorn uvicorn
启动命令(按CPU核心数调整工作进程数,建议为核心数*2):
gunicorn -w 4 -k uvicorn.workers.UvicornWorker main:app --bind 0.0.0.0:8000
4. 数据处理细节优化
原代码用np.append逐个添加元素,每次操作都会重新分配内存,效率极低。改用列表收集后转数组(批量处理后此步骤已被替代,此处仅作参考):
# 原低效写法 t_vals = np.array([]) for text in texts: t_vals = np.append(t_vals, t_val) # 优化写法 t_vals = [] for text in texts: t_vals.append(t_val) t_vals = np.array(t_vals)
5. 其他优化建议
- 用
nvidia-smi查看显存占用,若有剩余可进一步调大批次大小(如一次处理64/128条文本),最大化GPU利用率。 - 确保
AutoTokenizer.from_pretrained加载的是本地缓存文件,避免重复下载。
内容的提问来源于stack exchange,提问作者grey
相关产品推荐
相关产品推荐

