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

如何修改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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.09.29 11:57:04