INNER CODE UNIT · Python
training_generator
vietnh1009/QuickDraw · train.py:56
training_generator = DataLoader(training_set, **training_params)
print ("there are {} images for training phase".format(training_set.__len__()))
test_set = MyDataset(opt.data_path, opt.total_images_per_class, opt.ratio, "test")
test_generator = DataLoader(test_set, **test_params)
print("there are {} images for test phase".format(test_set.__len__()))
model = QuickDraw(num_classes=training_set.num_classes)
if os.path.isdir(opt.log_path):
shutil.rmtree(opt.log_path)
os.makedirs(opt.log_path)
writer = SummaryWriter(opt.log_path)
# writer.add_graph(model, torch.rand(opt.batch_size, 1, 28, 28))
if torch.cuda.is_available():
model.cuda()