Pre-generate pseudo labels via network forward or uniformly assignment
(dataloader, model, class_num, data_len, init_via_forward=False)
| 238 | return losses.avg |
| 239 | |
| 240 | def compute_labels(dataloader, model, class_num, data_len, init_via_forward=False): |
| 241 | '''Pre-generate pseudo labels via network forward or uniformly assignment''' |
| 242 | if args.verbose: |
| 243 | logger.info('Compute labels') |
| 244 | batch_time = AverageMeter() |
| 245 | data_time = AverageMeter() |
| 246 | end = time.time() |
| 247 | |
| 248 | if init_via_forward: |
| 249 | model.eval() |
| 250 | label_list = [] |
| 251 | for i, inputs in enumerate(dataloader): |
| 252 | data_time.update(time.time() - end) |
| 253 | |
| 254 | inputs = inputs.cuda() |
| 255 | with torch.no_grad(): |
| 256 | output = model(inputs) |
| 257 | batch_time.update(time.time() - end) |
| 258 | |
| 259 | label = output.argmax(dim=1).cpu() |
| 260 | if i == 0: |
| 261 | label_list = label |
| 262 | else: |
| 263 | label_list = torch.cat([label_list, label], dim=0) |
| 264 | if args.verbose and (i % 100) == 0: |
| 265 | logger.info('{0}/{1}\t' |
| 266 | 'Time: {batch_time.val:.3f} ({batch_time.avg:.3f})' |
| 267 | 'Data: {data_time.val:.3f} ({data_time.avg:.3f})\t' |
| 268 | .format(i, len(dataloader), batch_time=batch_time, data_time=data_time)) |
| 269 | end = time.time() |
| 270 | label_list = label_list.numpy() |
| 271 | model.train() |
| 272 | else: |
| 273 | label_list = np.array([int(np.random.uniform(0, class_num)) for _ in range(data_len)]) |
| 274 | |
| 275 | return label_list |
| 276 | |
| 277 | if __name__ == '__main__': |
| 278 | main() |
no test coverage detected