MCPcopy Create free account
hub / github.com/DanielShalam/BPA / main

Function main

train.py:115–199  ·  view source on GitHub ↗
()

Source from the content-addressed store, hash-verified

113
114
115def main():
116 args = get_args()
117 utils.set_seed(seed=args.seed)
118 print("\n".join("%s: %s" % (k, str(v)) for k, v in sorted(dict(vars(args)).items())))
119 output_dir = utils.get_output_dir(args=args)
120
121 # define datasets and loaders
122 args.set_episodes = dict(train=args.train_episodes, val=args.eval_episodes, test=args.test_episodes)
123 if not args.eval:
124 train_dataloader = utils.get_dataloader(set_name='train', args=args, constant=False)
125 val_dataloader = utils.get_dataloader(set_name='val', args=args, constant=True)
126 else:
127 val_dataloader = utils.get_dataloader(set_name='test', args=args, constant=False)
128 train_dataloader = None
129
130 # define model and load pretrained weights if available
131 model = utils.get_model(args.backbone, args)
132 model = model.to(device)
133 utils.load_weights(model, args.pretrained_path)
134
135 # BPA and few-shot classification method (e.g. proto, pt-map...)
136 bpa = None
137 if 'bpa' in args.method.lower():
138 bpa = BPA(
139 distance_metric=args.distance_metric,
140 ot_reg=args.ot_reg,
141 mask_diag=args.mask_diag,
142 sinkhorn_iterations=args.sink_iters,
143 max_scale=args.max_scale
144 )
145 fewshot_method = utils.get_method(args=args, bpa=bpa)
146
147 # few-shot labels
148 train_labels = utils.get_fs_labels(args.method, args.train_way, args.num_query, args.num_shot)
149 val_labels = utils.get_fs_labels(args.method, args.val_way, args.num_query, args.num_shot)
150
151 # initialized wandb
152 if args.wandb:
153 utils.init_wandb(exp_name=output_dir.split('/')[-1] if output_dir[-1] != '/' else output_dir.split('/')[-2],
154 args=args)
155
156 # define loss
157 criterion = utils.get_criterion_by_method(method=args.method)
158
159 # Test-set evaluation
160 if args.eval:
161 print(f"Evaluate model for {args.test_episodes} episodes... ")
162 loss, acc = eval_one_epoch(model, val_dataloader, fewshot_method, criterion, val_labels, 0, args, set_name='test')
163 print("Final evaluation results:\nAccuracy={:.4f}, Loss={:.4f}".format(acc, loss))
164 exit(1)
165
166 # define optimizer and scheduler
167 optimizer, lr_scheduler = utils.get_optimizer_and_lr_scheduler(args=args, params=model.parameters())
168
169 # evaluate model before training
170 if args.eval_first:
171 print("Evaluate model before training... ")
172 eval_one_epoch(model, val_dataloader, fewshot_method, criterion, val_labels, -1, args, set_name='val')

Callers 1

train.pyFile · 0.70

Calls 4

BPAClass · 0.90
eval_one_epochFunction · 0.85
train_one_epochFunction · 0.85
get_argsFunction · 0.70

Tested by

no test coverage detected