基于PyTorch的高效样本对采样优化:150k数据集训练瓶颈解决
问题描述
我有15万张图像和15万条音频数据,每组图像对应一条音频,希望通过图像数据训练网络实现音频间的映射。
构建源-目标训练对时,目标样本需要从训练数据子集采样且满足阈值条件。目前在自定义Dataset的__getitem__方法中完成子集筛选与阈值校验,这种动态采样方式在大数据量下效率极低,成为训练瓶颈。
现有实现代码如下:
def _getitem_(idx): sourceimg, sourceaudio = self.data[idx] Target_dataset = create_subset(sourceimg, sourceaudio) Iterations = 0 While true: targetimg, targetaudio = np.random(Target_dataset) If cal_value(target_img) >= threshold: Return sourceimg, sourceaudio, targetimg, targetaudio Iterations+=1 If iterations >= 10: _getitem_(random.randint(0, len(self.dataset))
def cal_val(targetimg, sourceimg): sourceimg=np.load(source_mat) targetimg = np.load(target_mat) diff = np.abs(source_mat - target_mat) thresholded_difference = np.where(diff >0 , 1, 0) average_difference = np.mean(thresholded_difference) return average_difference
def create_subset(sourceimg): subarray=[] sourceindex1, sourceindex2 = os.path.basename(sourceimg).split("/")[1:3] for element in arr: elementindex1, elementindex2 = os.path.basename(element).split("/")[1:3] if(elementindex1 == sourceindex1 and elementindex2 == sourceindex2): subarray.append(element) return subarray
优化方案
1. 预构建分组字典,避免动态遍历筛选子集
原create_subset每次调用都要遍历全量数据找同组样本,这是最大的效率杀手。直接在Dataset初始化时按图像路径的index1和index2分组,用字典存储每个分组对应的样本列表,后续直接通过key取子集:
from collections import defaultdict import os import random from torch.utils.data import Dataset class CustomDataset(Dataset): def __init__(self, data, threshold): self.data = data self.threshold = threshold # 预构建分组字典:key是(index1, index2),value是同组的(图像路径, 音频路径)列表 self.group_dict = defaultdict(list) for img_path, audio_path in data: # 修正路径分割逻辑:os.path.basename只会取最后文件名,改用split拆分路径 path_parts = img_path.split("/") idx1, idx2 = path_parts[-3], path_parts[-2] # 根据实际路径层级调整 self.group_dict[(idx1, idx2)].append((img_path, audio_path))
2. 预计算并缓存阈值校验结果,消除重复IO与计算
原cal_val每次都要加载图像文件计算差异,IO开销极大。提前离线计算所有同组样本对的差异值,缓存到文件中,训练时直接读取缓存:
import pickle import numpy as np # 离线预处理脚本,只运行一次 def precompute_diff_cache(data): diff_cache = {} group_dict = defaultdict(list) # 先分组 for img_path, audio_path in data: path_parts = img_path.split("/") idx1, idx2 = path_parts[-3], path_parts[-2] group_dict[(idx1, idx2)].append(img_path) # 计算同组内所有样本对的差异 for key, img_paths in group_dict.items(): # 一次性加载同组所有图像 imgs = [np.load(path) for path in img_paths] # 两两计算差异并缓存 for i in range(len(img_paths)): src_img = img_paths[i] src_mat = imgs[i] for j in range(len(img_paths)): if i == j: continue tgt_img = img_paths[j] tgt_mat = imgs[j] diff = np.abs(src_mat - tgt_mat) thresholded_diff = np.where(diff > 0, 1, 0) avg_diff = np.mean(thresholded_diff) diff_cache[(src_img, tgt_img)] = avg_diff # 保存缓存到文件 with open("diff_cache.pkl", "wb") as f: pickle.dump(diff_cache, f) # 在Dataset初始化时加载缓存 def __init__(self, data, threshold): # ... 其他初始化代码 ... with open("diff_cache.pkl", "rb") as f: self.diff_cache = pickle.load(f)
3. 提前为每个源样本筛选候选目标,避免随机重试
结合预分组和缓存,在初始化时就为每个源样本生成符合阈值的目标样本列表,训练时直接随机选择:
def __init__(self, data, threshold): # ... 之前的初始化代码 ... # 预构建每个源样本的候选目标列表 self.candidate_targets = [] for src_img, src_audio in data: path_parts = src_img.split("/") key = (path_parts[-3], path_parts[-2]) same_group_samples = self.group_dict[key] valid_targets = [] for tgt_img, tgt_audio in same_group_samples: if tgt_img == src_img: continue avg_diff = self.diff_cache[(src_img, tgt_img)] if avg_diff >= threshold: valid_targets.append((tgt_img, tgt_audio)) self.candidate_targets.append(valid_targets) # 优化后的__getitem__ def __getitem__(self, idx): src_img, src_audio = self.data[idx] valid_targets = self.candidate_targets[idx] # 处理无有效目标的情况(替代原逻辑的重试10次) if not valid_targets: rand_idx = random.randint(0, len(self.data)-1) return self.__getitem__(rand_idx) # 随机选一个符合条件的目标 tgt_img, tgt_audio = random.choice(valid_targets) return src_img, src_audio, tgt_img, tgt_audio
4. 修正原代码的语法错误
原代码存在多处语法问题,会影响运行效率和正确性:
_getitem_需改为__getitem__(双下划线)While、If、Return要小写np.random(Target_dataset)应改为random.choice(Target_dataset)cal_value与cal_val函数名不一致,需统一source_mat、arr等变量未定义,需补充逻辑- 递归调用
_getitem_时缺少闭合括号
内容的提问来源于stack exchange,提问作者mahnoor.fatima
相关产品推荐
相关产品推荐

