(data_set, backbone, batch_size, nfolds=10)
| 225 | |
| 226 | @torch.no_grad() |
| 227 | def test(data_set, backbone, batch_size, nfolds=10): |
| 228 | print('testing verification..') |
| 229 | data_list = data_set[0] |
| 230 | issame_list = data_set[1] |
| 231 | embeddings_list = [] |
| 232 | time_consumed = 0.0 |
| 233 | for i in range(len(data_list)): |
| 234 | data = data_list[i] |
| 235 | embeddings = None |
| 236 | ba = 0 |
| 237 | while ba < data.shape[0]: |
| 238 | bb = min(ba + batch_size, data.shape[0]) |
| 239 | count = bb - ba |
| 240 | _data = data[bb - batch_size: bb] |
| 241 | time0 = datetime.datetime.now() |
| 242 | img = ((_data / 255) - 0.5) / 0.5 |
| 243 | net_out: torch.Tensor = backbone(img) |
| 244 | _embeddings = net_out.detach().cpu().numpy() |
| 245 | time_now = datetime.datetime.now() |
| 246 | diff = time_now - time0 |
| 247 | time_consumed += diff.total_seconds() |
| 248 | if embeddings is None: |
| 249 | embeddings = np.zeros((data.shape[0], _embeddings.shape[1])) |
| 250 | embeddings[ba:bb, :] = _embeddings[(batch_size - count):, :] |
| 251 | ba = bb |
| 252 | embeddings_list.append(embeddings) |
| 253 | |
| 254 | _xnorm = 0.0 |
| 255 | _xnorm_cnt = 0 |
| 256 | for embed in embeddings_list: |
| 257 | for i in range(embed.shape[0]): |
| 258 | _em = embed[i] |
| 259 | _norm = np.linalg.norm(_em) |
| 260 | _xnorm += _norm |
| 261 | _xnorm_cnt += 1 |
| 262 | _xnorm /= _xnorm_cnt |
| 263 | |
| 264 | embeddings = embeddings_list[0].copy() |
| 265 | embeddings = sklearn.preprocessing.normalize(embeddings) |
| 266 | acc1 = 0.0 |
| 267 | std1 = 0.0 |
| 268 | embeddings = embeddings_list[0] + embeddings_list[1] |
| 269 | embeddings = sklearn.preprocessing.normalize(embeddings) |
| 270 | print(embeddings.shape) |
| 271 | print('infer time', time_consumed) |
| 272 | _, _, accuracy, val, val_std, far = evaluate(embeddings, issame_list, nrof_folds=nfolds) |
| 273 | acc2, std2 = np.mean(accuracy), np.std(accuracy) |
| 274 | return acc1, std1, acc2, std2, _xnorm, embeddings_list |
| 275 | |
| 276 | |
| 277 | def dumpR(data_set, |
nothing calls this directly
no test coverage detected