(args)
| 39 | |
| 40 | |
| 41 | def main(args): |
| 42 | config = Config("train", binary=True, only_det=True) # need to change |
| 43 | config_global = ConfigGlobal("train", binary=True, only_det=True) |
| 44 | |
| 45 | num_epochs = args.nepoch |
| 46 | need_log = args.log |
| 47 | num_workers = args.nworker |
| 48 | start_epoch = 1 |
| 49 | batch_size = args.batch |
| 50 | num_agent = args.num_agent |
| 51 | |
| 52 | # Specify gpu device |
| 53 | device = torch.device("cuda" if torch.cuda.is_available() else "cpu") |
| 54 | device_num = torch.cuda.device_count() |
| 55 | print("device number", device_num) |
| 56 | |
| 57 | if args.com in {"mean", "max", "cat", "sum", "v2v", "ind_mae", "joint_mae", "late", "vqvae", "vqstar"}: |
| 58 | flag = args.com |
| 59 | else: |
| 60 | raise ValueError(f"com: {args.com} is not supported") |
| 61 | |
| 62 | config.flag = flag |
| 63 | |
| 64 | agent_idx_range = range(1, num_agent) if args.no_cross_road else range(num_agent) |
| 65 | |
| 66 | test_dataset = MultiTempV2XSimDet( |
| 67 | dataset_roots=[f"{args.data}/agent{i}" for i in agent_idx_range], |
| 68 | config=config, |
| 69 | config_global=config_global, |
| 70 | split="train", |
| 71 | bound="both", |
| 72 | kd_flag=args.kd_flag, |
| 73 | no_cross_road=args.no_cross_road, |
| 74 | time_stamp = args.time_stamp |
| 75 | ) |
| 76 | test_data_loader = DataLoader( |
| 77 | test_dataset, shuffle=False, batch_size=batch_size, num_workers=num_workers |
| 78 | ) |
| 79 | print("Testing dataset size:", len(test_dataset)) |
| 80 | |
| 81 | # logger_root = args.logpath if args.logpath != "" else "logs" |
| 82 | |
| 83 | if args.no_cross_road: |
| 84 | num_agent -= 1 |
| 85 | |
| 86 | if args.com == "joint_mae" or args.com == "ind_mae": |
| 87 | # Juexiao added for mae |
| 88 | model = multiagent_mae.__dict__[args.mae_model](norm_pix_loss=args.norm_pix_loss, time_stamp=args.time_stamp, mask_method=args.mask) |
| 89 | # also include individual reconstruction: reconstruct then aggregate |
| 90 | elif args.com == "late": |
| 91 | model = CNNNet( |
| 92 | config, |
| 93 | layer=args.layer, |
| 94 | kd_flag=args.kd_flag, |
| 95 | num_agent=num_agent, |
| 96 | train_completion=True, |
| 97 | ) |
| 98 | elif args.com == "vqvae": |
no test coverage detected