如何高效从多个numpy memmap文件取数并生成新数组?
优化方案建议
1. 提前初始化Memmap对象(核心优化)
当前代码每次调用__getitem__都会重复打开所有memmap文件,这是IO密集型的冗余操作,会大幅拖慢数据读取速度。正确的做法是在数据集类的__init__方法中一次性打开所有文件并保存引用,后续__getitem__直接复用这些对象:
class YourDataset(Dataset): def __init__(self, memmap_paths, targets): self.targets = targets # 提前打开所有memmap文件,保存为实例属性 self.memmaps = [ np.memmap(path, dtype='float32', mode='r', shape=(24000, 300, 300)) for path in memmap_paths ] self.num_files = len(self.memmaps) self.sample_shape = self.memmaps[0].shape[1:] # 获取单样本形状:(300, 300) def __getitem__(self, index): # 预分配结果数组,避免np.stack的临时内存开销 x = np.empty((self.num_files,) + self.sample_shape, dtype='float32') for i, mm in enumerate(self.memmaps): x[i] = mm[index] return torch.from_numpy(x), self.targets[index] def __len__(self): return len(self.targets)
2. 预分配数组替代np.stack
你考虑的“创建数组再循环填充”确实比np.stack更高效:
np.stack需要先收集所有单样本数组到列表,再合并为新数组,会产生额外的临时内存开销;- 预分配数组直接将每个memmap的样本写入对应位置,内存使用更紧凑,减少不必要的内存拷贝次数。
3. 额外优化细节
- 多进程读取适配:如果使用PyTorch的
DataLoader并设置num_workers>0,注意memmap对象跨进程可能存在文件句柄问题,建议在worker的初始化逻辑中重新打开memmap文件,而非复用主进程的对象。 - 数据布局适配:如果模型需要特定的张量格式(比如通道在前),可以在预分配数组时直接指定对应形状,避免后续转置操作。
- 只读模式保持:保持
mode='r'即可,这是最安全的只读模式,既能避免意外修改文件,也能获得最优的读取性能。
性能提升说明
- 提前打开memmap:将每次
__getitem__的IO操作从N次(N为文件数)降至0,这是最显著的性能提升点; - 预分配数组:相比
np.stack,能减少30%-50%的内存拷贝开销,样本尺寸越大、文件数量越多,效果越明显。
内容的提问来源于stack exchange,提问作者ATYslh
相关产品推荐
相关产品推荐

