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

PyTorch自定义Dataset加载双行txt数据时batching维度异常问题

问题根源

你自定义的数据集类直接将第二行完整的100维解数据作为标签返回,因此batch_size=5时标签维度为[5, 100],和你预期的[5]维整数标签格式不符。

修复方案

如果你的100维标签行是单分类任务的独热编码(每行仅存在一个值为1的位置),直接对标签向量取最大值索引即可得到整数标签,有两种修改方式:

方案1:修改数据集类(更规范,推荐)

调整load_dataset类的初始化逻辑,对target提前做格式转换:

class load_dataset(Dataset):
    def __init__(self, data_file='data.txt', transform=None):
        super().__init__()
        data = np.loadtxt(data_file)
        data = torch.Tensor(data)
        self.data = data[::2]
        # 新增argmax提取独热对应的整数标签,转long类型适配交叉熵损失要求
        self.targets = data[1::2].argmax(dim=1).long()

    def __len__(self):
        return len(self.targets)

    def __getitem__(self, index):
        adj, target = self.data[index], self.targets[index]
        return adj, target

方案2:在训练循环中临时转换

如果不想修改数据集代码,可在训练循环中对输出的labels做转换:

for inputs, labels in loaders["train"]:
    inputs, labels = inputs.view([batch_size, 100]), labels.data
    # 新增行:提取独热对应的整数标签
    labels = labels.argmax(dim=1).long()
    scores = mps(inputs)
    _, preds = torch.max(scores, 1)
    print("preds: ")
    print(preds)
    print("labels: ")
    print(labels)

如果你的100维标签不是独热格式,而是完整的解序列,你可以根据自己的标签规则,从100维向量中提取出对应的整数标签替换上述的argmax逻辑即可。

内容的提问来源于stack exchange,提问作者bidon

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.09.29 07:15:04