From 77abac65464cb056b02804f37c1bfc7054c3b574 Mon Sep 17 00:00:00 2001 From: ZardLi1115 <262531544+ZardLi1115@users.noreply.github.com> Date: Sat, 16 May 2026 21:17:15 +0000 Subject: [PATCH] Fix LFW eval distance shape --- utils/utils_fit.py | 4 ++-- utils/utils_metrics.py | 7 ++++++- 2 files changed, 8 insertions(+), 3 deletions(-) diff --git a/utils/utils_fit.py b/utils/utils_fit.py index 3b61c6a..2b87433 100644 --- a/utils/utils_fit.py +++ b/utils/utils_fit.py @@ -7,7 +7,7 @@ from tqdm import tqdm from utils.utils import get_lr -from utils.utils_metrics import evaluate +from utils.utils_metrics import calculate_lfw_distance, evaluate def fit_one_epoch(model_train, model, loss_history, loss, optimizer, epoch, epoch_step, epoch_step_val, gen, gen_val, Epoch, cuda, test_loader, Batch_size, lfw_eval_flag, fp16, scaler, save_period, save_dir, local_rank): @@ -115,7 +115,7 @@ def fit_one_epoch(model_train, model, loss_history, loss, optimizer, epoch, epoc if cuda: data_a, data_p = data_a.cuda(local_rank), data_p.cuda(local_rank) out_a, out_p = model_train(data_a), model_train(data_p) - dists = torch.sqrt(torch.sum((out_a - out_p) ** 2, 1)) + dists = calculate_lfw_distance(out_a, out_p) distances.append(dists.data.cpu().numpy()) labels.append(label.data.cpu().numpy()) diff --git a/utils/utils_metrics.py b/utils/utils_metrics.py index 110fc2b..24fb14f 100644 --- a/utils/utils_metrics.py +++ b/utils/utils_metrics.py @@ -95,6 +95,11 @@ def calculate_val_far(threshold, dist, actual_issame): far = float(false_accept) / float(n_diff) return val, far +def calculate_lfw_distance(out_a, out_p): + out_a = out_a.reshape(out_a.size(0), -1) + out_p = out_p.reshape(out_p.size(0), -1) + return torch.sqrt(torch.sum((out_a - out_p) ** 2, 1)) + def test(test_loader, model, png_save_path, log_interval, batch_size, cuda): labels, distances = [], [] pbar = tqdm(enumerate(test_loader)) @@ -111,7 +116,7 @@ def test(test_loader, model, png_save_path, log_interval, batch_size, cuda): # 获得预测结果的距离 #--------------------------------------# out_a, out_p = model(data_a), model(data_p) - dists = torch.sqrt(torch.sum((out_a - out_p) ** 2, 1)) + dists = calculate_lfw_distance(out_a, out_p) #--------------------------------------# # 将结果添加进列表中