(projectors, student, train_loaders, test_loader, unlabelled_train_loader, whole_train_test_loader, args)
| 177 | |
| 178 | |
| 179 | def train(projectors, student, train_loaders, test_loader, unlabelled_train_loader, whole_train_test_loader, args): |
| 180 | |
| 181 | optimizer = SGD(list(projectors.parameters()) + list(student.parameters()), lr=args.lr, momentum=args.momentum, |
| 182 | weight_decay=args.weight_decay) |
| 183 | |
| 184 | exp_lr_scheduler = lr_scheduler.CosineAnnealingLR( |
| 185 | optimizer, |
| 186 | T_max=args.epochs, |
| 187 | eta_min=args.lr * 1e-3, |
| 188 | ) |
| 189 | |
| 190 | sup_con_crit = SupConLoss() |
| 191 | best_test_acc_lab = 0 |
| 192 | best_acc_lab = 0 |
| 193 | |
| 194 | for epoch in range(args.epochs): |
| 195 | |
| 196 | loss_record = AverageMeter() |
| 197 | train_acc_record = AverageMeter() |
| 198 | |
| 199 | student.train() |
| 200 | projectors.train() |
| 201 | |
| 202 | loaders =[iter(l) for l in train_loaders] |
| 203 | |
| 204 | max_len = max([len(loader) for loader in loaders]) |
| 205 | |
| 206 | # Load data from each dataloaders, the total number of dataloader is expert_num + 1 |
| 207 | for _ in tqdm(range(max_len)): |
| 208 | fine_loss = 0 |
| 209 | expert_images = [] |
| 210 | expert_class_labels = [] |
| 211 | expert_mask_lab = [] |
| 212 | |
| 213 | for idx, loader in enumerate(loaders): |
| 214 | try: |
| 215 | item = next(loader) |
| 216 | except StopIteration: |
| 217 | loaders[idx] = iter(train_loaders[idx]) |
| 218 | item = next(loaders[idx]) |
| 219 | |
| 220 | images, class_labels, uq_idxs, mask_lab = item |
| 221 | mask_lab = mask_lab[:, 0] |
| 222 | images = torch.cat(images, dim=0) # [B*2/num_experts,3,224,224] |
| 223 | if args.use_global_con: |
| 224 | # Load subset data |
| 225 | if idx < args.experts_num: |
| 226 | expert_images.append(images) |
| 227 | expert_class_labels.append(class_labels) |
| 228 | expert_mask_lab.append(mask_lab.bool()) |
| 229 | |
| 230 | # Load whole dataset |
| 231 | else: |
| 232 | all_images = images.to(device) |
| 233 | all_class_labels = class_labels.to(device) |
| 234 | all_mask_lab = mask_lab.to(device).bool() |
| 235 | else: |
| 236 | expert_images.append(images) |
no test coverage detected