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")
在自定义数据读取和网络后,用mindspore.Model.train时,发生如下报错:
我能够理解报错的意思,但是我应该在何处对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")