Skip to content
Open
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
4 changes: 2 additions & 2 deletions utils/utils_fit.py
Original file line number Diff line number Diff line change
Expand Up @@ -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):
Expand Down Expand Up @@ -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())

Expand Down
7 changes: 6 additions & 1 deletion utils/utils_metrics.py
Original file line number Diff line number Diff line change
Expand Up @@ -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))
Expand All @@ -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)

#--------------------------------------#
# 将结果添加进列表中
Expand Down