(args)
| 329 | |
| 330 | |
| 331 | def main(args): |
| 332 | if args.log_wandb: |
| 333 | if has_wandb: |
| 334 | wandb.init(project=args.dataset, |
| 335 | name = args.experiment, |
| 336 | entity="spikingtransformer", |
| 337 | config=args) |
| 338 | |
| 339 | else: |
| 340 | print("You've requested to log metrics to wandb but package not found. " |
| 341 | "Metrics not being logged to wandb, try `pip install wandb`") |
| 342 | |
| 343 | max_test_acc1 = 0. |
| 344 | test_acc5_at_max_test_acc1 = 0. |
| 345 | |
| 346 | utils.init_distributed_mode(args) |
| 347 | print(args) |
| 348 | |
| 349 | output_dir = os.path.join(args.output_dir, f'{args.model}_b{args.batch_size}_T{args.T}') |
| 350 | |
| 351 | if args.T_train: |
| 352 | output_dir += f'_Ttrain{args.T_train}' |
| 353 | |
| 354 | if args.weight_decay: |
| 355 | output_dir += f'_wd{args.weight_decay}' |
| 356 | |
| 357 | if args.opt == 'adamw': |
| 358 | output_dir += '_adamw' |
| 359 | else: |
| 360 | output_dir += '_sgd' |
| 361 | |
| 362 | if not os.path.exists(output_dir): |
| 363 | utils.mkdir(output_dir) |
| 364 | |
| 365 | output_dir = os.path.join(output_dir, f'lr{args.lr}') |
| 366 | if not os.path.exists(output_dir): |
| 367 | utils.mkdir(output_dir) |
| 368 | |
| 369 | device = torch.device(args.device) |
| 370 | print("device:", device) |
| 371 | data_path = args.data_path |
| 372 | dataset_train, dataset_test, train_sampler, test_sampler = load_data(args.dataset, data_path, args.distributed, args.T) |
| 373 | data_loader = torch.utils.data.DataLoader( |
| 374 | dataset=dataset_train, |
| 375 | batch_size=args.batch_size, |
| 376 | shuffle=True, |
| 377 | num_workers=args.workers, |
| 378 | drop_last=True, |
| 379 | pin_memory=True) |
| 380 | |
| 381 | data_loader_test = torch.utils.data.DataLoader( |
| 382 | dataset=dataset_test, |
| 383 | batch_size=args.batch_size, |
| 384 | shuffle=False, |
| 385 | num_workers=args.workers, |
| 386 | drop_last=False, |
| 387 | pin_memory=True) |
| 388 |
no test coverage detected