如何修改MNIST数据加载代码 输出相邻奇偶数字对供多模态VAE使用
修改逻辑很简单,你当前代码的问题是取第二个模态样本时,直接用第一个模态的数字作为key去查索引,所以只能得到相同数字的配对。按照你设定的(偶,奇)配对规则,偶数x对应的配对奇数为x+1,只需修改__getitem__里的索引查找key即可。
以下是修改后的完整代码:
import random import torch class JointDataset(torch.utils.data.Dataset): def __init__(self, mnist_pt_path_1, mnist_pt_path_2): self.mnist_pt_path_1 = mnist_pt_path_1 self.mnist_pt_path_2 = mnist_pt_path_2 # 加载MNIST pt文件 self.mnist_data_1, self.mnist_targets_1 = torch.load(self.mnist_pt_path_1) self.mnist_data_2, self.mnist_targets_2 = torch.load(self.mnist_pt_path_2) self.mnist_target_idx_mapping = self.process_mnist_labels() def process_mnist_labels(self): numbers_dict = {0: [], 1: [], 2: [], 3:[], 4:[], 5:[], 6:[], 7: [], 8:[], 9:[]} for i in range(len(self.mnist_targets_2)): mnist_target = self.mnist_targets_2[i].item() numbers_dict[mnist_target].append(i) return numbers_dict def __len__(self): return len(self.mnist_data_1) def __getitem__(self, index: int): """ Args: index (int): 索引 模态1: 偶数数字 模态2: 奇数数字 """ mnist_img_1, mnist_target_1 = self.mnist_data_1[index], int(self.mnist_targets_1[index]) # 仅修改这一行:查找当前偶数+1对应的奇数的索引列表 indices_list = self.mnist_target_idx_mapping[(mnist_target_1 + 1)] # 从索引列表中随机选一个 idx = random.choice(indices_list) mnist_img_2 = self.mnist_data_2[idx] mnist_target_2 = int(self.mnist_targets_2[idx]) return mnist_img_1/255, mnist_img_2/255, mnist_target_1, mnist_target_2
如果要额外加鲁棒性校验,避免第一个数据集中混入奇数导致报错,可以在获取mnist_target_1后加一段校验:
# 校验模态1是否为偶数 if mnist_target_1 % 2 != 0: raise ValueError(f"模态1预期输入偶数数字,实际得到数字:{mnist_target_1}")
内容的提问来源于stack exchange,提问作者Achint Kumar
相关产品推荐
相关产品推荐

