大规模文本数据高效预测方法及PyTorch推理代码性能优化咨询
大规模文本批量推理效率优化指南
你猜测的频繁torch.cat操作拖慢效率的判断是完全正确的,该操作每次执行都会重新申请内存、拷贝所有已有的embedding数据,批次量级越大,额外耗时越高。以下是可落地的优化方案:
核心优化方案
1. 替换动态拼接逻辑,减少冗余内存拷贝
最高优先级的优化是避免每轮循环都执行拼接操作,可选两种实现:
- 已知总样本数和embedding维度的场景:预分配固定大小的张量,每轮循环直接将批次结果写入对应位置,全程无额外拷贝
- 样本数不确定的场景:先将每批次的embedding存入Python列表,所有批次推理完成后只做一次全局拼接,性能远高于逐次拼接
2. 数据加载侧优化
- 给
DataLoader开启pin_memory=True,同时设置non_blocking=True的异步数据传输参数,减少CPU到GPU的数据搬运耗时 - 调整
DataLoader的num_workers参数,通常设置为CPU核心数的1/2到2/3,避免数据加载成为瓶颈 - 文本token化等预处理逻辑提前完成并缓存到磁盘,避免每次推理都重复执行预处理
3. 推理侧性能优化
- 用
torch.inference_mode()替代torch.no_grad(),推理场景下该模式会禁用更多不需要的计算逻辑,性能更高 - 若无特殊精度要求,将
float64替换为float32甚至float16,低精度推理速度可提升30%以上,显存占用也会大幅降低 - 若使用Transformer类模型,可在tokenization阶段开启
pad_to_multiple_of=8参数,对齐张量维度适配GPU TensorCore加速 - 条件允许的情况下可通过ONNX Runtime、TensorRT等工具对模型做量化、编译优化,推理速度可提升2-10倍
优化后参考代码
预分配张量版本(性能最优)
def get_embeddings(model, data_loader, device): model.eval() # 提前计算总样本数与embedding维度 total_samples = len(data_loader.dataset) test_batch = next(iter(data_loader)) test_emb = model.predict( test_batch["input_ids"].to(device), attention_mask=test_batch["attention_mask"].to(device) ) emb_dim = test_emb.shape[1] # 预分配固定大小张量 embeddings = torch.empty((total_samples, emb_dim), dtype=torch.float32, device=device) with torch.inference_mode(): current_idx = 0 for d in tqdm(data_loader): input_ids = d["input_ids"].to(device, non_blocking=True) attention_mask = d["attention_mask"].to(device, non_blocking=True) batch_emb = model.predict(input_ids, attention_mask=attention_mask) batch_size = batch_emb.shape[0] # 直接写入对应位置,无拼接操作 embeddings[current_idx:current_idx+batch_size] = batch_emb current_idx += batch_size return embeddings.cpu().numpy()
列表暂存版本(适配不确定样本数场景)
def get_embeddings(model, data_loader, device): model.eval() emb_list = [] with torch.inference_mode(): for d in tqdm(data_loader): input_ids = d["input_ids"].to(device, non_blocking=True) attention_mask = d["attention_mask"].to(device, non_blocking=True) emb_list.append(model.predict(input_ids, attention_mask=attention_mask)) # 仅执行一次全局拼接 return torch.cat(emb_list).cpu().numpy()
内容的提问来源于stack exchange,提问作者Thomas
相关产品推荐
相关产品推荐

