训练报错:For primitive[SoftmaxCrossEntropyWithLogits], the dimension of logits must be equal to 2, but got 4
收藏回复举报
训练报错:For primitive[SoftmaxCrossEntropyWithLogits], the dimension of logits must be equal to 2, but got 4
新人帖
发表于2022-12-04 20:06:10
0 查看

在自定义数据读取和网络后,用mindspore.Model.train时,发生如下报错:

true

我能够理解报错的意思,但是我应该在何处对logits进行修改呢

我的主程序代码如下: if name == "main": parser = get_command_line_parser_tmp() #读取命令行输入的参数 args = parser.parse_args() set_seed(args.seed) import mindspore as md from dataloader.data_utils import set_up_datasets

args = set_up_datasets(args)
network = MYNET(args)
ls = md.nn.SoftmaxCrossEntropyWithLogits(sparse=True)

# optimization definition
opt = md.nn.Momentum(filter(lambda x: x.requires_grad, network.get_parameters()), 0.01, 0.9)
model = md.Model(network, loss_fn=ls, optimizer=opt, metrics={'acc'})

#自定义数据集
dataroot = 'dataloader/'
batch_size_base = 128
class_index = np.arange(60)
trainset = CIFAR100(root=dataroot, train=True, download=False, transform=None, index=class_index,
                    base_sess=True)
testset = CIFAR100(root=dataroot, train=False, download=False, index=class_index, base_sess=True)

trainloader = GeneratorDataset(source=trainset, column_names=["image", "label"]).shuffle(buffer_size=10000).batch(
    128)
testloader = GeneratorDataset(source=testset, column_names=["image", "label"]).shuffle(buffer_size=10000).batch(100)

model.train(1, trainloader)

自定义数据读取的代码如下:

from PIL import Image import os import os.path import numpy as np import pickle

import torchvision.transforms as transforms from torchvision.datasets.vision import VisionDataset from torchvision.datasets.utils import check_integrity, download_and_extract_archive from mindspore.dataset import GeneratorDataset

class CIFAR100(VisionDataset): """CIFAR100 <https://www.cs.toronto.edu/~kriz/cifar.html>_ Dataset.

This is a subclass of the `CIFAR10` Dataset.
"""
base_folder = 'cifar-100-python'
url = "https://www.cs.toronto.edu/~kriz/cifar-100-python.tar.gz"
filename = "cifar-100-python.tar.gz"
tgz_md5 = 'eb9058c3a382ffc7106e4002c42a8d85'
train_list = [
    ['train', '16019d7e3df5f24257cddd939b257f8d'],
]

test_list = [
    ['test', 'f0ef6b0ae62326f3e7ffdfab6717acfc'],
]
meta = {
    'filename': 'meta',
    'key': 'fine_label_names',
    'md5': '7973b15100ade9c7d40fb424638fde48',
}

def __init__(self, root, train=True, transform=None, target_transform=None,
             download=False, index=None, base_sess=None):

    super(CIFAR100, self).__init__(root, transform=transform,
                                  target_transform=target_transform)
    self.root = os.path.expanduser(root)
    self.train = train  # training set or test set

    if download:
        self.download()

    if not self._check_integrity():
        raise RuntimeError('Dataset not found or corrupted.' +
                           ' You can use download=True to download it')

    if self.train:
        downloaded_list = self.train_list
        self.transform = transforms.Compose([
            transforms.RandomCrop(32, padding=4),
            transforms.RandomHorizontalFlip(),
            transforms.ToTensor(),
            transforms.Normalize(mean=[0.507, 0.487, 0.441], std=[0.267, 0.256, 0.276])
        ])
    else:
        downloaded_list = self.test_list
        self.transform = transforms.Compose([
            transforms.ToTensor(),
            transforms.Normalize(mean=[0.507, 0.487, 0.441], std=[0.267, 0.256, 0.276])
        ])

    self.data = []
    self.targets = []

    # now load the picked numpy arrays
    for file_name, checksum in downloaded_list:
        file_path = os.path.join(self.root, self.base_folder, file_name)
        with open(file_path, 'rb') as f:
            entry = pickle.load(f, encoding='latin1')
            self.data.append(entry['data'])
            if 'labels' in entry:
                self.targets.extend(entry['labels'])
            else:
                self.targets.extend(entry['fine_labels'])

    self.data = np.vstack(self.data).reshape(-1, 3, 32, 32)
    self.data = self.data.transpose((0, 2, 3, 1))  # convert to HWC

    self.targets = np.asarray(self.targets)

    if base_sess:
        self.data, self.targets = self.SelectfromDefault(self.data, self.targets, index)
    else:  # new Class session
        if train:
            self.data, self.targets = self.NewClassSelector(self.data, self.targets, index)
        else:
            self.data, self.targets = self.SelectfromDefault(self.data, self.targets, index)

    self._load_meta()

def SelectfromDefault(self, data, targets, index):
    data_tmp = []
    targets_tmp = []
    for i in index:
        ind_cl = np.where(i == targets)[0]
        if data_tmp == []:
            data_tmp = data[ind_cl]
            targets_tmp = targets[ind_cl]
        else:
            data_tmp = np.vstack((data_tmp, data[ind_cl])) #数组叠加
            targets_tmp = np.hstack((targets_tmp, targets[ind_cl]))

    return data_tmp, targets_tmp

def NewClassSelector(self, data, targets, index):
    data_tmp = []
    targets_tmp = []
    ind_list = [int(i) for i in index]
    ind_np = np.array(ind_list)
    index = ind_np.reshape((5,5))
    for i in index:
        ind_cl = i
        if data_tmp == []:
            data_tmp = data[ind_cl]
            targets_tmp = targets[ind_cl]
        else:
            data_tmp = np.vstack((data_tmp, data[ind_cl]))
            targets_tmp = np.hstack((targets_tmp, targets[ind_cl]))

    return data_tmp, targets_tmp

def _load_meta(self):
    path = os.path.join(self.root, self.base_folder, self.meta['filename'])
    if not check_integrity(path, self.meta['md5']):
        raise RuntimeError('Dataset metadata file not found or corrupted.' +
                           ' You can use download=True to download it')
    with open(path, 'rb') as infile:
        data = pickle.load(infile, encoding='latin1')
        self.classes = data[self.meta['key']]
    self.class_to_idx = {_class: i for i, _class in enumerate(self.classes)}

def __getitem__(self, index):
    """
    Args:
        index (int): Index

    Returns:
        tuple: (image, target) where target is index of the target class.
    """
    img, target = self.data[index], self.targets[index]

    # doing this so that it is consistent with all other datasets
    # to return a PIL Image
    img = Image.fromarray(img)

    if self.transform is not None:
        img = self.transform(img)

    if self.target_transform is not None:
        target = self.target_transform(target)

    return img, target

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

def _check_integrity(self):
    root = self.root
    for fentry in (self.train_list + self.test_list):
        filename, md5 = fentry[0], fentry[1]
        fpath = os.path.join(root, self.base_folder, filename)
        if not check_integrity(fpath, md5):
            print("the root is:",root,fpath)
            return False
    return True

def download(self):
    if self._check_integrity():
        print('Files already downloaded and verified')
        return
    download_and_extract_archive(self.url, self.root, filename=self.filename, md5=self.tgz_md5)

def extra_repr(self):
    return "Split: {}".format("Train" if self.train is True else "Test")

本帖最后由 匿名用户2022/12/28 14:33:58 编辑

我要发帖子