MCPcopy Create free account
hub / github.com/coperception/star / main

Function main

completion/train_completion.py:41–354  ·  view source on GitHub ↗
(args)

Source from the content-addressed store, hash-verified

39
40
41def 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 )

Callers 1

Calls 14

resume_from_cpuMethod · 0.95
step_mae_completionMethod · 0.95
step_vae_completionMethod · 0.95
step_completionMethod · 0.95
MultiTempV2XSimDetClass · 0.90
CNNNetClass · 0.90
VQVAENetClass · 0.90
CoModuleClass · 0.90
printFunction · 0.85
stepMethod · 0.80
state_dictMethod · 0.80

Tested by

no test coverage detected