如何在Python中高效批量读取PyTorch .pt文件并追加张量到列表
优化Colab读取大量PyTorch .pt文件的方案
针对你在Colab上处理40000个小.pt文件的IO瓶颈问题,以下是几个直接落地的优化方案:
1. 合并小文件为批量大文件(最有效)
小文件的频繁远程IO是速度骤降的核心原因。建议先一次性将多个小文件合并为少数大文件,后续读取时直接处理大文件,彻底减少IO次数:
import os import torch from glob import glob # 按1000个文件为一组合并 pt_files = sorted(glob('/content/drive/MyDrive/your_pt_files/*.pt')) batch_size = 1000 save_dir = '/content/drive/MyDrive/merged_pt_files/' os.makedirs(save_dir, exist_ok=True) for i in range(0, len(pt_files), batch_size): batch_files = pt_files[i:i+batch_size] batch_tensors = [] for f in batch_files: data = torch.load(f, map_location='cpu', weights_only=True) tensor = data['mean_representations']['MY_LAYER'] batch_tensors.append(tensor) # 保存批量张量 torch.save(torch.stack(batch_tensors), f'{save_dir}/batch_{i//batch_size}.pt')
后续读取时直接加载批量文件,速度会提升几十倍:
merged_files = sorted(glob('/content/drive/MyDrive/merged_pt_files/*.pt')) all_tensors = [] for f in merged_files: batch_tensor = torch.load(f, map_location='cpu') all_tensors.append(batch_tensor) all_tensors = torch.cat(all_tensors, dim=0)
2. 多进程并行读取IO密集型任务
Colab的单线程IO无法充分利用带宽,用多进程并行读取可以缓解远程IO的延迟:
import torch from glob import glob from concurrent.futures import ProcessPoolExecutor def load_tensor(file_path): # 加载时只读取需要的张量,跳过其他内容 with open(file_path, 'rb') as f: data = torch.load(f, map_location='cpu', weights_only=True) return data['mean_representations']['MY_LAYER'] pt_files = sorted(glob('/content/drive/MyDrive/your_pt_files/*.pt')) # 限制进程数,Colab一般支持4-8个进程,避免资源耗尽 with ProcessPoolExecutor(max_workers=6) as executor: # 并行加载所有文件的张量 all_tensors = list(executor.map(load_tensor, pt_files)) # 转为大张量(可选,比列表更省内存) all_tensors = torch.stack(all_tensors)
注意:如果出现进程崩溃,可尝试添加torch.multiprocessing.set_start_method('spawn')(放在代码开头)。
3. 优化torch.load参数与内存使用
- 使用
weights_only=True(PyTorch 1.13+):只加载权重张量,跳过模型结构、优化器状态等无关内容,减少加载时间。 - 指定
map_location='cpu':避免不必要的GPU内存占用,同时加载速度更快。 - 预分配大张量替代列表追加:提前计算总尺寸,直接分配内存,避免列表动态扩容的开销:
import torch from glob import glob pt_files = sorted(glob('/content/drive/MyDrive/your_pt_files/*.pt')) total_num = len(pt_files) # 预分配40000x1280的张量 all_tensors = torch.zeros((total_num, 1280), dtype=torch.float32, device='cpu') for idx, f in enumerate(pt_files): data = torch.load(f, map_location='cpu', weights_only=True) all_tensors[idx] = data['mean_representations']['MY_LAYER']
4. 临时下载到Colab本地磁盘
Colab本地磁盘(/content/)的读写速度远快于挂载的Google Drive。可以批量下载文件到本地,读取后再删除:
import os import shutil import torch from glob import glob source_dir = '/content/drive/MyDrive/your_pt_files/' temp_dir = '/content/temp_pt_files/' os.makedirs(temp_dir, exist_ok=True) # 批量复制文件到本地(按批次处理,避免占满本地磁盘) batch_size = 2000 pt_files = sorted(glob(f'{source_dir}*.pt')) all_tensors = [] for i in range(0, len(pt_files), batch_size): batch_files = pt_files[i:i+batch_size] # 复制到本地 for f in batch_files: shutil.copy(f, temp_dir) # 读取本地文件 local_files = glob(f'{temp_dir}*.pt') for f in local_files: data = torch.load(f, map_location='cpu', weights_only=True) all_tensors.append(data['mean_representations']['MY_LAYER']) # 清理本地文件 for f in local_files: os.remove(f) all_tensors = torch.stack(all_tensors)
内容的提问来源于stack exchange,提问作者Andrija
相关产品推荐
相关产品推荐

