You need to enable JavaScript to run this app.
优惠活动
大模型
产品
解决方案
定价
更多

基于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

相关产品推荐
方舟 Agent Plan

超全模态模型 × Harness 升级,最新支持 Deepseek-V4.1-Flash、GLM-5.3 系列、Doubao-Seedream-5.0-pro、Kimi-K3 (部分), 限时 9.9 元起

最近更新时间:2026.06.26 15:33:15