INNER CODE UNIT · Python
test_generator
vietnh1009/QuickDraw · train.py:59
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()
criterion = nn.CrossEntropyLoss()
if opt.optimizer == "adam":
optimizer = torch.optim.Adam(model.parameters(), lr=opt.lr)