(model, data_loader, log_step=10, logging=logger.info)
| 70 | |
| 71 | |
| 72 | def encode_data(model, data_loader, log_step=10, logging=logger.info): |
| 73 | |
| 74 | batch_time = AverageMeter() |
| 75 | val_logger = LogCollector() |
| 76 | |
| 77 | # switch to evaluate mode |
| 78 | model.eval() |
| 79 | |
| 80 | end = time.time() |
| 81 | |
| 82 | # np array to keep all the embeddings |
| 83 | img_embs = None |
| 84 | cap_embs = None |
| 85 | |
| 86 | # compute the number of max word |
| 87 | max_n_word = model.opt.max_word |
| 88 | |
| 89 | for i, data_i in enumerate(data_loader): |
| 90 | |
| 91 | # make sure val logger is used |
| 92 | images, captions, lengths, ids, img_ids = data_i |
| 93 | |
| 94 | model.logger = val_logger |
| 95 | |
| 96 | # compute the embeddings |
| 97 | img_emb, cap_emb, lengths = model.forward_emb(images, captions, lengths) |
| 98 | |
| 99 | if img_embs is None: |
| 100 | # for local visual features |
| 101 | img_embs = torch.zeros((len(data_loader.dataset), img_emb.size(1), img_emb.size(2))) |
| 102 | # for local textual features |
| 103 | cap_embs = torch.zeros((len(data_loader.dataset), max_n_word, cap_emb.size(2))) |
| 104 | |
| 105 | cap_lens = torch.zeros(len(data_loader.dataset)).long() |
| 106 | |
| 107 | # cache embeddings |
| 108 | img_embs[ids] = img_emb.cpu() |
| 109 | |
| 110 | n_word = min(max(lengths), max_n_word) |
| 111 | |
| 112 | cap_embs[ids, :n_word, :] = cap_emb[:, :n_word, :].cpu() |
| 113 | cap_lens[ids] = lengths.cpu() |
| 114 | |
| 115 | # measure elapsed time |
| 116 | batch_time.update(time.time() - end) |
| 117 | end = time.time() |
| 118 | |
| 119 | if i % log_step == 0: |
| 120 | logging('Test: [{0}/{1}]\t' |
| 121 | '{e_log}\t' |
| 122 | 'Batch-Time {batch_time.val:.3f} ({batch_time.avg:.3f})\t' |
| 123 | .format(i, len(data_loader.dataset) // data_loader.batch_size + 1, batch_time=batch_time, e_log=str(model.logger))) |
| 124 | del images, captions |
| 125 | |
| 126 | return img_embs, cap_embs, cap_lens |
| 127 | |
| 128 | |
| 129 | def evalrank(model_path, model=None, data_path=None, split='dev', fold5=False, save_path=None): |
no test coverage detected