图像分类任务中如何高效堆叠/集成预训练模型?
多预训练模型堆叠的预处理&推理耗时优化方案
以下是3种可落地的优化思路,可根据你的使用场景组合选择:
方案1:离线预生成所有基模型的表征缓存
这是成本最低、提速效果最明显的方案,适合基模型不需要和元学习器一起微调的场景:
- 提前跑一次全数据集的表征提取流程,把每个基模型对每张图的输出结果按唯一ID存到本地(可选择h5py、parquet或者numpy二进制格式存储),后续训练元学习器的时候直接读取缓存的表征拼接即可,完全不需要重复跑预处理和基模型推理。
- 新增基模型时不需要重新跑全流程,只需要单独提取新模型的表征,和原有缓存拼接即可。
参考实现代码:
import h5py import numpy as np # 第一步:预生成表征缓存,只需要跑一次 with h5py.File("feat_cache.h5", "w") as f: for idx, img in enumerate(images): # 提取各模型的隐藏层表征 input_1 = processor_1(img) out_1 = model_1(input_1).detach().cpu().numpy() input_2 = processor_2(img) out_2 = model_2(input_2).detach().cpu().numpy() # 按图片索引存储 f.create_dataset(f"model1/{idx}", data=out_1) f.create_dataset(f"model2/{idx}", data=out_2) # 第二步:训练元学习器时直接读缓存,速度提升10倍以上 X, y = [], [] with h5py.File("feat_cache.h5", "r") as f: for idx in range(len(images)): out1 = f[f"model1/{idx}"][()] out2 = f[f"model2/{idx}"][()] X.append(np.concatenate([out1, out2], axis=0)) X = np.array(X) # 直接喂给XGBoost等元学习器训练即可
方案2:批处理+并行化加速在线提取
如果基模型需要和元学习器联动微调,必须在线提取表征,可通过以下方式提速:
- 把逐图循环改为批处理,同时用多进程/多线程并行执行不同processor的预处理逻辑,避免单进程串行处理的耗时。
- 可把不同的基模型部署到不同的GPU/CPU核心上,同时执行批次推理,进一步压缩耗时。
参考实现代码:
from torch.utils.data import DataLoader, Dataset import concurrent.futures # 自定义数据集返回原始图和标签 class RawImgDataset(Dataset): def __getitem__(self, idx): return images[idx], labels[idx] def __len__(self): return len(images) dataloader = DataLoader(RawImgDataset(), batch_size=32, num_workers=8) # 可把两个模型放到不同GPU上并行推理 model_1 = model_1.to("cuda:0") model_2 = model_2.to("cuda:1") for batch_imgs, batch_labels in dataloader: # 多线程并行处理两个模型的输入 with concurrent.futures.ThreadPoolExecutor() as executor: input1_batch = list(executor.map(processor_1, batch_imgs)) input2_batch = list(executor.map(processor_2, batch_imgs)) input1 = torch.stack(input1_batch).to("cuda:0") input2 = torch.stack(input2_batch).to("cuda:1") # 并行推理得到表征 out1 = model_1(input1).detach().cpu() out2 = model_2(input2).detach().cpu() batch_feat = torch.cat([out1, out2], dim=1)
方案3:预处理逻辑统一对齐(仅适用于兼容场景)
如果多个预训练模型的预处理差异只是resize尺寸、归一化参数这类不改变图像语义的操作,可以把预处理逻辑做统一对齐,避免重复做图像解码、裁剪等重操作:
- 举例:model1要求输入224224,model2要求384384,你可以统一把原始图预处理到384*384,给model1的输入额外加一个Resize层降到224即可,比两次从原始图做全流程预处理快很多。
- 如果模型的预处理包含不同的裁剪逻辑(比如一个中心裁剪一个随机裁剪),该方案不适用。
内容的提问来源于stack exchange,提问作者skidjoe
相关产品推荐
相关产品推荐

