(args)
| 39 | |
| 40 | |
| 41 | def main(args): |
| 42 | config = Config("train", binary=True, only_det=True) |
| 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 | auto_resume_path = args.auto_resume_path |
| 52 | |
| 53 | |
| 54 | # Specify gpu device |
| 55 | device = torch.device("cuda" if torch.cuda.is_available() else "cpu") |
| 56 | device_num = torch.cuda.device_count() |
| 57 | print("device number", device_num) |
| 58 | |
| 59 | if args.com in {"mean", "max", "cat", "sum", "v2v", "ind_mae", "joint_mae", "late", "vqvae", "vqstar"}: |
| 60 | flag = args.com |
| 61 | else: |
| 62 | raise ValueError(f"com: {args.com} is not supported") |
| 63 | |
| 64 | config.flag = flag |
| 65 | |
| 66 | agent_idx_range = range(1, num_agent) if args.no_cross_road else range(num_agent) |
| 67 | |
| 68 | training_dataset = MultiTempV2XSimDet( |
| 69 | dataset_roots=[f"{args.data}/agent{i}" for i in agent_idx_range], |
| 70 | config=config, |
| 71 | config_global=config_global, |
| 72 | split="train", |
| 73 | bound="both", |
| 74 | kd_flag=args.kd_flag, |
| 75 | no_cross_road=args.no_cross_road, |
| 76 | time_stamp = args.time_stamp |
| 77 | ) |
| 78 | training_data_loader = DataLoader( |
| 79 | training_dataset, shuffle=True, batch_size=batch_size, num_workers=num_workers |
| 80 | ) |
| 81 | print("Training dataset size:", len(training_dataset)) |
| 82 | |
| 83 | logger_root = args.logpath if args.logpath != "" else "logs" |
| 84 | |
| 85 | if args.no_cross_road: |
| 86 | num_agent -= 1 |
| 87 | |
| 88 | if args.com == "joint_mae" or args.com == "ind_mae": |
| 89 | # Juexiao added for mae |
| 90 | model = multiagent_mae.__dict__[args.mae_model](norm_pix_loss=args.norm_pix_loss, time_stamp=args.time_stamp, mask_method=args.mask, |
| 91 | encode_partial=args.encode_partial, no_temp_emb=args.no_temp_emb, decode_singletemp=args.decode_singletemp) |
| 92 | # also include individual reconstruction: reconstruct then aggregate |
| 93 | elif args.com == "vqstar": |
| 94 | model = VQSTAR.vqstar( |
| 95 | norm_pix_loss=args.norm_pix_loss, time_stamp=args.time_stamp, mask_method=args.mask, |
| 96 | decay=args.decay, commitment_cost=args.commitment_cost, |
| 97 | num_vq_embeddings=args.num_vq_embeddings, vq_embedding_dim=args.vq_embedding_dim |
| 98 | ) |
no test coverage detected