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
相关产品推荐
相关产品推荐

